ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

从零手撕Transformer:深入原理与PyTorch实现

从零手撕Transformer:深入原理与PyTorch实现 1. 为什么我劝你手撕一遍Transformer而不是直接调库刚入门深度学习那会儿我和很多人一样nn.TransformerEncoderLayer一包forward一调模型跑起来就算完事。直到有一次线上推理延迟飙高我盯着torch.profiler的输出发愣——明明只改了nhead和d_model为什么显存涨了将近一倍那一刻我才意识到我对这个“已经封装好的黑盒”其实一无所知。Transformer 这个架构从 2017 年那篇《Attention Is All You Need》开始已经统治了 NLP、视觉、语音、时序预测几乎所有序列建模场景。你刷到的 ViT、Swin Transformer、Point Transformer、YOLO 里塞的注意力模块本质上都是它的变体。但真正能把它讲清楚、写清楚的人并不多大部分教程要么停留在画架构图要么直接甩一段官方实现让你自己悟。这篇东西就是写给那些“想真正搞懂”的人。我会从位置编码、多头注意力、Encoder、Decoder四个核心模块出发一块一块拆原理一行一行写代码最后拼成一个能跑通的最小可用 Transformer。适合已经会 PyTorch 基础、但对手撕模型还没底的新手也适合想回头补一补底层细节的老手。看完你至少能做到三件事能手写多头注意力不查文档、能说清位置编码为什么用 sin/cos、能自己搭一个 Encoder-Decoder 做序列到序列的任务。提示本文所有代码基于 PyTorch不依赖torch.nn.Transformer的任何封装全部从nn.Linear、nn.LayerNorm这种最基础的组件搭起。这样你才能真正看清每个张量是怎么流动的。2. 整体架构拆解先看清骨架再动手2.1 Transformer 到底解决了什么问题在 Transformer 之前序列建模基本是 RNN 和 CNN 的天下。RNN 的问题是串行计算第 t 步必须等第 t-1 步算完长序列上根本没法并行而且梯度在时间维度上传播长距离依赖容易衰减。CNN 虽然能并行但感受野要靠堆层数来扩大捕捉全局依赖的效率很低。Transformer 的核心思路很直接用注意力机制直接建模任意两个位置之间的关系一步到位拿到全局信息而且整个计算过程可以完全并行。代价是注意力本身对位置不敏感——你把序列打乱注意力算出来的结果是一样的。所以必须额外注入位置信息这就是位置编码存在的意义。理解这一点很关键Transformer 注意力建模内容关系 位置编码注入顺序信息 前馈网络做非线性变换 残差和归一化保证训练稳定。后面所有模块都是围绕这四个东西展开的。2.2 整体数据流从输入到输出走一遍一个标准的 Encoder-Decoder Transformer数据流大致是这样的源序列经过词嵌入Embedding变成向量再加上位置编码进入 Encoder 堆叠层每层包含多头自注意力 前馈网络每部分都有残差连接和 LayerNormEncoder 输出一组上下文表示作为 Decoder 的 Key 和 Value目标序列同样经过嵌入和位置编码进入 DecoderDecoder 每层有三个子模块带掩码的自注意力、交叉注意力Query 来自 DecoderKey/Value 来自 Encoder、前馈网络最后经过线性层映射到词表维度接 Softmax 得到概率分布我习惯把这个流程画成一张“张量形状追踪表”因为手撕的时候最容易出错的就是维度对不上。下面这张表是我自己调试时常用的你可以对照着看阶段张量形状说明输入 token(batch, seq_len)整数索引Embedding 后(batch, seq_len, d_model)d_model 通常 512加位置编码(batch, seq_len, d_model)逐元素相加多头拆分(batch, nhead, seq_len, d_k)d_k d_model / nhead注意力输出(batch, seq_len, d_model)多头拼接后FFN 输出(batch, seq_len, d_model)中间层 d_ff 通常 2048最终输出(batch, seq_len, vocab_size)映射到词表注意d_model必须能被nhead整除否则多头拆分时会报维度错误。这是新手最常踩的坑之一我后面还会专门讲。2.3 为什么我要按“模块化”的方式手撕很多人手撕 Transformer 喜欢一口气写一个几百行的类结果一出错根本不知道是哪一层的问题。我的建议是分模块实现、分模块测试先把位置编码单独写出来用一个小张量验证数值再写多头注意力用随机输入检查输出形状和数值范围最后才拼装成完整模型。这样做的好处是每个模块都能独立验证出错时排查范围小。而且这些模块本身是可复用的——比如你以后要做视觉任务位置编码换成二维的注意力模块几乎不用改。这也是为什么工业界很多代码库都是把 Attention、FFN、PositionalEncoding 拆成独立文件的原因。3. 位置编码让模型知道谁在前谁在后3.1 为什么不能只用可学习的位置嵌入最朴素的想法是给每个位置分配一个可学习的向量和词嵌入一样训练。GPT 系列早期就是这么做的。但它有个明显问题训练时没见过的长度推理时就没法处理。比如你训练时最大长度 512推理时来了个 600 的序列第 513 到 600 个位置根本没有对应的嵌入向量。正弦位置编码Sinusoidal Positional Encoding解决了这个问题。它用一组固定的三角函数生成位置向量不参与训练因此可以外推到任意长度。虽然外推效果在长序列上会打折扣但至少不会直接崩掉。这也是原论文选择它的核心原因。3.2 sin/cos 编码的数学原理与直觉公式长这样$$PE_{(pos, 2i)} \sin\left(\frac{pos}{10000^{2i/d_{model}}}\right)$$ $$PE_{(pos, 2i1)} \cos\left(\frac{pos}{10000^{2i/d_{model}}}\right)$$其中pos是位置i是维度索引。偶数维用 sin奇数维用 cos。这个设计背后的直觉是不同维度对应不同频率的波。低维度i 小频率高变化快能区分相邻位置高维度i 大频率低变化慢能编码长距离信息。这有点像二进制编码——低位变化快高位变化慢组合起来能表示任意整数。更妙的是任意位置pos k的编码可以表示为pos编码的线性变换。这意味着模型可以通过线性操作学到相对位置关系而不只是绝对位置。这一点在理解为什么 sin/cos 比随机初始化更好时非常关键。3.3 手写位置编码模块import torch import torch.nn as nn import math 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: (max_len, 1) position torch.arange(0, max_len, dtypetorch.float).unsqueeze(1) # div_term: (d_model/2,) 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) # 增加 batch 维度: (1, max_len, d_model) pe pe.unsqueeze(0) # 注册为 buffer不参与训练但会随模型保存 self.register_buffer(pe, pe) def forward(self, x): # x: (batch, seq_len, d_model) x x self.pe[:, :x.size(1), :] return self.dropout(x)几个细节值得说清楚。第一register_buffer很关键它让pe成为模型状态的一部分保存和加载时不会丢但不会被优化器更新。第二div_term用exp和log组合计算比直接写10000 ** (2i/d_model)数值更稳定。第三pe[:, 0::2]这种切片赋值要保证维度对齐position * div_term的形状是(max_len, d_model/2)正好填满偶数列。实操心得我一开始写的时候忘了unsqueeze(0)结果x pe广播出错。调试时打印形状是最快的定位方法别嫌麻烦。3.4 位置编码的常见变体与选择建议正弦编码不是唯一选择。实际项目里你会遇到这几种可学习位置嵌入BERT、GPT 早期用的简单直接但外推能力差相对位置编码T5、Transformer-XL 用的编码的是位置差而非绝对位置长序列更友好旋转位置编码RoPELLaMA、Qwen 系列用的通过旋转矩阵注入位置信息外推能力强现在是大模型主流ALiBi直接给注意力分数加一个和距离相关的偏置实现简单选哪个取决于你的场景。如果是固定长度分类任务可学习嵌入完全够用如果是变长序列生成正弦编码或 RoPE 更稳。我个人的经验是做研究和入门手撕先用正弦编码把流程跑通再根据任务换更高级的。4. 多头注意力Transformer 的心脏4.1 从单头到多头为什么要“分头”单头注意力的计算是Query 和 Key 做点积得到注意力分数Softmax 归一化后对 Value 加权求和。公式是$$\text{Attention}(Q, K, V) \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$那为什么要多头核心原因是单头只能学到一种“关注模式”。比如在翻译任务里有的头可能关注语法依赖有的头关注语义相似有的头关注位置邻近。多头让模型在不同子空间里并行学习多种关系最后拼接起来表达能力更强。你可以这样类比单头注意力像一个人用一个视角看问题多头注意力像一组专家各自从不同角度看然后汇总意见。d_model512、nhead8时每个头的维度是 64计算量和单头 512 维差不多但表达能力强得多。4.2 缩放点积注意力那个根号 d_k 到底干嘛的QK^T的结果是点积维度是d_k。当d_k很大时点积的方差会随维度线性增长导致 Softmax 的输入值很大梯度趋近于零也就是 Softmax 饱和。除以sqrt(d_k)就是把方差拉回到 1 附近保证梯度健康。我做过一个小实验d_k64时不除根号Softmax 输出几乎变成 one-hot梯度小得可怜除以 8 之后分布平滑很多训练也稳定。这个细节很多人写代码时会漏但它是原论文明确强调的。4.3 手撕多头注意力完整代码class MultiHeadAttention(nn.Module): def __init__(self, d_model, nhead, dropout0.1): super().__init__() assert d_model % nhead 0, d_model 必须能被 nhead 整除 self.d_model d_model self.nhead nhead self.d_k d_model // nhead # 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.dropout nn.Dropout(pdropout) def forward(self, query, key, value, maskNone): batch_size query.size(0) # 1. 线性投影并拆分成多头 # (batch, seq_len, d_model) - (batch, nhead, seq_len, d_k) Q self.w_q(query).view(batch_size, -1, self.nhead, self.d_k).transpose(1, 2) K self.w_k(key).view(batch_size, -1, self.nhead, self.d_k).transpose(1, 2) V self.w_v(value).view(batch_size, -1, self.nhead, self.d_k).transpose(1, 2) # 2. 缩放点积注意力 scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) # scores: (batch, nhead, seq_len_q, seq_len_k) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn torch.softmax(scores, dim-1) attn self.dropout(attn) # 3. 加权求和 context torch.matmul(attn, V) # context: (batch, nhead, seq_len_q, d_k) # 4. 多头拼接 context context.transpose(1, 2).contiguous().view( batch_size, -1, self.d_model ) # 5. 输出投影 return self.w_o(context)这段代码有几个容易出错的地方我逐个说。第一view和transpose的顺序。w_q(query)输出是(batch, seq_len, d_model)view(batch, -1, nhead, d_k)把它变成(batch, seq_len, nhead, d_k)再transpose(1, 2)变成(batch, nhead, seq_len, d_k)。这个顺序不能反否则头与头之间的数据会串。第二contiguous()不能省。transpose之后张量在内存里不连续直接view会报错。contiguous()把它拷贝成连续内存虽然有一点开销但正确性优先。第三mask 的用法。masked_fill(mask 0, -inf)里mask 为 0 的位置被填成负无穷Softmax 后变成 0相当于屏蔽掉。Decoder 的自注意力需要因果掩码上三角为 0Encoder 的 padding 掩码需要屏蔽 padding 位置。这两种掩码形状不同后面会细讲。4.4 多头注意力的三种使用场景多头注意力在 Transformer 里有三种用法区别只在 Q、K、V 的来源类型Query 来源Key/Value 来源用途自注意力当前层输入当前层输入Encoder、Decoder 内部掩码自注意力当前层输入当前层输入带因果掩码Decoder 防止看到未来交叉注意力Decoder 输入Encoder 输出Decoder 对齐源序列交叉注意力是 Encoder-Decoder 架构的关键。Decoder 在生成每个词时通过 Query 去“询问” Encoder 的输出找到源序列中最相关的部分。这也是机器翻译里对齐关系的来源。注意交叉注意力的 mask 通常来自 Encoder 的 padding 掩码形状是(batch, 1, 1, seq_len_k)需要广播到(batch, nhead, seq_len_q, seq_len_k)。5. Encoder 与 Decoder把模块拼成完整架构5.1 前馈网络看似简单但不能省每个 Encoder/Decoder 层里注意力后面都跟着一个前馈网络FFN。结构是两层线性加 ReLUclass PositionwiseFeedForward(nn.Module): def __init__(self, d_model, d_ff, dropout0.1): super().__init__() self.linear1 nn.Linear(d_model, d_ff) self.linear2 nn.Linear(d_ff, d_model) self.dropout nn.Dropout(pdropout) self.relu nn.ReLU() def forward(self, x): return self.linear2(self.dropout(self.relu(self.linear1(x))))d_ff通常是d_model的 4 倍比如 512 对 2048。为什么先升维再降维我的理解是注意力层做的是线性加权表达能力有限FFN 提供非线性变换让模型能拟合更复杂的函数。升维到 4 倍是为了给非线性变换足够的空间再压回原维度保持残差连接的形状一致。5.2 残差连接与 LayerNorm训练稳定的关键每个子模块注意力、FFN外面都包了一层LayerNorm(x Sublayer(x))。残差连接解决深层网络梯度消失LayerNorm 稳定每层的输入分布。这里有个细节原论文用的是Post-LN先加残差再归一化但后来很多实现改成Pre-LN先归一化再进子模块因为 Pre-LN 训练更稳定不需要 warmup 也能收敛。我手撕时两种都写过实测 Pre-LN 在小数据集上确实更省心。class EncoderLayer(nn.Module): def __init__(self, d_model, nhead, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, nhead, dropout) self.ffn PositionwiseFeedForward(d_model, d_ff, dropout) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout1 nn.Dropout(pdropout) self.dropout2 nn.Dropout(pdropout) def forward(self, x, maskNone): # Pre-LN 写法 attn_out self.self_attn(self.norm1(x), self.norm1(x), self.norm1(x), mask) x x self.dropout1(attn_out) ffn_out self.ffn(self.norm2(x)) x x self.dropout2(ffn_out) return x5.3 Decoder 层三个子模块的协作Decoder 层比 Encoder 多一个交叉注意力而且自注意力要加因果掩码class DecoderLayer(nn.Module): def __init__(self, d_model, nhead, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, nhead, dropout) self.cross_attn MultiHeadAttention(d_model, nhead, dropout) self.ffn PositionwiseFeedForward(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(pdropout) self.dropout2 nn.Dropout(pdropout) self.dropout3 nn.Dropout(pdropout) def forward(self, x, enc_out, src_maskNone, tgt_maskNone): # 1. 掩码自注意力 attn_out self.self_attn(self.norm1(x), self.norm1(x), self.norm1(x), tgt_mask) x x self.dropout1(attn_out) # 2. 交叉注意力Q 来自 DecoderK/V 来自 Encoder cross_out self.cross_attn(self.norm2(x), enc_out, enc_out, src_mask) x x self.dropout2(cross_out) # 3. 前馈网络 ffn_out self.ffn(self.norm3(x)) x x self.dropout3(ffn_out) return x因果掩码的生成很简单就是一个上三角为 0 的矩阵def generate_causal_mask(seq_len): # (1, 1, seq_len, seq_len) mask torch.tril(torch.ones(seq_len, seq_len)).bool() return mask.unsqueeze(0).unsqueeze(0)torch.tril保留下三角含对角线上三角为 0。这样位置 i 只能看到位置 0 到 i看不到未来。5.4 完整模型拼装与形状验证把上面所有模块拼起来class Transformer(nn.Module): def __init__(self, src_vocab, tgt_vocab, d_model512, nhead8, num_enc_layers6, num_dec_layers6, d_ff2048, max_len5000, dropout0.1): super().__init__() self.src_embed nn.Embedding(src_vocab, d_model) self.tgt_embed nn.Embedding(tgt_vocab, d_model) self.pos_enc PositionalEncoding(d_model, max_len, dropout) self.encoder_layers nn.ModuleList([ EncoderLayer(d_model, nhead, d_ff, dropout) for _ in range(num_enc_layers) ]) self.decoder_layers nn.ModuleList([ DecoderLayer(d_model, nhead, d_ff, dropout) for _ in range(num_dec_layers) ]) self.fc_out nn.Linear(d_model, tgt_vocab) self.d_model d_model def forward(self, src, tgt, src_maskNone, tgt_maskNone): # 嵌入 位置编码乘 sqrt(d_model) 是原论文的做法 src_emb self.pos_enc(self.src_embed(src) * math.sqrt(self.d_model)) tgt_emb self.pos_enc(self.tgt_embed(tgt) * math.sqrt(self.d_model)) enc_out src_emb for layer in self.encoder_layers: enc_out layer(enc_out, src_mask) dec_out tgt_emb for layer in self.decoder_layers: dec_out layer(dec_out, enc_out, src_mask, tgt_mask) return self.fc_out(dec_out)嵌入后乘sqrt(d_model)是原论文的细节目的是让嵌入的数值范围和位置编码匹配避免位置信息被淹没。这个操作很多人会忽略但对训练稳定性有影响。验证形状假设batch2, src_len10, tgt_len8, d_model512, tgt_vocab1000输入src(2,10)、tgt(2,8)输出应该是(2, 8, 1000)。我第一次跑通时盯着这个形状看了半天确认无误才敢往下做训练。6. 常见问题与排查技巧实录6.1 维度不匹配类问题速查手撕 Transformer 时90% 的报错都是维度问题。我整理了一张速查表报错信息常见原因解决方法view size is not compatibletranspose 后没 contiguous加.contiguous()The size of tensor a must match位置编码长度不够增大 max_lend_model % nhead ! 0头数不能整除调整 d_model 或 nheadmasked_fill shape mismatchmask 维度没广播对检查 mask 的 unsqueezeExpected 3D input忘了 batch 维度输入加 unsqueeze(0)6.2 训练不收敛的排查思路模型能跑通不代表能训好。我遇到过几次 loss 不降的情况排查下来通常是这几个原因学习率太大。Transformer 对学习率敏感原论文用了 warmup 策略前几千步线性升温再衰减。如果你直接上1e-3很可能一开始就发散。建议从1e-4起步配合 warmup。没有 mask padding。如果 batch 里序列长度不一padding 位置参与注意力会引入噪声。一定要生成 padding 掩码把 padding 位置的注意力分数屏蔽掉。LayerNorm 位置不对。Post-LN 需要 warmupPre-LN 相对宽容。如果你用的是 Post-LN 又没 warmup训练前期会很难看。初始化太随意。PyTorch 默认初始化对 Transformer 来说偏大原论文用了 Xavier 初始化。我一般会手动初始化一遍def init_weights(m): if isinstance(m, nn.Linear): nn.init.xavier_uniform_(m.weight) if m.bias is not None: nn.init.zeros_(m.bias) elif isinstance(m, nn.Embedding): nn.init.normal_(m.weight, mean0, std0.02)6.3 显存爆炸的优化技巧Transformer 的显存占用主要来自注意力矩阵复杂度是O(seq_len^2)。序列长度翻倍显存涨四倍。几个实用的优化手段梯度检查点用时间换空间torch.utils.checkpoint可以省一半显存混合精度torch.cuda.amp自动把部分计算转成 fp16显存和速度都有提升减小 batch size最直接但要注意梯度累积保持等效 batchFlash Attention如果 PyTorch 版本够新F.scaled_dot_product_attention会自动用上速度和显存都优化明显实操心得我调试时习惯先用小模型d_model64, nhead4, 2 层跑通流程确认无误再放大。这样单次迭代几秒钟改起来快不会浪费时间等大模型跑完才发现 bug。6.4 从手撕到实战的过渡建议手撕一遍之后你会发现官方nn.Transformer的封装其实很好用但你已经知道它内部在干什么了。这时候可以尝试几个进阶方向换位置编码把正弦编码换成 RoPE感受一下外推能力的差异改注意力试试稀疏注意力或线性注意力理解效率与表达的权衡迁移到视觉把序列换成图像 patch就是 ViT 的雏形做时序预测把输入换成时间序列输出换成预测值就是一个回归模型我自己在做一个时序预测项目时就是把 Transformer 的 Decoder 去掉只留 Encoder输出接一个线性层做回归。改动很小但效果比 LSTM 好不少尤其是长序列上。最后分享一个我踩过的坑别在第一次手撕时就追求和论文完全一致。原论文的很多细节比如 Post-LN、warmup 的具体步数是为特定数据集调的你照搬未必适合你的任务。先把主干跑通再根据实际情况微调这才是工程上更靠谱的路径。
返回列表