从零实现PyTorch多头注意力:原理、代码与调试避坑指南
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_{i1}^{d_k} q_i k_i$的均值是0方差是$d_k$。除以$\sqrt{d_k}$之后方差重新回到1softmax的输入就稳定了。这个推导我在第一次读论文时没在意后来自己手写实现时发现不缩放确实训练不动才回头把这个细节补上。提示如果你自己实现注意力缩放这一步千万别省。我见过有人直接把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$$h8$那么每个头的维度是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.13.9 - 3.1111.8 / 12.1pip install torch --index-url https://download.pytorch.org/whl/cu1181.133.8 - 3.1011.6 / 11.7conda install pytorch pytorch-cuda11.7 -c pytorch -c nvidia1.113.7 - 3.910.2 / 11.3pip install torch1.11.0cu113如果你用的是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, maskNone): # 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, dropout0.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, maskNone): 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, devicecpu): mask torch.tril(torch.ones(seq_len, seq_len, devicedevice)) 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_model16, num_heads4) x torch.randn(2, 5, 16) mask generate_causal_mask(5).unsqueeze(0).unsqueeze(0) out, attn mha(x, x, x, maskmask) 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 memorybatch或seq_len太大减小batch或用梯度累积loss不下降忘记缩放或mask位置错误检查是否除以sqrt(d_k)mask是否在softmax前nan in lossmask全为负无穷导致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_idx0, sample_idx0): # attn_weights: (batch, heads, seq_len, seq_len) weights attn_weights[sample_idx, head_idx].detach().cpu().numpy() plt.imshow(weights, cmapviridis) plt.colorbar() plt.title(fHead {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 AttentionPyTorch 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_maskmask)这个函数会自动选择最优的注意力实现在支持的硬件上能快好几倍。我实测下来在A100上比手写版本快3倍左右显存也省了不少。4.4 几个容易忽略的细节初始化注意力层的权重初始化对训练稳定性影响很大。nn.Linear默认用Kaiming初始化一般够用。但如果训练初期loss震荡厉害可以试试把$W_Q$、$W_K$的初始化标准差调小一点。Dropout的位置注意力dropout加在softmax之后、乘$V$之前这是Transformer论文里的做法。也有实现加在注意力权重上效果差不多。但不要加在$Q$、$K$、$V$上那样会破坏相似度计算。残差连接多头注意力的输出通常会加一个残差连接再过一个LayerNorm。残差连接让梯度能直接回传LayerNorm稳定分布。这两个组件虽然简单但少了任何一个深层网络都很难训练。位置编码自注意力本身没有位置概念打乱输入顺序结果不变。所以需要额外加位置编码。正弦位置编码是原始Transformer的做法现在也有很多用可学习的位置嵌入。如果你的任务对位置敏感比如翻译位置编码不能省。5. 从注意力到完整模型组装一个迷你Transformer5.1 编码器层的搭建有了多头注意力就可以搭一个完整的编码器层了。一个标准的编码器层包含多头自注意力、残差连接、LayerNorm、前馈网络、再一个残差连接和LayerNorm。class EncoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout0.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, maskNone): 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, 2i1)} \cos(pos / 10000^{2i/d_{model}}) $$实现起来很简单class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len5000): super().__init__() pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.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_len512, dropout0.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, maskNone): x self.embedding(x) x self.pos_encoding(x) for layer in self.layers: x layer(x, mask) # 用平均池化代替[CLS] token x x.mean(dim1) 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_norm1.0) optimizer.step() total_loss loss.item() return total_loss / len(dataloader)梯度裁剪这一步在Transformer训练里几乎是必须的。注意力层的梯度有时候会突然变得很大不裁剪的话loss会直接飞掉。max_norm1.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但推理时单条输入没有paddingmask逻辑不一致导致注意力分布偏移。后来统一了mask生成逻辑问题才解决。这个教训是训练和推理的预处理逻辑必须完全一致尤其是mask相关的部分。注意力机制的内容远不止这些从自注意力到交叉注意力从编码器到解码器从文本到图像到语音它的应用还在不断扩展。但核心思想始终没变让模型学会“看哪里”。把这个思想吃透剩下的都是工程问题。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →