ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

PyTorch手写多头注意力:原理、维度变换与避坑实践

PyTorch手写多头注意力:原理、维度变换与避坑实践 做Transformer相关的项目时我自己踩过最深的坑不是模型不收敛而是用 Torch 手写多头注意力Multi-Head Attention的时候Q、K、V 的维度一不留神就拼错一训练就报警告甚至直接抛出 shape 对不上的异常。后来我把 PyTorch 源码里的实现拆开看才真正理解 head 到底是怎么从张量里被“切”出来的也把很多平时只停留在公式上的细节给补上了。这篇就把多头注意力机制的实现过程、背后的设计原因以及实操中那些特别容易被忽略的细节一次性讲清楚。这篇文章适合这几类人正在读 Transformer 源码但被维度变换绕晕的人准备面试、需要从原理到实现都能讲明白多头注意力的人以及想自己改注意力层、做性能优化但不想被 PyTorch 的底层行为坑到的人。我会从原理讲到完整可运行的 Torch 代码再给出一套最小实验验证方法最后把常见的维度错误、安装和编译问题、KV Cache 优化过一次梳理完。1. 多头注意力机制从“一个问题”到“多个观察视角”1.1 一个头能捕捉的关系是有限的先回到注意力机制最基础的定义。对一个序列中的每个 tokenSelf-Attention 会计算它和其他所有 token 之间的相关度然后用这个相关度去加权聚合其他 token 的信息。公式长这样[ Attention(Q, K, V) softmax\left(\frac{QK^T}{\sqrt{d_k}}\right)V ]其中 (Q) 是查询(K) 是键(V) 是值。直观理解就是每个 token 发出一条查询商品有自己的“键”通过点积算相似度再用 softmax 得到权重最后按权重拿回“值”里的信息。问题在于如果只用一组 Q、K、V整个模型从一而终只有一个“观察视角”。这个视角要么更关注局部语法关系要么更关注远程指代要么更关注位置关系很难同时把多种关系都抓到位。比如“苹果不便宜但口感很好”这句话里既要识别到“苹果”和“口感”的修饰关系又要捕捉到“不便宜”和“口感好”之间的转折关系。单头注意力很容易顾此失彼因为它只有一套权重去权衡所有信息。1.2 多头注意力让每个头专注一种关系多头注意力的做法是不把 Q、K、V 当一个整体用而是先把它们的特征维度切分成多个子空间每个子空间对应一个“头”。每个头独立做一次注意力计算最后再把结果拼起来做一次线性投影输出。原论文里给的是这样的形式[ head_i Attention(QW_i^Q, KW_i^K, VW_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) 是输出投影矩阵。这样做的核心好处在于不同的头可以学习到不同类型的依赖关系。实际训练中你会发现有的头的 attention 分布特别集中基本聚焦在相邻词上有的头分布很均匀更像是某种平滑还有的头会形成比较清晰的“句法模式”。这就是多头带来的能力上限提升——模型的表达能力不再被单一视角限制住。别把多头理解成“模型大了、参数多了所以模型更强”。多头本质上给了模型一组并行的“特征子空间”每个头可以单独盯一类关系最后融合的时候再互相补充。这也是为什么很多实验里在总维度不变的情况下多头确实比单头收敛更快、效果更稳。1.3 参数与维度约定别把 head_dim 和 embed_dim 搞混实现之前必须先搞清楚几个关键维度embed_dim或d_model每个 token 的输入特征维度。num_heads头的数量一般记为 (h)。head_dim每个头分到的维度通常等于embed_dim // num_heads。以标准配置为例d_model 512h 8那么head_dim 64。每个头拿到 512 维里的其中 64 维做属于自己的注意力计算。核心的维度变化流程是输入 (B, T, d_model) - 线性投影得到 Q/K/V仍然是 (B, T, d_model) - 拆成 (B, T, h, head_dim) - 转置为 (B, h, T, head_dim) - 计算注意力 - 转置回 (B, T, h, head_dim) - 重新拼成 (B, T, d_model) - 输出投影这里最容易搞混的一点是每个头的 Q/K/V 并不是“单独用一个小 Linear 层算出来的”而是先用一个大 Linear 把输入映射到d_model维再通过view和transpose在张量里切分。后面实现时会看到这个细节对整个张量布局影响非常大。2. 用 Torch 手写多头注意力完整代码与逐步拆解2.1 模块初始化投影矩阵怎么放先给出一个完整的、可以直接放进项目里的MultiHeadAttention模块。我没有直接调torch.nn.MultiheadAttention而是把整个过程拆开写这样更容易看到里面的每一步。import math import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadAttention(nn.Module): def __init__(self, embed_dim, num_heads, dropout0.0, biasTrue): super().__init__() assert embed_dim % num_heads 0, \ fembed_dim({embed_dim}) must be divisible by num_heads({num_heads}) self.embed_dim embed_dim self.num_heads num_heads self.head_dim embed_dim // num_heads self.dropout nn.Dropout(dropout) # 四个投影矩阵每个都是 embed_dim - embed_dim self.q_proj nn.Linear(embed_dim, embed_dim, biasbias) self.k_proj nn.Linear(embed_dim, embed_dim, biasbias) self.v_proj nn.Linear(embed_dim, embed_dim, biasbias) self.out_proj nn.Linear(embed_dim, embed_dim, biasbias) self._reset_parameters() def _reset_parameters(self): # Transformer 里常见的初始化方式保证训练初期更稳定 for proj in [self.q_proj, self.k_proj, self.v_proj, self.out_proj]: nn.init.xavier_uniform_(proj.weight) if proj.bias is not None: nn.init.constant_(proj.bias, 0.0)初始化中有几个值得注意的点。第一四个投影矩阵都是embed_dim - embed_dim没有把 Q/K/V 到每个头的投影拆成独立小矩阵。这样处理是出于实现方便和效率考虑PyTorch 的nn.Linear内部是矩阵乘法大矩阵相乘的 GPU 利用率通常好过一堆小矩阵。拆开写也不难但没必要。第二embed_dim能否整除num_heads必须开头断言。如果不整除后面的view(B, T, num_heads, head_dim)会直接报错而且错误信息非常不友好。提前断言至少能让问题更早暴露。第三Q/K/V 的偏置bias是不是需要实际项目里有争议。原版 Transformer 的 Q/K/V 投影是不带 bias 的但 PyTorch 官方实现里默认带。我这里保留了 bias使用时可以通过参数关闭。2.2 forward 中的 Q/K/V 变换为什么是 view transpose前向计算是整个模块的核心代码里每一步都值得仔细看def forward(self, query, key, value, attn_maskNone, need_weightsFalse): B, T_q, _ query.shape T_k key.size(1) H self.num_heads D self.head_dim # 1. 线性投影 q self.q_proj(query) # (B, T_q, embed_dim) k self.k_proj(key) # (B, T_k, embed_dim) v self.v_proj(value) # (B, T_k, embed_dim) # 2. 拆分成多头并调整张量维度顺序 # view: (B, T, embed_dim) - (B, T, H, D) # transpose: (B, T, H, D) - (B, H, T, D) q q.view(B, T_q, H, D).transpose(1, 2) k k.view(B, T_k, H, D).transpose(1, 2) v v.view(B, T_k, H, D).transpose(1, 2) # 3. 计算缩放点积注意力分数 scores q k.transpose(-2, -1) / math.sqrt(D) # scores: (B, H, T_q, T_k) # 4. 处理 mask if attn_mask is not None: if attn_mask.dtype torch.bool: # True 表示保留False 表示 mask 掉 scores scores.masked_fill(~attn_mask, float(-inf)) else: # float 类型作为 additive mask直接加到分数上 scores scores attn_mask # 5. softmax dropout attn_weights torch.softmax(scores, dim-1, dtypetorch.float32) attn_weights attn_weights.to(scores.dtype) attn_weights self.dropout(attn_weights) # 6. 加权聚合 out attn_weights v # (B, H, T_q, D) # 7. 把头维度合并回特征维度 out out.transpose(1, 2).contiguous().view(B, T_q, self.embed_dim) # 8. 输出投影 out self.out_proj(out) if need_weights: return out, attn_weights return out很多人第一次看这段代码时不理解为什么不能直接reshape非要viewtranspose最后还要contiguous。实际上view不会改变张量在内存中的数据顺序它只是重新解释这个张量。transpose则是交换维度交换后张量的内存布局变成了非连续的non-contiguous。所以当你transpose之后直接reshape就会碰到 “shape is invalid for input size” 之类的问题。最稳妥的做法就是transpose - contiguous - view先让内存布局重新规整再把它重新映射成想要的形状。这一步非常容易踩坑尤其是从 PyTorch 的nn.MultiheadAttention源码里抄代码时少写一个contiguous()训练立刻崩。另外为什么要把(B, T, H, D)转成(B, H, T, D)因为点积计算q k.transpose(-2, -1)是在最后两个维度上做的我们要让每个头分别计算自己那个子空间内所有 token 之间的注意力分数。如果把头维度放在 1 维batch在 0 维head在 1 维seq在 2 维这样通过广播机制所有头的注意力计算可以一次完成不需要写 for 循环。2.3 注意力权重计算与 scaling为什么要除 sqrt(head_dim)注意力分数是q和k的点积。在实现里我没有直接去掉缩放因子而是用了scores q k.transpose(-2, -1) / math.sqrt(D)这个D是head_dim。为什么要除以 (\sqrt{d_k})原因和 softmax 的梯度性质有关。如果 Q、K 的每个元素大致是均值 0、方差 1 的分布那么点积 (q \cdot k) 的均值是 0但方差大约是 (d_k)。当 (d_k) 变大时点积的绝对值也会变大softmax 的输入会被推到饱和区。softmax 进入饱和区之后对输入的梯度会非常小模型很难继续学下去。除以 (\sqrt{d_k}) 之后点积的方差被拉回到 1 左右softmax 的输入分布更温和梯度能正常回传。这就是为什么原论文里要特意加这个缩放而不是单纯为了让数值好看。从我自己的测试经验来看这个缩放因子对训练稳定性影响很大。去掉它之后小型模型在训练早期非常容易 loss 直接变 NaN尤其是用 FP16 混合精度训练的时候。所以不要觉得“反正 softmax 会归一化乘不乘无所谓”这个缩放必须先做。2.4 Mask 的正确用法padding mask 与 causal mask多头注意力里并不是所有 token 对都允许互相 attend。常见的两种 mask 是 padding mask 和 causal mask。Padding mask的作用是让某些 key 位置是 padding token不应该参与注意力计算。它的形状通常是(B, 1, 1, T_k)然后广播到(B, H, T_q, T_k)这个四维权重的维度上。注意 pad mask 的维度必须和scores能广播建议直接写成 4 维。# key_padding_mask: (B, T_k)True 表示是 padding key_padding_mask torch.zeros(B, T_k, dtypetorch.bool) key_padding_mask[:, -5:] True # 假设最后 5 个位置是 padding # 扩展维度用于广播 mask ~key_padding_mask.unsqueeze(1).unsqueeze(2) # (B, 1, 1, T_k) # 在 forward 里传入 attn_mask maskCausal mask则用于自回归场景保证当前位置只能看到当前位置和之前的位置看不到未来。实现方式很简单causal_mask torch.tril(torch.ones(T_q, T_k, dtypetorch.bool)) # causal_mask.shape (T_q, T_k)然后传入attn_mask causal_mask在 forward 里布尔 mask 分支会执行masked_fill(~attn_mask, -inf)把右上角的位置全部置为负无穷。这里有一个很关键的细节mask 的处理必须在 softmax 之前。如果你先 softmax 再 mask得到的结果根本不是“这个位置不参与”而是“这个位置被分配了很小但非零的权重”等于 mask 失效。还有一个容易被忽略的问题-inf在注意力分数里会不会导致 NaN。理论上不会因为 softmax 会把-inf映射成 0 概率。但如果某一行的所有位置都被 mask 成-infsoftmax 会得到nan因为这一步在数学上是 0/0。你在实现时一定要保证每个 query 位置至少能 attend 到一个有效的 key 位置。如果自己的数据里出现了全 padding 的行最好在数据侧先过滤掉。3. 跑一个最小实验验证多头注意力实现3.1 造一个“字符预测”玩具数据集写一个模块之后不能只看 forward 能跑通就算完。我一般会先用一个特别小的任务验证梯度是否正常、loss 是否能降下去。这里推荐用“字符级预测”任务给出一段离散符号序列模型要预测下一个符号。这个任务本身很难但作为注意力模块的验证非常合适因为数据可以随机生成不依赖外部数据集。我把整个训练代码精简成下面这样class CharDataset(torch.utils.data.Dataset): def __init__(self, seq_len, num_samples, vocab_size): self.seq_len seq_len self.num_samples num_samples self.vocab_size vocab_size self.data torch.randint(0, vocab_size, (num_samples, seq_len 1)) def __len__(self): return self.num_samples def __getitem__(self, idx): x self.data[idx, :-1] y self.data[idx, 1:] return x, y这里每个样本是一串随机的离散 ID输入前seq_len个预测后移一位的seq_len个。虽然不像真实文本那样有语义结构但足以验证多头注意力的前向、反向、mask 和投影层整个过程是否正确。3.2 组装一个极简 Transformer 块验证用的模型不需要完整 Transformer。我把一个多头注意力层和一层 FFN 拼起来组成一个最简的 Transformer 块就能看出问题。class MinimalFormerBlock(nn.Module): def __init__(self, d_model, num_heads, ff_dim, dropout0.1): super().__init__() self.attn MultiHeadAttention(d_model, num_heads, dropoutdropout) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.ffn nn.Sequential( nn.Linear(d_model, ff_dim), nn.GELU(), nn.Linear(ff_dim, d_model), nn.Dropout(dropout), ) def forward(self, x, attn_maskNone): x x self.attn(self.norm1(x), self.norm1(x), self.norm1(x), attn_maskattn_mask) x x self.ffn(self.norm2(x)) return x训练时用 CrossEntropyLoss优化器选择比较简单model MinimalFormerBlock(d_model64, num_heads4, ff_dim128, dropout0.1) model.train() optimizer torch.optim.AdamW(model.parameters(), lr5e-4) loss_fn nn.CrossEntropyLoss() for step, (x, y) in enumerate(dataloader): logits model(x) # (B, T, d_model) logits model.out_proj(logits) loss loss_fn(logits.reshape(-1, vocab_size), y.reshape(-1)) optimizer.zero_grad() loss.backward() optimizer.step()注意out_proj是我追加的一个输出层负责把d_model映射到词表大小。实测下来当d_model64、num_heads4、seq_len32、batch_size64时这个模型几十步内 loss 就能明显下降说明多头注意力的前向和反向都是通的。如果跑下来 loss 保持不动或者训练中直接报错那大概率不是数据集的锅而是模块内部的张量维度或者 mask 处理有问题。3.3 用 gradcheck 和 loss 曲线双重检查除了看 loss 是否下降还可以用torch.autograd.gradcheck做数值梯度校验。这个工具会拿解析梯度和数值梯度做对比如果差异太大就说明反向传播实现有问题。from torch.autograd import gradcheck x torch.randn(2, 4, 16, dtypetorch.float64, requires_gradTrue) model MultiHeadAttention(embed_dim16, num_heads4).double() # gradcheck 对浮点精度要求很高需要用 float64 test gradcheck(model, (x, x, x), eps1e-6, atol1e-4) print(test)我自己使用时有一个习惯先只测不带 mask 的情况通过之后再测带 causal mask 的情况。一旦加了-inf数值梯度的计算可能会因为掩码边界导致较大的误差这时可以把atol调大一点或者直接用fast_modeTrue。如果gradcheck返回 False优先检查是不是在 forward 里用了in-place操作比如scores mask这种写法很容易让 autograd 的梯度路径断裂。4. 常见问题速查与性能优化经验4.1 维度不对、device 不一致等报错怎么破多头注意力实现里的大多数报错翻来覆去就那么几类。我列了一个排查表碰见问题可以先对着查。报错信息常见原因解决办法shape ... is invalid for input size ...张量不是 contiguous 就 view/reshape先.contiguous()再viewExpected tensor to be on the same deviceQ/K/V 或 mask 没有移动到同一设备统一.to(device)mask 也要送 GPUDimension out of range把只适用于 3D 张量的操作放在 4D 张量上检查当前张量形状是否为(B, H, T, D)AssertionError: embed_dim must be divisibleembed_dim不能整除num_heads调整num_heads或embed_dimexpected scalar type Long but found Floatmask 类型混用布尔 mask 用torch.booladditive mask 用floatloss.backward() 之后梯度是 None某个输入没有参与计算图或者 mask 把整行都 mask 掉了检查 mask 是否让某行全为-inf这里最值得强调的是 mask 类型。我建议在自己模块里把 mask 分为两个入口一个接收attn_mask一个接收key_padding_mask然后统一转换成 additive mask。这样调用方不容易混。PyTorch 官方nn.MultiheadAttention就是这么设计的目的就是避免 mask 语义互相干扰。4.2 torch 安装失败 / FlashAttention 编译失败怎么解决这个问题几乎每个项目都遇到过。最常见的是命令行执行pip install torch却提示ERROR: Could not find a version that satisfies the requirement torch。这类错误十有八九不是 torch 本身的问题而是 Python 版本和你当前环境能匹配到的 torch 版本对不上。比如太新的 Python 3.12、3.13 在某些旧 torch 版本里还没有预编译包pip就会直接说找不到版本。我现在的固定做法是先执行python -V确认 Python 版本。到 PyTorch 官方安装命令页面选对应的安装命令不要只敲一个裸的pip install torch。安装时用pip install torch2.4.0这种带版本号的方式避免装到最新版导致依赖冲突。如果要匹配 CUDA优先用 PyTorch 官方给出的源不要在来源不明的代码仓库里下载压缩包自己解压。接着是 FlashAttention 的编译问题。如果你的环境是 Python 3.12、CUDA 12.9、Torch 2.4 这种比较新的组合FlashAttention 的预编译 wheel 很可能还没跟上。此时强行源码编译经常会卡在ninja、gcc或 CUDA toolkit 版本不匹配的问题上。我的建议是如果项目不需要极致的长序列性能先退回 PyTorch 自带的F.scaled_dot_product_attention它内部已经能自动选择高效的注意力内核。等 FlashAttention 的 wheel 版本覆盖到你的环境之后再考虑切换没必要死磕编译。4.3 推理阶段的 KV Cache 与长序列优化训练和推理对注意力实现的关注点是不一样的。训练时我们通常把完整序列一次性喂进去计算所有 token 的注意力。但推理是逐 token 生成的如果每次都重新计算历史 token 的 K、V浪费会非常大。KV Cache 的思路是把已经计算过的 key 和 value 缓存下来每次只计算当前 token 的 Q、K、V然后把当前 K、V 拼到缓存后面再做注意力。我用伪代码表示一下这个逻辑def infer_step(model, token_ids, past_k, past_v): x embedding(token_ids) # 对每一层注意力 # 1. 只对当前 token 做 k_proj/v_proj # 2. past_k cat([past_k, k], dimseq) # 3. past_v cat([past_v, v], dimseq) # 4. attention 只拿 past_k/past_v 和当前 q 计算 return logits, new_k, new_v但注意KV Cache 会让 mask 的形状也变化。过去完整的(T_q, T_k)下三角 mask 已经不再适用因为推理时 T_q 通常是 1T_k 会越来越大。此时只需要保证当前 query 能看到所有缓存的 key不再需要额外的 causal mask——因为你只生成未来 token当前 query 天然就只能“看到”历史。如果你做的是长序列任务又不方便用 FlashAttention还可以考虑这几个方向使用F.scaled_dot_product_attentionPyTorch 2.x 会自动选择内存高效的 attention 实现。减小head_dim比如保持embed_dim不变增加num_heads但注意总计算量会变。推理时把不需要的中间张量及时释放或者用torch.no_grad()包裹生成循环。4.4 关于多头数量不是越多越好最后聊一个很容易被忽略的细节。很多初学者以为num_heads越大越好但实际并不是。头数增加之后每个头的head_dim会相应变小如果head_dim太小比如只有 8 甚至 4单个头的表达能力会受限注意力分布很容易变得不稳定。反过来头数太少又会让多头机制失去意义。常见的配置是d_model512, num_heads8或者d_model768, num_heads12大致保持head_dim在 64 左右。我自己在调模型时会做一个小实验固定d_model分别跑 1 头、2 头、4 头、8 头观察 loss 下降速度和最终指标。多头的收益通常在头数从 1 提升到 4 时非常明显再往上就是锦上添花甚至还会引入训练不稳定的问题。如果你在项目里调试多头注意力建议也保留一个“可配置 head 数”的参数别写死。我自己在实际项目里最常犯的错误永远是transpose之后忘了contiguous()。现在养成了一个习惯只要在多维张量上做了维度交换下一步要view或reshape之前都会先确认张量是否连续。另一个能省很多时间的小技巧是调试多头注意力时先用seq_len8、d_model16这种极小的尺寸跑通 gradcheck再去碰真实数据。大模型里那些张量形状问题多数在小尺寸场景下一眼就能看出来。
返回列表