ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

多头注意力机制原理与PyTorch实现详解

多头注意力机制原理与PyTorch实现详解 自注意力机制本身并不复杂核心思想就是一句话让序列里的每个 token 都能和序列里的其他 token 做信息交互。但真正让 Transformer 在 NLP、CV、多模态任务里全面站稳脚跟的是那个看起来只加了一个字的模块——多头注意力机制Multi-Head Attention。它解决的是单头注意力“只能有一组注意力分布”的表达瓶颈而方案不是增加复杂度而是把 Q、K、V 投影到多个低维子空间并行计算。这篇文章把多头注意力的动机、数学原理、PyTorch 实现、因果掩码、MQA/GQA 变体以及它与残差、层归一化、FFN 的配合方式全部拆开讲一遍。如果你正在学 Transformer准备从零复现 BERT 或 GPT 系列模型如果你读懂了论文里的公式但一到了代码里就被view、transpose、contiguous绕晕或者你只是想知道为什么大模型推理时都在强调 KV Cache而 GQA 能加速那么多——这篇文章就是给你准备的。读完你不仅能手写一个可运行的多头注意力模块还能说清楚它为什么有效、有哪些容易踩的坑。1. 多头注意力机制核心速览先把关键信息放在最前面。这一节不做推导只给结论。下面的表格基本覆盖了多头注意力机制的理解坐标。项目说明机制名称多头注意力机制Multi-Head Attention, MHA提出论文Attention Is All You NeedTransformer 原论文解决的核心问题单头自注意力表达能力有限难以同时建模多种依赖关系核心思路将 Q/K/V 投影到多个低维子空间并行做注意力计算再把结果拼接起来是否增加参数量标准实现下不增加Q/K/V 总参数量与单头版本一致典型配置d_model512num_heads8每个头维度 d_kd_v64前置知识缩放点积注意力、Softmax、线性投影、矩阵维度变换主要应用Transformer、BERT、GPT、ViT、多模态模型等绝大多数现代架构常见变体MHA、MQA、GQA以及 FlashAttention 等高频实现学习门槛中等需要一定矩阵基础但代码复现并不难判断自己是否真正理解多头注意力可以拿下面三个问题自测当d_model512、num_heads8时每个头的维度是多少为什么是这个数多头注意力的总参数量为什么和单头注意力相同它到底“多”在哪里训练 GPT 这类自回归模型时为什么在注意力分数上要加一个上三角掩码这三个问题如果在读完后都能回答说明这章就真正通了。2. 为什么需要“多头”单头自注意力的局限单头自注意力的局限主要体现在三个层面。第一表示能力单一。自注意力输出是Attention(Q, K, V)它本质上是在一组 Softmax 权重下对所有 Value 向量做加权求和。一个注意力头只能输出一种加权方式的结果。但语言中一个词可能需要同时建模多种关系比如“苹果”这个词既和“红色”有颜色关系又和“水果”有类别关系还和“乔布斯”有品牌关系。单头注意力只能把这些关系全部揉在一起最终得到一个平均化的上下文表示。第二Softmax 存在“平均化”倾向。当序列长度变长时注意力分数经过 Softmax 后很容易变得平缓尤其是每个 token 的表示都比较接近的时候。这时候注意力头实际上退化成了一种近似平均池化操作没有真正突出某一个位置。解决思路有两个方向一是降低温度增大注意力分布的尖锐程度二是让模型同时尝试多组不同的注意力分布总有一组能学到关键依赖。第三优化的计算路径太单一。单头自注意力从一个全量矩阵运算中学习依赖关系一组 W_Q、W_K、W_V 只能覆盖一种语义空间。模型把所有的语法、语义、指代、位置信息全部塞进同一个低维投影里梯度更新时这些信息会互相干扰。多头注意力解决问题的思路很直接既然一个头不够那就并行跑多个头。每个头使用独立的投影矩阵把输入映射到不同的子空间学习不同类型的依赖关系。最后把多个头的输出拼接起来再经过一次线性投影融合成完整的表示。这样做既保留了注意力的全局交互能力又增加了模型的表达自由度而且参数总量不涨。3. 多头注意力机制原理拆解多头注意力机制的输入是一个序列表示矩阵X形状为(batch_size, seq_len, d_model)。整个计算过程分四步。3.1 生成 Q、K、V 投影输入通过三个可学习矩阵 W_Q、W_K、W_V 得到查询、键、值矩阵。在标准实现中这三个矩阵的维度都是(d_model, d_model)Q XW_Q, K XW_K, V XW_V3.2 按头拆分把 d_model 维度平均切分成 h 份每份维度 d_k d_model / h。拆分在代码中常见做法是先经过 Linear(d_model, d_model) 得到形状 (batch, seq_len, d_model)再通过 view 和 transpose 重排为 (batch, h, seq_len, d_k)。这一操作等价于把一个大矩阵切成了 h 个子矩阵每个子矩阵代表一个子空间中的投影。3.3 缩放点积注意力每个头独立计算注意力分数。缩放点积注意力的标准公式为$$ \text{Attention}(Q, K, V) \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)V $$其中除以 sqrt(d_k) 是关键细节。当 d_k 较大时QK^T 的结果会有较大方差Softmax 的梯度会变得非常小训练不稳定。除以 sqrt(d_k) 是为了把分数拉回到合理的数值区间。3.4 拼接并输出投影将 h 个头的输出在最后一个维度上拼接得到维度为 d_model 的向量再经过输出矩阵 W_O 完成一次线性变换$$ \text{MultiHead}(X) \text{Concat}(\text{head}_1, \ldots, \text{head}_h)W^O $$其中$$ \text{head}_i \text{Attention}(XW_i^Q, XW_i^K, XW_i^V) $$这里 W_i^Q、W_i^K、W_i^V 的维度是 (d_model, d_k)。从参数总量看h 个头总共 h * (3 * d_model * d_k) 3 * d_model^2 个参数和单头版本完全一致。区别在于单头是一个大的线性投影多头是把这个投影拆成了 h 份不同的子空间并通过输出投影重新融合。4. 为什么多头有效4 个关键原因多头注意力之所以成为 Transformer 最核心的组件不是因为它“听起来复杂”而是因为它在四个方面都有明确作用。4.1 子空间并行多头各司其职Transformer 原论文通过在机器翻译模型上的可视化实验观察到不同的注意力头确实在学习不同类型的关系有的头关注句法依赖比如动词和主语有的头关注指代关系比如代词和先行词有的头关注相邻位置的局部特征还有的头关注长距离的跨段依赖。多头机制本质上是让模型拥有 h 次机会去学习不同的注意力模式而不是强迫一个头把所有关系都学会。4.2 打破单头 Softmax 的“平均化”多头机制相当于把原来的一个 Softmax 分布变成了 h 个独立的 Softmax 分布。每个头只需要在自己的子空间里找到最重要的位置不需要承担所有信息的加权责任。即使某一个头出现退化成“平均池化”的情况其他头仍然可以保持尖锐的注意力分布。多个头的组合让模型更稳定。4.3 参数效率极高很多第一次接触多头注意力的人会误以为“多头”意味着 h 倍参数量。事实并非如此。多头通过拆分 d_model 维度来降低每个头的维度总参数量和单头完全一致。它增加的是“表征的自由度”而不是“参数的数量”。这也是为什么 Transformer 能在参数量不变的情况下获得更高的模型容量。4.4 训练更稳定梯度更平滑单头注意力的输出是一个大矩阵直接参与最后的加权求和所有信息集中在同一个路径上。多头输出经过拼接和线性投影后梯度可以通过多个分支回传到不同的子空间避免了单个注意力头的梯度主导整个模型更新的问题。多个头还可以配合 Dropout 机制使用不同头随机丢弃部分注意力权重相当于在注意力层面做了集成学习。5. 环境准备与代码复现理解公式之后必须动手写代码。这里给出一套完全可运行的 PyTorch 实现不需要 GPUCPU 环境即可验证维度逻辑。如果你本机已经有 PyTorch直接跳过安装步骤。首先确认 Python 版本建议 Python 3.8 以上然后安装 PyTorch。pip install torch安装完成后检查是否可以正常导入。import torch import torch.nn as nn import torch.nn.functional as F import math print(torch.__version__)本文代码的位置是在一个自包含的 Python 脚本里运行不依赖额外项目结构。建议把下面的代码保存为multi_head_attention.py后续修改参数方便调试。6. 手写多头注意力模块PyTorch这里给出一个最典型的实现方式。它严格按照上文公式展开重点在于理解 Q/K/V 的维度变换和 mask 的传递逻辑。import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads, dropout0.1): super().__init__() assert d_model % num_heads 0, d_model must be divisible by 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, x, maskNone): batch_size, seq_len, _ x.size() # 1. 生成 Q、K、V并拆分成多头形状 # 目标形状: (batch, num_heads, seq_len, d_k) Q self.W_Q(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) K self.W_K(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) V self.W_V(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) # 2. 计算缩放点积注意力分数 # scores shape: (batch, num_heads, seq_len, seq_len) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) # 3. 如果传入 mask则 mask 为 0 的位置置为负无穷 if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) # 4. Softmax 得到注意力权重再作用于 V attn F.softmax(scores, dim-1) attn self.dropout(attn) context torch.matmul(attn, V) # shape: (batch, num_heads, seq_len, d_k) # 5. 拼接所有头恢复 d_model 维度 context context.transpose(1, 2).contiguous().view(batch_size, seq_len, self.d_model) # 6. 输出投影 output self.W_O(context) return output接下来做一个小测试验证输出形状是否正确。x torch.randn(2, 10, 512) # batch2, seq_len10, d_model512 mha MultiHeadAttention(d_model512, num_heads8) y mha(x) print(input shape:, x.shape) print(output shape:, y.shape)预期输出input shape: torch.Size([2, 10, 512]) output shape: torch.Size([2, 10, 512])输出形状和输入形状完全一致这符合 Transformer 中残差连接的使用前提。这里最容易踩坑的是view和transpose的配合。view是把最后两个维度重组为(num_heads, d_k)transpose(1, 2)再把num_heads维度提前最终得到(batch, num_heads, seq_len, d_k)。这里的顺序一旦写错后续矩阵乘法的形状就会全部错位。如果你熟悉 einsum也可以写一个更紧凑的等价版本scores torch.einsum(bqhd,bkhd-bhqk, Q, K) / math.sqrt(self.d_k) context torch.einsum(bhqk,bkhd-bqhd, attn, V)两种写法计算逻辑完全一致。einsum 可读性稍差但不容易出现维度顺序错误。7. 因果自注意力与 Mask 实现在 GPT 等自回归模型里多头注意力不能直接使用普通版本必须加一个因果掩码causal mask所以这部分单独拿出来讲。在很多开源代码中你会看到它被写作 Causal Self-Attention也就是“因果自注意力”。因果掩码的核心逻辑生成任务中token 在位置 t 只能看到位置 t 的 token不能看到未来的 token。否则模型在训练时“偷看”了未来信息推理时就没有对应的未来 token造成训练和推理不一致。掩码的计算非常简单。先用torch.tril生成一个下三角矩阵再将掩码应用到注意力分数矩阵上。在 PyTorch 中实现如下def subsequent_mask(seq_len): mask torch.tril(torch.ones(seq_len, seq_len)).bool() return mask # shape: (seq_len, seq_len)测试一下mask subsequent_mask(5) print(mask)输出tensor([[ True, False, False, False, False], [ True, True, False, False, False], [ True, True, True, False, False], [ True, True, True, True, False], [ True, True, True, True, True]])在多头注意力 forward 中调用时mask 需要扩展为和 scores 相同的维度也就是(batch, num_heads, seq_len, seq_len)seq_len x.size(1) mask subsequent_mask(seq_len).unsqueeze(0).unsqueeze(0) # (1, 1, seq_len, seq_len) mask mask.expand(x.size(0), mha.num_heads, -1, -1) # (batch, num_heads, seq_len, seq_len) output mha(x, maskmask)在 forward 内部mask 为 False 的位置会被masked_fill替换为负无穷经过 Softmax 后这些位置的权重趋近于 0。这里有一个实现上的关键点mask 必须在 softmax 之前加而不是在 softmax 之后把权重置零。如果 softmax 之后直接置零所有权重之和不再等于 1会破坏概率分布的语义而 softmax 之前加负无穷是标准做法。除了因果掩码实际工程中还会用到 padding mask目的是让注意力忽略掉 padding token。对于自回归模型通常需要同时使用 padding mask 和 causal mask两者取交集。实现上可以通过torch.logical_and将两个掩码合并成一个布尔矩阵再一次性传给注意力模块。8. 多头注意力变体对比MHA / MQA / GQA随着大模型推理部署的发展多头注意力机制出现了几个重要变体。理解这些变体能帮你理解为什么新一代大模型都在提“减少 KV Cache”。变体全称核心思想参数量推理效率代表应用情况MHAMulti-Head Attention每个头都有自己的 K、V高一般Transformer、BERT、早期 GPTMQAMulti-Query Attention所有头共享一组 K、V只有 Q 独立低快部分早期大模型GQAGrouped-Query Attention若干个头共享一组 K、V中较快Llama 2、Llama 3、Mistral 等在自回归生成时模型每生成一个新 token 都需要用到之前所有 token 的 K、V 向量。如果不做缓存每一步都重新计算代价太高。因此引擎会将历史 K、V 缓存到显存中这部分缓存就是 KV Cache。MHA 因为每个头都需要各自缓存 K、V显存占用最高。MQA 让所有头共享一组 K、V缓存显著减少但会牺牲一部分模型表达能力。GQA 是折中方案把头分成若干组每组内部共享 K、V。它既减少了缓存量又保留了一定程度的表达多样性。这也是为什么 Llama 2 之后的很多开源模型都选择 GQA。如果你自己在实现 Transformer 推理可以从 MHA 开始跑通后再优化为 GQA。优化的第一步是理解“哪些 weight 需要缓存”Q 每次都重新生成不需要缓存K、V 需要跨 step 保留并拼接。9. 多头注意力与残差、层归一化、FFN 的配合多头注意力并不是单独工作的。在 Transformer 中它总是和残差连接、层归一化、前馈网络组合成一个完整的 Transformer Block。一个标准 Transformer Block 的计算过程如下x x MultiHeadAttention(LayerNorm(x)) x x FeedForward(LayerNorm(x))其中 FeedForward 通常是一个两层的多层感知机MLP先升维再降维中间用 ReLU 或 GELU 激活。这个结构的两个关键点第一多头注意力输出经过残差和 LayerNorm 后数值分布会更稳定。多头注意力内部的矩阵乘法和 Softmax 操作会让数值范围波动很大直接堆叠多层会出现训练不稳定的情况。层归一化LayerNorm在每个 token 维度上做归一化平均值拉到 0、方差拉到 1有效缓解梯度爆炸或消失。第二多头注意力本质上是线性投影和加权求和单靠它无法引入非线性。FFN 中的非线性激活函数承担了这部分工作。多头注意力负责在不同 token 之间交互信息FFN 负责在每个 token 内部做更高维的特征变换两者分工明确。Pre-LN 和 Post-LN 是实现上的一个重要差别。上面给出的写法是 Pre-LN先 LayerNorm 再进注意力。GPT 系列模型普遍使用 Pre-LN因为它可以让深层网络训练更稳定。原始 Transformer 论文中的结构更接近 Post-LN先注意力再 LayerNorm。理解这个差别有助于阅读不同开源模型的源码。10. 常见问题与排查方法实际写代码时最容易出的问题集中在维度变换和 mask 逻辑上。下面整理了一份排查清单。现象可能原因检查方式解决方案矩阵乘法维度对不上d_model 无法被 num_heads 整除打印 Q、K、V 的 shape调整 num_heads或修改 d_model输出 shape 与输入不一致view和transpose顺序写反在 forward 中逐步打印 shape按view - transpose顺序重新组织训练损失不下降或速度太慢忘记除以 sqrt(d_k)检查 scores 计算代码加上math.sqrt(self.d_k)mask 没有生效mask 维度与 scores 不一致打印 mask 和 scores 的 shape将 mask 扩展到 (batch, heads, seq_len, seq_len)Softmax 后取 mask 置零对概率直接置零分布不再归一检查 mask 是在 softmax 前还是后在 softmax 前通过负无穷屏蔽单头输出正常多头后结果异常拼接后没有调用 contiguous检查.view前的报错拼接前调用.contiguous()长序列显存溢出注意力分数矩阵为 O(n^2)监控显存和序列长度使用 FlashAttention、稀疏注意力或梯度检查点推理延迟过高MHA 的 KV Cache 占用过高观察显存占用与缓存大小切换到 GQA 或 MQA其中contiguous的问题尤其隐蔽。transpose操作不会让内存连续此时直接调用view会报错。代码中先transpose再contiguous再view这是正确顺序。如果你在实现中遇到view size is not compatible with input tensors size大概率就是这里出了问题。11. 最佳实践与使用建议学习多头注意力机制不需要一开始就追求复杂实现。建议按下面的顺序推进。第一先跑通小规模测试。d_model 设 128、num_heads 设 4序列长度设 16先用随机张量验证输出形状。形状全部正确后再加 mask最后再接入 LayerNorm 和 FFN。第二结合 loss 曲线判断实现是否正确。完全随机初始化时Transformer 的 loss 应该短暂下降且不会立即发散。如果 loss 在第一步就变成 NaN优先检查 scores 的缩放因子和 LayerNorm 的 eps 参数。第三头数并不是越大越好。常见工程经验是每个头的维度在 64 附近例如 d_model512 对应 8 个头d_model768 对应 12 个头。头数过少表达能力受限头数过多单个头维度太小能够学到的特征有限且矩阵乘法形状更碎GPU 利用率反而下降。第四长序列场景要主动优化。多头注意力的计算复杂度是 O(n^2)序列长度从 512 提升到 2048计算量会增长 16 倍。实际工程中可以考虑 FlashAttention、稀疏注意力、局部窗口注意力等方案而不是盲目堆算力。第五推理阶段要关注 KV Cache 的复用。自回归模型生成时对 QKV 的处理方式完全不同Q 只和当前 token 有关不缓存K、V 需要保存历史。如果只是拿模型做训练可以暂时忽略 KV Cache如果做部署和接口服务KV Cache 就是性能优化的核心。第六任何涉及真实数据训练或应用的项目要注意数据授权、隐私保护和内容合规。模型训练使用他人文本、图像、语音数据时需要确认是否有合法使用权生成内容对外发布前需要根据应用场景做好安全审核。12. 总结与下一步多头注意力机制的核心可以浓缩为一句话在参数总量不变的条件下把单一大矩阵投影拆成多个子空间并行学习再拼接融合让模型获得更多样、更稳定的注意力模式。它本身不是复杂机制但却是理解和复现几乎所有现代大模型的必经之路。建议下一步动手做三件事一是修改num_heads从 1 改成 4、8、16观察输出变化和显存波动二是给当前模块加入因果掩码跑一个简单的 n-gram 预测任务验证自回归逻辑三是继续学习位置编码和 FlashAttention位置编码解决的是“注意力本身不感知顺序”的问题FlashAttention 解决的是长序列下显存和速度的问题。把多头注意力这一步踩扎实后面再看 BERT、GPT、ViT 的源码会发现大量代码都是这一章的重复与扩展。
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进