1. 注意力机制到底解决了什么问题
1.1 从翻译任务里的一个尴尬现象说起
早些年做机器翻译的时候,我遇到过一个很典型的问题:输入一句中文“我爱吃苹果”,模型翻译成英文时,前面几个词都翻得挺准,到了“苹果”这里,有时候会翻成“apple”,有时候会翻成“fruit”,甚至偶尔翻成“phone”。当时我以为是词表不够大,后来把词表扩了一倍,问题依旧。真正的原因不在词表,而在于模型在生成“apple”这个目标词的时候,并没有“回头看”源句子里对应的那个词,它只是把整句话压成了一个固定长度的向量,然后凭这个向量去猜。
这个固定长度向量就是早期序列到序列模型的瓶颈。编码器把整句“我爱吃苹果”压成一个向量,解码器再从这个向量里恢复出“I love eating apples”。句子短的时候还行,句子一长,前面信息就被后面覆盖掉了。注意力机制要干的事情非常朴素:让解码器在生成每一个词的时候,能够直接去源句子的各个位置“看一眼”,并且根据当前需要决定看哪里看得重一些。生成“apple”时,就把注意力集中在“苹果”这个词上;生成“I”时,注意力就落在“我”上。
这个思路放到今天,已经不只是翻译在用。文本分类、情感分析、问答系统、新闻摘要、甚至时序预测,只要涉及“一串输入对应一个输出”的场景,注意力机制几乎都成了标配。你如果正在入门NLP,或者已经写过几行PyTorch但一直没搞明白nn.MultiheadAttention里那几个参数到底在干嘛,那这篇内容就是写给你的。我会从最朴素的原理讲起,一路推到多头注意力、因果自注意力,最后给出一份可以直接跑起来的PyTorch代码,并且把我在实际调试中踩过的坑一并交代清楚。
1.2 注意力机制的核心直觉:加权求和
把注意力机制拆到最底层,它其实就是一个加权求和。假设你有一组输入向量,每个向量代表一个词的信息,现在你要计算某个查询(query)对应的输出,做法是:拿这个query去和每一个输入向量算一个相似度分数,把分数归一化成权重,然后对所有输入向量做加权平均。权重大的位置,说明当前query更关注那里。
用生活里的例子类比:你在图书馆找一本关于“注意力机制”的书,管理员(query)会扫一遍书架上的每本书(key),判断哪本和你的需求最匹配,匹配度高的书(value)你就重点翻,匹配度低的就略过。最终你脑子里形成的“这本书讲了什么”的印象,就是所有书内容的加权综合。这就是注意力机制的全部精髓,剩下的都是工程上的优化。
用公式表达就是:
$$ \text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V $$
这里的$Q$、$K$、$V$分别对应查询、键、值。$QK^T$算的是相似度,除以$\sqrt{d_k}$是为了防止点积结果过大导致softmax梯度消失,这个细节后面会展开讲。softmax把分数变成概率分布,最后乘$V$得到加权结果。
1.3 为什么是“缩放点积”而不是别的
相似度计算有很多种方式,比如加性注意力(additive attention)、点积注意力(dot-product attention)。早期Bahdanau那篇论文用的是加性注意力,用一个前馈网络来算分数。后来Vaswani等人在Transformer里改成了点积,原因很实际:点积可以用矩阵乘法一次性算完,GPU上跑得快。加性注意力要过一层网络,计算量大,并行度还低。
但点积有个副作用:当维度$d_k$变大时,点积的方差会随之增大,softmax的输入可能落在梯度极小的饱和区,训练会变得困难。所以加了一个缩放因子$\frac{1}{\sqrt{d_k}}$。这个缩放不是拍脑袋来的,假设$q$和$k$的每个分量都是均值0、方差1的独立随机变量,那么它们的点积$q \cdot k = \sum_{i=1}^{d_k} q_i k_i$的均值是0,方差是$d_k$。除以$\sqrt{d_k}$之后,方差重新回到1,softmax的输入就稳定了。这个推导我在第一次读论文时没在意,后来自己手写实现时发现不缩放确实训练不动,才回头把这个细节补上。
提示:如果你自己实现注意力,缩放这一步千万别省。我见过有人直接把
QK^T丢进softmax,结果loss一直不降,排查了半天才发现是这里的问题。
2. 自注意力与多头注意力:把注意力用到极致
2.1 自注意力:自己查自己
前面说的注意力是解码器查编码器,query来自一边,key和value来自另一边,这叫交叉注意力(cross-attention)。而自注意力(self-attention)是query、key、value全部来自同一个序列。听起来有点奇怪:自己查自己有什么意义?
意义在于,自注意力让序列里每个位置都能直接和所有其他位置交互。比如“苹果”这个词,在自注意力里它会去和“我”“爱”“吃”分别算相似度,从而把上下文信息融合进自己的表示。这样“苹果”的向量就不再是孤立的词向量,而是带有“被吃”这个语境信息的向量。相比之下,RNN要靠隐藏状态一步步传递信息,距离远了就衰减;自注意力一步到位,任意两个位置之间的距离都是1。
自注意力的计算过程可以拆成四步:
- 对输入序列$X$做三个线性变换,得到$Q = XW_Q$、$K = XW_K$、$V = XW_V$。
- 计算$QK^T$,得到每个位置对其他位置的相似度分数矩阵。
- 除以$\sqrt{d_k}$后过softmax,得到注意力权重。
- 用权重对$V$加权求和,得到输出。
这里$W_Q$、$W_K$、$W_V$都是可学习参数,维度通常是$d_{model} \times d_k$。输入$X$的每一行是一个词的向量,整个序列并行计算,没有循环,所以训练时可以充分利用GPU。
2.2 多头注意力:多个视角看同一句话
单个自注意力有一个局限:它只能学到一种“关注模式”。但一句话里的关系是多样的,有的位置关注语法依赖,有的位置关注语义相似,有的位置关注位置邻近。一个注意力头很难同时兼顾。
多头注意力的做法是:把$d_{model}$维的输入切成$h$份,每份维度是$d_k = d_{model} / h$,每一份独立做一次自注意力,最后把$h$个结果拼接起来,再过一层线性变换。这样每个头可以在自己的子空间里学习不同的关注模式。
举个例子,$d_{model}=512$,$h=8$,那么每个头的维度是64。8个头各自算自己的$Q$、$K$、$V$,各自得到64维的输出,拼起来又是512维。计算量和单头512维差不多,但表达能力更强。
多头注意力的代码实现里,最常见的写法是把$h$个头的$W_Q$、$W_K$、$W_V$合并成一个大矩阵,一次矩阵乘法算完再reshape。这样比循环8次快得多。PyTorch的nn.MultiheadAttention内部就是这么做的。
2.3 因果自注意力:不能偷看未来
在翻译、文本生成这类任务里,解码器生成第$t$个词时,只能看到前$t-1$个词,不能看到后面的词。但自注意力默认是全局的,每个位置都能看到所有位置,这就“作弊”了。解决办法是加一个因果掩码(causal mask),把未来位置的注意力分数设成负无穷,softmax之后这些位置的权重就变成0。
掩码通常是一个上三角矩阵,对角线及以下为0,以上为负无穷。PyTorch里可以用torch.triu生成,也可以用torch.nn.Transformer.generate_square_subsequent_mask直接生成。这个掩码在训练时必须加,推理时因为是一个词一个词生成的,天然看不到未来,但为了代码统一,通常也会加上。
注意:因果掩码加的位置是在softmax之前,加在缩放之后。顺序是:先算$QK^T$,再除以$\sqrt{d_k}$,再加掩码,最后softmax。顺序错了结果就不对。
3. PyTorch实战:从零实现多头注意力
3.1 环境准备与版本对应
动手之前先把环境弄干净。PyTorch的版本和Python版本、CUDA版本之间有对应关系,装错了会各种报错。我整理了一份常见组合,供你参考:
| PyTorch版本 | 推荐Python版本 | CUDA版本 | 安装命令示例 |
|---|---|---|---|
| 2.1+ | 3.9 - 3.11 | 11.8 / 12.1 | pip install torch --index-url https://download.pytorch.org/whl/cu118 |
| 1.13 | 3.8 - 3.10 | 11.6 / 11.7 | conda install pytorch pytorch-cuda=11.7 -c pytorch -c nvidia |
| 1.11 | 3.7 - 3.9 | 10.2 / 11.3 | pip install torch==1.11.0+cu113 |
如果你用的是Windows下的WSL环境,建议直接在WSL里装Linux版的PyTorch,不要用Windows版再映射,性能损失明显。AMD显卡的话,ROCm版本的PyTorch支持有限,部分算子可能回退到CPU,训练速度会打折扣,这一点要有心理准备。
安装完成后用下面这段代码验证:
import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else "CPU only")如果cuda.is_available()返回False,先检查驱动和CUDA版本是否匹配,再检查是不是装成了CPU版。
3.2 手写单头自注意力
先从最简单的单头自注意力开始,把每一步都写清楚,方便理解。
import torch import torch.nn as nn import torch.nn.functional as F import math class SingleHeadAttention(nn.Module): def __init__(self, d_model, d_k): super().__init__() self.d_k = d_k self.W_q = nn.Linear(d_model, d_k) self.W_k = nn.Linear(d_model, d_k) self.W_v = nn.Linear(d_model, d_k) def forward(self, x, mask=None): # x: (batch, seq_len, d_model) Q = self.W_q(x) # (batch, seq_len, d_k) K = self.W_k(x) V = self.W_v(x) scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) # scores: (batch, seq_len, seq_len) if mask is not None: scores = scores.masked_fill(mask == 0, float('-inf')) attn = F.softmax(scores, dim=-1) output = torch.matmul(attn, V) # (batch, seq_len, d_k) return output, attn这段代码里几个关键点:transpose(-2, -1)是把最后两维转置,得到$K^T$;masked_fill把掩码为0的位置填成负无穷;softmax在最后一维做,也就是每个query对所有key归一化。
测试一下:
batch, seq_len, d_model, d_k = 2, 5, 16, 8 x = torch.randn(batch, seq_len, d_model) attn_layer = SingleHeadAttention(d_model, d_k) out, attn_weights = attn_layer(x) print(out.shape) # torch.Size([2, 5, 8]) print(attn_weights.shape) # torch.Size([2, 5, 5])注意力权重矩阵的每一行加起来应该是1,可以验证一下:
print(attn_weights.sum(dim=-1))3.3 多头注意力的完整实现
单头理解之后,多头就是把它并行化。下面这份实现把$h$个头的投影合并成一个大矩阵,效率更高。
class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads, dropout=0.1): super().__init__() assert d_model % num_heads == 0, "d_model必须能被num_heads整除" self.d_model = d_model self.num_heads = num_heads self.d_k = d_model // num_heads self.W_q = nn.Linear(d_model, d_model) self.W_k = nn.Linear(d_model, d_model) self.W_v = nn.Linear(d_model, d_model) self.W_o = nn.Linear(d_model, d_model) self.dropout = nn.Dropout(dropout) def forward(self, query, key, value, mask=None): batch_size = query.size(0) # 线性投影并拆分成多头 Q = self.W_q(query).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) K = self.W_k(key).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) V = self.W_v(value).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) # Q, K, V: (batch, num_heads, seq_len, d_k) scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores = scores.masked_fill(mask == 0, float('-inf')) attn = F.softmax(scores, dim=-1) attn = self.dropout(attn) context = torch.matmul(attn, V) # (batch, num_heads, seq_len, d_k) context = context.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model) output = self.W_o(context) return output, attnview和transpose的顺序很关键。先把d_model拆成num_heads × d_k,再把num_heads这一维换到前面,这样每个头的数据就是连续的。最后拼接时反过来操作,注意要加.contiguous(),否则view会报错。
3.4 因果掩码的生成与使用
因果掩码是一个下三角矩阵,对角线及以下为1,以上为0。生成方式:
def generate_causal_mask(seq_len, device='cpu'): mask = torch.tril(torch.ones(seq_len, seq_len, device=device)) return mask # (seq_len, seq_len)使用时需要扩展到(batch, num_heads, seq_len, seq_len),因为masked_fill要求形状能广播:
seq_len = 5 mask = generate_causal_mask(seq_len) mask = mask.unsqueeze(0).unsqueeze(0) # (1, 1, seq_len, seq_len)测试一下因果掩码的效果:
mha = MultiHeadAttention(d_model=16, num_heads=4) x = torch.randn(2, 5, 16) mask = generate_causal_mask(5).unsqueeze(0).unsqueeze(0) out, attn = mha(x, x, x, mask=mask) print(attn[0, 0]) # 打印第一个样本第一个头的注意力矩阵你会看到注意力矩阵是下三角的,每个位置只关注自己和前面的位置,未来位置权重为0。
4. 调试与排查:那些文档里不会写的问题
4.1 常见报错与解决思路
实际写代码时,报错是家常便饭。我整理了一份速查表,覆盖了大部分高频问题:
| 报错信息 | 可能原因 | 解决方法 |
|---|---|---|
RuntimeError: The size of tensor a must match... | mask形状和scores不匹配 | 检查mask是否扩展到(batch, heads, seq, seq) |
RuntimeError: view size is not compatible... | transpose后没加contiguous | 在view前加.contiguous() |
CUDA out of memory | batch或seq_len太大 | 减小batch,或用梯度累积 |
loss不下降 | 忘记缩放或mask位置错误 | 检查是否除以sqrt(d_k),mask是否在softmax前 |
nan in loss | mask全为负无穷导致softmax输出nan | 确保每行至少有一个位置可见 |
其中nan这个问题特别隐蔽。如果某个query对所有key都被mask掉,softmax的输入全是负无穷,输出就是nan。因果掩码不会出现这种情况,因为对角线总是可见的。但如果你自己写padding mask,把padding位置全mask掉,而某个query恰好全是padding,就会出问题。解决办法是给mask加一个极小值而不是负无穷,或者确保每行至少有一个可见位置。
4.2 注意力权重的可视化排查
训练不收敛的时候,把注意力权重打印出来看看,往往能发现端倪。正常情况下,注意力权重应该是一个比较分散的分布,如果某个头几乎全部权重都集中在一个位置,说明这个头可能“死”了,没有学到有效模式。
import matplotlib.pyplot as plt def plot_attention(attn_weights, head_idx=0, sample_idx=0): # attn_weights: (batch, heads, seq_len, seq_len) weights = attn_weights[sample_idx, head_idx].detach().cpu().numpy() plt.imshow(weights, cmap='viridis') plt.colorbar() plt.title(f'Head {head_idx} Attention') plt.xlabel('Key position') plt.ylabel('Query position') plt.show()我一般会在训练初期每隔几个epoch画一次,观察注意力模式有没有从均匀分布逐渐变得有结构。如果一直是均匀分布,可能是学习率太小或者初始化有问题。
4.3 性能优化:让注意力跑得更快
注意力机制的计算复杂度是$O(n^2 d)$,序列一长就吃不消。几个实用的优化方向:
- 混合精度训练:用
torch.cuda.amp把计算转成float16,显存占用减半,速度提升明显。注意softmax部分要保持float32,否则容易溢出。 - 梯度检查点:用
torch.utils.checkpoint把中间激活值不保存,反向传播时重算,用时间换显存。 - Flash Attention:PyTorch 2.0之后内置了
scaled_dot_product_attention,底层用了Flash Attention的实现,速度和显存都有大幅优化。如果你的PyTorch版本够新,直接用这个函数替代手写实现。
from torch.nn.functional import scaled_dot_product_attention # 替代手写的scores计算和softmax output = scaled_dot_product_attention(Q, K, V, attn_mask=mask)这个函数会自动选择最优的注意力实现,在支持的硬件上能快好几倍。我实测下来,在A100上比手写版本快3倍左右,显存也省了不少。
4.4 几个容易忽略的细节
初始化:注意力层的权重初始化对训练稳定性影响很大。nn.Linear默认用Kaiming初始化,一般够用。但如果训练初期loss震荡厉害,可以试试把$W_Q$、$W_K$的初始化标准差调小一点。
Dropout的位置:注意力dropout加在softmax之后、乘$V$之前,这是Transformer论文里的做法。也有实现加在注意力权重上,效果差不多。但不要加在$Q$、$K$、$V$上,那样会破坏相似度计算。
残差连接:多头注意力的输出通常会加一个残差连接,再过一个LayerNorm。残差连接让梯度能直接回传,LayerNorm稳定分布。这两个组件虽然简单,但少了任何一个,深层网络都很难训练。
位置编码:自注意力本身没有位置概念,打乱输入顺序结果不变。所以需要额外加位置编码。正弦位置编码是原始Transformer的做法,现在也有很多用可学习的位置嵌入。如果你的任务对位置敏感(比如翻译),位置编码不能省。
5. 从注意力到完整模型:组装一个迷你Transformer
5.1 编码器层的搭建
有了多头注意力,就可以搭一个完整的编码器层了。一个标准的编码器层包含:多头自注意力、残差连接、LayerNorm、前馈网络、再一个残差连接和LayerNorm。
class EncoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout=0.1): super().__init__() self.self_attn = MultiHeadAttention(d_model, num_heads, dropout) self.feed_forward = nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Dropout(dropout), nn.Linear(d_ff, d_model) ) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.dropout1 = nn.Dropout(dropout) self.dropout2 = nn.Dropout(dropout) def forward(self, x, mask=None): attn_out, _ = self.self_attn(x, x, x, mask) x = self.norm1(x + self.dropout1(attn_out)) ff_out = self.feed_forward(x) x = self.norm2(x + self.dropout2(ff_out)) return x前馈网络的隐藏层维度$d_{ff}$通常是$d_{model}$的4倍,这是Transformer论文里的设定。这个比例不是随便定的,4倍能在表达能力和计算量之间取得比较好的平衡。
5.2 位置编码的实现
位置编码的公式是:
$$ PE_{(pos, 2i)} = \sin(pos / 10000^{2i/d_{model}}) $$
$$ PE_{(pos, 2i+1)} = \cos(pos / 10000^{2i/d_{model}}) $$
实现起来很简单:
class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=5000): super().__init__() pe = torch.zeros(max_len, d_model) position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1) div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) pe = pe.unsqueeze(0) # (1, max_len, d_model) self.register_buffer('pe', pe) def forward(self, x): return x + self.pe[:, :x.size(1), :]用register_buffer把位置编码注册成buffer,这样它不会参与梯度更新,但会跟着模型一起保存和加载。
5.3 一个完整的文本分类模型
把编码器层堆几层,再加一个分类头,就是一个能用的文本分类模型:
class TextClassifier(nn.Module): def __init__(self, vocab_size, d_model, num_heads, d_ff, num_layers, num_classes, max_len=512, dropout=0.1): super().__init__() self.embedding = nn.Embedding(vocab_size, d_model) self.pos_encoding = PositionalEncoding(d_model, max_len) self.layers = nn.ModuleList([ EncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) self.classifier = nn.Linear(d_model, num_classes) self.dropout = nn.Dropout(dropout) def forward(self, x, mask=None): x = self.embedding(x) x = self.pos_encoding(x) for layer in self.layers: x = layer(x, mask) # 用平均池化代替[CLS] token x = x.mean(dim=1) x = self.dropout(x) return self.classifier(x)这个模型可以直接拿去做中文新闻分类。输入是token id序列,输出是类别logits。训练时用交叉熵损失,优化器用AdamW,学习率3e-4左右,配合warmup效果更稳。
5.4 训练循环与注意事项
训练循环本身不复杂,但有几个细节值得注意:
def train_epoch(model, dataloader, optimizer, criterion, device): model.train() total_loss = 0 for batch in dataloader: input_ids = batch['input_ids'].to(device) labels = batch['labels'].to(device) optimizer.zero_grad() logits = model(input_ids) loss = criterion(logits, labels) loss.backward() # 梯度裁剪,防止梯度爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() total_loss += loss.item() return total_loss / len(dataloader)梯度裁剪这一步在Transformer训练里几乎是必须的。注意力层的梯度有时候会突然变得很大,不裁剪的话loss会直接飞掉。max_norm=1.0是个比较安全的默认值。
学习率调度也很关键。Transformer论文里用的是warmup加逆平方根衰减,前4000步线性增加,之后按步数平方根倒数衰减。PyTorch里可以用LambdaLR实现:
def lr_lambda(step): d_model = 512 warmup_steps = 4000 return d_model ** (-0.5) * min(step ** (-0.5), step * warmup_steps ** (-1.5)) scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)这个调度策略能让训练初期稳定,后期收敛快。我试过直接用固定学习率,前期容易震荡,后期又降不下来,效果差不少。
6. 注意力机制的变体与扩展方向
6.1 通道注意力与空间注意力
注意力机制不只用在NLP,计算机视觉里也大量使用。通道注意力(如SE模块)是给每个通道学一个权重,重要的通道放大,不重要的抑制。空间注意力(如CBAM)是在特征图的每个位置上算权重,告诉网络“看哪里”。这两种注意力和NLP里的自注意力思路一致,只是作用维度不同。
如果你做多模态任务,比如图文匹配,可以把文本的自注意力和图像的通道注意力结合起来,让模型同时关注“哪些词重要”和“哪些区域重要”。
6.2 时序注意力
时间序列预测里,注意力机制用来捕捉不同时间步之间的依赖。和NLP不同的是,时序数据没有明确的“词”边界,而且往往需要处理多变量。做法通常是把每个时间步的特征向量当作一个token,做自注意力,再取最后一个时间步的输出做预测。因果掩码在这里同样重要,因为预测未来时不能看到未来的数据。
6.3 高效注意力的几个方向
标准注意力的$O(n^2)$复杂度在长序列上是个硬伤。几个主流的优化方向:
- 稀疏注意力:只计算部分位置的注意力,比如局部窗口注意力、膨胀注意力。
- 线性注意力:用核函数近似softmax,把复杂度降到$O(n)$。
- 低秩近似:把注意力矩阵分解成低秩矩阵的乘积。
- Flash Attention:不改变数学结果,通过分块计算和显存优化提升速度,是目前最实用的方案。
这些方法各有取舍,稀疏注意力实现简单但可能丢失全局信息,线性注意力理论优雅但实际效果有时不如标准注意力。选哪个取决于你的任务对精度和速度的要求。
6.4 我个人的选型建议
如果你刚开始做NLP项目,我的建议是:先用标准多头注意力把baseline跑通,再考虑优化。很多项目序列长度也就一两百,标准注意力完全够用,过早引入复杂变体反而增加调试成本。等baseline稳定了,发现推理速度是瓶颈,再针对性地上Flash Attention或者稀疏注意力。
另外,PyTorch 2.0之后的scaled_dot_product_attention已经自动做了很多优化,优先用它,不要自己手写。手写版本除了教学目的,生产环境里没有优势。
7. 写在最后的一些实操体会
注意力机制从2014年提出到现在,已经成了深度学习的基石之一。但我在带新人的时候发现,很多人能背出公式,却说不清楚$Q$、$K$、$V$各自代表什么,也不知道为什么要缩放。这其实是因为跳过了“自己手写一遍”这一步。你只要亲手实现一次单头注意力,再扩展到多头,把掩码加上去,把训练跑起来,那些公式自然就活了。
调试注意力模型时,我最常用的手段是打印注意力权重矩阵。它就像模型的“注意力地图”,能直观告诉你模型在看哪里。如果发现某个头始终关注[SEP]或者padding位置,那这个头基本没学到东西,可以考虑减掉。如果所有头都关注同一个位置,说明多头没有起到多视角的作用,可能需要调整初始化或者增加正则。
最后分享一个我踩过的坑:有一次做中文新闻分类,模型在验证集上表现很好,但上线后效果差很多。排查后发现是padding mask的问题。训练时batch内序列长度对齐用了padding,但推理时单条输入没有padding,mask逻辑不一致导致注意力分布偏移。后来统一了mask生成逻辑,问题才解决。这个教训是:训练和推理的预处理逻辑必须完全一致,尤其是mask相关的部分。
注意力机制的内容远不止这些,从自注意力到交叉注意力,从编码器到解码器,从文本到图像到语音,它的应用还在不断扩展。但核心思想始终没变:让模型学会“看哪里”。把这个思想吃透,剩下的都是工程问题。