多头注意力机制原理与PyTorch实现:从QKV拆分到掩码调试

📅 发布时间:2026/8/30 21:36:06
多头注意力机制原理与PyTorch实现:从QKV拆分到掩码调试 多头注意力机制简单说就是把自注意力中的 Q、K、V 分成多个子空间并行计算让每个头关注不同的关系。它解决的核心问题是单头注意力只能学到一种固定的词与词关联模式而真实文本里存在语法、指代、局部搭配、全局语义等多种关系单头很难兼顾。这篇文章适合正在学习 Transformer、自己动手实现注意力层或者用 PyTorch 搭模型时搞不清num_heads、embed_dim、掩码维度的人。我最想强调的是多头不是简单的“多复制几份注意力”而是先投影、再切分、最后拼接这中间任何一个维度写错都会让训练结果变得很怪。下面按我实际调试的顺序拆开讲。1. 为什么单头自注意力不够用先从注意力图说起1.1 单头注意力本质上只能学一种“关系”自注意力做的事情是计算句子中每个词和其他词的相关程度。比如“苹果”这个词它可能和“吃”有关系也可能和“公司”有关系还可能和“红色”有关系。如果只有一个注意力头那么模型只能学习一个注意力权重的分布这个分布是所有关系的混合。当信息本身包含多种关系时单头就要在多个目标之间做折中结果往往是每种关系都学得不够清楚。我早期实现自注意力时只用单头去跑一个简单的文本分类任务。模型能收敛但看注意力可视化时发现几乎所有词的注意力都集中在相邻词上很少出现长距离依赖。后来换成了 8 头模型在同样 epoch 下收敛更快注意力图也呈现出明显的多样性。这说明单头不是不能工作而是表达能力的上限比较低。1.2 不同头可以捕捉不同语义位置、语法、局部与全局多头注意力机制的核心意图就是让模型同时从多个子空间中学习不同种类的依赖关系。有的头可能学到“动词 宾语”这样的局部语法关系有的头可能学到“代词和它指代的名词”这种跨句长距离关系有的头则学到“位置相邻的词倾向高权重”这种位置关系。这些关系不一定都由人在设计时规定而是模型通过数据自动分化出来的。实验里常见的观察是低层 Transformer 的头更多学习局部、短距离的模式高层 Transformer 的头更容易关注长距离依赖某些头会呈现明显的语法行为比如 attention 集中到句号、逗号等标点上。这些都不是单头可以稳定学出来的。多头提供了一个更大的参数空间让不同头在梯度更新过程中逐渐分化。这个过程不是人为强制切分的而是通过随机初始化和不同投影矩阵的差异自然形成的。1.3 多头并不是把模型“变大”而是把注意力空间拆开有一个容易误解的地方多头注意力的参数量并没有变成原来的 N 倍。因为每个头的 Q、K、V 投影都是从一个完整的线性层里切分出来的。比如embed_dim7688 个头每个头的 Q 维度仍然是64 768 / 8总的 QKV 计算量和单头是近似相等的。多头只是把一个大矩阵拆成了 8 个小矩阵分别做注意力再把结果拼回去。这个“拆开再合并”的设计带来的是表达能力上的收益而不是参数规模的膨胀。理解这一点很重要否则你会误认为多头就是堆硬件从而在配置模型时把头数设得很大结果模型反而变慢变差。2. 多头注意力机制的原理拆解Q、K、V 和头数2.1 Query、Key、Value 的含义自注意力里每个 token 都会生成三个向量Query查询代表“我想找什么信息”Key键代表“我能提供什么信息”Value值代表“找到之后我实际要取到的内容”。注意力权重就是 Query 和 Key 的相似度Value 则按这个相似度进行加权求和。多头注意力就是在多头下各自做这个流程。很多人第一次接触时会把 Q、K、V 想得很玄。其实可以把它类比成一个检索过程你有一个问题Query然后去数据库里比对每个条目的索引Key算出一个相关分数最后根据分数把这些条目的内容Value加权取出来。这个类比在文本、图像和推荐系统里都适用。2.2 缩放点积注意力公式与为什么除以根号 d_k标准的注意力公式是Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) V除以根号 d_k 是为了平衡点积的方差。如果不缩放当维度 d_k 变大时点积的数值会迅速变大进入 softmax 的饱和区导致梯度非常小模型很难训练。缩放之后点积的方差大致保持在 1梯度可以正常回传。实际编码时需要特别注意这个缩放。很多人手写注意力时忘了除以sqrt(d_k)结果模型初始 loss 下降很慢训练不稳定。若加了缩放一开始的注意力权重分布会相对均匀训练过程也更好控制。2.3 头的概念怎么把一个 768 维向量切成 12 头假设输入序列中每个 token 的向量维度是d_model768num_heads12。多头做法是用三个权重矩阵分别把输入投影到 Q、K、V每个维度都是768再把 Q、K、V 的最后一维切成 12 份每份维度64对每一份分别计算缩放点积注意力得到 12 个输出把 12 个输出拼接回来变成 768 维最后经过一个输出投影矩阵再变回 768 维。切头的时候常见实现是view和transpose的组合。比如 PyTorch 中先view(batch_size, seq_len, num_heads, head_dim)然后transpose(1, 2)变成(batch_size, num_heads, seq_len, head_dim)。这里最容易出错的就是维度排列。如果顺序写错了模型不会直接报错但结果会乱掉训练 loss 会一直不降。3. 手推一遍多头注意力计算流程从输入到输出3.1 输入向量和词嵌入维度我们先设定一个具体场景句子长度为 10batch size 为 2d_model64num_heads4那么每个头维度d_h16。输入x的 shape 是(2, 10, 64)这是最常见的batch_first形式。如果使用 PyTorch 的nn.MultiheadAttention需要注意batch_first参数。默认是False输入是(seq_len, batch, embed_dim)设置batch_firstTrue后输入才是(batch, seq_len, embed_dim)。很多初学者在这个地方看文档不仔细导致维度不匹配报错或者 shape 对但结果完全不对。3.2 线性投影得到 Q、K、V输入x分别经过三个线性层得到q、k、v。这三个线性层的输入维度都是64输出维度也都是64。在单头的情况下这个过程非常简单。但在多头情况下这三个线性层的输出会被切分。实际代码里nn.MultiheadAttention内部把三个线性层合并成一个大的in_proj_weight方便一次计算。手写实现时为了方便理解可以分开写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)这里没有减少参数量只是把大矩阵拆开。对于d_model64, num_heads4每个头分到 16 维。投影之后Q、K、V 的形状仍然是(batch, seq_len, d_model)只是内部会被切分。3.3 按头拆分计算注意力权重拆分的方法是先 reshape再交换维度。batch, seq_len, d_model x.shape head_dim d_model // num_heads q q.view(batch, seq_len, num_heads, head_dim).transpose(1, 2) k k.view(batch, seq_len, num_heads, head_dim).transpose(1, 2) v v.view(batch, seq_len, num_heads, head_dim).transpose(1, 2)此时 q、k、v 的形状都是(batch, num_heads, seq_len, head_dim)。然后对每个头单独计算注意力分数scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(head_dim) attn torch.softmax(scores, dim-1) context torch.matmul(attn, v)得到的context形状是(batch, num_heads, seq_len, head_dim)这就是每个头各自的加权结果。这里 key 的转置要转最后两个维度不要转错成transpose(1,2)。之前我在手写时把k.transpose(-2,-1)写成了k.transpose(1,2)结果得分矩阵的 shape 不正确找了好一会才发现。3.4 加权求和、拼接、输出投影得到每个头的输出后需要把多个头的结果拼接回d_model维度。context context.transpose(1, 2).contiguous().view(batch, seq_len, d_model)注意transpose之后需要用contiguous()否则view会报错。这也是实际实现里非常常见的坑。拼接之后再过一层输出投影self.w_out nn.Linear(d_model, d_model) output self.w_out(context)到这一步一个完整的多头注意力层就结束了。总结下来流程是线性投影 - 切头 - 缩放注意力 - 拼接 - 输出投影。这个流程在几乎所有 Transformer 里都是一样的。4. PyTorch 实现多头注意力关键参数与掩码4.1 使用 nn.MultiheadAttention 的快速路径如果你只是想在模型里用多头注意力不关心内部细节直接用 PyTorch 封装好的接口最快。import torch.nn as nn mha nn.MultiheadAttention( embed_dim64, num_heads4, dropout0.1, batch_firstTrue, )调用时attn_output, attn_weights mha(query, key, value)如果query和key都是同一个输入这就是自注意力。attn_output是加了注意力之后的结果attn_weights是注意力权重矩阵可以用来可视化。这里有几个参数需要说清楚embed_dim输入输出的向量维度通常也是整个模型的隐藏层维度num_heads注意力头的数量dropout注意力权重上的 dropout 概率默认是 0.0batch_first控制输入输出维度的顺序。这种封装省事但不容易理解内部机制。我建议在学原理时还是手写一遍等理清楚了再回到封装接口。4.2 手写多头注意力的代码示例关键维度变化为了更清楚我给出一个最小可运行的多头注意力实现。这里只关注注意力核心部分省略了残差和层归一化。import math import torch import torch.nn as nn class MultiHeadAttention(nn.Module): def __init__(self, d_model64, num_heads4, dropout0.1): super().__init__() assert d_model % num_heads 0 self.d_model d_model self.num_heads num_heads self.head_dim 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_out nn.Linear(d_model, d_model) self.dropout nn.Dropout(dropout) def forward(self, x, maskNone): batch, seq_len, d_model x.shape q self.w_q(x) k self.w_k(x) v self.w_v(x) q q.view(batch, seq_len, self.num_heads, self.head_dim).transpose(1, 2) k k.view(batch, seq_len, self.num_heads, self.head_dim).transpose(1, 2) v v.view(batch, seq_len, self.num_heads, self.head_dim).transpose(1, 2) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn torch.softmax(scores, dim-1) attn self.dropout(attn) context torch.matmul(attn, v) context context.transpose(1, 2).contiguous().view(batch, seq_len, d_model) output self.w_out(context) return output, attn这个实现足够用于学习。需要说明的是真实 Transformer 中还会在输出投影后加 dropout这个可以按需调整。4.3 因果自注意力与 mask 处理所谓因果自注意力就是语言模型里每个 token 只能看到当前和之前的位置不能看到未来。实现方法是把上三角位置的注意力分数设置为一个非常大的负数比如-inf这样 softmax 之后这些位置的权重会变成 0。在nn.MultiheadAttention里attn_mask参数可以传入一个 shape 为(L, S)的布尔矩阵或浮点矩阵。在 PyTorch 里attn_mask如果类型是 boolTrue表示允许参与注意力如果类型是 float加法掩码被 mask 的位置会加上-inf。手写时可以先生成一个上三角掩码seq_len x.shape[1] mask torch.triu(torch.ones(seq_len, seq_len), diagonal1).bool() mask ~mask # 下三角和主对角线保留然后对 scores 执行scores scores.masked_fill(mask 0, float(-inf))这里的维度要注意。scores 是(batch, num_heads, seq_len, seq_len)而 mask 是(seq_len, seq_len)。PyTorch 的广播机制会先在后面补维度所以可以直接用。但是如果你传入的 mask 维度是(batch, seq_len, seq_len)则需要手动mask.unsqueeze(1)扩展到(batch, 1, seq_len, seq_len)。这个细节很容易忽略。4.4 训练和推理时 dropout 与 causal 的差异训练阶段注意力权重后面会接一个 dropout用来防止过拟合。推理阶段一般关掉 dropout并把模型设为eval()模式。如果你使用 PyTorch 的nn.Dropout它会在训练和 eval 之间自动切换不需要手动去处理。因果掩码在训练和推理时都需要加。训练时因为每个 token 都要同时作为预测目标因果掩码是必须的。推理时生成第 n 个 token同样不能让它看到未来所以也要加。换句话说因果掩码和训练还是推理无关它是模型结构的一部分。5. 多头之后为什么还要层归一化和残差连接5.1 残差避免梯度消失让深层网络可训练Transformer 的每一个多头注意力模块之后都有一段残差连接加上层归一化。残差连接的表达式是output x MultiHeadAttention(x)这样做的好处是梯度可以直接从深层传到浅层。如果去掉残差当 Transformer 层数增加时梯度在反向传播过程中会逐渐消失训练会变得极其困难。这也是从传统深层网络到 ResNet 再到 Transformer 一路沿用的经验。实际测试时我在一个 6 层 Transformer 上去掉残差训练 loss 基本不下降。加回残差后同样配置、同样数据loss 很快下降。这个差别非常明显所以残差不是可选项而是必备结构。5.2 层归一化稳定激活分布层归一化的作用是对一个 token 的所有特征维度做归一化使每个 token 的输出均值为 0、方差为 1然后通过可学习的缩放和平移参数恢复表达能力。它和 BatchNorm 的区别在于LayerNorm 不依赖 batch 内的其他样本所以在变长输入和 RNN、Transformer 场景下更稳定。在多头注意力后面接 LayerNorm可以防止注意力输出数值范围变化过大。注意力输出的量级受 softmax 影响虽然 softmax 把权重限制在 0 到 1 之间但 Value 本身的数值范围可能很大拼接之后的值域也不稳定。LayerNorm 把这些数值重新拉回到一个稳定的分布后续前馈网络会更安全。5.3 Post-Norm 和 Pre-Norm 的差异实际实现中层归一化和残差的先后顺序有两种Post-Normx LayerNorm(Attention(x))原始 Transformer 使用这种Pre-Normx Attention(LayerNorm(x))很多新模型使用这种。Post-Norm 更符合“先学习、再归一化”的思路但是深层模型训练时需要额外的 warmup。Pre-Norm 让梯度传播更稳定可以在更深的模型上不用 warmup但通常最终效果比 Post-Norm 略弱一点点。我实践下来如果是小模型、学习任务简单Post-Norm 问题不大。如果模型层数超过 12 层建议优先试 Pre-Norm。这篇文章只讲注意力层层归一化是它的直接配套所以放在这里一起说。6. 多头注意力的调试经验头数、模型维度和性能边界6.1 头数与维度的关系d_model % num_heads 0手写多头注意力时第一个硬性约束是d_model % num_heads 0。如果embed_dim768可选头数可以是 1、2、3、4、6、8、12 等。如果d_model800选 12 头就会报错因为 800 不能被 12 整除。这不是一个可以忽略的小问题。有些封装库会提示错误有些则会在初始化阶段就报错。但如果你自己写切分逻辑用view时维度不整除会直接抛异常。所以配置模型时我习惯先检查d_model % num_heads 0再开始写代码。6.2 头数不是越大越好过小欠拟合过大多样性下降头数过少模型表达能力不足注意力关系区分度不够头数过多每个头的维度太小单个头可能学不到足够的信息而且头之间容易冗余。比如d_model64如果设num_heads32每个头只有 2 维显然太极端。常见的比例是每个头维度 32 到 128 之间。对于d_model5128 头或 16 头都常见对于d_model76812 头是经典配置。经验判断标准如果训练 loss 很高、注意力可视化结果混乱可以增加头数如果模型变慢且注意力图看不出明显差异可能是头数过多。调头数时建议每次都重新初始化模型不要用上一次的 checkpoint 直接改头数。6.3 性能与显存多头注意力的计算成本多头注意力的计算复杂度主项是O(batch_size * num_heads * seq_len^2 * head_dim)整理后约等于O(batch_size * d_model * seq_len^2)。也就是说在d_model不变时头数改变不会显著改变总计算量。真正影响性能的是序列长度seq_len平方增长。显存方面注意力分数矩阵是(batch, num_heads, seq_len, seq_len)如果 batch 大、序列长、头数多显存会很快吃紧。此时可以考虑梯度检查点、混合精度训练、或者对长序列做稀疏注意力。但这些都是后续优化方向不要在第一次实现时就引入太多复杂度。6.4 判断注意力是否正常查看注意力权重分布和梯度训练完一个小模型后可以打印某个输入句子的注意力权重。正常的多头注意力应该呈现一定的模式比如有些头关注相邻词有些头关注特定的语法成分。如果所有头的注意力权重都非常均匀像均匀分布一样可能表示模型没训练好或者学习率设置太大。如果注意力权重很多是接近 0可能是 softmax 已经饱和要检查是否缺少缩放因子sqrt(d_k)。梯度方面最直接的办法是检查注意力层参数的梯度范数。如果梯度为 NaN通常是因为初始化、学习率或者输入数据本身含有 NaN。如果梯度很小可以适当调大学习率或检查残差连接。7. 常见报错与排查顺序维度、掩码、训练不稳定7.1 维度不匹配batch_first、embed_dim 和 head_dim遇到维度报错先看报错信息里的 shape再对照当前数据流分析输入 x 的 shape 是不是(batch, seq_len, d_model)线性层输入输出维度是否正确view时d_model能不能整除num_headstranspose之后是否需要用contiguous()。我见过最多的错误是view之前没有contiguous()。因为transpose之后 tensor 在内存里不是连续存储的直接view会抛出 “view size is not compatible” 之类的错误。解决方法是transpose(1,2).contiguous().view(...)。这个坑几乎每个人都会踩到。7.2 因果掩码失效forward 里 mask 要处理对在多头注意力里如果 mask 是二维的要确保它和 scores 维度兼容。scores是(batch, num_heads, seq_len_q, seq_len_k)mask通常需要是(seq_len_q, seq_len_k)或者(batch, seq_len_q, seq_len_k)。如果维度不对masked_fill可能不会报错但 mask 没有作用到正确的位置上。比如想实现因果自注意力直接用torch.triu(..., diagonal1)生成上三角矩阵这个矩阵中True代表需要被 mask 的位置。如果你把布尔值搞反了模型会看到未来训练 loss 表面上看不出来但推理时会崩。一个快速验证方法是构造一个简单的句子让模型预测下一个词如果生成时看到未来预测结果会“作弊”到令人惊讶的准确率但换到新数据就完全不行。7.3 训练不稳定检查 scale、dropout、初始化训练 loss 出现 NaN 或不下降时按照这个顺序排查检查输入数据是否包含 NaN 或极大值检查注意力里是否漏了1/sqrt(d_k)检查 softmax 的得分矩阵是否有-inf如果 mask 填错可能导致整行都是-infsoftmax 输出 NaN检查 dropout 是否在推理时没有关闭检查学习率是否太大Transformer 通常需要 warmup。如果所有设置都正常可以先用很小的模型、很小的随机数据跑一个过拟合测试。如果 loss 能降到很低说明注意力层本身能工作问题可能出在数据处理上。7.4 对比 nn.MultiheadAttention 和手写实现当手写实现行为异常时我通常把它和 PyTorch 官方的nn.MultiheadAttention输出做对比。给定同样的输入关闭所有随机性后两者的输出应该非常接近。如果不一致就去检查 QKV 初始化方式和缩放因子。官方实现里 Q、K、V 的投影是合并在一起的存在一个整体权重切分方式可能和手写默认的顺序不同。比如官方可能把 QKV 按顺序拼接成一个更大的矩阵切分时q取第一段k取第二段v取第三段。手写时如果分开线性层则不需要考虑这种切分顺序但要确保初始化分布一致。这种对比方法对定位 bug 非常有效。我自己写的多头注意力第一次输出和官方不一致原因就是漏了输出投影后的 dropout以及没有把注意力权重的维度调回正确顺序。逐个核对后两者数字基本一致我才放心继续往下搭 Transformer。最后再说一句我在实际项目里使用多头注意力时并不会盲目追求“头数多代表效果好”。更稳妥的做法是先用小模型把单条样本跑通检查输出 shape 和注意力权重分布再进入正式训练。如果这篇文章里的维度讲解和报错排查能让你少查半小时资料说明这些坑确实是普遍存在的。