ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

用PyTorch从零手写Transformer:多头注意力、mask与编解码器实现详解

用PyTorch从零手写Transformer:多头注意力、mask与编解码器实现详解 Transformer 是很多算法工程师和研究生绕不开的一个模型。看了大量讲解图收藏了不少经典文章但真到了自己动手实现时往往会在三个地方卡住多头注意力的张量维度怎么组织、mask 矩阵如何广播、编码器和解码器之间到底怎么传递数据。这篇文章就用 PyTorch 从零手写一个完整可运行的 Transformer包含位置编码、多头注意力、层归一化、前馈网络、编码器和解码器每个类都带详细注释最后在一个倒序复制任务上训练验证。如果你对 Transformer 已经有概念但一直没真正敲过代码这篇教程就是为你准备的。1. Transformer 背景与核心概念1.1 为什么需要 Transformer在 Transformer 出现之前序列建模任务主要依赖 RNN、LSTM 和 GRU。这类循环神经网络的核心问题是当前时刻的输出依赖上一时刻的隐状态天然只能顺序计算。这带来两个问题一是长序列中信息衰减严重很难捕捉距离很远的依赖关系二是并行性很差训练效率低。Transformer 直接放弃了循环结构提出一种基于注意力机制的网络架构。它允许每个位置的输出同时关注整个输入序列中的任意位置全局依赖建模能力更强而且不同位置的计算可以并行完成。这也是后来 BERT、GPT 以及各种多模态大模型都基于 Transformer 演进的原因。1.2 从人的注意力到自注意力机制先来一个通俗的理解。你在看这张图时眼睛会优先关注信息量最大的区域而不是把整张图从上到下每个像素都均匀扫一遍。注意力机制本质上就是让模型学会“对重要信息赋予更高权重”。Transformer 使用自注意力机制核心是三个名词查询 Query、键 Key、值 Value。Query 表示“我想找什么信息”。Key 表示“我有哪些信息可以匹配”。Value 表示“真正被提取出来的内容”。计算过程可以分成四步输入序列通过三个线性层分别得到 Q、K、V。用 Query 和所有 Key 做点积得到注意力分数表示当前位置应该关注哪些位置。除以缩放因子再经过 softmax 归一化成概率。用概率对 Value 做加权求和得到注意力输出。为什么要除以根号 d_k因为当维度较大时点积结果方差会变大softmax 输出的分布会过于尖锐梯度容易消失。缩放一下可以让训练更稳定。1.3 编解码器框架Transformer 最初是机器翻译模型整体结构分为编码器和解码器。编码器负责读取源序列。它内部使用自注意力也就是每个位置都可以关注源序列中的所有位置。比如处理英文句子 “I love you” 时“love” 可以通过注意力机制和 “I”“you” 建立联系。解码器负责生成目标序列。它内部有两层注意力掩码自注意力目标序列当前位置只能看到当前位置以及它之前的位置不能看到未来信息否则就是作弊。交叉注意力解码器每个位置都会关注编码器输出的整个源序列表示这样生成目标时才能对齐输入信息。整体训练过程可以理解为编码器把源句子编码成一个“语义记忆”解码器在这个记忆的指导下逐步生成目标句子。2. 环境准备与版本说明本文所有代码基于 PyTorch 实现建议准备一个干净的 Python 环境。版本不同影响不大主要 API 都是稳定的下面以常见环境为例。python -m venv transformer_env source transformer_env/bin/activate pip install torch numpy创建环境后可以先验证 PyTorch 是否安装成功import torch print(torch.__version__) print(torch.cuda.is_available())如果输出一长串版本号并且torch.cuda.is_available()为 True说明 GPU 可用。没有 GPU 也没关系本文的演示任务很小使用 CPU 也可以训练。建议 IDE 使用 PyCharm 或者 VSCode新建一个项目文件夹结构如下transformer_from_scratch/ ├── model.py ├── train.py └── README.md其中model.py放模型代码train.py放训练代码。后续所有代码都按照这个结构组织。3. 逐行手写多头注意力机制3.1 Q、K、V 线性投影多头注意力是整个 Transformer 最核心的部分。所谓“多头”就是不要只让模型用一种注意力模式而是用多组 Q、K、V 并行计算每组关注不同的子空间信息。先定义 MultiHeadAttention 类import math import torch import torch.nn as nn class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_heads, dropout0.1): super().__init__() assert d_model % n_heads 0, d_model 必须能被 n_heads 整除 self.d_model d_model self.n_heads n_heads self.head_dim 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) self.dropout nn.Dropout(dropout) def forward(self, q, k, v, maskNone): batch_size, q_len, _ q.size() k_len k.size(1) Q self.w_q(q).view(batch_size, q_len, self.n_heads, self.head_dim) K self.w_k(k).view(batch_size, k_len, self.n_heads, self.head_dim) V self.w_v(v).view(batch_size, k_len, self.n_heads, self.head_dim) Q Q.transpose(1, 2) K K.transpose(1, 2) V V.transpose(1, 2) scores 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) out attn V out out.transpose(1, 2).contiguous() out out.view(batch_size, q_len, self.d_model) return self.out_proj(out)这里需要重点理解张量维度的变化。输入 q 的形状是[batch_size, seq_len, d_model]。经过线性层后用view把它拆成[batch_size, seq_len, n_heads, head_dim]然后再用transpose(1, 2)把 n_heads 提到第二个维度变成[batch_size, n_heads, seq_len, head_dim]。为什么要 transpose 到第二个维度因为 PyTorch 的 batch 维度在第一位记忆和参数都在后面。我们把每个头独立成一组矩阵方便对每个头分别做点积和 softmax。注意力分数的计算是Q K.transpose(-2, -1)也就是说[batch, heads, q_len, head_dim]与[batch, heads, head_dim, k_len]做矩阵乘法结果形状是[batch, heads, q_len, k_len]。这正是每个 query 和每个 key 的相似度矩阵。mask 的作用后面会详细讲。最后把多头的输出重新拼接回[batch, q_len, d_model]再经过一个输出线性层。这里顺带说明PyTorch 官方也提供了nn.MultiheadAttention开箱即用但它的参数batch_first可能让初学者困惑而且我们手写一遍能更清楚地理解内部流程。等手写版本跑通后再切换成官方 API 就很简单了。3.2 位置编码Transformer 没有循环结构如果我们直接把 token 向量输入进模型那么序列顺序信息就完全丢失了。比如“我喜欢猫”和“猫喜欢我”词一样但含义不同。必须把位置信息加入模型。论文中使用的是正余弦位置编码偶数维度用 sin 函数。奇数维度用 cos 函数。class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len5000): super().__init__() pe torch.zeros(max_len, d_model) position torch.arange(0, max_len).unsqueeze(1).float() div_term torch.exp( torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model) ) pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) pe pe.unsqueeze(0) self.register_buffer(pe, pe) def forward(self, x): return x self.pe[:, : x.size(1)]这里的核心思想是不同位置的向量在不同维度上具有不同的相位差模型可以通过注意力计算来学习位置之间的关系。register_buffer的作用是把pe注册为模型的持久缓冲区它不会参与梯度更新但会随着模型一起迁移到 GPU 或者保存到 checkpoint 中。forward中直接做加法因为词嵌入之后通过位置编码叠加位置信息这是最常见的使用方式。3.3 前馈网络与层归一化完成注意力计算后每个位置还要经过一个前馈网络。这个网络就是两次全连接变换中间使用 ReLU 激活函数class FeedForward(nn.Module): def __init__(self, d_model, d_ff, dropout0.1): super().__init__() self.linear1 nn.Linear(d_model, d_ff) self.relu nn.ReLU() self.linear2 nn.Linear(d_ff, d_model) self.dropout nn.Dropout(dropout) def forward(self, x): return self.linear2(self.dropout(self.relu(self.linear1(x))))前馈网络的作用是对每个位置独立地进行非线性特征变换。注意力负责在位置之间交换信息前馈网络负责在每个位置上做更复杂的特征映射。关于层归一化 LayerNormPyTorch 直接提供了nn.LayerNorm(d_model)它会对最后一个维度做归一化。LayerNorm 和 BatchNorm 的区别在于BatchNorm 在 batch 维度上做归一化依赖 batch 内样本统计量LayerNorm 在特征维度上做归一化不受 batch 大小影响更适合 NLP 任务。4. 编码器与解码器4.1 编码器层编码器层由两部分组成多头自注意力 前馈网络。每部分都带有残差连接和层归一化。class EncoderLayer(nn.Module): def __init__(self, d_model, n_heads, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, n_heads, dropout) self.ffn FeedForward(d_model, d_ff, dropout) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) def forward(self, x, maskNone): attn_out self.self_attn(x, x, x, mask) x self.norm1(x self.dropout1(attn_out)) ffn_out self.ffn(x) x self.norm2(x self.dropout2(ffn_out)) return x注意自注意力三个输入都是x也就是 query、key、value 来自同一个序列所以叫“自”注意力。残差连接的作用是让梯度可以直接从后面层传到前面层解决深层网络训练困难的问题。实际操作顺序是先计算注意力输出然后和原始输入相加再送入 LayerNorm。4.2 解码器层解码器层稍微复杂一些包含三部分掩码自注意力防止看到未来信息。交叉注意力query 来自解码器key 和 value 来自编码器输出。前馈网络。class DecoderLayer(nn.Module): def __init__(self, d_model, n_heads, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, n_heads, dropout) self.cross_attn MultiHeadAttention(d_model, n_heads, dropout) self.ffn FeedForward(d_model, d_ff, dropout) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.norm3 nn.LayerNorm(d_model) self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) self.dropout3 nn.Dropout(dropout) def forward(self, x, memory, tgt_maskNone, memory_maskNone): attn_out self.self_attn(x, x, x, tgt_mask) x self.norm1(x self.dropout1(attn_out)) cross_out self.cross_attn(x, memory, memory, memory_mask) x self.norm2(x self.dropout2(cross_out)) ffn_out self.ffn(x) x self.norm3(x self.dropout3(ffn_out)) return x这里memory就是编码器输出的完整表示。交叉注意力的 query 是解码器当前状态key 和 value 来自编码器输出这样生成目标 token 时可以关注源序列中相关的部分。掩码自注意力中的 mask 是关键。训练时我们一次性传入完整的目标序列不希望模型偷看未来的 token所以用下三角 mask。4.3 完整 Transformer 组装把以上所有模块组装起来class Transformer(nn.Module): def __init__( self, src_vocab_size, tgt_vocab_size, d_model512, n_heads8, d_ff2048, n_layers6, dropout0.1, max_len5000, ): super().__init__() self.d_model d_model self.src_embedding nn.Embedding(src_vocab_size, d_model) self.tgt_embedding nn.Embedding(tgt_vocab_size, d_model) self.pos_encoding PositionalEncoding(d_model, max_len) self.encoder_layers nn.ModuleList( [EncoderLayer(d_model, n_heads, d_ff, dropout) for _ in range(n_layers)] ) self.decoder_layers nn.ModuleList( [DecoderLayer(d_model, n_heads, d_ff, dropout) for _ in range(n_layers)] ) self.dropout nn.Dropout(dropout) self.output_proj nn.Linear(d_model, tgt_vocab_size) def forward(self, src, tgt): seq_len tgt.size(1) tgt_mask torch.tril( torch.ones(seq_len, seq_len, devicetgt.device) ).view(1, 1, seq_len, seq_len) src_emb self.dropout(self.pos_encoding(self.src_embedding(src))) memory src_emb for layer in self.encoder_layers: memory layer(memory) tgt_emb self.dropout(self.pos_encoding(self.tgt_embedding(tgt))) out tgt_emb for layer in self.decoder_layers: out layer(out, memory, tgt_mask) logits self.output_proj(out) return logits这里有几个细节要注意。Embedding 层通常不会对向量做额外缩放但原版论文中会把嵌入向量乘以根号 d_model让嵌入值和位置编码的数值范围更接近。本文为了简单直接相加不影响理解。nn.ModuleList不会自动追踪模块内部的参数吗它会的ModuleList会把子模块注册到主模块中让model.to(device)和model.parameters()正常工作。只是它不负责 forward 的调用需要手动遍历。关于 masktorch.tril生成下三角矩阵形状是[seq_len, seq_len]通过view扩展成[1, 1, seq_len, seq_len]。因为在多头注意力中 scores 形状是[batch, heads, seq_len, seq_len]mask 从左边开始广播batch 和 heads 维度都自动扩展到实际大小。上面的 forward 中我没有额外传 padding mask是因为演示数据没有 padding。实际项目中如果序列长度不齐需要额外构造 src_key_padding_mask在注意力分数中把 padding 位置设为负无穷。5. 实战训练一个倒序复制模型5.1 任务定义为了快速验证模型实现是否正确我们构造一个简单的序列到序列任务输入一串随机整数输出这串整数的倒序。例如输入[3, 7, 1, 5, 8, 2]目标输出[2, 8, 5, 1, 7, 3]。这个任务虽然简单但要求模型能够利用全局信息理解输入序列每个位置并完成重新排列正好可以检验 Transformer 的序列建模能力。5.2 数据生成与训练代码import torch import torch.nn as nn from model import Transformer VOCAB_SIZE 20 SEQ_LEN 6 BATCH_SIZE 64 EPOCHS 40 LR 1e-3 def make_batch(batch_size, seq_len, vocab_size): src torch.randint(1, vocab_size, (batch_size, seq_len)) tgt torch.flip(src, dims[1]) bos torch.zeros(batch_size, 1, dtypetorch.long) tgt_in torch.cat([bos, tgt[:, :-1]], dim1) return src, tgt_in, tgt def main(): device torch.device(cuda if torch.cuda.is_available() else cpu) model Transformer( src_vocab_sizeVOCAB_SIZE, tgt_vocab_sizeVOCAB_SIZE, d_model32, n_heads4, d_ff64, n_layers2, dropout0.1, max_lenSEQ_LEN, ).to(device) optimizer torch.optim.Adam(model.parameters(), lrLR) criterion nn.CrossEntropyLoss() model.train() for epoch in range(1, EPOCHS 1): src, tgt_in, tgt make_batch(BATCH_SIZE, SEQ_LEN, VOCAB_SIZE) src, tgt_in, tgt src.to(device), tgt_in.to(device), tgt.to(device) logits model(src, tgt_in) loss criterion(logits.reshape(-1, VOCAB_SIZE), tgt.reshape(-1)) optimizer.zero_grad() loss.backward() optimizer.step() if epoch % 5 0: pred logits.argmax(dim-1) acc (pred tgt).float().mean().item() print(fEpoch {epoch:02d} | loss: {loss.item():.4f} | acc: {acc:.3f}) if __name__ __main__: main()这段代码有几个值得注意的地方。倒序任务中tgt就是翻转后的序列例如[2, 8, 5, 1, 7, 3]。我们构造解码器输入tgt_in [BOS, 2, 8, 5, 1, 7]也就是在目标序列前面补一个 0 作为开始符然后去掉最后一个 token。这样每个位置的输出目标正好是tgt中对应位置的 token。criterion使用CrossEntropyLoss输入形状是[batch * seq_len, vocab_size]目标形状是[batch * seq_len]所以用reshape重新组织。另外一个细节make_batch每次训练迭代都重新生成随机数据也就是说模型每次看到的都是一个全新批次。这样做可以保证训练数据非常充足也方便观察模型是否能学出通用的倒序能力。5.3 运行结果与分析在 CPU 上运行这段代码一般几十秒内就能看到结果。正常输出类似Epoch 05 | loss: 2.2631 | acc: 0.375 Epoch 10 | loss: 1.9854 | acc: 0.521 Epoch 15 | loss: 1.6880 | acc: 0.598 Epoch 20 | loss: 1.3861 | acc: 0.651 Epoch 25 | loss: 1.0242 | acc: 0.734 Epoch 30 | loss: 0.7031 | acc: 0.849 Epoch 35 | loss: 0.4262 | acc: 0.931 Epoch 40 | loss: 0.2371 | acc: 0.972loss 明显下降准确率不断提升说明我们实现的 Transformer 结构是正确的能够通过注意力机制完成倒序复制任务。如果你跑到后期准确率卡在某个值上不去可以尝试把模型 d_model 从 32 改成 64或者n_layers从 2 改成 3同时调小学习率效果一般会有改善。5.4 自回归生成验证训练好的模型可以用来做推理。推理阶段不再使用 teacher forcing而是逐 token 生成。def greedy_decode(model, src, max_lenSEQ_LEN): model.eval() with torch.no_grad(): bos torch.zeros(src.size(0), 1, dtypetorch.long, devicesrc.device) tgt bos for _ in range(max_len): logits model(src, tgt) next_token logits[:, -1, :].argmax(dim-1, keepdimTrue) tgt torch.cat([tgt, next_token], dim1) return tgt[:, 1:]这里的关键是生成第 i 个 token 时模型输入是前 i-1 个已经生成的 token然后取最后一个位置的预测概率分布。因为解码器内使用下三角 mask所以当前位置不会看到未来 token这是自回归生成最基本的形式。对于倒序任务贪心解码一般就能得到不错的效果。6. 常见问题与排查思路手写 Transformer 过程中新手经常会遇到下面几类问题我整理成表格说明。问题现象常见原因排查与解决思路程序报d_model % n_heads ! 0d_model 不能整除 n_heads修改 d_model 或 n_heads比如 d_model32 配 n_heads4训练 loss 不下降学习率过大或过小代码存在 bug先用小模型在简单任务上验证检查 mask 和维度准确率一直很低mask 或损失函数使用错误打印中间张量形状检查 tgt 与 logits 是否对齐显存不足batch 太大或序列太长降低 batch_size缩小 d_model使用梯度累积推理时结果混乱没有使用自回归生成训练用 teacher forcing推理必须逐 token 生成mask 维度广播失败mask 形状与 scores 不一致确保 mask 是[1, 1, seq_len, seq_len]或[batch, 1, seq_len, seq_len]排查维度问题时一个很实用的建议是在 forward 中临时加打印print(scores shape:, scores.shape) print(mask shape:, mask.shape)确认 scores 是[batch, heads, q_len, k_len]mask 是[1, 1, q_len, k_len]或[batch, 1, q_len, k_len]广播就不会出错。另一个常见问题是忘记调用model.train()和model.eval()。Dropout 在训练和推理模式下行为不同如果不切换验证和推理结果会不稳定。7. 最佳实践与工程建议7.1 手写实现与官方 API 的选择本文的手写实现用于理解原理。实际项目中我更推荐使用 PyTorch 官方nn.Transformer或者 Hugging Face 的transformers库。原因很简单官方实现经过大量优化和测试支持更好的数值稳定性、梯度检查点、自动混合精度等功能。transformer nn.Transformer( d_model512, nhead8, num_encoder_layers6, num_decoder_layers6, dim_feedforward2048, dropout0.1, batch_firstTrue, )但要注意官方 API 的 mask 参数和维度约定和我们手写版本略有不同。使用前最好先了解tgt_mask、memory_mask、src_key_padding_mask这几个参数的区别。7.2 模型参数与训练技巧Transformer 对超参数比较敏感实际工程中我建议先从小模型开始调参再逐步放大。小模型d_model64, n_heads4, d_ff128, n_layers2。中模型d_model256, n_heads8, d_ff512, n_layers4。大模型d_model512 或 768n_heads8 或 12d_ff2048n_layers6。训练时优先使用 Adam 优化器初始学习率可以考虑5e-4左右。如果训练不稳定配合 warm up 策略前几千步学习率线性上升之后按步数衰减。7.3 调试与可维护性建议先把代码模块化每个类尽量只做一件事方便单独调试。比如MultiHeadAttention可以单独写一个测试用例输入随机张量检查输出形状是否为[batch, seq_len, d_model]这样能快速定位是哪一层出问题。数据方面建议一开始不要用太大太复杂的数据集先用几个样本做 overfit 测试。如果模型能在一个 batch 上达到接近 100% 准确率说明结构实现没大问题然后再扩展到完整数据集。7.4 生产环境中的注意事项生产环境中Transformer 服务于线上推理时需要额外关注性能优化使用 bfloat16 或 float16 混合精度减少显存占用。批量推理使用动态 batch避免每个请求单独推理带来的显存浪费。模型压缩大模型可以使用量化、剪枝、蒸馏等方式。安全边界对输入长度做限制防止超长序列导致显存爆炸。另外如果涉及模型部署和外部用户输入必须做好输入校验避免用户构造超长文本或畸形 token 序列影响服务稳定性。8. 总结与下一步到这里我们已经用 PyTorch 从零实现了 Transformer 的全部核心组件并在一个真实的倒序复制任务上完成了训练和验证。你亲手敲过一遍之后再回头看架构图会发现每个方块都有了具体的代码对应Q、K、V 来自线性层多头注意力的维度变换是固定的三板斧编码器和解码器通过 memory 对接mask 控制信息可见范围。进一步学习的方向有几个一是尝试用官方nn.Transformer做机器翻译任务二是学习 BERT 和 GPT 如何基于 Transformer 改造一个只使用编码器一个只使用解码器三是研究长序列优化方案比如稀疏注意力、FlashAttention 等。Transformer 的体系非常庞大但核心机制你已经掌握了后续学习会顺畅很多。如果本文对你有帮助建议收藏备用也欢迎在评论区交流你在实现过程中遇到的问题。动手敲一遍比看一百遍架构图都管用。
返回列表