
你有没有过这种体验读 Transformer 源码时单头注意力一眼就能看懂——Q 和 K 做点积softmax 得到权重再和 V 加权求和。可代码里一旦出现多头各种 view、transpose、reshape 扑面而来你就开始怀疑自己是不是真的理解了。这种困惑不是你的问题。多头注意力是 Transformer 里“概念一句话、实现好几层”的典型代表。绝大多数人对它的理解停留在“就是多头多个头并行”但真正到了上手实现、写训练代码、调精度、部署推理的时候才会发现其中的维度细节、掩码处理、数值稳定性问题都藏着坑。本文从核心原理讲到完整代码再讲到部署阶段容易踩的数值问题一次把多头注意力讲透。读完你会得到三样东西不会忘的概念框架、可以直接跑的 PyTorch 实现、以及实际工程里替换和调试多头注意力时的排查清单。1. 这篇文章真正要解决的问题很多深度学习初学者学到 Transformer 时会经历这样三个阶段第一个阶段看论文知道注意力公式是softmax(QK^T / sqrt(d_k))V但不知道为什么 Q、K、V 要拆成多份。第二个阶段看别人的源码发现代码里全是维度变换看不懂view、transpose到底在做什么。第三个阶段自己写模型发现训练时 loss 不降、推理时结果不对却又不知道问题出在多头注意力的哪一个环节。这篇文章要解决的就是这三个阶段的问题而且不打算停留在“会用”的层面。我们会讨论一个经常被忽略的技术判断多头注意力真正提升的不是“能力上限”而是“优化的效率”。单头注意力在理论上也可以拟合复杂函数但训练不稳定、收敛慢多头注意力通过多个子空间并行让不同的头分别学习不同类型的关系从而让梯度更容易在不同“语义维度”上传播。下面这些读者最适合精读本文正在学习 Transformer 源码看到多头注意力就想跳过的人。需要自己实现或修改注意力机制的算法工程师。在部署场景中遇到 fp16、bf16 精度问题想理解数值根源的推理工程师。如果只是想调用现成的nn.MultiheadAttention本文也能帮上忙因为我们会专门讲官方接口的输入输出约定和最容易踩的坑。2. 注意力机制的核心原理Q、K、V 与缩放点积2.1 Q、K、V 到底是什么学习注意力机制第一个绕不开的问题是Query、Key、Value 这三个词到底是什么意思。用一个最常见的搜索场景来理解。你在百度搜索“多头注意力”时Query是你输入到搜索框里的查询词代表“我想关注什么”。Key是每篇网页的标题、标签代表“我有什么可被匹配的信息”。Value是网页的正文内容代表“一旦匹配上了实际要提取的信息”。注意力机制做的事情就是把 Query 和每个 Key 做相似度计算用 softmax 把相似度转成权重再用权重对 Value 做加权求和。在实际代码里Q、K、V 通常是同一个输入经过三个不同的线性层得到的。因为三个线性层的参数不同所以同一份输入会被映射到三个不同的空间中分别承担“查询”“键”“值”的角色。2.2 缩放点积注意力的公式与细节单头缩放点积注意力的公式如下Attention(Q, K, V) softmax(Q K^T / sqrt(d_k)) V其中Q 的形状是(..., T_Q, d_k)K 的形状是(..., T_K, d_k)V 的形状是(..., T_V, d_v)通常d_k d_v这里初学者最容易问的问题是为什么除以sqrt(d_k)答案是为了稳定 softmax 的梯度。当向量维度 d_k 增大时Q 和 K 里每个元素都是均值为 0、方差为 1 的随机变量时它们的点积结果的方差会变成 d_k。比如 d_k 64 时点积值的方差是 64标准差是 8。这意味着很多点积值会落在 [-8, 8] 甚至更远范围经过指数运算后softmax 的输入会非常大导致 softmax 函数进入饱和区梯度变得非常小模型难以训练。除以sqrt(d_k)之后点积结果的方差恢复到 1 左右softmax 的输入保持在一个合理范围内梯度可以稳定回传。2.3 单头注意力的局限单头注意力并非不能工作但在面对复杂语言结构时它存在一个明显问题一次注意力只计算一种加权求和关系。例如在句子“小明把苹果放在桌上然后拿起它咬了一口”中“它”需要关注“苹果”来获得语义指代关系同时也需要关注“小明”来理解谁是动作主体。单头注意力只能给出一个归一化后的权重分布它在权衡这些不同关系时会互相干扰。这不是说单头注意力无法表示复杂关系而是它把多种关系压缩进一组权重里模型的表示能力和训练效率都会受到限制。3. 多头注意力到底“多”在哪里3.1 多头的本质多套投影矩阵多头注意力的思想非常简单与其只用一组 Q、K、V 线性投影不如用多组投影每组投影计算一次注意力最后把结果拼起来再投影一次。用公式表示head_i Attention(Q W_i^Q, K W_i^K, V W_i^V) MultiHead(Q, K, V) Concat(head_1, ..., head_h) W^O这里 W_i^Q、W_i^K、W_i^V 是第 i 个头的投影矩阵W^O 是输出的投影矩阵。注意一个关键设计每个头使用的 d_k 并不是完整的 d_model而是d_k d_model / h。所以多头注意力总体的参数量和计算量和单头完整注意力几乎一样。这是一个经常被误解的点多头不是把注意力重复堆几遍而是在同样的参数量预算下把模型切分成多个更小维度的子空间。3.2 为什么多头比单头更有效从原理上有两个层面的解释第一层是“表示空间多样性”。不同的投影矩阵会把输入映射到不同的特征子空间每个头关注的模式因此不同。有的头可能倾向于关注相邻词之间的局部依赖有的头可能学习到长距离的句法关系有的头可能主要负责指代消解。第二层是“梯度传播路径的多样性”。多个头给梯度提供了多条并行传播的通道。即便某一个头的梯度被某种饱和效应削弱其他头依然可以提供有效的学习信号。需要强调的是不同头的“分工”并不是人为设定出来的而是训练中自动涌现的。近年也有研究指出多头中的某些头存在冗余甚至可以被剪枝而不影响最终效果。这说明多头并不保证每个头一定学到独特信息但它在统计上大大提高了好特征被学到的概率。4. 多头注意力的完整计算过程为了把代码看懂我们必须把矩阵的维度变化掰碎。下面以一个典型配置为例batch_size 2序列长度 T 10d_model 512num_heads 8d_k d_v 512 / 8 644.1 第一步三个线性投影输入 X 的形状是(2, 10, 512)。分别经过三个线性层 W_Q、W_K、W_V得到Q: (2, 10, 512) K: (2, 10, 512) V: (2, 10, 512)4.2 第二步拆分成多头把最后一维 512 拆成 8 个 64 维的向量形状变成(2, 10, 8, 64)然后通过 transpose 把头的维度移到序列维度之前(2, 8, 10, 64)此时每个头的 Q、K、V 形状是(2, 8, 10, 64)表示 8 个不同的注意力头并行计算。4.3 第三步每个头独立计算注意力对每个头执行缩放点积注意力Q 和 K 的转置做矩阵乘法得到权重矩阵(2, 8, 10, 10)。除以sqrt(64) 8。如果存在 mask则在 softmax 之前把需要遮蔽的位置替换为-inf。softmax 后得到归一化权重。权重与 V 相乘得到每个头的输出(2, 8, 10, 64)。4.4 第四步合并并输出把头的维度换回后面形状变成(2, 10, 8, 64)再 reshape 为(2, 10, 512)。最后经过一个输出线性层 W_O形状不变得到多头注意力的最终输出。输入 (2, 10, 512) - 线性投影 (2, 10, 512) - 拆分头 (2, 8, 10, 64) - 注意力计算 (2, 8, 10, 64) - 合并头 (2, 10, 512) - 输出投影 (2, 10, 512)5. PyTorch 完整代码实现5.1 从零实现多头注意力这是理解内部机制最重要的一份代码。建议不要直接复制到项目里而是手敲一遍并逐行注释。# 文件路径mha_from_scratch.py import math import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadAttention(nn.Module): def __init__(self, d_model: int, num_heads: int, dropout: float 0.0): 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(pdropout) def forward(self, query, key, value, maskNone): 参数形状说明 query: (B, T_Q, d_model) key: (B, T_K, d_model) value: (B, T_V, d_model) mask: 可选形状 (B, 1, 1, T_K) 或 (B, num_heads, T_Q, T_K) batch_size query.size(0) # 1. 线性投影 Q self.w_q(query) # (B, T_Q, d_model) K self.w_k(key) # (B, T_K, d_model) V self.w_v(value) # (B, T_V, d_model) # 2. 拆分成多头 # view: (B, T, d_model) - (B, T, num_heads, d_k) # transpose: (B, T, num_heads, d_k) - (B, num_heads, T, d_k) Q Q.view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) K K.view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) V V.view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) # 3. 缩放点积注意力 scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) # scores: (B, num_heads, T_Q, T_K) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn_weights F.softmax(scores, dim-1) attn_weights self.dropout(attn_weights) context torch.matmul(attn_weights, V) # context: (B, num_heads, T_Q, d_k) # 4. 合并多头 context context.transpose(1, 2).contiguous() context context.view(batch_size, -1, self.d_model) # 5. 输出投影 output self.w_o(context) return output, attn_weights if __name__ __main__: batch_size 2 seq_len 10 d_model 512 num_heads 8 x torch.randn(batch_size, seq_len, d_model) mha MultiHeadAttention(d_model, num_heads) output, attn mha(x, x, x) print(output shape:, output.shape) print(attn shape:, attn.shape)代码里最容易看晕的部分是第二步的view加transpose。这里解释一下view把最后一维 512 拆成(8, 64)矩阵逻辑上变成了四维。transpose把(num_heads, T)的顺序换到(T, num_heads)前面这样模型在计算注意力时每个头在 batch 和序列两个维度上都是独立的。如果你在实现时漏掉contiguous()后续调用view时会报错因为transpose返回的张量在内存中不是连续的。这是一个很容易被忽略的细节。5.2 使用 nn.MultiheadAttention 快速实现PyTorch 官方提供了封装好的nn.MultiheadAttention用法如下# 文件路径mha_official_api.py import torch import torch.nn as nn batch_size 2 seq_len 10 d_model 512 num_heads 8 # 官方接口默认输入是 (sequence_length, batch_size, embed_dim) x torch.randn(seq_len, batch_size, d_model) mha nn.MultiheadAttention( embed_dimd_model, num_headsnum_heads, dropout0.1, batch_firstFalse, # 默认为 False ) attn_output, attn_weights mha(x, x, x) print(attn_output:, attn_output.shape) # (10, 2, 512) print(attn_weights:, attn_weights.shape) # (2, 10, 10)如果你习惯使用(batch, sequence, feature)格式把batch_firstTrue即可# 推荐在实际项目中使用 batch_firstTrue mha_batch_first nn.MultiheadAttention( embed_dimd_model, num_headsnum_heads, dropout0.1, batch_firstTrue, ) x_batch_first torch.randn(batch_size, seq_len, d_model) attn_output, attn_weights mha_batch_first(x_batch_first, x_batch_first, x_batch_first) print(attn_output:, attn_output.shape) # (2, 10, 512) print(attn_weights:, attn_weights.shape) # (2, 10, 10)官方接口返回的注意力权重形状是(B, T_Q, T_K)它已经把多头权重做了平均。如果你需要拿到每个头单独的权重可以传入average_attn_weightsFalse参数。一个常见的坑是在同一个项目中有的模块用了batch_firstTrue有的模块没有设置导致输入输出的维度不一致拼接时出现 bug。建议在项目入口统一所有模块的batch_first配置。5.3 一个最小可运行的 Transformer 编码器层把多头注意力放回 Transformer 编码器层中是理解它实际作用的最佳方式。这个例子包含残差连接、LayerNorm 和前馈网络。# 文件路径transformer_encoder_layer_min.py import torch import torch.nn as nn class TransformerEncoderLayerMin(nn.Module): def __init__(self, d_model: int, num_heads: int, dim_feedforward: int, dropout: float 0.1): super().__init__() self.mha nn.MultiheadAttention( embed_dimd_model, num_headsnum_heads, dropoutdropout, batch_firstTrue, ) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.ffn nn.Sequential( nn.Linear(d_model, dim_feedforward), nn.ReLU(), nn.Dropout(dropout), nn.Linear(dim_feedforward, d_model), ) self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) def forward(self, x): # x: (B, T, d_model) # 多头注意力 残差 attn_out, _ self.mha(x, x, x) x x self.dropout1(attn_out) x self.norm1(x) # 前馈网络 残差 ffn_out self.ffn(x) x x self.dropout2(ffn_out) x self.norm2(x) return x if __name__ __main__: d_model 512 num_heads 8 dim_feedforward 2048 layer TransformerEncoderLayerMin(d_model, num_heads, dim_feedforward) x torch.randn(2, 10, d_model) out layer(x) print(encoder layer output:, out.shape)这里遵循 Transformer 原始论文的规范Attention 和 FFN 之后都使用残差连接并且每个块的输出经过 LayerNorm。前馈网络的中间维度通常是 d_model 的 4 倍比如 512 - 2048。运行以上三个代码块你会得到预期的输出维度。如果程序能跑通说明你已经完成了多头注意力的基本实现闭环。6. 运行结果与效果验证6.1 正确性的基本验证运行 5.1 的代码正常输出如下output shape: torch.Size([2, 10, 512]) attn shape: torch.Size([2, 8, 10, 10])这说明输入经过多头注意力后输出和输入的序列长度保持一致。attn的形状中第二维是 8恰好对应 8 个头。6.2 验证多头是否有差异一个值得做的实验是观察不同头的注意力权重是否真的有差异。你可以对同一句输入分别取出 8 个头的注意力矩阵做可视化。# 文件路径visualize_heads.py import matplotlib.pyplot as plt import torch from mha_from_scratch import MultiHeadAttention d_model 64 num_heads 4 seq_len 6 torch.manual_seed(42) x torch.randn(1, seq_len, d_model) mha MultiHeadAttention(d_model, num_heads) _, attn_weights mha(x, x, x) # attn_weights: (1, num_heads, T_Q, T_K) fig, axes plt.subplots(1, num_heads, figsize(12, 3)) for i in range(num_heads): axes[i].imshow(attn_weights[0, i].detach().numpy(), cmapviridis) axes[i].set_title(fhead {i}) axes[i].axis(off) plt.tight_layout() plt.savefig(mha_heads.png, dpi150)如果训练过一轮真实数据你会发现不同头的注意力图中高亮区域明显不同。这说明多头确实在输入的不同特征子空间上执行了不同模式的匹配。6.3 如何判断实现结果是否正确在实际训练中验证多头注意力模块是否正常工作的最直接方法是看 loss 曲线的下降趋势。如果 loss 在前几步没有明显下降先确认 Q、K、V 的初始化是否正确线性层默认初始化通常没有问题。如果 loss 直接变成 NaN需要检查两点一是 softmax 之前是否出现了过大的数值二是 fp16 训练下注意力矩阵是否溢出。如果 loss 下降正常但模型效果远低于预期可以考虑调整 num_heads 和 d_k而不是盲目增大 d_model。7. 常见问题与排查方法问题现象可能原因排查方式解决方案view报错提示 shape 不匹配输入形状不是(B, T, d_model)打印 Q/K/V 的 shape确认线性层输出维度使用x.reshape或先contiguous()再viewtranspose后view报内存不连续错误transpose返回的是非连续张量检查是否调用了contiguous()在 view 前调用.contiguous()mask 传进去后结果全变成 NaNmask 中被遮蔽的位置写了-inf但某些行全被 mask打印 mask 形状和 scores 最大值/最小值检查 mask 生成逻辑保证每行至少有一个有效位置训练时 loss 不下降dropout 过大或学习率不合适先关闭 dropout 测试再看 loss 曲线适当降低学习率确认 model.train() 和 model.eval() 正确推理和训练输出不一致忘记切换 model.eval()dropout 仍然生效检查模型中是否存在 dropout推理前调用model.eval()或使用torch.no_grad()包裹fp16 训练时 loss 波动剧烈注意力分数在 fp16 下溢出查看 scores 的取值范围是否超过 65504使用 bf16 代替 fp16或对 attention logits 做额外缩放多头注意力输出维度错误合并多头时 view 的维度写错打印 context 在 view 前后的 shape合并前先transpose(1, 2)再view(B, T, d_model)这里单独强调一下 fp16 的坑。在混合精度训练中softmax的输入如果数值过大fp16 能表示的最大值是 65504一旦超过就变成 inf后续计算会出现 NaN。bf16的范围和 fp32 相同只是精度更低所以很多大模型训练任务会选择 bf16 而不是 fp16。如果部署环境不支持 bf16就要额外检查注意力 logits 的数值范围。8. 多头注意力的最佳实践与工程建议8.1 选择合适的头数原始 Transformer 论文中使用的是 d_model 512num_heads 8。这已经成为一种默认配置但不是所有任务都必须照搬。实际选择时可以遵循几条经验d_model 必须能被 num_heads 整除。如果 d_model 768可以选择 num_heads 8 或 12。头数过少子空间的维度 d_k 过大每个头的表达能力粗难以学到细粒度关系。头数过多每个头的 d_k 过小单头信息容量不足且训练和推理开销都会增加。从工程角度看如果算力有限保持 num_heads 8 是一个稳妥的选择。模型规模变大时可以按比例增大头数但 d_k 维持在 64 左右通常是合理的。8.2 与位置编码的配合多头注意力本身是“无序”的它不会自动感知 token 的位置。没有位置编码时把句子“我打你”和“你打我”输入模型多头注意力计算出的语义模式完全一样因为三个 token 都在同一组位置上互换。实际实现中务必在进入多头注意力之前加入位置编码。常见的做法有绝对位置编码把正弦位置编码或可学习位置编码直接加到输入 embedding 上。相对位置编码在注意力分数计算时额外加入相对位置偏差很多新模型使用这个方法。如果任务对位置顺序非常敏感比如代码生成、SQL 生成、时间序列预测建议优先考虑相对位置编码它对序列长度的泛化能力更强。8.3 掩码的正确设计在 Transformer 中掩码分为两类对应多头注意力里的两个不同位置Padding Mask用于遮蔽 padding 位置。形状一般为(B, 1, 1, T_K)会在注意力机制的 scores 矩阵上把无效位置置为-inf。Causal Mask用于自回归生成保证当前位置只能看到当前及之前的位置。形状是上三角矩阵通常为(T_Q, T_K)。一个高频错误是解码器中同时用到两种掩码时直接把两个 mask 相加导致所有位置都被遮蔽。正确做法是先做 padding mask再做 causal mask确保两个条件同时满足。8.4 部署阶段的精度选择在训练阶段多使用 fp16 或 bf16 混合精度来加速。在推理阶段如果模型已经充分收敛适当降低精度可以有效减少显存占用但要注意 dropout 层在推理时已经关闭。当模型部署到不支持 bf16 的硬件上时推荐先对输入做标准化并检查注意力 logits 的数值范围。如果 logits 偏大可以考虑在 embedding 层后面接一层 LayerNorm或者使用动态缩放。另一个常见的部署优化是使用 FlashAttention。它把注意力计算重构成分块计算减少显存读写并且在 fp16 下数值稳定性更好。如果你的模型是在长序列场景下运行替换成 FlashAttention 往往能带来明显的加速效果但需要确认硬件兼容性。8.5 注意力权重可视化与调试在多语言翻译、文本分类、视觉 Transformer 等任务中可视化注意力权重可以帮助定位模型学到了什么。但不要过度依赖可视化。新的研究不断证明注意力权重和模型决策之间不一定是简单的因果关系。它更适合用于发现明显异常比如所有头都关注同一个位置、某个头完全没有多样性、或者 mask 之外出现异常高权重。9. 总结与后续学习方向多头注意力是 Transformer 架构中最具有代表性的设计之一。它解决的问题很明确单头注意力只能用一组权重去分配对所有 token 的关注程度难以同时表达多种语义关系。多头机制通过多套线性投影和多个子空间并行让模型可以在一次前向计算中同时学习多种关系模式同时整体参数量和单头注意力基本持平。这篇文章里我重点讲述了三个层面的内容概念层Q、K、V 的含义以及缩放点积注意力中为什么要除以sqrt(d_k)。实现层从零实现多头注意力到 PyTorch 官方 API再到完整编码器层的代码闭环。工程层掩码设计、头数选择、fp16/bf16 数值稳定性、FlashAttention 优化方向。下一步的学习建议非常明确不要停留在看文章。打开编辑器手写一份多头注意力然后做一个小实验——把模型里的多头换成单头对比在同一个任务上的 loss 下降速度和最终效果。这个对比会让你对“多头为什么有效”有比任何文章都深的体会。之后可以继续学习位置编码的数学原理、Transformer Encoder-Decoder 的整体结构、以及 KV Cache 在推理阶段如何复用 K 和 V从而加速生成。这些都是基于注意力机制的重要工程主题。如果你在尝试实现时遇到奇怪的维度问题或数值问题回看第 7 节的排查表大概率能在几分钟内定位到原因。