
Transformer 这个词我最早是在 2017 年的那篇《Attention Is All You Need》里看到的当时第一反应是“又一个涨点的模型”没想到后来它几乎重塑了整个深度学习版图。无论是做 NLP、CV 还是时间序列预测只要你还在接触模型就一定绕不开 Transformer。这篇博文我想从一个实践者的角度把“初见 Transformer”时需要搞清楚的架构、原理、代码和踩坑经验一次性讲透。我不打算堆公式而是用做项目的思路把每个模块为什么存在、怎么实现、有哪些坑讲明白适合刚入门深度学习的同学也适合想快速上手 Transformer 做实验的工程师。1. 初见Transformer它到底解决了什么问题1.1 从RNN的痛点说起在 Transformer 出现之前序列建模的主流工具是 RNN、LSTM 和 GRU。RNN 的核心逻辑是按时间步逐个处理输入当前时刻的隐藏状态依赖于上一个时刻的输出。这个设计天然适合序列但也带来两个致命问题一是并行性差因为每一步都要等前一步算完训练速度上不去二是长距离依赖问题虽然 LSTM 通过门控机制缓解了梯度消失但信息在传递过程中还是会衰减句子一长前面的关键信息就容易被丢掉。我当年用 LSTM 做文本分类时最头疼的就是长文本。输入 500 个字以后模型基本只记得后半段的内容前面的重要实体经常被忽略。当时的解决办法无非是加大隐层维度、做双向编码、加注意力机制但这些都是治标不治本。Transformer 的思路是彻底抛弃循环结构一次性看到整个序列然后用注意力机制直接建模任意两个位置之间的关系。这样既解决了并行问题也让长距离依赖变得不再困难。1.2 Transformer的核心思想注意力机制注意力机制的本质是“按相关性加权提取信息”。在处理一个词时模型不是只看这个词本身而是根据它与序列中其他词的相关程度把整个序列的信息加权汇总。这个“相关程度”就是通过 Query、Key、Value 三个向量计算出来的。你可以把注意力理解为在公司里开评审会Query 是你当前需要解决的问题Key 是每个参会者的擅长领域Value 是每个参会者能提供的具体建议。你会先根据 Query 和每个 Key 的匹配度决定该听谁的再把所有人的建议按匹配度加权汇总。自注意力就是让序列中的每个元素都作为“提问者”去询问其他元素从而获得全局上下文。这一步是 Transformer 所有能力的根源。它不像 CNN 那样只能看到局部感受野也不像 RNN 那样靠循环逐步传递信息而是直接建立全连接的关系。代价是计算复杂度是 O(n²) 的也就是序列越长计算量增长越快。这也是后续很多优化工作的核心突破口比如稀疏注意力、窗口注意力等后面讲到视觉 Transformer 时会再提。2. 架构拆解编码器与解码器的秘密2.1 输入嵌入与位置编码PE计算详解Transformer 的输入首先是 token 序列。在 NLP 里token 通常是一个词或子词在图像领域token 可能是一个图像块。每个 token 会通过一个嵌入层映射成 d_model 维的向量这个向量就是模型能处理的语义表示。但纯粹的嵌入向量没有位置信息。自注意力是“对顺序不敏感”的它把序列当成一个集合来处理把“我爱你”和“你爱我”看成完全一样的输入。为了打破这种对称性Transformer 引入了位置编码Positional Encoding。位置编码有两种常见方式一种是让模型自己学习一套位置嵌入Learned Positional Embedding另一种是使用固定的三角函数公式Sinusoidal Positional Encoding。论文里用的是后者公式为import numpy as np def sinusoidal_positional_encoding(max_len, d_model): pe np.zeros((max_len, d_model)) for pos in range(max_len): for i in range(0, d_model, 2): pe[pos, i] np.sin(pos / (10000 ** (2 * i / d_model))) if i 1 d_model: pe[pos, i 1] np.cos(pos / (10000 ** (2 * i / d_model))) return pe这个公式看似绕实际思路是用不同频率的正弦和余弦波来编码位置。偶数维度用 sin奇数维度用 cos。为什么用三角函数而不是直接用 0、1、2、3 这样的整数因为归一化后的数值范围有限而且三角函数可以通过线性变换表达相对位置关系。比如 PE(posk) 可以由 PE(pos) 的某个线性组合近似得到这有助于模型学习位置之间的相对关系。实际工程中很多预训练模型也直接采用可学习位置嵌入效果差异不大但三角函数方案不需要训练参数且能外推到比训练时更长的序列。这里的 d_model 是模型的隐藏维度max_len 是最大序列长度。在图像任务中ViT 沿用了可学习位置嵌入因为它要处理的是固定大小的图像块序列不需要外推。2.2 多头自注意力与前馈网络自注意力Self-Attention的计算流程可以分成三步先把嵌入向量通过权重矩阵映射成 Query、Key、Value然后计算 Q 和 K 的点积并缩放再通过 softmax 得到注意力权重最后加权 Value。公式写作import torch import torch.nn.functional as F def scaled_dot_product_attention(Q, K, V, maskNone): d_k Q.size(-1) scores torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k, dtypetorch.float32)) if mask is not None: scores scores.masked_fill(mask 0, -1e9) attn_weights F.softmax(scores, dim-1) output torch.matmul(attn_weights, V) return output, attn_weights除以 sqrt(d_k) 这个缩放操作经常被忽略但它非常关键。如果 d_k 很大Q 和 K 点积的方差也会很大导致 softmax 的梯度极小、模型难以训练。缩放后注意力分布会更平滑训练更稳定。多头注意力就是把刚才的流程并行做 h 次每次使用不同的权重矩阵得到 h 个不同的表示子空间。之所以用多头而不是仅仅加大单头维度是因为不同的头可以关注不同的关系模式。翻译任务里有的头关注语法依赖有的头关注指代关系有的头关注相邻词。项目实践中我最常用的配置是 8 个头。增加头数能让模型更灵活但头数过多也会导致每个头分到的维度太窄反而学不到有效信息一般确保每个头的维度是 64 左右比较稳妥。自注意力之后是一个前馈网络Feed-Forward NetworkFFN通常是两层全连接加一个 ReLU 激活函数。FFN 的作用是对每个 token 的位置进行非线性变换增强模型的表达能力。这里有个容易忽略的点FFN 是逐位置共享的也就是同一个 FFN 会独立作用于序列中的每个 token用代码实现时一般是一个 Conv1d 或者两个 Linear 层。2.3 残差连接、层归一化与掩码每个子层注意力、FFN外面都会接一个残差连接和层归一化LayerNorm。残差连接帮助梯度直接流过深层网络层归一化则保证每一层的输入分布稳定。Transformer 中的 LayerNorm 是对每个 token 的 d_model 维做归一化而不是对 batch 或 channel 做归一化这一点与 BN 不同原因是序列长度会动态变化LayerNorm 不受 batch size 影响处理变长输入更稳定。解码器中还有两个关键掩码。第一个是 Padding Mask用来屏蔽掉输入中补齐的无效 token第二个是 Look-Ahead Mask也叫 causal mask保证模型在预测第 i 个 token 时只能看到前 i-1 个 token不能看到未来的信息。这个掩码通常实现为一个上三角矩阵在计算注意力分数时把未来位置置为负无穷。我见过不少新手在实现 masked multi-head attention 时忘记把 mask 传给所有头只加在了一个头上结果训练时 loss 一直不降。排查了很久才发现是 mask 广播维度出了问题。建议把 mask 的 shape 设计成 [batch, 1, seq_len, seq_len]这样就能自动广播到所有头。3. 不止NLPTransformer的视觉版图3.1 Vision TransformerViT怎么把图片变成序列ViT 的出发点很简单既然 Transformer 能处理序列那为什么不把图片也切成小方块当成序列输入具体做法是把一张 H×W×C 的图片切成长宽为 P 的 patch得到 N 个图像块每个图像块展平后通过线性映射变成 d_model 维的嵌入向量。为了让模型知道每个 patch 的位置还要加上位置嵌入并在序列开头加一个特殊的 [CLS] token它最终对应的输出向量就用来做分类。我之前第一次跑 ViT 时觉得这个映射有点粗暴但实验效果确实好。在 ImageNet 这类大数据集上ViT 能超过同期 CNN因为它有全局感受野。不过这个模型也很吃数据在小数据集上直接训练效果不如 ResNet原因是它缺少 CNN 内置的归纳偏置。解决办法是先在大规模数据上预训练再迁移到小数据集上微调。如果你想自己实现核心是把Rearrange操作加入数据流在 PyTorch 里可以用einopsfrom einops import rearrange # x: [batch, channels, height, width] patches rearrange(x, b c (h p1) (w p2) - b (h w) (p1 p2 c), p1patch_size, p2patch_size)这一步会把图片变成类似文本序列的形式后面的 Transformer 处理就可以完全复用 NLP 的代码。这也是为什么我说学 Transformer 一定要先吃透编码器结构因为视觉模型只是把输入换成了 patch核心模块没变。3.2 Swin Transformer与层级化设计ViT 的全局注意力在大尺寸图片上计算量太大因为 patch 数量多O(n²) 的复杂度无法接受。Swin Transformer 的思路是引入窗口注意力window attention只在局部窗口内做自注意力然后通过 shift window 的方式让不同窗口之间交换信息。这样既保持了 Transformer 的表达能力又能构建类似 CNN 的金字塔层级结构特征图尺寸逐层缩小适合做检测、分割和密集预测。Swin Transformer 的另一个启发是架构设计要尊重输入数据的物理结构。文本一维所以用一维位置编码图像二维所以要设计二维相对位置偏置relative position bias。相对位置偏置是 Swin 的一个关键技巧它让注意力分数在计算时加上一个可学习的偏置项这个偏置只与两个 token 之间的相对位置有关。这样做比绝对位置嵌入更高效而且泛化更强实际做视觉任务时可以直接借鉴这个思路。我做目标检测任务时经常把 Swin 作为骨干网络替换掉 ResNet在 COCO 数据集上 mAP 有明显提升。但要注意Swin 的窗口划分和 shift 实现比普通 ViT 复杂代码不好调试建议在理解论文的基础上先跑官方开源代码再动手改不要一上来就重写。3.3 目标检测与多模态应用Transformer 在目标检测领域最有代表性的工作是 DETR 和它的后续变体。DETR 把检测当成集合预测问题直接用 Transformer 输出一组目标框和类别而不需要手工设计锚框和后处理。它利用二分图匹配Hungarian Algorithm把预测框和真实框一一对应然后用 Transformer 编码器和解码器完成目标查询。虽然 DETR 收敛慢但它的整体流程非常优雅省去了大量工程细节。最近两年Transformer 也大量应用于多模态场景比如高光谱图像分类、RGB-T可见光与热红外行人检测、无人机感知等。核心思路是用不同模态的编码器分别提取特征再通过跨模态注意力cross-modal attention融合。以 RGB-T 检测为例可见光图像偏纹理细节热红外图像提供温度信息两者在某些条件下互补性很强。用变形可变形交叉注意力deformable cross-attention可以在弱对齐数据下依然保持较好的融合效果。多模态 Transformer 的难点不是模型定义而是数据对齐。如果两个模态的图片没有精确对齐简单的拼接或相加效果很差。实操中我一般先用特征级对齐模块做粗对齐再输入 Transformer 做融合这样比直接在原始像素上做跨模态注意力稳定得多。4. 手写一个Mini Transformer从零开始的前向传播4.1 数据准备与超参数设定为了讲清楚前向传播我直接写了一个极简 Transformer 编码器用 PyTorch 实现用来做序列预测。这个例子可以去掉了复杂的数据预处理方便看到模型骨架。我先定义超参数d_model 512 # 嵌入维度 n_heads 8 # 多头注意力头数 n_layers 6 # 编码器层数 d_ff 2048 # 前馈网络隐藏维度 max_len 128 # 最大序列长度 vocab_size 10000 # 词表大小 batch_size 32这里的 d_ff 通常是 d_model 的 4 倍左右为什么因为 Transformer 论文里的 FFN 维度就是 2048对应的 d_model 是 512。更大的 d_ff 能增加模型容量但参数量和计算量也会明显上升。很多轻量级 Transformer 会把 d_ff 压缩到 2 倍或 3 倍比如 Restormer 就针对计算效率做了不少结构上的精简。输入数据我直接用随机整数序列模拟目的是走通流程x torch.randint(0, vocab_size, (batch_size, max_len))如果你要做时间序列预测这里的“token”就是数值型特征可以先把序列切片成固定窗口然后做归一化。后面我会专门讲时间序列的注意事项。4.2 代码实现核心模块先写嵌入和位置编码class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len128): 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() * (-np.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) # shape [1, max_len, d_model] self.register_buffer(pe, pe) def forward(self, x): return x self.pe[:, :x.size(1)]注意这里用了register_buffer这样位置编码会随模型移动到 GPU但不会参与训练。接着是多头注意力class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_heads): super().__init__() 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.w_o nn.Linear(d_model, d_model) def forward(self, x, maskNone): batch_size, seq_len, _ x.size() Q self.w_q(x).view(batch_size, seq_len, self.n_heads, self.d_k).transpose(1, 2) K self.w_k(x).view(batch_size, seq_len, self.n_heads, self.d_k).transpose(1, 2) V self.w_v(x).view(batch_size, seq_len, self.n_heads, self.d_k).transpose(1, 2) scores Q K.transpose(-2, -1) / math.sqrt(self.d_k) if mask is not None: scores scores.masked_fill(mask 0, -1e9) attn F.softmax(scores, dim-1) out attn V out out.transpose(1, 2).contiguous().view(batch_size, seq_len, -1) return self.w_o(out)这里我习惯先contiguous()再view()因为transpose之后张量不是连续的直接view会报错。这是很常见的坑。接下来是 FFN 和编码器层class FeedForward(nn.Module): def __init__(self, d_model, d_ff): super().__init__() self.net nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model) ) def forward(self, x): return self.net(x) 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) self.ffn FeedForward(d_model, d_ff) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout nn.Dropout(dropout) def forward(self, x, maskNone): x x self.dropout(self.self_attn(self.norm1(x), mask)) x x self.dropout(self.ffn(self.norm2(x))) return x这里我用了“Pre-LN”结构也就是先做 LayerNorm再做子层计算。原始论文是“Post-LN”子层计算后再归一化。Post-LN 在深层网络中训练不稳定经常需要 warmup 配合Pre-LN 在深模型中更稳定也是现代大多实现的标准做法。4.3 训练与预测示例定义完整编码器和预测头class TransformerEncoder(nn.Module): def __init__(self, vocab_size, d_model, n_heads, d_ff, n_layers, max_len, num_classes1): super().__init__() self.embed nn.Embedding(vocab_size, d_model) self.pe PositionalEncoding(d_model, max_len) self.layers nn.ModuleList([ EncoderLayer(d_model, n_heads, d_ff) for _ in range(n_layers) ]) self.norm nn.LayerNorm(d_model) self.fc_out nn.Linear(d_model, num_classes) def forward(self, x, maskNone): x self.pe(self.embed(x)) for layer in self.layers: x layer(x, mask) x self.norm(x) return self.fc_out(x[:, 0]) # 取第一个token的输出类似CLS训练代码和普通 MLP 没有区别model TransformerEncoder(vocab_size, d_model, n_heads, d_ff, n_layers, max_len) optimizer torch.optim.Adam(model.parameters(), lr1e-4) loss_fn nn.MSELoss() for epoch in range(10): optimizer.zero_grad() output model(x) loss loss_fn(output, y) loss.backward() optimizer.step() print(fepoch {epoch}, loss: {loss.item():.4f})如果你只想跑通前向传播这个代码足够了。但真正做预测任务时我建议把学习率设置成动态的比如前 N 步 warmup后面再衰减。Transformer 对学习率比较敏感固定学习率训练时稍不留神就 loss 震荡甚至发散。5. 实战经验调参、踩坑与常见问题5.1 训练不收敛怎么办Transformer 训练不收敛是最常见的问题。我遇到的情况无非几种一是学习率太大。Transformer 对 Adam 优化器比较友好但 lr 一般在 1e-4 到 1e-5 区间。使用 warmup 策略可以显著提升稳定性先让学习率从 0 线性升到峰值再按指数或余弦衰减。二是数据和标签不对齐。我在做时间序列需要数据窗口时经常出现把未来数据当标签输入的情况模型看起来在“学”实际上是偷看了未来信息测试效果极差。还有分类任务里如果标签从 1 开始而不是 0 开始也可能导致损失异常。三是多头注意力维度整除问题。d_model 必须是 n_heads 的整数倍。我会在定义模型时用断言检查assert d_model % n_heads 0, d_model must be divisible by n_heads四是没有加 LayerNorm 位置放错了。如果你用的是 Post-LN后续网络层多了就梯度爆炸直接把 norm 调整成 Pre-LN 能解决大多数不稳定问题。5.2 显存爆炸与序列长度Transformer 最大的痛点是显存开销。自注意力分数矩阵的大小是 [batch_size, n_heads, seq_len, seq_len]序列长度翻倍显存占用就变成四倍。输入长度 1024 时单样本还能扛一旦到 4096一般单卡就爆了。我之前处理长文档时最直接的办法是截断把超过 512 的部分直接丢掉效果损失很大。后来用了两个技巧第一个是窗口注意力或者局部注意力只让每个 token 关注临近的一部分 token计算量降到 O(n×w)w 是窗口大小。第二个是梯度检查点gradient checkpointing以时间换空间。前向传播时不保存中间激活值反向传播时重新计算能省下大量显存代价是训练速度变慢。这个很适合单卡调参的阶段。如果只是推理还能用 Flash Attention 这类高效实现它在 GPU 上做 IO 优化不仅显存省速度还更快。现在 PyTorch 已经内置了scaled_dot_product_attention可以直接替换手写注意力建议优先用这个。5.3 时间序列预测的注意事项Transformer 用于时间序列预测时很多人直接套 NLP 的代码结果效果不如 LSTM于是就说 Transformer 不适合时间序列。实际上问题往往出在数据处理上。时间序列和文本不同它的趋势性和季节性会影响模型性能。我一般会先做差分处理把非平稳序列变成平稳序列再输入模型。其次位置编码只表达了位置顺序没有表达时间间隔。如果你的数据是不等间隔采样建议加入时间戳特征作为额外信息或者把时间间隔编码进注意力。另外预测长度和输入长度要保持合理的比例。我测试下来的经验是输入窗口至少 2 到 5 倍于预测长度效果比较好。如果要做长序列预测可以考虑专门设计的 Time Series Transformer或者在注意力里加入稀疏化设计而不是把序列硬塞进去。6. 从看懂到玩转学习路线与变体速览6.1 从论文到代码的学习路径第一次接触 Transformer 的人我建议按这条路径走能省不少时间先看《Attention Is All You Need》论文原文只抓核心图不看附录公式再找一个注释详尽的开源实现把前向传播每一步的 shape 打印出来盯着张量维度的变化看一遍然后自己动手写一个 mini 版本只需要支持前向传播不需要训练最后跑一个小任务比如文本分类或简单的序列预测把反向传播一跑通整个模型就真正属于你了。很多人一上来就刷各种讲解视频只看不动手结果看完还是不会写代码。Transformer 是一个工程性很强的模型只靠看是学不会的。哪怕是把别人代码抄一遍也比只看图解强。如果想深入理解可解释性可以试试 Transformer Explainer 这类可视化工具它能把注意力权重、embedding 变化展示出来。我每次调模型时也会用可视化看看注意力矩阵很多时候能直观发现某个 head 是否崩塌成了单点注意力。6.2 Transformer变体速览Transformer 的变体非常多我梳理几条主线方便你按需选择NLP 方向有 BERT编码器主导、GPT解码器主导、T5编码器-解码器。BERT 适合做理解类任务GPT 适合做生成类任务。视觉方向有 ViT、Swin、DeiT 等DeiT 用蒸馏方法解决了 ViT 需要超大训练集的问题。轻量化方向有 Restormer 等针对图像复原任务大幅优化了计算量。多模态方向有 CLIP、ALBEF 等用对比学习对齐图像和文本特征。目标检测有 DETR、Deformable DETR用 Transformer 做集合预测。还有各种 Transformer 改进版比如 Longformer、BigBird 用稀疏注意力处理长文本Performer 用核近似方法降低复杂度的。我建议不要盲目追新还是回到自己的任务需求。先分析输入数据的特点是局部相关为主还是全局相关为主再决定用标准注意力还是局部注意力然后选一个成熟的预训练权重开始微调比从零训练省太多资源。我个人做项目时最常用的组合是“标准 Transformer 编码器 针对性位置编码 合适的注意力模式”大部分任务都能在这个骨架上跑通。真正困难的从来不是搭模型而是认清数据、理解任务、找到适合的归纳偏置。这也是我见过了这么多变体之后回头再看“初见 Transformer”时最深的感受把基础架构吃透剩下的都是围绕它做加减法。