用PyTorch手写注意力机制:从单头到多头完整实现 📅 发布时间:2026/9/1 2:18:50 👁 浏览次数: 注意力机制是 Transformer 的核心也是很多人在读完论文后第一个想手动实现的部分。但论文里的公式和实际代码之间有一层“工程间隙”维度怎么变换、掩码怎么传、多头怎么合并、训练时和推理时有什么区别。这篇文章不准备讲过于理论化的内容就用 PyTorch 从零搭建一个可用于分类或序列任务的注意力模块把 QKV、缩放点积、多头合并、掩码这些关键点拆开讲清楚。适合已经了解 Transformer 基本结构但还没手写过完整代码的读者。我建议你不要一上来就复制大段模型代码先按下面这几步走先理解自注意力的数据流向再写一个最简单的单头版本跑通之后改成多头最后再处理掩码和批量输入。这样每一步出问题都能定位到具体环节而不是面对几百行代码无从下手。1. 自注意力到底在计算什么先别急着写代码很多人第一次看 Transformer 论文时最困惑的是 Q、K、V 这三个矩阵。其实可以把自注意力理解成一种“带权重的信息汇总”序列里的每个位置都会发出一个查询然后在其他所有位置上做匹配匹配度高的位置贡献更多信息最后把这些信息按权重混合起来。整个过程不依赖循环结构所以可以并行计算长距离依赖这也是 Transformer 相比 RNN 的优势所在。1.1 Q、K、V 不是三个独立的东西而是同一输入的三种投影假设输入序列的形状是 (batch_size, seq_len, d_model)。Q、K、V 本质上是对同一组输入分别做线性变换后得到的结果只是它们的用途不同QQuery代表“当前位置想找什么”。KKey代表“每个位置能提供什么的标签”。VValue代表“每个位置实际携带的信息”。计算注意力分数时用 Q 和 K 做点积。点积越大说明这个位置在当前上下文中越相关。随后用 softmax 归一化成权重再对 V 做加权求和得到当前层输出。这里最容易踩坑的点是缩放因子。为什么点积之后要除以 sqrt(d_k)因为当 d_k 较大时点积数值会很大softmax 的输入落在梯度很小的区域反向传播时梯度容易消失。除以 sqrt(d_k) 可以让注意力分数的方差保持在合理范围训练更稳定。这个细节在论文里写了但很多实现里容易忽略尤其是从别的框架迁移代码时。1.2 单头注意力是整个实现的最小单元先把多头放一边写一个只有单个注意力头的模块。这个模块只需要四步对输入做 Q、K、V 线性投影。计算 Q 与 K 的点积并缩放。对注意力分数做 softmax。用注意力权重对 V 加权求和。用 PyTorch 实现时可以通过torch.matmul一次完成批量矩阵乘法不需要写循环。这个版本的代码虽然简单但对理解数据张量的形状变化很有帮助。import torch import torch.nn as nn import torch.nn.functional as F class SingleHeadAttention(nn.Module): def __init__(self, d_model, d_k, d_v): super().__init__() self.d_k d_k self.d_v d_v self.w_q nn.Linear(d_model, d_k) self.w_k nn.Linear(d_model, d_k) self.w_v nn.Linear(d_model, d_v) def forward(self, x, maskNone): # x: (batch_size, seq_len, d_model) Q self.w_q(x) # (batch_size, seq_len, d_k) K self.w_k(x) # (batch_size, seq_len, d_k) V self.w_v(x) # (batch_size, seq_len, d_v) scores torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5) # scores: (batch_size, seq_len, seq_len) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn_weights F.softmax(scores, dim-1) output torch.matmul(attn_weights, V) # output: (batch_size, seq_len, d_v) return output先跑通这个版本。你可以在任意一批形状为 (2, 10, 64) 的随机数据上验证一下正常情况输出形状应该是 (2, 10, d_v)。如果这一步输出尺寸不对说明矩阵乘法维度没对齐。1.3 为什么先写单头而不是直接写多头单头版本代码量小心智负担低。它可以让你把注意力机制本身和数据流看清楚Q、K、V 的张量形状变化、softmax 沿哪个维度计算、mask 在哪个时机生效。我认为学习路线应该是单头理解原理多头用于实战。如果你直接看多头代码很容易被维度变换绕晕尤其是当 d_model 和 n_heads 不完全整除的时候拼接和 reshape 的顺序会让人头大。先单头后多头可以避免“功能跑通了但根本不知道代码在干什么”的问题。2. 从单头扩展到多头维度拆分与合并是关键多头注意力的核心思想不是让模型只关注一个“匹配关系”而是让它在不同的表示子空间里分别做匹配最后把结果拼接起来。通俗点说多个头可以分别关注不同的词与词之间的关联类型有的头可能更关注语法关系有的头可能更关注语义相似度。2.1 多头实现的标准顺序多头注意力在代码上一般有两种实现方式。一种是把 Q、K、V 都拆成多头分别计算注意力再合并结果另一种是先把线性层做完再在特征维度上拆分成多个头。第二种方式在实践中更常见也方便批量计算。关键点在于张量形状的变化。假设 d_model 是 512n_heads 是 8那么每个头的维度是 64。输入 x 经过线性层得到 Q形状为 (batch_size, seq_len, 512)然后把最后那个维度拆成 8 组每组 64 维再把这 8 组看成独立的批次维度以便一次性完成所有头的矩阵运算。具体的拆分操作是# Q: (batch_size, seq_len, d_model) batch_size, seq_len, _ Q.size() Q Q.view(batch_size, seq_len, n_heads, d_k).transpose(1, 2) # Q: (batch_size, n_heads, seq_len, d_k)view将最后一维拆成 n_heads 和 d_ktranspose把 n_heads 放到 batch 之后。这样后面计算注意力分数时每个头独立计算互不干扰。2.2 完整的多头注意力代码下面是一个不带掩码的多头注意力实现适合做序列分类和一般特征提取任务。class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_heads): super().__init__() assert d_model % n_heads 0, d_model must be divisible by n_heads self.d_model d_model self.n_heads n_heads self.d_k d_model // n_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.out_proj nn.Linear(d_model, d_model) def forward(self, x, maskNone): batch_size, seq_len, _ x.size() Q self.w_q(x) K self.w_k(x) V self.w_v(x) # 拆分多头 Q Q.view(batch_size, seq_len, self.n_heads, self.d_k).transpose(1, 2) K K.view(batch_size, seq_len, self.n_heads, self.d_k).transpose(1, 2) V V.view(batch_size, seq_len, self.n_heads, self.d_k).transpose(1, 2) scores torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn_weights F.softmax(scores, dim-1) context torch.matmul(attn_weights, V) # context: (batch_size, n_heads, seq_len, d_k) # 合并多头 context context.transpose(1, 2).contiguous().view( batch_size, seq_len, self.d_model ) output self.out_proj(context) return output注意最后合并多头时要先transpose把 n_heads 换回原位置再做view。如果没有加contiguous()部分 PyTorch 版本会报错因为transpose后的张量在内存中不一定连续。这是一个很小的点但很多人第一次手写时都会卡在这里。2.3 头数怎么选不是越多越好头数过少模型的表示能力不足头数过多每个头分到的维度变小模型可能学不到足够丰富的特征。在常见的 Transformer 配置里d_model 为 512 时头数常取 8d_model 为 768 时头数常取 12。这样每个头的维度基本维持在 64 左右。如果你的任务是文本分类这类中等规模任务d_model 取 128 或 256n_heads 取 4 或 8通常是够用的。不要一上来就仿照大型模型的参数配置在小数据集上很容易过拟合而且训练速度会明显变慢。3. 掩码机制padding mask 和 causal mask 的差异很多人在看注意力代码时最难理解的是 mask 的作用。其实 mask 的用途只有两类一类是让模型忽略无效位置一类是防止模型看到未来信息。3.1 padding mask让模型不关注填充位序列数据经常长度不一致实际使用时需要对短句子做 padding填充到相同长度。填充位置没有真实语义如果不处理模型会把注意力分配到这些填充位上影响最终特征提取。padding mask 一般是一个形状为 (batch_size, seq_len) 的布尔张量有效位置为 True填充位置为 False。在计算 scores 时用masked_fill(mask 0, float(-inf))把填充位置变成负无穷softmax 之后这些位置的权重就会变成接近 0。这里要注意 mask 的维度。在多头注意力中scores 的形状是 (batch_size, n_heads, seq_len, seq_len)所以 mask 需要是 (batch_size, 1, 1, seq_len) 的形状才能正确广播。如果你的 mask 维度不对训练时不一定报错但结果会完全错误。我曾经在这个问题上排查了很久最后发现是 mask 少扩展了一维。3.2 causal mask自回归任务的关键如果任务是语言模型、文本生成这类自回归任务每个位置只能关注它之前的位置不能看到未来。做法是在 scores 矩阵的上三角部分填充负无穷。实现方式很简单def create_causal_mask(seq_len): return torch.triu(torch.ones(seq_len, seq_len), diagonal1).bool()这个矩阵的右上角为 True含义是“第 i 行不能看到第 j 列当 j i 时”。把它传给masked_fill之后未来位置的注意力分数会变成负无穷。需要注意的是如果你同时处理 padding 和因果掩码要把两个 mask 合并后再传入。3.3 为什么 softmax 之前的负无穷处理不能省有人可能觉得把无效位置先乘一个很小的权重不就行了实际上不行。softmax 是对所有分数做指数归一化如果无效位置不是负无穷它的指数值仍然是一个有限正数会占据一部分概率质量。只有使用负无穷softmax 后这些位置的权重才会变成 0。但这里有个数值稳定性问题。如果 scores 中同时出现了很大的正数和负无穷softmax 在计算exp(x)时可能溢出。PyTorch 的softmax内部会做数值稳定处理一般情况下没问题。不过如果你在自定义 CUDA 内核或低层框架里实现就要自己先减去最大值再做指数运算。4. 把注意力模块放进一个可训练的小模型里单看注意力模块是没法直接做任务的你需要把它和输入嵌入层、位置编码、前馈网络组合起来才能形成一个完整的可训练模型。这里我建议先做一个“最小可用版”嵌入层 位置编码 一层多头注意力 前馈网络 分类头。4.1 位置编码不能省略但可以用简单的实现Transformer 没有循环结构如果不加位置信息模型会认为输入序列是一个无序集合。位置编码最常见的方式是使用 sin 和 cos 函数但为了减少代码复杂度也可以直接使用可学习的位置嵌入。在中小规模任务上可学习位置嵌入效果通常不比三角函数差。PyTorch 里可以直接使用nn.Embedding(max_len, d_model)然后加上序列索引。这样做的好处是实现简单坏处是序列长度超过训练时的最大长度时需要插值或重新训练。所以生产环境里如果你要处理不定长文本还是建议用官方 Transformer 里的静态位置编码实现。4.2 一个用于文本分类的最小模型class TransformerEncoderBlock(nn.Module): def __init__(self, d_model, n_heads, d_ff, dropout0.1): super().__init__() self.attention MultiHeadAttention(d_model, n_heads) self.norm1 nn.LayerNorm(d_model) self.ffn nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model) ) self.norm2 nn.LayerNorm(d_model) self.dropout nn.Dropout(dropout) def forward(self, x, maskNone): # 先注意力再残差再 LayerNorm attn_out self.attention(x, mask) x self.norm1(x self.dropout(attn_out)) # 再前馈再残差再 LayerNorm ffn_out self.ffn(x) x self.norm2(x self.dropout(ffn_out)) return x这个结构里残差连接和 LayerNorm 都有各自的理由。残差连接保证深层网络梯度能顺畅回传LayerNorm 稳定每层输出的分布。如果你去掉残差只堆叠多层注意力训练时 loss 下降会很慢甚至不收敛。4.3 前馈网络层 d_ff 怎么设置Transformer 论文里 d_ff 通常取 d_model 的 4 倍。d_model 是 512 时d_ff 是 2048。这个参数不是固定不变的任务简单时可以缩小到 2 倍任务复杂时可以适当放大。但要记住d_ff 越大参数量增长越明显。前馈网络里的两个线性层是 Transformer 中参数量的大头占比往往超过注意力部分。在低显存环境下做实验优先减小 d_ff而不是减少头数这样对模型表现影响更可控。5. 完整实现一个基于注意力机制的文本分类例子前面几节已经把所有模块拆分讲完这一节把它们组装成一个能跑通训练流程的文本分类模型。我以简单的中文文本分类任务为例但数据量很小主要用于验证代码路径而不是真正追求准确率。5.1 训练一个最小可运行的分类模型import torch import torch.nn as nn class SimpleTransformerClassifier(nn.Module): def __init__(self, vocab_size, d_model128, n_heads4, d_ff512, num_classes2, max_len128): super().__init__() self.embedding nn.Embedding(vocab_size, d_model) self.pos_embedding nn.Embedding(max_len, d_model) self.encoder TransformerEncoderBlock(d_model, n_heads, d_ff) self.pooler nn.AdaptiveAvgPool1d(1) self.classifier nn.Linear(d_model, num_classes) def forward(self, input_ids, maskNone): seq_len input_ids.size(1) positions torch.arange(seq_len, deviceinput_ids.device).unsqueeze(0) x self.embedding(input_ids) self.pos_embedding(positions) x self.encoder(x, mask) x x.transpose(1, 2) # (batch_size, d_model, seq_len) x self.pooler(x).squeeze(-1) # (batch_size, d_model) return self.classifier(x)我这里只叠加了一层编码器。如果你想要更好的效果可以循环堆叠多层TransformerEncoderBlock但每一层的内存占用都会线性增加。第一次跑通时先不要超过两层。5.2 训练时的几个判断标准训练开始后你要重点观察几个指标而不是只看最终准确率第一个 batch 的 loss 有没有明显下降。如果没有下降先检查学习率是否过大以及数据有没有正常流入模型。训练到第几个 step 时 loss 开始下行。一般第一个 epoch 内就应该看到明显趋势如果完全不动大概率是 mask 维度或标签类型有问题。验证集上的 loss 和训练集差距是否拉大。如果训练 loss 快速下降但验证 loss 上升说明过拟合可以先加 dropout或者减少模型层数。我一般会在训练前先做一次“单 batch 过拟合测试”拿一个 batch 的少量样本训练几十步如果模型能够在训练集上快速记住这批样本说明代码路径是通的。如果连单 batch 都学不动那可能是模型结构、学习率或数据预处理的问题这时候不要在更大数据集上浪费时间。5.3 验证输出是否正确的简单方法你也可以用 PyTorch 自带的nn.MultiheadAttention作为参照对比输出维度。但要注意PyTorch 内置实现默认输入格式是 (seq_len, batch_size, d_model)和很多自定义实现不同。你需要先把输入改成对应的维度顺序或者直接对比一个较小矩阵的数值结果。更常见的验证方式是自己构造一组简单数据比如只有两个句子每个句子长度不一样加上 padding mask观察注意力权重矩阵。打印出各个位置在句子内部的注意力分布看看填充位置是否真的被忽略。这个检查在调试阶段非常有效。6. 常见问题与排查顺序手写注意力机制时大部分报错都不是“模型设计错了”而是张量形状、掩码维度、设备不一致或参数初始化问题。下面按实际排查顺序整理一下。6.1 形状错误优先看 QKV 的变换最常见的形状错误包括拆分多头后view无法执行通常是因为 d_model 不能整除 n_heads。解决方法是先把 n_heads 改成 d_model 的约数或者在代码开头加断言。scores的维度不是 (batch_size, n_heads, seq_len, seq_len)而是多了一维或少了一维。打印Q.size()、K.size()、scores.size()就知道问题在哪。输出维度恢复不了 d_model。这是因为合并多头时view的目标形状不对检查一下batch_size, seq_len, self.d_model是否写反。6.2 训练 loss 不下降按这个顺序排查先确认数据输入 id 是否从 0 开始是否包含超出vocab_size的 token。标签是否从 0 开始类别数是否和分类头匹配。mask 是否对应了 padding 位置而不是全部为 True。再确认模型是否做了残差连接。如果没有残差深层模型很容易梯度消失。是否使用了 LayerNorm。不用 LayerNorm 时训练稳定性会差很多。学习率是否过大。Transformer 对学习率比较敏感一般需要配合 warmup 或较小的初始学习率。最后再考虑是不是数据量太小。如果训练数据只有几百条模型很难学出通用特征这属于数据问题不是代码问题。6.3 显存不够从三个角度降占用注意力机制的显存占用和序列长度的平方成正比。如果 seq_len 是 1024单层的注意力分数矩阵大小就是 1024×1024批大小稍微增大就会占用大量显存。降低显存占用的常见做法减小批大小。最简单直接但会影响训练稳定性可以同时调低学习率。使用梯度累积。把一个大 batch 拆成多个小 batch梯度累加后再更新参数效果接近大 batch。使用稀疏注意力或窗口注意力。这是更高级的优化适合长序列任务但不适合第一次实现。先把普通注意力跑通再考虑优化。如果你只是做学习验证不要一上来就处理几千字的长文本先用 64 或 128 的序列长度跑通所有逻辑都是一样的但显存压力小很多。6.4 不同 PyTorch 版本之间的兼容性手写代码时尽量避免使用已经废弃的 API。比如旧版本常见torch.bmm新版本推荐torch.matmul。masked_fill的张量形状校验在不同版本上可能有细微差异如果报错提示 mask 维度不匹配就手动用mask.unsqueeze(1).unsqueeze(1)扩展维度。还有一个实际经验如果你的代码里同时用了transpose和view在transpose之后必须调用contiguous()。这个错误在老版本 PyTorch 里经常出现新版本部分情况下会自动处理但你不应该依赖这种自动行为。7. 从手写学习到实际工程还要注意什么手写一遍注意力机制最大的收益是理解了 Transformer 内部的张量流转。但真正做工程时我通常不会从零手写而是直接使用 PyTorch 内置的nn.TransformerEncoderLayer或nn.MultiheadAttention。原因很简单内置实现经过大量测试数值稳定性、性能优化和 backward 正确性都有保障。手写代码更适合学习、研究和定制化需求。不过如果你要改注意力机制比如加相对位置编码、做局部窗口注意力、或者实现稀疏注意力那手写版本就是必需的。你只有在理解了标准实现之后才知道从哪里切入修改而不是把整个模型重写一遍。7.1 手写版与内置版怎么选维度手写版PyTorch 内置版学习价值高适合理解原理中适合快速使用定制能力高任意改注意力逻辑较低需要绕内置接口稳定性需要自己测试久经验证性能未优化可能较慢有优化如 flash attention 路径代码量较多少所以我的建议是学习阶段手写业务阶段优先内置版只有内置版满足不了需求时再换手写版。不要为了“炫技”在项目里强行手写所有模块工程价值并不高。7.2 后续可以继续扩展的方向当你把基础多头注意力写完之后可以尝试扩展几个方向加入缩放点积之外的注意力打分函数比如加性注意力。实现带有相对位置编码的注意力比如 Transformer-XL 里的相对位置编码。把多头注意力改成稀疏注意力控制在长序列上的计算开销。加入 attention dropout。在 softmax 之后、对 V 加权之前对注意力权重做 dropout可以有效缓解过拟合。这些方向都有开源实现可以参考但前提是你已经能熟练看懂基础版。否则你连“多出来的参数在哪里”都找不到。7.3 最后想提醒的一句话很多人在学习注意力机制时喜欢直接把整份源码复制下来跑通之后就觉得“会了”。但一到面试、改代码或换任务就暴露问题。真正有效率的做法是先跑通最小例子然后自己改掉一个设计比如去掉缩放因子、去掉残差、改成单头观察训练行为变化。这样亲手做过一遍你才算真正理解了注意力机制。如果你现在刚开始动手建议先拿一个短的文本序列把上面的代码逐行跑一遍打印每一层的输出形状。不要急着一口气写完所有模块。每一步都有明确的观察结果之后再进入下一步你会发现自己对 Transformer 的理解会扎实很多。