Transformer核心:多头注意力机制原理与PyTorch实现 这次我们来看 Transformers 里最核心也最容易在阅读源码时绕晕的一个模块多头注意力。它属于 Transformer 章节 7.1.2 的内容前接 7.1 注意力机制基础后面直接通向 BERT、GPT、ViT 这些实际模型。很多同学在看结构图时觉得 Q、K、V 三条分支懂了但一旦落到 PyTorch 代码里就不清楚每个张量到底是几维、为什么要算点积、为什么拆成多头后又要拼回去。这篇文章要做的就是把这条链路彻底打通先拆注意力机制的公式和直觉再讲多头注意力到底在做什么然后给出一版可运行的 PyTorch 从零实现并和官方nn.MultiheadAttention做对比验证。先说结论。多头注意力的核心价值不是简单增加参数量而是把注意力计算切到多个子空间并行执行让模型能够同时关注不同位置的多种依赖关系。从工程角度看它只是在缩放点积注意力前后做几次 reshape 和线性变换计算量没有本质增加但表达能力更强。全程不需要高端显卡CPU 就能跑通。文章按这个顺序展开第一部分给出多头注意力的核心认知速览第二部分从缩放点积注意力讲起这是所有后续内容的基础第三部分拆解多头注意力的完整流程第四部分给环境准备和 PyTorch 实现第五部分用代码验证维度、掩码和注意力可视化第六部分看它在真实模型中的应用最后是常见问题排查和最佳实践。1. 多头注意力核心认知速览先给一张速览表把多头注意力这个模块的关键信息放在最前面方便你判断这篇文章是否值得往下读。维度说明所属模块Transformer 编码器和解码器的核心子层核心公式Attention(Q,K,V) softmax(QKᵀ / √d_k)V多头计算形式MultiHead(Q,K,V) Concat(head₁, …, head_h)W^O核心思想将 d_model 维空间切成 h 个 d_k 维子空间并行计算注意力可训练参数Q/K/V 三个线性层 输出线性层共 4 组权重典型 head 数Transformer 论文默认 8BERT-base 为 12GPT-2 为 12硬件门槛理解原理无需 GPU小规模验证 CPU 可运行与单头区别多头能同时捕获多种关系单头只能做一种加权聚合这里先解释几个常见符号后面代码和公式都会用到d_model输入向量的维度也是每个 token 的表示维度。n_head多头数量。d_k每个 head 的查询/键向量维度通常d_k d_model / n_head。d_v每个 head 的值向量维度通常在标准实现中d_v d_k。理解多头注意力不需要先读完整篇 Transformer只需要知道它接受一个形状为(batch_size, seq_len, d_model)的张量经过 Q、K、V 三个线性映射和若干次 reshape 后输出和输入同形状的张量同时内部还产出一组注意力权重。2. 注意力机制基础从加权求和说起2.1 Query、Key、Value 的直觉注意力机制可以这样理解把一份文本切成长度为 n 的 token 序列后每个 token 被表示成一个向量。当我们处理第 i 个 token 时希望模型能动态地决定“应该重点关注序列里的哪些位置”。这里的 Query 可以理解为当前查询向量Key 是其他所有位置的索引向量Value 是其他位置的内容向量。模型先用当前 Query 和所有 Key 做相似度计算得到一组权重再把这些权重应用到 Value 上最终得到当前 token 的上下文表示。这就是一次“加权求和”的过程。Query 来自“我要查什么”Key 是“别人能提供什么索引”Value 是“别人实际提供的内容”。注意力机制就是在给定 Query 的情况下从所有 Key 中找出相关度然后按相关度聚合 Value。2.2 缩放点积注意力的数学形式论文《Attention Is All You Need》中给出的注意力函数是缩放点积注意力公式如下Attention(Q, K, V) softmax(QKᵀ / √d_k)V这个公式看起来简单但每个矩阵的维度必须清楚Q 的形状是(batch_size, seq_len_q, d_k)。K 的形状是(batch_size, seq_len_k, d_k)。V 的形状是(batch_size, seq_len_k, d_v)。QKᵀ 计算后得到(batch_size, seq_len_q, seq_len_k)表示每个 Query 和每个 Key 之间的相似度。softmax 作用在最后一个维度上让每一行的注意力权重之和为 1。最后乘 V得到(batch_size, seq_len_q, d_v)。在自注意力场景里seq_len_q 和 seq_len_k 相等Q、K、V 都来自同一个输入序列。在编码器-解码器注意力场景里Q 来自解码器K 和 V 来自编码器输出。2.3 为什么要除以根号 d_k缩放因子 √d_k 是公式里最容易忽略但最关键的部分。如果不做缩放两个 d_k 维向量的点积结果会随着维度增加而变大。当 d_k 较大时点积结果的数值会很大softmax 的输入进入梯度饱和区表现为梯度非常小训练不稳定。除以 √d_k 后点积的方差被拉回 1 附近softmax 的梯度能保持在合理范围。从实现角度看这个缩放不需要额外学习参数只是一个常数操作。但它在训练稳定性上非常重要。手写多头注意力时容易把这个缩放漏掉导致训练时 loss 不下降这是常见问题之一。2.4 单头注意力的局限如果只做一次注意力计算模型只能学到一组加权模式。但一句话里往往同时存在多种关系相邻词之间的语法关系、远距离的指代关系、否定词的作用范围、句法结构中的父子节点关系。单头注意力只能把这些关系混合在同一个加权平均里无法分别建模。多头注意力要解决的正是这个问题让不同的头去学习不同类型的依赖关系最终把多个子空间的信息拼接起来交给输出投影层融合。这也是为什么不是“多一份参数”这么简单而是“多一份表达能力”。3. 多头注意力的原理拆解3.1 整体思路多头注意力可以拆成五个步骤对输入做 Q、K、V 三个线性投影。把 Q、K、V 按头数 reshap 并转置拆成 n_head 个子空间。在每个子空间里独立执行缩放点积注意力。把所有头的输出拼接回d_model维。过输出线性投影 W^O。整体公式为MultiHead(Q,K,V) Concat(head₁, …, head_h)W^O其中每个头为head_i Attention(QW_i^Q, KW_i^K, VW_i^V)注意 Q、K、V 并不是一开始就拆好的而是先通过三组线性层映射到同样维度再在计算时切成多段。这个“先映射再切分”的做法在工程实现上非常高效。3.2 张量形状变化为了不抽象这里把张量形状变化列成一张表以batch_size2、seq_len10、d_model512、n_head8为例每个头的维度d_kd_v64。步骤输入形状输出形状Q/K/V 线性投影(2, 10, 512)(2, 10, 512)拆分为多头(2, 10, 512)(2, 8, 10, 64)每个头单独注意力(2, 8, 10, 64)(2, 8, 10, 64)拼接所有头(2, 8, 10, 64)(2, 10, 512)输出投影(2, 10, 512)(2, 10, 512)reshape 的细节要特别注意。一个(2, 10, 512)的张量先变成(2, 10, 8, 64)然后用transpose(1, 2)变成(2, 8, 10, 64)。这样才能保证每个头都看到完整的序列而不是把序列拆成多段。3.3 为什么有效不同头关注不同关系论文中提到训练完毕后观察不同头的注意力权重会发现不同头分布在不同区域有些头主要关注相邻词有些头关注远距离依赖还有些头关注特定的语法关系。这说明多头并不是简单的重复而是模型在训练中自动把不同的“关系查找任务”分配给了不同的子空间。不过也要说明并非每个头都一定学到可解释的模式有部分头可能互相冗余。这也催生了后续的 MQAMulti-Query Attention和 GQAGrouped Query Attention等优化它们核心思想都是减少 Key 和 Value 的冗余头数降低推理显存和带宽开销。理解标准多头注意力是理解这些优化方案的前提。4. 环境准备与 PyTorch 从零实现4.1 环境准备多头注意力是理论模块不需要特殊硬件。建议用 Python 3.8 以上版本和 PyTorch 1.10 或 2.x。安装命令如下具体 PyTorch 版本请以官方安装页为准pip install torch numpy matplotlib安装完成后可以用下面的命令确认 PyTorch 是否可用python -c import torch; print(torch.__version__)如果你的环境里已经有 PyTorch可以直接跳过安装步骤。下面所有实验在 CPU 上就能运行。4.2 手写多头注意力模块从零实现一版多头注意力核心代码并不长。这里给出一个适合学习的最小实现没有封装太多复杂细节方便对照公式。import math import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadAttention(nn.Module): 从零实现的多头注意力模块。 d_model: Transformer 模型的宽度 n_head: 头数 def __init__(self, d_model, n_head, dropout0.1): super().__init__() assert d_model % n_head 0, d_model 必须能被 n_head 整除 self.d_model d_model self.n_head n_head self.d_k d_model // n_head self.d_v d_model // n_head 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) # 1. 线性投影 Q self.w_q(query) # (batch, seq_len, d_model) K self.w_k(key) # (batch, seq_len, d_model) V self.w_v(value) # (batch, seq_len, d_model) # 2. 拆分多头 # 先 view 成 (batch, seq_len, n_head, d_k) # 再 transpose 成 (batch, n_head, seq_len, d_k) Q Q.view(batch_size, -1, self.n_head, self.d_k).transpose(1, 2) K K.view(batch_size, -1, self.n_head, self.d_k).transpose(1, 2) V V.view(batch_size, -1, self.n_head, self.d_k).transpose(1, 2) # 3. 缩放点积注意力 # scores: (batch, n_head, seq_len, seq_len) 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_weights F.softmax(scores, dim-1) attn_weights self.dropout(attn_weights) # context: (batch, n_head, seq_len, d_k) context torch.matmul(attn_weights, V) # 4. 拼接多头 # 先转置回 (batch, seq_len, n_head, d_k) # 再 contiguous view 回 (batch, seq_len, d_model) context context.transpose(1, 2).contiguous().view( batch_size, -1, self.d_model ) # 5. 输出投影 output self.w_o(context) return output, attn_weights代码中几个容易出错的地方view之后必须注意维度的排列顺序。先view(batch_size, -1, n_head, d_k)再transpose(1, 2)得到的才是(batch, n_head, seq_len, d_k)。transpose之后张量内存可能不连续拼接前需要调用contiguous()否则view会报错。masked_fill(mask 0, float(-inf))会把被 mask 掉的位置变成负无穷softmax 后这些位置的权重接近 0。4.3 与 nn.MultiheadAttention 对比测试PyTorch 官方已经提供了nn.MultiheadAttention我们可以并行跑一遍验证形状是否一致。import torch import torch.nn as nn d_model 512 n_head 8 batch_size 2 seq_len 10 dropout 0.1 custom_mha MultiHeadAttention(d_model, n_head, dropout) builtin_mha nn.MultiheadAttention(d_model, n_head, dropoutdropout, batch_firstTrue) x torch.randn(batch_size, seq_len, d_model) out_custom, attn_custom custom_mha(x, x, x) out_builtin, attn_builtin builtin_mha(x, x, x) print(自定义多头注意力输出形状:, out_custom.shape) print(内置多头注意力输出形状:, out_builtin.shape) print(自定义注意力权重形状:, attn_custom.shape) print(内置注意力权重形状:, attn_builtin.shape)预期输出是自定义多头注意力输出形状: torch.Size([2, 10, 512]) 内置多头注意力输出形状: torch.Size([2, 10, 512]) 自定义注意力权重形状: torch.Size([2, 8, 10, 10]) 内置注意力权重形状: torch.Size([2, 10, 8, 10])两边的输出张量形状完全一致只是注意力权重的维度排布不同。自定义实现里注意力权重是(batch, n_head, seq_len, seq_len)内置模块默认返回(batch, seq_len, n_head, seq_len)。这是因为官方接口里把 seq_len 放在了前面语义上没有差别。数值上两边不会完全一致因为线性层初始化参数不同。需要验证的是变换逻辑而不是数值相等。如果你希望严格对齐可以手动把内置模块的in_proj_weight和in_proj_bias拷贝到自定义实现中再对比输出但一般学习阶段不需要做这一步。5. 功能验证维度、掩码与注意力可视化5.1 维度变化验证用一个小例子逐步打印每个阶段的张量形状是最快理解多头注意力的方式。下面脚本把上一节的 Q、K、V 中间变量分别打印出来import math import torch import torch.nn as nn d_model 64 n_head 4 batch_size 1 seq_len 6 class DebugMHA(nn.Module): def __init__(self): super().__init__() self.n_head n_head self.d_k d_model // n_head 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) def forward(self, x): B x.size(0) Q self.w_q(x) K self.w_k(x) V self.w_v(x) print(Q shape:, Q.shape) Q Q.view(B, -1, self.n_head, self.d_k).transpose(1, 2) K K.view(B, -1, self.n_head, self.d_k).transpose(1, 2) V V.view(B, -1, self.n_head, self.d_k).transpose(1, 2) print(Q after split:, Q.shape) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) print(scores shape:, scores.shape) attn torch.softmax(scores, dim-1) context torch.matmul(attn, V) print(context shape:, context.shape) context context.transpose(1, 2).contiguous().view(B, -1, d_model) print(context after concat:, context.shape) out self.w_o(context) print(output shape:, out.shape) return out x torch.randn(batch_size, seq_len, d_model) model DebugMHA() model(x)运行这段代码你会清楚看到每一步的维度变化。判断实现是否正确标准就是最终输出形状和输入形状一致以及scores的形状符合(batch, n_head, seq_len, seq_len)。5.2 mask 掩码对注意力的影响在 Transformer 中mask 主要有两种用途padding mask把无效位置填充为很小的数避免模型关注填充 token。因果 mask解码器中避免模型看到未来 token。以 padding mask 为例假设序列长度为 4其中最后一个 token 是填充项mask 向量为[1, 1, 1, 0]。在注意力计算里这个 mask 会被广播到所有 head。用上一节的手写模块测试import torch from your_implementation import MultiHeadAttention model MultiHeadAttention(d_model64, n_head4) x torch.randn(2, 4, 64) # 假设第二批次的最后一个 token 是 padding mask torch.tensor([ [1, 1, 1, 1], [1, 1, 1, 0] ]).unsqueeze(1).unsqueeze(2) # (2, 1, 1, 4) out, attn model(x, x, x, maskmask) print(attn)mask 形状需要能广播到(2, 4, 4, 4)的 scores 上所以这里是(batch, 1, 1, seq_len)。可以看到被 mask 位置对应的注意力权重几乎为 0因为masked_fill将分数设置为-inf后softmax 会把它们压到 0。常见错误是把 mask 的形状搞错。如果 mask 少了维度masked_fill会广播失败或者没有按预期遮挡。建议在实现里写成scores.masked_fill(mask 0, float(-inf))并通过打印 shape 确认广播后形状。5.3 注意力权重可视化注意力权重是理解模型行为的直接入口。以一句话为例子可以画出每个 token 到其他 token 的热力图。import matplotlib.pyplot as plt import torch from your_implementation import MultiHeadAttention model MultiHeadAttention(d_model64, n_head4) tokens [我, 爱, 深度学习, 和, 自然语言处理] x torch.randn(1, len(tokens), 64) _, attn model(x, x, x) head_idx 0 plt.figure(figsize(6, 5)) plt.imshow(attn[0, head_idx].detach().numpy(), cmapBlues) plt.xticks(range(len(tokens)), tokens, rotation45) plt.yticks(range(len(tokens)), tokens) plt.colorbar() plt.title(Head 0 Attention Weights) plt.tight_layout() plt.show()如果使用随机初始化的模型热力图通常比较平滑没有明显规律。如果模型已经训练好你会看到不同 head 的关注点有明显差异。这也是验证“多头有效”最直观的方式。5.4 资源与性能观察思路多头注意力的主要计算开销来自注意力矩阵(batch_size, n_head, seq_len, seq_len)序列长度增大时注意力矩阵按平方增长这是 Transformer 被称为“二次复杂度”模型的原因。观察资源占用可以分两部分CPU 或 GPU 推理耗时可以用time模块粗略统计。如果使用 GPU可以用torch.cuda.max_memory_allocated()查看峰值显存。实际数值取决于序列长度、batch size 和 head 数不能一概而论。重点要记住head 数增加会加大注意力矩阵的数量但每个 head 的维度变小总的参数量和计算量并不会成倍增长因为每个头只负责d_model / n_head维子空间。6. 多头注意力在真实模型中的应用6.1 Transformer 编码器与解码器在 Transformer 编码器中每层包含两个子层多头注意力和前馈网络。输入经过多头注意力之后会经过残差连接和 LayerNorm再进入前馈网络。在解码器中多头注意力出现两次第一次是自注意力用因果 mask 屏蔽未来 token。第二次是编码器-解码器注意力Query 来自解码器Key 和 Value 来自编码器输出帮助解码器获取输入序列的信息。这两种场景下多头注意力的计算逻辑完全相同区别只在于 mask 和输入来源。6.2 BERT、GPT 等预训练模型BERT-base 使用 12 层 Transformer每层 12 个 headd_model768每个头的维度是 64。GPT-2 同样使用 12 个 head更大规模版本会增加层数和 head 数。阅读 BERT 和 GPT 源码时你会发现它们对多头注意力的实现有两种风格一种是显式做 QKV 线性投影后再 reshape 拆头另一种是使用nn.MultiheadAttention封装。前者的好处是可控性强后者更简洁。理解了手写版本后再去读这两类源码都会轻松很多。6.3 视觉 TransformerViTViT 把图片切分成固定大小的 patch每个 patch 展开成 token 后送入标准 Transformer。图像里的多头注意力同样按 patch 位置计算相关性因此不同 head 可能学到不同尺度或方向的图像特征。这是近两年视觉领域大量使用 Transformer 结构的直接原因之一。6.4 head 数量怎么选论文中默认 8 个 head之后大量模型沿用 12 或 16 个 head。head 数增加能提升模型表达能力但显存和训练时间也会增加。更关键的是很多研究发现并非所有 head 都对最终效果有贡献剪掉部分 head 性能下降有限。这也是多查询注意力MQA和分组查询注意力GQA的出发点在推理阶段减少 K、V 头数降低显存带宽开销同时保持模型输出质量。理解这些优化的前提还是把标准多头注意力吃透。7. 常见问题与排查方法问题现象可能原因排查方式解决方案代码报错无法viewtranspose后张量内存不连续检查报错信息确认是否提示contiguous拼接前调用contiguous()d_model无法整除n_head维度配置不合理打印d_model和n_head调整d_model或n_head让两者整除训练 loss 不下降忘记除以√d_ksoftmax 梯度饱和检查注意力代码中是否有math.sqrt(self.d_k)补上缩放因子注意力权重几乎均匀分布模型未训练或初始化不合理打印注意力矩阵观察是否集中在少数 token先跑小数据集验证再调学习率mask 不生效mask 形状不对广播后无法匹配 scores打印 mask 和 scores 的形状把 mask 扩展到(batch, 1, 1, seq_len)形式解码器看到未来信息因果 mask 设置错误检查 mask 是否为上三角矩阵使用torch.triu(..., diagonal1)生成上三角掩码序列稍长就显存不足注意力矩阵 O(n²) 占用过大查看峰值显存缩短序列、使用 FlashAttention 或分块注意力自定义实现和官方输出差异大初始化参数不同对比两者的逻辑和形状不对比具体数值如需严格对齐复制官方权重初始化方式8. 最佳实践与学习建议8.1 实现与调试建议第一次写多头注意力时不要上来就跑大模型。建议先用d_model64、n_head4、seq_len8这样的小参数把链路跑通确认每一步的维度都如预期。调试时可以沿用三件套在每个阶段打印张量形状。写一个和nn.MultiheadAttention的对比脚本校验形状。用一个你熟悉的小任务比如简单的序列复制或情感分类让模型训练 20 步左右观察 loss 是否能下降。还有一个容易被忽略的点是初始化。PyTorch 的nn.Linear默认初始化通常能直接使用但如果你复现论文时要严格对齐需要关注权重初始化的细节。学习阶段不需要过度纠结这一点。8.2 从理解到工程理解标准多头注意力后就可以继续往下读这些方向位置编码注意力本身不包含位置信息需要靠位置编码补充。完整 Transformer Encoder/Decoder多头注意力只是其中一个子层。FlashAttention通过分块计算和 IO 优化把注意力计算变得更快更省显存。MQA / GQA在推理阶段减少 K、V 头数的优化方案。不同变体的源码实现读 transformers 库或 Fairseq 的注意力实现看工程化封装思路。建议动手做一个小实验用标准多头注意力实现一个最小字符级语言模型在几十万字符的语料上训练几轮然后观察生成结果。这个实验能让你把公式、代码和实际效果连通。9. 总结与下一步多头注意力是整个 Transformer 架构中最需要熟练掌握的模块之一。它建立在缩放点积注意力之上通过拆接多个子空间让模型可以同时表达多种依赖关系。这篇文章给出了完整公式、形状变化、PyTorch 从零实现和与官方模块的对比验证重点在于你能否跟着代码跑一遍。如果你已经看懂了这版实现下一步建议做三件事把代码里的 mask 改成因果 mask实现一个最小解码器。找一个已经训练好的小型 BERT 或 GPT 模型提取某层注意力权重并可视化。对比 FlashAttention 的原理理解标准实现里哪些地方可以优化。在跑通上述内容之前不需要急着读大模型源码更不用一开始就在大批量数据上训练。先把多头注意力的每个张量形状烂熟于心再去碰完整 Transformer 会顺畅得多。