如果你最近在学深度学习,不管是看 NLP 的论文还是刷视觉模型的博客,大概率都会撞见“Transformer”这个名字。第一次听到“转换器”这个翻译的时候,我其实挺懵的——它到底要把什么东西转换成什么?后来自己动手复现、训练、踩坑过几轮之后才发现,这个结构之所以能统治大半个深度学习领域,核心原因就一句话:它让模型学会“看全局”和“找关系”。
这篇文章不是把原论文的公式再抄一遍给你看,而是从我的实际使用体验出发,把 Transformer 的原理、代码、训练技巧和避坑经验一次讲透。内容包括:为什么 RNN 会被它取代、自注意力机制到底在算什么、怎么用 PyTorch 从零实现一个能跑的小模型、Vision Transformer 在图像任务里怎么用、以及我训练时踩过的各种坑。适合刚入门深度学习、想做 NLP 或 CV 项目、或者面试前想系统梳理 Transformer 知识的朋友。
1. 项目概述:为什么要专门写一篇 Transform?
1.1 这个东西到底是什么,能干什么
Transformer 最直白的理解,就是一个“序列到序列”的转换工具。给它一串输入,比如一个句子、一段语音特征、一组图像切片,它输出另一串信息。输出可以是翻译后的句子、加了理解语义的向量表示、或者分类标签。
这个结构最早出现在 2017 年 Google 那篇经典的 Attention Is All You Need 论文里,原本是用来做机器翻译的。后来大家发现,它不仅能翻译,还能做文本分类、命名实体识别、问答系统、语音识别、图像分类、目标检测,甚至最近这几年火得一塌糊涂的大语言模型,底层基本都是 Transformer 堆出来的。
我最早用它是在文本分类任务上。当时公司项目需要做一批长文档的分类,之前用 TextCNN 和 BiLSTM,效果总差口气,尤其是遇到那种五六百字长、关键信息散在各处的文本,RNN 类模型抓不住重点。后来换成了基于 Transformer 的模型,效果提升非常明显。
1.2 它到底解决了什么问题
要理解 Transformer 的价值,就得先看它之前的模型有什么毛病。
RNN 处理序列是一个字一个字往过看,每一步的隐状态依赖上一步的输出。这种串行方式带来两个痛点:一是速度慢,没法并行化;二是长距离信息容易丢,因为信息要一步一步往后传,时间长了前面的内容要么被遗忘,要么梯度消失传到后面什么都没了。
LSTM 和 GRU 缓解了梯度问题,但串行的本质没变。CNN 可以并行处理,但感受野有限,想看到很远的依赖关系就得堆很多层,或者用膨胀卷积,始终不够优雅。
Transformer 换了个思路:既然我们要找一句话里词和词之间的关系,那干脆让每个词直接去和其他所有词做“相关性计算”。你跟我关系大,我就多参考你的信息;关系小,就少看两眼。这种机制叫做自注意力,也就是 Transformer 的核心。它把序列建模彻底并行化,同时天然能捕捉长距离依赖。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 核心原理:撑起 Transformer 的三块基石
2.1 自注意力机制:让每个词看见整句话
自注意力机制是我见过的深度学习结构里,思路最“直觉”的设计之一。可以把它想象成一个会议室:里面每个人都在发言,但每个人注意力有限,只能重点听跟自己议题相关的人说话,其他人顺带听一耳朵。
具体到公式层面,每个词会生成三个向量:Query(查询)、Key(键)、Value(值)。Query 代表“我现在想找什么信息”,Key 代表“我能提供什么信息”,Value 是“我真正给出的内容”。当前词会和序列里所有词的 Key 做点积,得到一个相关性分数,再用 softmax 归一化成权重,最后把所有权重和 Value 加权求和,就得到了这个词的新表示。
用代码实现核心的计算过程大概是这样的:
python复制import torch
import torch.nn.functional as F
def scaled_dot_product_attention(Q, K, V, mask=None):
d_k = Q.size(-1)
scores = torch.matmul(Q, K.transpose(-2, -1)) / (d_k ** 0.5)
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
注意这里有个细节:除以根号 d_k。为什么要除?因为当向量维度变高时,点积的数值会变得很大,大的数值丢进 softmax 之后,梯度区域会变得非常平缓,模型训练起来会很费劲。除以根号 d_k 就是把数值拉回到一个合适的范围,保证梯度能正常传播。这个细节我一开始没注意,结果模型训练得很慢,后来查了论文才发现这个设计是有讲究的。
2.2 多头注意力:从多个角度看问题
一个注意力只能给出一组相关性权重,但句子里的关系往往不止一种:有的注意力要关注语法关系,有的要关注语义关联,有的可能关注指代关系。让模型一次只用一个视角去理解,明显不够用。
多头注意力的思路很简单:把 Query、Key、Value 分别拆成多个头,每个头独立地做注意力计算,最后拼起来再过一层线性变换。这样每个头可以学到不同的关系模式,就像团队里分成几个小组,每个小组负责一种分析维度,最后汇总结果。
看下面的多头注意力核心实现,就是在单头注意力的基础上加了一个“拆分再合并”的过程:
python复制import torch
import torch.nn as nn
import math
class MultiHeadAttention(nn.Module):
def __init__(self, d_model, n_head, dropout=0.1):
super().__init__()
assert d_model % n_head == 0
self.d_model = d_model
self.n_head = n_head
self.d_k = d_model // n_head
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, x):
batch_size, seq_len, _ = x.size()
Q = self.w_q(x).view(batch_size, seq_len, self.n_head, self.d_k).transpose(1, 2)
K = self.w_k(x).view(batch_size, seq_len, self.n_head, self.d_k).transpose(1, 2)
V = self.w_v(x).view(batch_size, seq_len, self.n_head, self.d_k).transpose(1, 2)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
attn = self.dropout(F.softmax(scores, dim=-1))
out = torch.matmul(attn, V)
out = out.transpose(1, 2).contiguous().view(batch_size, seq_len, -1)
return self.w_o(out)
多头注意力后面通常还跟着残差连接和层归一化。残差连接解决的是深层网络退化问题,让梯度能直接流过;层归一化则是把每一层的数据分布拉回稳定区间。这两样东西组合在一起,堆几十层 Transformer 才能稳定训练。
2.3 位置编码:给序列补上顺序信息
自注意力机制本身根本不关心词的先后顺序。你把句子从头到尾排列和从尾到头排列,算出来的相关性分数一模一样。这对自然语言来说是不可接受的,因为“我爱你”和“你爱我”完全不是一个意思。
位置编码就是干这个事的。最常见的做法是用一组正弦和余弦函数,给每个位置生成一个固定长度的向量,然后把它加到词的嵌入向量上。这个向量的特点是:相邻位置之间的编码是连续的,而且任意两个位置的编码可以通过线性变换关联起来,模型就能感知到相对位置关系。
我当时第一次看位置编码公式的时候觉得有点玄乎,后来动手画了一下不同位置的编码曲线才明白:位置 0 和位置 1 的编码向量很相似,位置 50 和位置 100 的编码向量差异就很大。这种“越近越相似,越远越不同”的性质,刚好是模型理解序列顺序所需要的。
为什么不用简单的数字编码?比如位置 1 就加 1,位置 2 就加 2?因为数字编码会让位置信息淹没在嵌入向量里,而且数字本身的数值范围没有上限,不好控制。正余弦编码的好处是数值始终在 -1 到 1 之间,而且能外推到比训练时更长的序列。
3. 实操:用 PyTorch 实现一个能跑的小 Transformer
3.1 环境配置要点
动手之前先说环境。Transformer 对硬件的要求确实比传统模型高一些,但跑小 demo 不需要多贵的显卡。我常用的配置是 PyTorch 2.x、Python 3.9 以上、CUDA 11.8 或更高版本。
没有 GPU 的话,小模型在 CPU 上也能跑,就是慢不少。我的建议是先用小模型、小数据把代码逻辑捋通,再上 GPU 跑全量。
3.2 核心代码:手写一个小型转换器
为了让你直观感受 Transformer 是怎么工作的,我写了一个特别小的演示项目:输入一串数字序列,模型要输出这串数字的倒序。比如输入 [3, 1, 4, 2],就要输出 [2, 4, 1, 3]。这个任务本身没有实际用途,但它能很好地考验模型对序列顺序和整体结构的理解能力。
先定义词表、嵌入和位置编码:
python复制import torch
import torch.nn as nn
import math
import random
VOCAB_SIZE = 12 # 0-9 数字 + 起始符 + 结束符
PAD_IDX = 0
BOS_IDX = 10
EOS_IDX = 11
D_MODEL = 64
N_HEAD = 4
NUM_LAYERS = 2
D_FF = 128
MAX_LEN = 20
BATCH_SIZE = 32
EPOCHS = 30
DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=5000):
super().__init__()
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.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)
self.register_buffer('pe', pe)
def forward(self, x):
return x + self.pe[:, :x.size(1)]
接着定义完整的 Transformer 模型。这里用 PyTorch 自带的 nn.Transformer 接口,足够稳定,也能避免自己实现 Encoder、Decoder 堆叠的重复代码:
python复制class Seq2SeqTransformer(nn.Module):
def __init__(self, vocab_size, d_model, n_head, num_layers, d_ff, max_len):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
self.pos_encoder = PositionalEncoding(d_model, max_len)
self.transformer = nn.Transformer(
d_model=d_model,
nhead=n_head,
num_encoder_layers=num_layers,
num_decoder_layers=num_layers,
dim_feedforward=d_ff,
dropout=0.1,
batch_first=True
)
self.fc_out = nn.Linear(d_model, vocab_size)
def forward(self, src, tgt):
src_emb = self.pos_encoder(self.embedding(src))
tgt_emb = self.pos_encoder(self.embedding(tgt))
output = self.transformer(src_emb, tgt_emb)
return self.fc_out(output)
数据生成部分非常简单,随机生成长度为 3 到 8 的数字序列,倒序作为目标序列。这里做了 Padding,让一个 batch 里的序列长度一致:
python复制def make_batch(batch_size):
src_list, tgt_in_list, tgt_out_list = [], [], []
for _ in range(batch_size):
length = random.randint(3, 8)
seq = [random.randint(1, 9) for _ in range(length)]
src = seq + [PAD_IDX] * (MAX_LEN - length)
rev = seq[::-1]
tgt_in = [BOS_IDX] + rev + [PAD_IDX] * (MAX_LEN - length - 1)
tgt_out = rev + [EOS_IDX] + [PAD_IDX] * (MAX_LEN - length - 1)
src_list.append(src)
tgt_in_list.append(tgt_in)
tgt_out_list.append(tgt_out)
return (torch.tensor(src_list, device=DEVICE),
torch.tensor(tgt_in_list, device=DEVICE),
torch.tensor(tgt_out_list, device=DEVICE))
训练循环里最核心的是两个东西:teacher forcing 和 attention mask。Teacher forcing 指的是解码时每一步都输入上一步的真实目标值,而不是模型的预测值。这样能加速收敛,模型不用从一开始就承担“错了就越错越离谱”的风险。Attention mask 的作用是让模型在预测位置 i 的时候,只能看到位置 i 及之前的内容,看不到未来信息。
python复制model = Seq2SeqTransformer(VOCAB_SIZE, D_MODEL, N_HEAD, NUM_LAYERS, D_FF, MAX_LEN).to(DEVICE)
# 生成解码器掩码:一个上三角矩阵,右上角全为 -inf
tgt_mask = nn.Transformer.generate_square_subsequent_mask(MAX_LEN).to(DEVICE)
criterion = nn.CrossEntropyLoss(ignore_index=PAD_IDX)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
for epoch in range(EPOCHS):
model.train()
total_loss = 0
for _ in range(200):
src, tgt_in, tgt_out = make_batch(BATCH_SIZE)
logits = model(src, tgt_in)
loss = criterion(logits.reshape(-1, VOCAB_SIZE), tgt_out.reshape(-1))
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += loss.item()
print(f'epoch {epoch+1}, loss {total_loss/200:.4f}')
这个小模型在我的笔记本 CPU 上,30 轮大概几分钟就跑完了。跑完以后 loss 能降到 0.05 以下,生成的倒序序列基本全对。
3.3 关键参数选择心得
关于参数设置,我给几个比较通用的经验值。d_model 是嵌入向量的维度,一般取 64、128、256 或 512。任务越复杂、数据量越大,维度可以设得越高。n_head 是注意力头数,常见取 8,但小任务取 4 也完全够,d_model 必须能被 n_head 整除。num_layers 是 Encoder 和 Decoder 的层数,小任务 1 到 2 层就够,大语言模型动辄几十层。
训练 Transformer 最容易出问题的还不是结构,而是学习率。这类模型对学习率极其敏感,我见过很多新手上来就设 1e-3 甚至更高,结果 loss 要么直接飞掉,要么震荡得厉害。比较可控的做法是先用 1e-4 到 5e-4 的较低学习率跑几轮看看趋势,再决定要不要往上调。
4. 从序列到图像:Vision Transformer 到底怎么用
4.1 ViT 做了什么改动
Transformer 在 NLP 领域站稳脚跟之后,大家自然想问:图像能不能用?直接把图像像素拉平当序列送进去是不现实的,一张 224×224 的图就有 5 万个像素点,计算量太大,而且像素和像素之间的空间关系也被破坏了。
ViT 的解决办法是先把图像切成固定大小的 patch,比如 16×16 像素一块,224×224 的图就能得到 196 个 patch。每个 patch 拉平并经过一个线性变换,变成一个嵌入向量,再拼上位置信息,然后送入标准的 Transformer Encoder。后面再接一个分类头,就能做图像分类。
这个设计非常简洁,但它和 CNN 有一个本质区别:CNN 天生自带“局部优先”的归纳偏置,它默认相邻像素相关性高,所以用小卷积核滑动提取特征。ViT 没这个先验,模型要从零开始学习“图像里哪些区域应该关注”。这就导致一个非常现实的问题:数据量不够的时候,ViT 很难打过 CNN。
我当时在一个只有几万张图片的小数据集上做过对比实验,ResNet 明显优于 ViT。但后来用千亿级别的数据预训练再迁移过来,ViT 的表现就反超了。简单来说,ViT 是一个“数据饥饿”的模型,它上限更高,但胃口也更大。
4.2 视觉 Transformer 的实践建议
如果你要在自己的图像项目里用 Transformer,我有几个经验供参考。
数据增强是必需品,而且要比 CNN 时代做得更狠。我自己的项目里用了 RandomResizedCrop、RandAugment、Mixup 和 CutMix 组合拳,效果提升非常明显。原因还是那句话:模型缺少归纳偏置,就得用数据增强把多样性补回来。
预训练权重一定要用。别自己从零训 ViT,尤其在中小数据集上。用 ImageNet 上预训练好的权重做迁移学习,能省下大量算力,效果也会好很多。PyTorch 官方和 timm 库里都有现成的预训练模型,接入成本很低。
最后一招是蒸馏。如果算力不够又想要 ViT 的效果,可以用一个训练好的 CNN 当老师,让 ViT 模仿 CNN 的输出分布。这叫知识蒸馏,DeiT 模型就干过这事,效果非常好,能在中等数据集上训练出超越直接训练 CNN 的模型。
5. 常见问题与排查技巧实录
5.1 训练不收敛或者 loss 震荡
这个问题我碰到的次数最多。排第一的原因是学习率没配好,Transformer 对学习率的忍耐范围很窄。建议用学习率热身策略:前几千步线性地从很小的值升到目标值,后面按步数衰减。我曾经用一个固定学习率怎么都训练不出来,换上 warmup 之后很快就收敛了。
第二个原因是数据没做好 Padding Mask。如果 padding 部分参与注意力计算,模型会去关注一堆没有意义的填充位,干扰真实的语义关系,loss 就降不下去。一定要在注意力计算的时候把 padding 位置屏蔽掉。
第三个原因是模型太深但数据量太小。小数据集没必要堆十几层 Transformer,先用小模型把任务跑通,再考虑放大。
5.2 显存溢出怎么办
Transformer 的显存消耗比 CNN 大得多,尤其是注意力矩阵的计算,会随序列长度平方级增长。512 长度的序列还好,到了 2048 长度,光注意力矩阵就能吃掉几个 G 显存。
我的排查顺序是:先减小 batch size,这个最直接;再开混合精度训练,PyTorch 里用 torch.cuda.amp 就能实现,显存直接砍一半,速度还能快不少;如果序列特别长,用梯度累积模拟较大的 batch size,或者用 gradient checkpointing 以时间换显存。
工业落地场景中,如果序列实在太长,一般会把长文本切成多个段处理,或者用稀疏注意力等变体结构,比如 Longformer、BigBird 这类能处理超长序列的模型。
5.3 效果怎么都不如 CNN
Transformer 不是万能药,我见过不少人抱怨 ViT 效果不如 CNN,结果一看训练集只有几千张图。这种情况下,把模型换成 DeiT 小模型,或者直接退回 CNN,都是合理的选择。
还有一种情况是任务本身并不需要全局建模。比如识别图像里的纹理、颜色这类局部特征占主导的任务,CNN 的归纳偏置让它在小数据上更占优。只有当输入中存在明显的长距离依赖,比如图像里有大目标跨越多个区域、文本有关键信息分散各处时,Transformer 的优势才能发挥出来。
下面把我常踩的坑整理成一个速查表,方便你排查问题的时候对照:
| 现象 | 可能原因 | 解决方向 |
|---|---|---|
| loss 不下降 | 学习率不合适 | 尝试 1e-4 到 5e-4,加 warmup |
| loss 飞到 NaN | 学习率过大,或梯度爆炸 | 降低学习率,加梯度裁剪 |
| 训练集效果好,验证集差 | 模型过拟合 | 加 dropout、数据增强、权重衰减 |
| 序列稍长就显存溢出 | 注意力矩阵太大 | 开混合精度、梯度检查点、减小 batch |
| 小数据上不如 CNN | 缺少归纳偏置 | 用 CNN 做老师蒸馏、用预训练权重 |
| 推理速度太慢 | 模型太大 | 模型蒸馏、量化、剪枝 |
这些坑概括起来就是一句话:Transformer 是个家族,不是单一模型。遇到问题别硬扛,换个变体、调调参数、试试蒸馏,很多时候就解开了。
我个人在实际操作中最深的体会是:理解 Transformer 最好的方式不是反复读公式,而是亲手实现一个最小的模型,再一点一点往里面加东西。先跑通,再优化,最后再回头研究那些细节,你会发现每个设计都不白给。这个小倒序任务也就几十行代码,但跑完之后你对注意力、位置编码、mask 的理解,会比看十遍论文都扎实。
