Transformer注意力机制从零实现:原理详解与PyTorch代码实战
之前用现成的nn.TransformerEncoderLayer搭模型时总觉得注意力机制是一个“黑盒”知道它叫 QKV也知道有 softmax但一旦需要改 mask、调多头数、或者把注意力权重拿出来可视化就只能去翻源码。后来在项目里需要把 Transformer 接入自己的序列特征提取流程我才真正动手把注意力机制从零实现了一遍。这篇文章就围绕“在 Transformer 中实现注意力机制”这条主线展开先讲清楚注意力机制的核心原理再逐行实现缩放点积注意力、多头注意力、位置编码和编码器层最后用一个可运行的文本分类小项目验证整个流程。无论你是刚接触 Transformer 的初学者还是想深入底层细节的开发者都可以参考这套实现思路。1. 背景与核心概念在正式写代码之前先回答一个问题为什么 Transformer 能取代 RNN、LSTM成为 NLP 和很多序列任务的主流架构核心答案就是注意力机制。RNN 类模型需要沿着时间步一步一步处理序列第 t 个时刻的隐藏状态依赖于前一个时刻的输出。这种串行结构有两个明显短板长距离依赖问题序列很长时前面的信息会被逐步“稀释”模型很难直接关注到离当前位置很远的关键信息。并行计算困难前后时间步之间存在依赖训练时很难对整条序列同时做计算。注意力机制改变了这种思路。它允许模型在计算某个位置的表示时直接“关注”序列中所有其他位置并且根据它们与当前位置的相关性来加权聚合信息。换句话说输入序列中任意两个位置之间都可以建立直接联系不再受距离限制。Transformer 正式把这种思路应用到了极致。它抛弃了循环结构完全依靠自注意力机制Self-Attention来捕捉序列内部的关系再配合位置编码Positional Encoding注入顺序信息。后来这套架构也被迁移到计算机视觉领域衍生出 Vision TransformerViT、Swin Transformer 等模型注意力机制因此成为深度学习中最基础、最核心的模块之一。不过要注意注意力机制本身有很多变体。常见的有自注意力机制Self-AttentionQ、K、V 来自同一个输入序列用于建模序列内部关系。多头注意力机制Multi-Head Attention把 Q、K、V 投影到多个子空间分别计算注意力再拼接起来。交叉注意力机制Cross-AttentionQ 来自一个序列K、V 来自另一个序列常用于解码器中对编码器输出做注意力。通道注意力机制SENet、CBAM、ECA主要用于卷积神经网络对特征图的通道维度做加权。本文会聚焦在 Transformer 中最核心的自注意力机制和多头注意力机制并且用 PyTorch 手写一遍。2. 环境准备与版本说明本文代码基于 Python 和 PyTorch 编写。示例中只会用到 PyTorch 的基础 API不依赖额外的模型库可以避免版本兼容问题。操作系统Windows / Linux / macOS 均可。Python 版本建议使用 3.9 及以上。PyTorch建议使用 PyTorch 2.x 版本CPU 环境即可运行本文示例。IDE推荐 PyCharm、VS Code 或 Jupyter Notebook。其他依赖matplotlib用于绘制注意力可视化可选。如果没有安装 PyTorch可以在命令行执行pip install torch torchvision如果只需要 CPU 版本可以根据 PyTorch 官网选择对应的安装命令。版本需要根据你的实际环境调整本文重点演示实现思路代码对 PyTorch 版本的依赖不强。建议按照下面的目录结构组织项目文件transformer_attention/ ├── attention.py # 缩放点积注意力、多头注意力 ├── transformer_layer.py # 位置编码、编码器层 ├── train.py # 文本分类训练示例 └── requirements.txt # 依赖清单3. 注意力机制核心原理拆解3.1 从 Q、K、V 三个向量说起注意力机制中的 Q、K、V 分别代表 Query查询、Key键、Value值。可以类比一个检索过程Query 是你要查的内容。Key 是资料库中每份文档的索引标签。Value 是资料库中每份文档的实际内容。注意力计算的过程就是拿 Query 和所有 Key 做相似度计算得到一组权重再用这组权重对 Value 做加权求和。在 Transformer 中Q、K、V 并不是输入本身而是输入通过三个权重矩阵映射得到的向量。假设输入是 (X)那么[ Q XW_Q ][ K XW_K ][ V XW_V ]其中 (W_Q)、(W_K)、(W_V) 都是可学习的参数矩阵。这样做的好处是模型可以通过训练调整这三个映射让注意力机制适配不同的任务。3.2 缩放点积注意力公式Transformer 中使用的是缩放点积注意力Scaled Dot-Product Attention公式如下[ Attention(Q, K, V) softmax(\frac{QK^T}{\sqrt{d_k}})V ]其中 (d_k) 是 Key 的维度。除以 (\sqrt{d_k}) 是为了防止点积结果过大导致 softmax 进入饱和区梯度变得极小。如果输入序列长度是 (n)每个位置的向量维度是 (d_k)那么 (QK^T) 的形状是 ((n, n))其中第 (i) 行第 (j) 列的值表示第 (i) 个 Query 和第 (j) 个 Key 的相似度。经过 softmax 之后每一行的权重之和为 1再用这些权重对 V 加权求和。3.3 Mask 的作用在实际实现中注意力机制还经常搭配 Mask 一起使用。Mask 主要解决两类问题Padding Mask文本序列长短不一batch 中短句会用pad符号补齐。这些 padding 位置是无效信息注意力计算时应该忽略。Causal Mask因果 Mask解码器生成时只能看到当前位置之前的信息不能看到未来。这时需要把未来位置遮住。实现方式通常是把需要屏蔽的位置的注意力分数设置为负无穷这样 softmax 之后权重趋近于 0。3.4 为什么需要多头注意力单头注意力只能学习一种“注意力模式”例如关注词与词之间的语法关系、指代关系、语义相似性等。不同任务可能同时需要多种不同的关系。多头注意力机制的做法是把 Q、K、V 投影到 (h) 个不同的子空间每个子空间独立计算注意力最后把所有头的输出拼接起来再经过一个输出投影。这样模型可以同时关注不同位置、不同语义子空间的信息。3.5 位置编码的意义Self-Attention 本身是不带位置信息的。如果把一句话的词序打乱注意力计算结果完全一样因为计算过程只关心两两之间的相似度不关心它们出现在哪个位置。为了让模型感知词的顺序Transformer 在输入中加入了位置编码。原始论文使用正弦余弦函数来生成位置向量[ PE_{(pos, 2i)} sin(pos / 10000^{2i/d_{model}}) ][ PE_{(pos, 2i1)} cos(pos / 10000^{2i/d_{model}}) ]其中 (pos) 是位置索引(i) 是维度索引。4. 在 Transformer 中实现注意力机制完整实战下面开始进入代码环节。我会按照“分层实现、逐块讲解”的方式把注意力机制完整实现一遍。4.1 创建项目结构先创建项目目录mkdir transformer_attention cd transformer_attention再创建三个 Python 文件attention.py transformer_layer.py train.py4.2 实现缩放点积注意力文件路径attention.py缩放点积注意力是整个 Transformer 中最底层的计算单元。它接收 Q、K、V 和可选的 Mask输出注意力聚合结果和注意力权重。import math import torch import torch.nn as nn import torch.nn.functional as F class ScaledDotProductAttention(nn.Module): def __init__(self, dropout0.1): super().__init__() self.dropout nn.Dropout(dropout) def forward(self, query, key, value, maskNone): # query, key, value 形状: (batch_size, num_heads, seq_len, d_k) d_k query.size(-1) # 1. 计算 Q 和 K 的点积并缩放 scores torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(d_k) # 2. 如果提供了 mask将需要屏蔽的位置设为负无穷 if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) # 3. 在最后一个维度上做 softmax attn_weights F.softmax(scores, dim-1) attn_weights self.dropout(attn_weights) # 4. 用注意力权重对 Value 做加权求和 output torch.matmul(attn_weights, value) return output, attn_weights这里有几个关键点需要重点说明。为什么要做缩放当 (d_k) 较大时Q 和 K 的点积结果方差也会变大。如果直接进 softmax会导致 softmax 输出非常接近 one-hot 分布梯度很小不利于训练。除以 (\sqrt{d_k}) 可以让分数分布保持在合理区间。mask 的格式是什么对于自注意力来说scores 的形状是(batch_size, num_heads, seq_len, seq_len)。如果我们传入的 mask 形状是(batch_size, 1, 1, seq_len)也就是每个样本标记哪些位置是有效 token那么它会在广播机制下自动应用到所有 head 和所有 query 位置。masked_fill(mask 0, float(-inf))的含义是mask 中为 0 的位置对应的分数被替换为负无穷。这样 softmax 后这些位置的注意力权重约等于 0。4.3 实现多头注意力文件路径attention.py多头注意力Multi-Head Attention是 Transformer 中的核心模块。它的做法是先把 Q、K、V 通过线性层投影再拆分成多个头分别计算注意力最后拼接。class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads, dropout0.1): super().__init__() assert d_model % num_heads 0, d_model 必须能被 num_heads 整除 self.d_model d_model self.num_heads num_heads self.d_k d_model // num_heads # Q、K、V 的投影矩阵 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.attention ScaledDotProductAttention(dropoutdropout) def forward(self, query, key, value, maskNone): batch_size query.size(0) # 1. 线性投影然后拆分为 num_heads 个头 Q self.w_q(query).view( batch_size, -1, self.num_heads, self.d_k ).transpose(1, 2) K self.w_k(key).view( batch_size, -1, self.num_heads, self.d_k ).transpose(1, 2) V self.w_v(value).view( batch_size, -1, self.num_heads, self.d_k ).transpose(1, 2) # 2. 如果 mask 是 (batch_size, seq_len) 形式扩展成可广播形状 if mask is not None: if mask.dim() 2: # (batch_size, 1, 1, seq_len) mask mask.unsqueeze(1).unsqueeze(1) elif mask.dim() 3: # (batch_size, seq_len, seq_len) - (batch_size, 1, seq_len, seq_len) mask mask.unsqueeze(1) # 3. 计算注意力 attn_output, attn_weights self.attention(Q, K, V, maskmask) # 4. 合并多头结果 attn_output attn_output.transpose(1, 2).contiguous().view( batch_size, -1, self.d_model ) # 5. 输出投影 output self.w_o(attn_output) return output, attn_weights这段代码中view和transpose的配合是容易出错的地方。原始 Q 的形状是(batch_size, seq_len, d_model)经过view后变成(batch_size, seq_len, num_heads, d_k)再transpose(1, 2)把num_heads提到第二个维度最终形状是(batch_size, num_heads, seq_len, d_k)。这里有一个小细节transpose之后张量的内存布局不是连续的直接view会报错。所以合并多头结果时要先transpose(...).contiguous()再view(...)。4.4 实现位置编码文件路径transformer_layer.py按照正弦余弦公式可以这样实现位置编码import torch import torch.nn as nn class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len5000, dropout0.1): super().__init__() self.dropout nn.Dropout(pdropout) # 初始化一个 (max_len, d_model) 的位置编码矩阵 pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.float).unsqueeze(1) # 计算不同维度的频率 div_term torch.exp( torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model) ) # 偶数维使用 sin奇数维使用 cos pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) # 增加 batch 维度(1, max_len, d_model) pe pe.unsqueeze(0) # 注册为 buffer不参与梯度更新会随模型保存和加载 self.register_buffer(pe, pe) def forward(self, x): # x 形状: (batch_size, seq_len, d_model) seq_len x.size(1) x x self.pe[:, :seq_len, :] return self.dropout(x)这里需要注意div_term中参数10000.0是原始论文中的设定也是常见的默认值。register_buffer的作用是把位置编码矩阵作为模型的一部分移动到 GPU 时不需要手动处理。4.5 实现 Transformer 编码器层文件路径transformer_layer.py标准 Transformer 编码器层由两部分组成多头自注意力子层和前馈神经网络子层。每个子层都带有残差连接和 Layer Normalization。class TransformerEncoderLayer(nn.Module): def __init__(self, d_model, num_heads, dim_feedforward, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, num_heads, dropout) # 前馈网络 self.linear1 nn.Linear(d_model, dim_feedforward) self.linear2 nn.Linear(dim_feedforward, d_model) # 归一化 self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) # Dropout self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) self.dropout3 nn.Dropout(dropout) self.activation nn.ReLU() def forward(self, src, maskNone): # 第一子层多头自注意力 残差 LayerNorm attn_output, attn_weights self.self_attn(src, src, src, maskmask) src src self.dropout1(attn_output) src self.norm1(src) # 第二子层前馈网络 残差 LayerNorm ffn_output self.linear2(self.dropout2(self.activation(self.linear1(src)))) src src self.dropout3(ffn_output) src self.norm2(src) return src, attn_weights注意力子层中Q、K、V 全部来自同一个输入src所以这是标准的自注意力机制。返回的attn_weights可以用来做可视化分析排查模型到底关注了哪些位置。4.6 文本分类训练示例文件路径train.py现在把上面的模块组装起来完成一个简单的文本情感分类任务。为了便于演示这里使用一个非常小的自定义数据集。首先定义模型import torch import torch.nn as nn from attention import MultiHeadAttention from transformer_layer import PositionalEncoding, TransformerEncoderLayer class TransformerClassifier(nn.Module): def __init__( self, vocab_size, d_model, num_heads, num_layers, dim_feedforward, num_classes, max_len, dropout0.1, ): super().__init__() self.embedding nn.Embedding(vocab_size, d_model) self.pos_encoder PositionalEncoding(d_model, max_len, dropout) self.encoder_layers nn.ModuleList([ TransformerEncoderLayer(d_model, num_heads, dim_feedforward, dropout) for _ in range(num_layers) ]) self.classifier nn.Linear(d_model, num_classes) self.d_model d_model def forward(self, tokens, maskNone): # 1. 词嵌入 x self.embedding(tokens) # 2. 缩放嵌入向量与位置编码保持量级一致 x x * torch.sqrt(torch.tensor(self.d_model, dtypetorch.float)) # 3. 加入位置编码 x self.pos_encoder(x) # 4. 逐层通过编码器层 attn_weights_list [] for layer in self.encoder_layers: x, attn_weights layer(x, maskmask) attn_weights_list.append(attn_weights) # 5. 取序列第一个 token 的表示作为分类特征 cls_feat x[:, 0, :] logits self.classifier(cls_feat) return logits, attn_weights_list然后构造一个极小的样本集from torch.utils.data import Dataset, DataLoader texts [ this movie is great and wonderful, i love this film very much, what a fantastic performance, the story is boring and dull, i hate this movie, this film is too bad to watch, ] labels [1, 1, 1, 0, 0, 0] # 构建词表 vocab {pad: 0, unk: 1} for text in texts: for token in text.split(): if token not in vocab: vocab[token] len(vocab) def encode(text, max_len10): tokens text.split()[:max_len] ids [vocab.get(token, vocab[unk]) for token in tokens] ids ids [vocab[pad]] * (max_len - len(ids)) return torch.tensor(ids, dtypetorch.long) class TextDataset(Dataset): def __init__(self, texts, labels): self.texts texts self.labels labels def __len__(self): return len(self.texts) def __getitem__(self, idx): return encode(self.texts[idx]), torch.tensor(self.labels[idx], dtypetorch.long)再写训练循环model TransformerClassifier( vocab_sizelen(vocab), d_model32, num_heads4, num_layers2, dim_feedforward64, num_classes2, max_len10, dropout0.1, ) optimizer torch.optim.Adam(model.parameters(), lr0.001) criterion nn.CrossEntropyLoss() dataset TextDataset(texts, labels) dataloader DataLoader(dataset, batch_size2, shuffleTrue) model.train() for epoch in range(30): total_loss 0.0 for src, label in dataloader: optimizer.zero_grad() logits, _ model(src) loss criterion(logits, label) loss.backward() optimizer.step() total_loss loss.item() if (epoch 1) % 10 0: print(fEpoch {epoch 1:3d}, Loss: {total_loss / len(dataloader):.4f})运行效果类似Epoch 10, Loss: 0.6742 Epoch 20, Loss: 0.5801 Epoch 30, Loss: 0.4827由于数据集非常小这里主要验证前向传播、反向传播和训练流程是否通畅。如果希望数据更真实可以换成公开的影评数据集、新闻分类数据集并配合更大的d_model、num_layers和更长的训练轮次。5. 常见问题与排查思路手写注意力机制时最常见的坑集中在维度匹配、Mask 广播、训练不收敛这几个方面。下面整理了一份排查清单。问题现象常见原因解决思路view报错提示形状不匹配Q、K、V 最后一维不是d_model检查输入维度是否为(batch, seq_len, d_model)检查d_model % num_heads 0transpose后view报错张量内存不连续先执行.contiguous()再执行.view()注意力分数全为nanmask 维度错误导致 softmax 输入全是-inf打印 scores 和 mask 的形状确认 mask 能正确广播训练 loss 不下降学习率过大或过小尝试3e-4或1e-4也可以加入学习率预热训练 loss 始终很高没有位置编码或没有 LayerNorm检查编码器层是否包含残差和归一化显存不足序列长度过大注意力矩阵是 (n^2) 复杂度缩小 batch、缩短序列或使用 FlashAttention多头注意力结果和官方实现不一致多头合并顺序错误检查view和transpose的顺序以及contiguous()的使用使用 mask 后padding 位置仍参与更新mask 没有传入注意力层检查 mask 参数是否一路从编码器层传到了ScaledDotProductAttention其中mask 维度是非常容易踩坑的点。下面给出一个标准的 padding mask 生成方法def make_padding_mask(ids, pad_idx0): # ids 形状: (batch_size, seq_len) mask (ids ! pad_idx).unsqueeze(1).unsqueeze(2) # 形状变为 (batch_size, 1, 1, seq_len) return mask在训练时可以这样传入src torch.tensor([[...], [...]]) # (batch, seq_len) mask make_padding_mask(src) logits, attn_weights_list model(src, maskmask)如果你想绘制注意力分布图只需要拿到attn_weightsimport matplotlib.pyplot as plt # 取第一个编码器层、第一个 batch、第一个 head weights attn_weights_list[0][0, 0].detach().numpy() plt.imshow(weights, cmaphot) plt.colorbar() plt.show()6. 最佳实践与工程建议6.1 从官方实现做对照自定义实现的目的是理解原理但在工程项目中除非有特殊需求否则优先使用 PyTorch 官方提供的nn.TransformerEncoderLayer、nn.MultiheadAttention。这些模块经过验证性能更好且支持 FlashAttention 等加速方案。建议的做法是先用本文实现跑通流程再用官方模块替换对比输出和训练指标确认自己的实现是否正确。6.2 注意力可视化是调试利器训练 Transformer 模型时loss 不下降并不一定代表模型完全坏了也有可能是注意力模式不合理。通过可视化注意力权重可以直观看到模型在关注什么。如果所有注意力权重都差不多说明模型没有学到有效的关系。如果注意力集中在pad位置说明 mask 没有生效。如果某些 head 明显集中在句首或句尾可能是位置编码的影响过强。6.3 数值稳定性注意力机制中使用了 softmax 和矩阵乘法在训练初期很容易出现数值不稳定的情况。建议始终保留缩放因子 ( \sqrt{d_k} )。Mask 使用-inf而不是0。使用torch.nn.utils.clip_grad_norm_做梯度裁剪。混合精度训练时注意-inf和nan的传播。6.4 数据安全与合规在真实项目中训练数据往往包含用户文本、日志等敏感信息。使用注意力机制处理这些数据时需要注意训练数据要经过脱敏处理不包含手机号、身份证号等隐私信息。模型发布前要进行数据审查避免生成或复述敏感内容。如果模型部署到生产环境需要关注输入数据的合规授权。6.5 可复现性实验阶段可以设定随机种子确保每次训练结果一致import random import numpy as np def set_seed(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed)同时在训练时保存最优的模型权重torch.save(model.state_dict(), transformer_attention.pth)7. 总结与学习路线本文从注意力机制的背景讲起完整实现了缩放点积注意力、多头注意力、位置编码和 Transformer 编码器层并用一个文本分类小项目验证了模型的可训练性。到这一步你已经掌握了 Transformer 中最核心的模块自注意力机制和多头注意力机制。如果想继续深入可以沿着下面几个方向扩展实现解码器中的交叉注意力机制把编码器的输出作为 K、V解码器自身的输出作为 Q。加入因果 Mask实现一个简单的自回归语言模型。替换位置编码方式尝试可学习位置编码、旋转位置编码RoPE、ALiBi。比较 CNN、RNN、Transformer 在相同任务上的表现差异。阅读 Vision TransformerViT源码了解注意力机制如何应用到图像 Patch 序列上。注意力机制没有想象中那么神秘把 QKV 的流动过程理清楚再动手写一遍很多疑惑都会自然消除。如果你在实践过程中遇到了其他问题欢迎对照本文的排查思路逐步定位也建议多打印中间张量的形状这是理解这类模型最有效的方法。如果这篇文章对你有帮助可以收藏备用。后续我也会继续整理 Transformer 相关实战内容欢迎一起交流。