ARTICLE DETAIL

资讯详情

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

从零手写Transformer:原理拆解与PyTorch实现

从零手写Transformer:原理拆解与PyTorch实现 Transformer 这个架构我从 2019 年开始反复啃从最开始看《Attention Is All You Need》一头雾水到后来自己用 NumPy 从零手写一遍、再用 PyTorch 复现一遍前后踩过的坑能写满一个笔记本。网上讲 Transformer 的文章很多但大部分要么停留在画框图上要么直接甩一段官方代码让你自己悟。这篇我想做的是把这两件事接起来——每一个模块先讲清楚它到底在解决什么问题、为什么这么设计然后立刻给出能跑的最小实现最后再拼成完整的 Encoder-Decoder。位置编码、多头注意力、残差加 LayerNorm、前馈网络、掩码机制这些热搜词里反复出现的东西我都会一个个拆开揉碎。这篇适合谁看如果你已经知道 Transformer 大概长什么样但一看到nn.MultiheadAttention就想跳过或者面试被问位置编码为什么用 sin/cos答不上来那这篇就是给你写的。我会尽量用生活化的类比解释原理代码部分保证可以直接复制运行参数计算过程也会写清楚。全程不依赖任何高级封装核心模块全部手写这样你才能真正理解每一行在干什么。1. 整体架构拆解与设计思路1.1 Transformer 到底想解决什么问题在 Transformer 出现之前处理序列任务的主流是 RNN 和 LSTM。它们的工作方式像一个字一个字读句子的人读完第一个字把记忆传给第二个字依次往下。这种方式有两个硬伤一是没法并行句子有多长就得串行算多少步GPU 再强也白搭二是长距离依赖会衰减句子开头的信息传到结尾时已经所剩无几。Transformer 的核心思路是把逐字传递换成全局直接连接。句子里的每个词都同时和所有其他词算一个相关性分数谁跟我关系近我就多吸收谁的信息。这样一来任意两个位置之间的距离都是 1 步长距离依赖问题自然消失而且所有位置的计算可以完全并行。这就是自注意力机制Self-Attention的威力。但天下没有免费的午餐。全局连接意味着计算量随序列长度平方增长这是 Transformer 最大的代价。另外把序列当成一个集合来处理就丢失了词的顺序信息所以必须额外注入位置编码。这两点贯穿了整个架构的设计。1.2 Encoder-Decoder 的分工与数据流原始 Transformer 是为机器翻译设计的结构上分两大块Encoder编码器负责理解输入句子把每个词变成一个富含上下文信息的向量。它由 N 层堆叠而成原论文 N6。每一层里有两个子层——多头自注意力和前馈网络。Decoder解码器负责生成目标句子每生成一个词都要参考两部分信息已经生成的历史通过带掩码的自注意力和 Encoder 的输出通过交叉注意力。每层里有三个子层。数据从输入到输出的完整路径是这样的输入词 → 词嵌入 → 加位置编码 → Encoder 多层处理 → 得到记忆矩阵 → Decoder 逐层处理自注意力 交叉注意力 前馈→ 线性层 → Softmax → 输出概率分布。我建议你在读后面每一节时脑子里都挂着这条数据流知道当前这个模块处在链条的哪个位置它的输入从哪来、输出给谁。这样零散的模块才能串成一个整体。1.3 为什么选择手写而不是直接调库PyTorch 的nn.Transformer一行就能建好整个模型为什么还要手写我的理由有三个。第一调库你永远不知道batch_first到底怎么影响维度、src_mask和tgt_mask有什么区别出错了只能瞎试。第二面试和实际调优时你需要改注意力头数、改位置编码方式、加自定义掩码这些都得动底层。第三手写一遍之后你对张量形状的敏感度会质变看到(batch, seq, dim)就知道下一步该转置还是该广播。提示手写不代表生产环境也要手写。实际项目里用官方实现或成熟框架更稳妥手写的价值在于理解和调试能力。2. 核心模块原理与逐行实现2.1 词嵌入与位置编码给词加上座位号词嵌入很直白就是把每个词 ID 映射成一个d_model维的稠密向量。原论文d_model512。这里有个热搜里常问的问题词嵌入矩阵是随机的吗答案是初始化时是随机的通常用正态分布但它会随着训练更新最终学出语义相近的词向量也相近。另外原论文里嵌入层会乘一个sqrt(d_model)目的是让嵌入的数值范围和位置编码匹配避免位置信息被淹没。位置编码才是重点。因为自注意力本身对顺序不敏感你把句子打乱算出来的注意力结果只是跟着换位置语义上完全等价。所以必须显式告诉模型谁在前谁在后。原论文用的是正弦余弦函数import numpy as np import torch import torch.nn as nn import math 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, dtypetorch.float).unsqueeze(1) 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) # (1, max_len, d_model) self.register_buffer(pe, pe) def forward(self, x): # x: (batch, seq_len, d_model) return x self.pe[:, :x.size(1), :]为什么用 sin/cos 而不是简单的 0,1,2,3我理解有三层原因。第一正弦函数的取值范围在 [-1,1]和归一化后的嵌入量级接近相加不会破坏原有信息。第二不同维度用不同频率低维变化快、高维变化慢相当于给每个位置生成了一个独特的频率指纹。第三也是最美的一点PE(posk)可以表示成PE(pos)的线性变换这意味着模型能通过线性操作推断相对位置对训练时没见过的更长序列也有一定泛化能力。div_term那个指数写法是数值稳定性的技巧等价于1 / (10000 ** (2i/d_model))但直接算幂容易溢出用 explog 更稳。这个细节很多教程不讲但你自己写的时候如果发现数值异常八成是这里的问题。2.2 缩放点积注意力注意力的最小内核注意力机制的本质是一次加权查询。你有三个矩阵Query我要找什么、Key我有什么标签、Value我的实际内容。用 Q 和 K 算相似度得到权重再用权重对 V 加权求和。公式是Attention(Q,K,V) softmax(QK^T / sqrt(d_k)) V。这里的sqrt(d_k)缩放是必须的。原因在于当d_k很大时Q 和 K 的点积方差会随维度线性增长导致 softmax 的输入值很大输出会变得极其尖锐接近 one-hot梯度几乎消失。除以sqrt(d_k)把方差拉回 1 附近softmax 才能工作在梯度健康的区间。def scaled_dot_product_attention(q, k, v, maskNone): # q: (batch, heads, seq_q, d_k) # k: (batch, heads, seq_k, d_k) # v: (batch, heads, seq_k, d_v) d_k q.size(-1) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_k) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn torch.softmax(scores, dim-1) output torch.matmul(attn, v) return output, attn掩码这里用-inf而不是 0是因为 softmax 里exp(-inf)0这样被屏蔽的位置权重严格为零。如果你填 0softmax 之后那些位置还是会有非零权重等于没屏蔽干净。这是新手最容易犯的错之一。2.3 多头注意力让模型从多个角度看问题单个注意力头只能学到一种关注模式。多头注意力的做法是把d_model切成 h 份每份独立做一次注意力最后拼回来再过一个线性层。原论文 h8每个头d_k d_v d_model/h 64。为什么多头有效打个比方读一句话时你可能同时关心语法主谓关系、指代关系、语义相似度。单个头被迫把这些混在一起多头则让每个头专注一个子空间。实测下来不同的头确实会分化出不同的功能有的盯相邻词有的盯句法结构。class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads, dropout0.1): super().__init__() assert d_model % num_heads 0 self.d_model d_model self.num_heads num_heads self.d_k d_model // num_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.w_o nn.Linear(d_model, d_model) self.dropout nn.Dropout(dropout) def forward(self, q, k, v, maskNone): batch_size q.size(0) # 线性投影后拆头: (batch, seq, d_model) - (batch, heads, seq, d_k) q self.w_q(q).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) k self.w_k(k).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) v self.w_v(v).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) out, attn scaled_dot_product_attention(q, k, v, mask) # 拼头: (batch, heads, seq, d_k) - (batch, seq, d_model) out out.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model) return self.w_o(out), attn维度变换是这里最容易出错的地方。view之后必须transpose因为我们要让 heads 维度独立出来参与批量矩阵乘。最后拼回去时先transpose再contiguous再viewcontiguous不能省否则view会报错。我当初就是漏了这个debug 了半小时。2.4 前馈网络与残差连接稳定深层的两大支柱每个子层后面都跟着Add Norm也就是残差连接加层归一化。残差连接解决的是深层网络梯度消失问题让梯度能通过捷径直接回传。LayerNorm 则是对每个样本的特征维度做归一化稳定训练。这里有个细节Transformer 用的是 Post-LN先残差后归一化还是 Pre-LN先归一化后残差原论文是 Post-LN但后来大家发现 Pre-LN 训练更稳定不需要 warmup 也能收敛。现在主流实现多用 Pre-LN。我两种都试过小模型差别不大层数一深 Pre-LN 明显更省心。前馈网络本身很简单两层线性加 ReLU中间维度通常是d_model的 4 倍class PositionwiseFeedForward(nn.Module): def __init__(self, d_model, d_ff2048, dropout0.1): super().__init__() self.linear1 nn.Linear(d_model, d_ff) self.linear2 nn.Linear(d_ff, d_model) self.dropout nn.Dropout(dropout) def forward(self, x): return self.linear2(self.dropout(torch.relu(self.linear1(x))))为什么中间要放大 4 倍我的理解是给模型足够的非线性容量去变换特征同时保持输入输出维度一致以便残差相加。这个 4 倍是经验值不是铁律小数据集上可以调小。3. 完整 Encoder-Decoder 搭建与实操3.1 Encoder 层与 Encoder 堆叠一个 Encoder 层 多头自注意力 前馈网络每个子层都包一层 AddNorm。堆叠 N 层就得到完整的 Encoder。注意自注意力里 Q、K、V 都来自同一个输入这就是自的含义。class EncoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, num_heads, 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(dropout) self.dropout2 nn.Dropout(dropout) 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 x class Encoder(nn.Module): def __init__(self, vocab_size, d_model, num_heads, d_ff, num_layers, max_len5000, dropout0.1): super().__init__() self.embedding nn.Embedding(vocab_size, d_model) self.pos_encoding PositionalEncoding(d_model, max_len) self.layers nn.ModuleList([ EncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) self.norm nn.LayerNorm(d_model) self.d_model d_model def forward(self, x, maskNone): x self.embedding(x) * math.sqrt(self.d_model) x self.pos_encoding(x) for layer in self.layers: x layer(x, mask) return self.norm(x)热搜里有人问Transformer 编码部分有多少编码器答案是原论文 6 层但这是超参可以调。层数越多表达能力越强但计算量和显存也线性增长。小任务 2-4 层往往就够。3.2 Decoder 层与两种掩码Decoder 比 Encoder 多一个交叉注意力子层它的 Q 来自 Decoder 自身K 和 V 来自 Encoder 的输出。这就是 Decoder 看输入句子的方式。Decoder 的自注意力必须加因果掩码causal mask保证预测第 t 个词时只能看到前 t-1 个词不能偷看未来。这个掩码是一个上三角为 0 的矩阵。另外还有 padding 掩码用来屏蔽填充位置。class DecoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, num_heads, dropout) self.cross_attn MultiHeadAttention(d_model, num_heads, 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(dropout) self.dropout2 nn.Dropout(dropout) self.dropout3 nn.Dropout(dropout) def forward(self, x, enc_out, self_maskNone, cross_maskNone): attn_out, _ self.self_attn(self.norm1(x), self.norm1(x), self.norm1(x), self_mask) x x self.dropout1(attn_out) cross_out, _ self.cross_attn(self.norm2(x), enc_out, enc_out, cross_mask) x x self.dropout2(cross_out) ffn_out self.ffn(self.norm3(x)) x x self.dropout3(ffn_out) return x def make_causal_mask(seq_len, device): # 下三角为 1上三角为 0 mask torch.tril(torch.ones(seq_len, seq_len, devicedevice)).bool() return mask.unsqueeze(0).unsqueeze(0) # (1, 1, seq, seq)掩码的形状要能广播到(batch, heads, seq_q, seq_k)。因果掩码是(1,1,seq,seq)padding 掩码通常是(batch,1,1,seq_k)两者可以相乘合并。我踩过的坑是掩码维度对不上导致广播错误结果注意力算出来是错的但程序不报错loss 就是降不下去。所以每次改掩码我都会打印形状确认。3.3 完整模型组装与一次前向把 Encoder、Decoder、输出层拼起来就是完整的 Transformer。输出层是一个线性层把d_model映射到词表大小接 softmax 得到概率。class Transformer(nn.Module): def __init__(self, src_vocab, tgt_vocab, d_model512, num_heads8, d_ff2048, num_layers6, max_len5000, dropout0.1): super().__init__() self.encoder Encoder(src_vocab, d_model, num_heads, d_ff, num_layers, max_len, dropout) self.decoder nn.ModuleList([ DecoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) self.tgt_embedding nn.Embedding(tgt_vocab, d_model) self.pos_encoding PositionalEncoding(d_model, max_len) self.fc_out nn.Linear(d_model, tgt_vocab) self.d_model d_model self.dropout nn.Dropout(dropout) def forward(self, src, tgt, src_maskNone, tgt_maskNone): enc_out self.encoder(src, src_mask) x self.tgt_embedding(tgt) * math.sqrt(self.d_model) x self.pos_encoding(x) for layer in self.decoder: x layer(x, enc_out, tgt_mask, src_mask) return self.fc_out(x)跑一次前向验证形状model Transformer(src_vocab1000, tgt_vocab1000, d_model64, num_heads4, d_ff256, num_layers2) src torch.randint(0, 1000, (2, 10)) tgt torch.randint(0, 1000, (2, 8)) tgt_mask make_causal_mask(8, src.device) out model(src, tgt, tgt_masktgt_mask) print(out.shape) # torch.Size([2, 8, 1000])形状对上了说明整个数据流是通的。这一步很重要很多人模型写完直接开训结果维度错了半天找不到原因。先用小维度、小 batch 跑通形状再上真实数据。3.4 参数计算模型到底有多大热搜里transformer 架构模型参数计算是个高频问题。我以d_model512, num_layers6, d_ff2048, h8为例算一下。单个多头注意力有 4 个d_model×d_model的权重矩阵Q、K、V、O参数量4 × 512 × 512 ≈ 1.05M。前馈网络是512×2048 2048×512 ≈ 2.1M。所以一个 Encoder 层约1.05M 2.1M 3.15M6 层约 18.9M。Decoder 层多一个交叉注意力约3.15M 1.05M 4.2M6 层约 25.2M。加上嵌入层vocab×512如果词表 3 万就是 15.4M源和目标各一份。粗算下来原论文规模在 60-65M 参数左右。这个计算能力在面试里很实用能快速估算显存需求推理时约参数量 × 4 字节训练时还要加上优化器状态和激活值通常是参数量的 3-4 倍。4. 常见问题排查与避坑实录4.1 训练不收敛的排查顺序Transformer 训练不收敛我总结了一套排查顺序按这个走基本能定位问题。排查项常见症状检查方法掩码方向loss 卡住不降打印掩码矩阵确认因果掩码是下三角位置编码输出与顺序无关打乱输入看输出是否变化学习率loss 震荡或爆炸用 warmup峰值 1e-4 到 5e-4维度变换报错或结果异常每步打印 tensor.shape梯度梯度为 nan加梯度裁剪 clip_grad_norm_学习率 warmup 是 Transformer 的标配。原论文的公式是lr d_model^(-0.5) × min(step^(-0.5), step × warmup^(-1.5))。前 warmup 步线性增长之后按步数平方根衰减。我实测下来不加 warmup 很容易在训练初期就发散尤其是 Post-LN 结构。4.2 位置编码的常见误区第一个误区是以为位置编码必须可学习。其实正弦编码是固定的效果和可学习编码差不多而且能外推到更长序列。第二个误区是忘了乘sqrt(d_model)。嵌入层不缩放的话位置编码的幅度可能盖过词嵌入模型学不到词义。第三个误区是位置编码加在了错误的位置。它必须在嵌入之后、进入 Encoder 之前加而且 Encoder 和 Decoder 各加各的。注意如果你用的是 batch_first 的库实现位置编码的维度顺序要跟着调整别直接套用(seq, batch, dim)的写法。4.3 多头注意力的调试技巧多头注意力出问题最隐蔽的是头塌缩——所有头学到了一样的模式等于白搭。判断方法是把注意力权重可视化看不同头的分布是否雷同。如果雷同可能是初始化不好或者学习率太大。另一个技巧是检查注意力权重的行和是否为 1。softmax 之后每行应该严格归一化如果不是说明掩码或 softmax 维度用错了。我习惯在调试时加一句assert torch.allclose(attn.sum(-1), torch.ones_like(attn.sum(-1)))能快速发现维度问题。4.4 显存不够怎么办Transformer 显存杀手是注意力矩阵复杂度O(seq²)。序列长度翻倍注意力显存翻四倍。几个实用手段一是减小 batch size用梯度累积模拟大 batch二是用混合精度训练显存能省近一半三是序列太长时考虑分块或稀疏注意力。我处理长文本时通常先把 batch 降到 8 以下再开混合精度基本能撑住 1024 长度。5. 从理解到扩展的几条路手写完这个基础版你其实已经拿到了继续深入的钥匙。想往视觉走可以看 Vision Transformer 怎么把图像切成 patch 当 token想往时序预测走把 Decoder 去掉只留 Encoder 就是不错的基线想研究高效注意力可以对比 Linformer、Performer 这些变体怎么把O(seq²)降下来。我自己在实际操作中的体会是Transformer 的难点从来不在单个模块而在于把十几个模块的维度、掩码、残差路径全部对齐。你手写一遍踩过的每一个维度错误都会变成以后调模型时的直觉。第一次跑通完整前向的那一刻比看十篇讲解都管用。后面再遇到什么 Swin、Point Transformer 这些变体你会发现它们的创新点往往只在某一个模块骨架还是这套东西。
返回列表