从全局走到重点:AI注意力机制笔记
测试时我们发现:当关键信息出现在 token 序列的前 20% 时,准确率会比出现在后 20% 低 15 个百分点。这篇笔记就是围绕这个偏差,记录我们怎么从 LSTM 一路试到 Self-Attention。
代码搜索项目里卡住的点
当时我们在做一个代码搜索项目,需要理解整个函数的语义来匹配查询。代码函数动辄几百行,前面定义的变量、函数名、注释往往在后面才用到。
最开始用的是标准的 LSTM,效果还行但总觉得"差点意思"。比如函数开头有个 user_id 变量,中间一百行各种逻辑,最后才用到这个变量做过滤。LSTM 的隐藏状态在传递过程中不断被新的信息覆盖,到了最后那一步,前面 user_id 的记忆已经淡得差不多了。
这个问题在长文本上更明显。测试时我们发现:当关键信息出现在 token 序列的前 20% 时,准确率会比出现在后 20% 低 15 个百分点。这不是随机误差,是系统性的位置偏差。
我当时的判断是:模型记性不算差,问题是它分不清该记什么、该丢什么。注意力机制想解决的正是这个——给它一把"聚光灯",让它知道哪里才是重点。
为什么选择注意力机制
注意力机制的核心想法很直白:生成每个输出时,对整个输入序列做加权汇聚,不同位置权重不同,不完全依赖上一步的隐藏状态。
举个不恰当的例子:你在看一篇论文时,不会每个词都花同样的精力。你会对标题、公式、结论多看几眼,对连接词、例句一扫而过。注意力机制就是让模型学会这种"扫一眼但知道重点在哪"的能力。
从工程角度看,它还有一个好处:解耦了信息位置和距离。不管关键信息离当前处理位置多远,只要给它足够的权重,模型就能"看"到。这比 RNN 那种"步步传递"的方式更适合长序列。
但当时我还有个顾虑:注意力机制的计算复杂度是 O(n²),当序列长度 n 增加时,计算量和内存占用会二次增长。在代码函数这种动辄上千 token 的场景里,会不会跑不动?
这个问题只能试出来才知道。
从 Seq2Seq Attention 开始
最经典的是 Bahdanau Attention,用在 Seq2Seq 的 decoder 端。每一步解码时,先计算当前状态和所有 encoder 隐藏状态的相似度,得到一组权重,再把 encoder 的输出按权重加权求和,作为这一步的上下文向量。
这个过程可以拆成几步:计算能量分数、归一化得到权重、加权求和得到上下文。
我先写了个简单的版本测试效果:
import torch
import torch.nn as nn
import torch.nn.functional as F
class BahdanauAttention(nn.Module):
def __init__(self, hidden_size):
super().__init__()
self.attn = nn.Linear(hidden_size * 2, hidden_size)
self.v = nn.Linear(hidden_size, 1, bias=False)
def forward(self, hidden, encoder_outputs):
# hidden: (batch, hidden_size)
# encoder_outputs: (batch, seq_len, hidden_size)
batch_size = encoder_outputs.size(0)
seq_len = encoder_outputs.size(1)
# 扩展 hidden 以匹配 seq_len
hidden = hidden.unsqueeze(1).repeat(1, seq_len, 1)
# 计算注意力分数
energy = torch.tanh(self.attn(torch.cat((hidden, encoder_outputs), dim=2)))
scores = self.v(energy).squeeze(2)
# 归一化得到注意力权重
attn_weights = F.softmax(scores, dim=1)
# 加权求和
context = torch.bmm(attn_weights.unsqueeze(1), encoder_outputs).squeeze(1)
return context, attn_weights
测试时发现,这个版本的注意力确实让模型"记住"了前面出现的关键词。有个函数开头定义了 RETRY_MAX=3,中间几百行逻辑,最后用这个常量做循环边界。加上注意力后,模型能准确识别出这个函数是"重试逻辑"。
但问题也很明显:还是不够快。每次解码都要计算整个序列的注意力分数,当序列长度超过 500 时,推理延迟已经到了 200ms,这在实时搜索场景里是不可接受的。
而且还有个更深的问题:这种注意力只在 decoder 端起作用,encoder 内部还是"各扫门前雪"。encoder 在编码长序列时,前面的信息在传递过程中还是会损失。
转向 Self-Attention
这时候 Transformer 出现了,直接把注意力机制搬到了 encoder 内部,而且做得更彻底——每个位置都能看到整个序列的所有位置,然后自己决定哪些信息该关注。
Self-Attention 的数学形式比 Seq2Seq Attention 稍微复杂一点,但核心思想差不多:
- 对每个位置,生成三个向量:Query(查询)、Key(键)、Value(值)
- 用 Query 和所有位置的 Key 计算相似度,得到注意力分数
- 用这些分数对 Value 做加权求和,得到该位置的输出
这个过程可以理解成:每个位置拿着自己的 Query 去"询问"其他位置,看谁的 Key 跟自己最相关,然后收集这些相关位置的 Value 作为自己的信息。
Multi-Head Self-Attention 把这个过程并行化多份,每份在不同的"子空间"里工作,最后把结果拼接起来。
class SelfAttention(nn.Module):
def __init__(self, embed_size, heads):
super().__init__()
self.embed_size = embed_size
self.heads = heads
self.head_dim = embed_size // heads
assert (self.heads * self.head_dim == embed_size), "Embed size needs to be divisible by heads"
self.values = nn.Linear(self.head_dim, self.head_dim, bias=False)
self.keys = nn.Linear(self.head_dim, self.head_dim, bias=False)
self.queries = nn.Linear(self.head_dim, self.head_dim, bias=False)
self.fc_out = nn.Linear(heads * self.head_dim, embed_size)
def forward(self, values, keys, query, mask):
N = query.shape[0]
value_len, key_len, query_len = values.shape[1], keys.shape[1], query.shape[1]
# Split embedding into self.heads pieces
values = values.reshape(N, value_len, self.heads, self.head_dim)
keys = keys.reshape(N, key_len, self.heads, self.head_dim)
query = query.reshape(N, query_len, self.heads, self.head_dim)
values = self.values(values)
keys = self.keys(keys)
queries = self.queries(query)
# Calculate attention scores
energy = torch.einsum("nqhd,nkhd->nhqk", [queries, keys])
# Masking for padded positions
if mask is not None:
energy = energy.masked_fill(mask == 0, float("-1e20"))
attention = torch.softmax(energy / (self.embed_size ** (1/2)), dim=3)
# Weighted sum of values
out = torch.einsum("nhql,nlhd->nqhd", [attention, values])
out = out.reshape(N, query_len, self.heads * self.head_dim)
return self.fc_out(attention)
用 Multi-Head Self-Attention 的好处是:让模型在不同的"子空间"里关注不同类型的信息。比如在代码里,某个 head 可能关注变量名和类型,另一个 head 可能关注控制流结构,还有一个 head 可能关注注释和文档。
踩坑实录
真正跑起来的时候,问题比预想的要多。
内存爆炸
第一个问题是内存。长序列的注意力矩阵是 n×n 的,当序列长度到 1024 时,这个矩阵已经占了几百 MB。在批量推理时,显存直接爆掉。
当时的临时方案是梯度检查点(gradient checkpointing):在前向传播时不保存中间结果,等到反向传播时再重新计算。这样用计算换内存,但训练速度慢了三倍。
后来试了几个长文本优化的方案:
分块注意力(Blockwise Attention):把序列分成小块,每块内部做 full attention,块之间只关注少量关键位置。这在一定程度上保留了局部信息,同时降低了计算量。
稀疏注意力(Sparse Attention):每个位置只关注固定数量的其他位置,比如最近的几个和几个固定间隔的位置。这把复杂度从 O(n²) 降到了 O(n),但效果会有损失。
Linformer:用低秩矩阵近似原来的长序列,把 n×n 的矩阵压缩成 n×k,k 是一个较小的常数。这个方案在长序列上效果不错,但短序列反而不如 full attention。
最终我们的选择是:在训练时用 blockwise attention,推理时对较短的序列(< 512)用 full attention,对长序列用 sparse attention。这样在效果和性能之间做了个折中。
数值不稳定
第二个问题是数值稳定性。在计算注意力分数时,softmax 内部的值如果太大,梯度会变得极小,几乎消失;如果太小,又会导致梯度爆炸。
标准做法是在 softmax 前除以 √d_k(d_k 是 key 的维度),这是为了让点积结果的方差保持在合理范围。
# 原始写法
scores = torch.matmul(query, key.transpose(-2, -1))
attention = torch.softmax(scores, dim=-1)
# 加缩放
scores = torch.matmul(query, key.transpose(-2, -1)) / (self.head_dim ** 0.5)
attention = torch.softmax(scores, dim=-1)
但即使加了缩放,在某些极端情况下还是有问题。比如当某个位置的所有注意力分数都特别大时,softmax 输出会趋近于 one-hot,梯度几乎消失。这会导致某些 head 的注意力变得"僵硬",只关注固定的几个位置。
我们的解决方案是加了个 temperature 参数,在实际应用中动态调整:
def attention_with_temperature(q, k, v, temperature=1.0):
scores = torch.matmul(q, k.transpose(-2, -1)) / (self.head_dim ** 0.5)
attention = torch.softmax(scores / temperature, dim=-1)
return torch.matmul(attention, v)
训练初期 temperature 设得大一点,让注意力更平滑;训练后期逐渐降低 temperature,让注意力更集中。这在一定程度上缓解了注意力僵硬的问题。
无法忽略噪声
第三个问题是注意力"太聪明"了——它会关注序列中的所有信息,包括噪声。
比如代码里有很多 boilerplate,像重复的导入、错误处理的模板代码、大量相似的 if-else 结构。这些信息对理解函数核心逻辑帮助不大,但模型还是会花注意力去看。
这在可视化注意力权重时特别明显:某些 head 会给大量的错误处理语句分配高权重,而真正重要的业务逻辑反而被稀释了。
我们试过几种方案:
预训练位置编码:在位置编码中加入先验信息,让模型倾向于关注函数的开头和结尾,因为这些位置通常是关键信息所在。
结构感知掩码:对于代码这种有明确结构的文本,利用 AST(抽象语法树)信息构造掩码,让某些 head 只关注语法树的特定层级,比如只关注函数定义层,只关注变量引用层。
对比学习:在训练时让模型区分"重要位置"和"噪声位置",通过对比学习让注意力更聚焦。
最终结构感知掩码的效果最好。它既利用了代码的结构特性,又不需要改变模型架构,实现成本相对较低。
结果与思考
折腾一圈,数字大概是这样:
- 准确率相比原始 LSTM 提升了 8 个百分点,长文本(> 500 tokens)上能到 15 个百分点
- 推理延迟相比 full attention 降了 40%,普通 GPU 上能压到 100ms 以内
- 训练成本大约是原来的 2.5 倍,这块没省下来
回过头看,前面那个位置偏差确实缓解了。模型能主动回头找 user_id、RETRY_MAX 这类变量,不再完全依赖 RNN 一步步传。
注意力机制也没那么玄。本质上是用更多计算和显存,换关键信息的"可追溯性"——不用一步步传,随时能扫一眼全序列。代价也摆在那:
- 相似度是静态算的,时序上动态变化的信息,attention 只能给一个固定权重
- 噪声位置照样参与计算,长文本里 boilerplate 会分走权重
- O(n²) 复杂度,整本书、整个代码库这种长度还是扛不住
Longformer、BigBird、Performer 这些长序列变体,还有 Memory-Augmented Transformer,都是在这个瓶颈上长出来的。我们这次用的 blockwise + sparse 折中,算是在效果和成本之间摸出来的一个点。
如果序列不长、特征工程已经到位,LSTM 未必输给 attention。反过来,关键信息散在几百 token 里、RNN 明显"忘事"的场景,attention 值得试。问题还是那句:你到底卡在哪。
版权声明: 本文首发于 指尖魔法屋-从全局走到重点:AI注意力机制笔记(https://blog.thinkmoon.cn/post/341-ai-attention-mechanism-global-focus-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。