如果你正在做 NLP、图像分类、时序预测,或者只是刷到“Transformer 涨点”“手撕 Transformer”这类词,却还不太清楚它内部到底怎么运转,这篇内容就是给你准备的。
先给一个明确判断:Transformer 不是一个只属于 NLP 的模型结构,它已经从语言模型走向了视觉、语音、时序预测、推荐系统等几乎所有深度学习领域。真正理解它,不是背下来“自注意力 + 多头 + 位置编码”这几个名词,而是能回答下面三个问题:为什么 RNN 会被它替代?它的每一层结构到底在处理什么?如果让我从零实现一个最小可用版本,我该怎么写?
这篇文章不会只停留在概念层面。我会从问题切入,讲清楚 Transformer 的核心原理,然后用 PyTorch 从零实现一个可以训练的分类模型,跑通训练和验证流程。接着再把视野扩展到 Vision Transformer、Swin Transformer 这些热门变体,最后给出工程落地和调试建议。
如果你之前看公式看得头大,或者复制过别人的代码却不知道怎么改,这篇文章会尽量让你读完后能自己动手。
1. 为什么最后是 Transformer
很多人第一次接触 Transformer,是在学习 NLP 的时候。当时主流的序列建模工具是 RNN、LSTM、GRU。它们有一个天然问题:按时间步顺序处理序列。这意味着第 100 个词要等前 99 个词算完才能开始,训练速度慢,而且长距离依赖容易丢失。虽然 LSTM 通过门控机制缓解了梯度消失,但本质上仍然受限于“按顺序”这个约束。
CNN 在 NLP 里也被用过。TextCNN 通过不同尺寸的卷积核提取 n-gram 特征,优点是能并行计算,缺点是感受野有限。想要捕捉长距离关系,就必须堆很多层或者用很大的卷积核,效率不高。
Transformer 换了一个思路:不再依赖顺序处理,而是让序列中的每个元素直接和所有元素计算相关性。这个机制叫自注意力。它带来两个关键变化:
- 计算可以并行,训练速度大幅提升。
- 任意两个位置之间只隔一次计算,长距离依赖不再是难题。
所以“为什么最后是 Transformer”这个问题的答案,可以概括为:它同时解决了 RNN 的串行瓶颈和 CNN 的局部感受野限制,而且随着数据量和算力增大,它的扩展性远好于前两者。更重要的是,Transformer 的架构足够通用,输入不一定非得是文本,只要你能把数据变成一组向量,就能用 Transformer 处理。
从工程角度看,Transformer 还带来了一个隐性优势:统一建模。以前做文本用 RNN,做图像用 CNN,做语音用专门的模型。现在 Transformer 提供了统一的基础结构,不同模态的数据经过适当编码后,都能塞进同一个架构。这也是 GPT、BERT、ViT、Swin Transformer 等模型真正重要的原因——它们共享同一套底层设计逻辑。
当然,Transformer 不是没有代价。它的计算复杂度是序列长度的平方,显存占用大,训练需要更多数据。这也是后面 Swin Transformer 这类模型尝试优化的方向之一。了解它的优点和局限,才算真正理解它。
2. 核心机制:自注意力与多头注意力
2.1 自注意力想解决什么问题
先看一个具体场景。假设输入一句话:“小明把苹果放在桌上,然后拿走了它。”要让模型知道“它”指代的是“苹果”还是“小明”,就需要让“它”这个位置的向量,能参考其他位置的向量。自注意力做的事情,就是让每个 token 根据自己的 Query 向量,去所有 token 的 Key 向量上做匹配,再用匹配结果对 Value 向量加权求和。
用大白话说:每个词都发出一条查询,问“谁和我相关?”,然后根据收到的答案,从其他词那里汇总信息。这个汇总结果就是当前词的新表示。整个过程可以做一次矩阵运算,完全并行。
2.2 缩放点积注意力公式
自注意力的核心公式是:
Attention(Q, K, V) = softmax(Q K^T / sqrt(d_k)) V其中:
- Q 是 Query 矩阵,代表当前元素要查询的信息。
- K 是 Key 矩阵,代表其他元素能被匹配的特征。
- V 是 Value 矩阵,代表其他元素实际提供的内容。
- d_k 是 Key 向量的维度,除以 sqrt(d_k) 是为了防止点积结果过大,导致 softmax 进入饱和区。
从实现角度,Q、K、V 通常来自同一个输入序列 X,经过不同的线性变换得到:
Q = X W_Q K = X W_K V = X W_V三种变换使用不同的权重矩阵,才能让模型在同一份输入上学到不同视角的表示。
2.3 多头注意力:不是一种注意力,而是多套并行
如果只做一次自注意力,模型只能学到一种相关性模式。但真实语言中的关系是多样的,可能是语法关系、指代关系、语义相似关系等。多头注意力把 Q、K、V 拆成 h 份,每份独立计算注意力,最后把所有头的结果拼接起来,再经过一个线性层。
公式如下:
MultiHead(Q, K, V) = Concat(head_1, ..., head_h) W_O head_i = Attention(Q W_Q^i, K W_K^i, V W_V^i)多头的好处有三个:
- 不同头可以关注不同位置的关系。
- 每个头运行在更低维空间,计算成本不会成倍增加。
- 增加模型并行度,表达能力更强。
实际项目中,BERT-base 使用 12 个头,GPT 使用 12 个头,ViT 的 large 版本使用 16 个头。头数不是越大越好,头数过大会导致每个头的维度太小,表达能力下降,也会增加训练开销。
2.4 注意力机制的一般视角
还有一点值得理解:注意力机制不是 Transformer 独有的。早年机器翻译中的 Bahdanau Attention 和 Luong Attention 就已经用注意力来对齐源语言和目标语言。Transformer 的贡献在于把它从辅助模块变成了主架构,并且用自注意力取代了所有循环结构。所以理解 Transformer,本质上就是理解自注意力在深层网络中的组织和实现方式。
3. Transformer 总体架构拆解
3.1 标准架构:编码器和解码器
原始论文《Attention Is All You Need》中,Transformer 采用编码器-解码器结构,两条线分别处理后输出。
编码器由 N 个相同的层堆叠,每层包含两个子层:
- 多头自注意力层。
- 逐位置前馈网络。
每个子层后面都接一个残差连接和层归一化。用公式表示就是:
x = LayerNorm(x + Sublayer(x))解码器与编码器类似,但有两点不同:
- 解码器使用带掩码的自注意力,防止当前位置看到未来位置。
- 解码器额外插入一个交叉注意力子层,让解码器能关注编码器的输出。
很多实际任务不一定需要完整的编码器-解码器结构。比如 BERT 只用编码器,适合理解任务;GPT 只用解码器,适合生成任务。这个取舍在工程上非常常见,后面讲视觉变体时还会看到。
3.2 位置编码:给并行模型一个顺序概念
自注意力本身不关心 token 的先后顺序,因为它同时计算所有两两关系。要引入顺序信息,必须在输入向量里注入位置信号。
原始 Transformer 使用正弦位置编码:
PE(pos, 2i) = sin(pos / 10000^(2i/d_model)) PE(pos, 2i + 1) = cos(pos / 10000^(2i/d_model))其中 pos 是位置索引,i 是维度索引,d_model 是模型维度。每个维度的频率不同,模型可以从角度关系里学到相对位置信息。
工程上,还有几种常见位置编码:
- 可学习位置编码:把位置向量当作普通参数训练,BERT 和 ViT 都采用了这种方案。
- 相对位置编码:建模两个 token 之间的距离而非绝对位置。
- 旋转位置编码(RoPE):在 LLaMA 等大模型中被广泛使用,能更好地处理长序列。
理解位置编码的重要性,是因为很多从零实现 Transformer 的人,最容易忽略的就是这里。如果你的模型在训练时表现不错,但序列一变长效果就很差,可以先检查位置编码的方案。
3.3 前馈网络与层归一化
每个 Transformer 块中的前馈网络(Feed-Forward Network,FFN)由两个线性层和一个激活函数组成。常用的设置为:
FFN(x) = max(0, x W_1 + b_1) W_2 + b_2第一个线性层将维度从 d_model 扩展到 4 * d_model,第二个线性层再降回 d_model。中间用一个 ReLU 激活函数,有些实现会换成 GELU。
层归一化(LayerNorm)对每个 token 的特征维度做归一化。它与 BatchNorm 的主要区别:
- BatchNorm 对一个 batch 的同一特征维度做归一化,依赖 batch 大小。
- LayerNorm 对单个样本的所有特征做归一化,不依赖 batch 大小,在 NLP 和 Transformer 中更稳定。
公式:
LayerNorm(x) = (x - mean) / sqrt(var + eps) * gamma + beta这里的 gamma 和 beta 是可学习参数。
3.4 残差连接的意义
Transformer 层数一般很深,BERT-base 有 12 层,GPT-3 有 96 层。如果没有残差连接,梯度很难传到浅层。残差连接让每一层的输出变为 x + Sublayer(x),相当于把原始信息沿着网络直接传递。这既缓解了梯度消失,也保证了模型不会因为层数增加而明显退化成恒等映射。
4. 环境准备与前置条件
在动手写代码前,先说明运行环境。本文核心演示使用 Python 和 PyTorch,具体版本以你实际安装为准,思路在不同版本下都适用。
建议环境:
- 操作系统:Windows / Linux / macOS 均可。
- Python 版本:3.8 或更高。
- PyTorch:2.0 或更高。
- CUDA:如果你有 NVIDIA 显卡,建议安装 CUDA 版 PyTorch,训练会快很多。
- Jupyter Notebook 或 VS Code 均可。
如果还没安装 PyTorch,可以用下面的命令安装 CPU 版本:
pip install torch torchvision需要 GPU 支持的话,建议到 PyTorch 官网选择对应的 CUDA 版本安装命令,这里不写死某个版本的 CUDA 号,以免因为显卡驱动不匹配导致安装失败。
安装完成后,可以运行一段代码验证环境:
import torch print(torch.__version__) print(torch.cuda.is_available()) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(device)如果打印的版本正常,并且torch.cuda.is_available()在有 GPU 的机器上返回 True,就说明环境没问题。
5. 手撕 Transformer:PyTorch 从零实现
这一章是全文核心。我们不用现成的nn.Transformer,而是手动实现每个组件,这样你能真正理解内部机制。
5.1 项目结构
为了方便维护,建议用下面的文件结构:
transformer-tutorial/ ├── data.py # 数据准备 ├── model.py # Transformer 模型定义 ├── train.py # 训练脚本 └── utils.py # 辅助函数这里为了控制篇幅,把关键代码放在 model.py 和 train.py 中,方便组合运行。
5.2 模型定义:model.py
先引入依赖:
import torch import torch.nn as nn import math然后是缩放点积注意力:
class ScaledDotProductAttention(nn.Module): def __init__(self, dropout=0.1): super().__init__() self.dropout = nn.Dropout(dropout) def forward(self, q, k, v, mask=None): 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_weights = torch.softmax(scores, dim=-1) attn_weights = self.dropout(attn_weights) output = torch.matmul(attn_weights, v) return output, attn_weights这段代码有几个关键细节:
- scores 的维度是 [batch_size, num_heads, seq_len, seq_len]。
- 除以 sqrt(d_k) 是为了稳定梯度。
- mask 中为 0 的位置会被替换成负无穷,softmax 后这些位置的概率趋近于 0。
- 返回 attn_weights 是为了方便可视化注意力权重。
然后是单头注意力模块:
class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_head, dropout=0.1): super().__init__() assert d_model % n_head == 0, "d_model must be divisible by n_head" 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.attention = ScaledDotProductAttention(dropout) self.dropout = nn.Dropout(dropout) def _split_heads(self, x): batch_size, seq_len, _ = x.size() x = x.view(batch_size, seq_len, self.n_head, self.d_k) x = x.transpose(1, 2) return x def forward(self, q, k, v, mask=None): batch_size = q.size(0) q = self._split_heads(self.w_q(q)) k = self._split_heads(self.w_k(k)) v = self._split_heads(self.w_v(v)) x, attn_weights = self.attention(q, k, v, mask) x = x.transpose(1, 2).contiguous().view(batch_size, -1, self.n_head * self.d_k) output = self.w_o(x) return output这里有一个很容易出错的地方:把多头拆开计算后,要记得把维度重新拼接回 [batch_size, seq_len, d_model]。同时,view之前要确保 tensor 在内存中是连续的,所以需要调用contiguous()。
接下来是位置编码:
class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=512, dropout=0.1): super().__init__() self.dropout = nn.Dropout(dropout) 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).transpose(0, 1) self.register_buffer("pe", pe) def forward(self, x): x = x + self.pe[: x.size(0), :] return self.dropout(x)位置编码使用正弦余弦函数,好处是它能外推到更长序列,并且相对位置信息隐含在相位差中。
下面是编码器层:
class FeedForward(nn.Module): def __init__(self, d_model, d_ff, dropout=0.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): x = self.linear1(x) x = torch.relu(x) x = self.dropout(x) x = self.linear2(x) return x class EncoderLayer(nn.Module): def __init__(self, d_model, n_head, d_ff, dropout=0.1): super().__init__() self.self_attn = MultiHeadAttention(d_model, n_head, dropout) self.ffn = FeedForward(d_model, d_ff, dropout) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.dropout = nn.Dropout(dropout) def forward(self, x, mask=None): x = x + self.dropout(self.self_attn(x, x, x, mask)) x = self.norm1(x) x = x + self.dropout(self.ffn(x)) x = self.norm2(x) return x注意这里先做残差,再 LayerNorm,这种写法叫 Post-Norm,是原始 Transformer 的实现方式。实际工程中很多模型改用 Pre-Norm,即先 LayerNorm 再做残差,训练更稳定,后面会细说。
最后是完整编码器模型:
class TransformerEncoder(nn.Module): def __init__(self, vocab_size, d_model, n_head, n_layers, d_ff, max_len, num_classes, dropout=0.1): super().__init__() self.embedding = nn.Embedding(vocab_size, d_model) self.positional_encoding = PositionalEncoding(d_model, max_len, dropout) self.layers = nn.ModuleList([ EncoderLayer(d_model, n_head, d_ff, dropout) for _ in range(n_layers) ]) self.norm = nn.LayerNorm(d_model) self.fc_out = nn.Linear(d_model, num_classes) def forward(self, x, mask=None): x = self.embedding(x) x = self.positional_encoding(x) for layer in self.layers: x = layer(x, mask) x = self.norm(x) cls_rep = x[:, 0, :] # 取第一个 token 的表示作为分类结果 logits = self.fc_out(cls_rep) return logits这里我们采用一个常见做法:在每个序列开头加一个特殊的 [CLS] token,最后用它的表示做分类。这个思路来自 BERT,在文本分类任务里非常实用。
5.3 训练脚本:train.py
再写一个最小训练脚本。数据部分用一个小型文本分类演示,你可以替换成自己的数据集。
import torch import torch.nn as nn from torch.utils.data import DataLoader, Dataset from model import TransformerEncoder # 假设有 4 条简单样本 texts = [ "this movie is great", "the film is terrible", "I love this book", "what a waste of time" ] labels = [1, 0, 1, 0] # 构造词表 def build_vocab(texts): vocab = {"<pad>": 0, "<cls>": 1} for text in texts: for word in text.split(): if word not in vocab: vocab[word] = len(vocab) return vocab vocab = build_vocab(texts) max_len = 6 def encode(text, vocab, max_len): tokens = ["<cls>"] + text.split()[: max_len - 1] ids = [vocab.get(w, 0) for w in tokens] ids = ids + [0] * (max_len - len(ids)) return torch.tensor(ids, dtype=torch.long) class TextDataset(Dataset): def __init__(self, texts, labels, vocab, max_len): self.texts = texts self.labels = labels self.vocab = vocab self.max_len = max_len def __len__(self): return len(self.texts) def __getitem__(self, idx): x = encode(self.texts[idx], self.vocab, self.max_len) y = torch.tensor(self.labels[idx], dtype=torch.long) return x, y dataset = TextDataset(texts, labels, vocab, max_len) dataloader = DataLoader(dataset, batch_size=2, shuffle=True) model = TransformerEncoder( vocab_size=len(vocab), d_model=32, n_head=4, n_layers=2, d_ff=64, max_len=max_len, num_classes=2, dropout=0.1 ) criterion = nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4) for epoch in range(30): total_loss = 0 for x, y in dataloader: optimizer.zero_grad() logits = model(x) loss = criterion(logits, y) loss.backward() optimizer.step() total_loss += loss.item() if (epoch + 1) % 5 == 0: print(f"Epoch {epoch + 1}, Loss: {total_loss / len(dataloader):.4f}")这里使用AdamW而不是传统 Adam,因为它把权重衰减和梯度更新解耦,是 Transformer 训练中最常用的优化器。
5.4 这份实现缺了什么
可以看到,上面的实现是一个编码器模型,适合做分类,但它不是完整的编码器-解码器 Transformer。如果你想做翻译或生成任务,还需要实现:
- DecoderLayer,加入掩码自注意力和交叉注意力。
- masked_fill 的因果掩码,确保位置 i 只能看到它之前的位置。
- 训练时用 teacher forcing,推理时逐步生成。
不过,对大多数想理解 Transformer 核心机制的人来说,理解编码器已经建立了一个非常重要的基础。
6. 运行结果与效果验证
把上面两个文件放在同一目录下,然后运行:
python train.py预期输出类似:
Epoch 5, Loss: 0.6123 Epoch 10, Loss: 0.4371 Epoch 15, Loss: 0.2789 Epoch 20, Loss: 0.1682 Epoch 25, Loss: 0.1094 Epoch 30, Loss: 0.0731如何判断模型训练成功了?
- Loss 是否在持续下降。如果 loss 停滞或升高,说明学习率可能太大,或代码存在 bug。
- 训练完成后,可以看预测结果:
model.eval() with torch.no_grad(): sample = encode("this film is terrible", vocab, max_len).unsqueeze(0) logits = model(sample) pred = torch.argmax(logits, dim=-1).item() print("Predicted label:", pred)如果训练顺利,这条样本应该输出 0。如果输出 1,说明模型还没收敛,可以增加训练轮数。
从实操角度看,这里第一个要检查的地方不是 loss,而是:
- 词表构建是否正确,有没有把词映射到正确的索引。
- 输入 batch 是否完整填充到相同的长度。
- 位置编码的维度是否和 embedding 维度一致。
7. 从 NLP 到视觉:Vision Transformer 与 Swin Transformer
Transformer 真正让人惊讶的地方,是它从语言模型进入了计算机视觉领域,并且开始挑战 CNN 的地位。
7.1 Vision Transformer 的核心思路
Vision Transformer 的做法很直接。把一张图片切成一堆 patch,比如 224x224 的图片切成 16x16 的 patch,会得到 196 个 patch。每个 patch 展平成一个向量,经过线性映射后变成 token。再加上一个 [CLS] token 和位置编码,然后送入标准的 Transformer 编码器。最后用 [CLS] token 的表示做分类。
这个设计最吸引人的一点是:几乎没有为视觉任务定制专门的架构,直接把图像变成了序列,就取得了很好的效果。它证明 Transformer 的建模能力不限于文本。
但 ViT 有一个明显弱点:需要大规模数据预训练。因为 patch 之间没有 CNN 那种天然的先验知识,在小数据集上容易过拟合。
7.2 Swin Transformer 的改进
Swin Transformer 为了解决 ViT 的问题,引入了两个关键设计:
- 层次化结构:不同阶段逐渐降低分辨率、增加通道数,类似 CNN 的 pyramid 结构。
- 窗口注意力:在局部窗口内计算自注意力,限制计算复杂度。
窗口注意力虽然减少了计算量,但窗口之间信息无法交互。Swin Transformer 为此引入了 shifted window 机制,交替移动窗口,让不同窗口之间的 token 有机会互相看到。这样既保留了 Transformer 的建模能力,又让计算复杂度从 O(n^2) 降到了 O(n)。
从工程角度看,Swin Transformer 的一个价值在于:它更容易复用到检测、分割等稠密预测任务上。ViT 处理这类任务时需要额外设计,而 Swin 的层次化结构天然适配。
7.3 视觉 Transformer 的落地选择
做图像分类时,到底该选 CNN 还是 Transformer?
一个务实的建议是:
- 如果数据集很小,比如几千张图片,建议先从 ResNet 这类 CNN 入手。
- 如果有大规模数据,或者能用预训练权重,ViT 和 Swin Transformer 值得优先考虑。
- 如果要部署到边缘设备,CNN 在速度、显存占用、推理优化上通常更省心。
这里并不是说 Transformer 一定比 CNN 好,只能说它的架构选择面更宽、上限更高,但需要的数据和算力也更多。
8. Transformer 的改进方向与工程实践
8.1 训练稳定性
原始 Transformer 的 Post-Norm 结构在深层网络下容易出现训练不稳。现在的主流做法是 Pre-LayerNorm,即:
x = x + Sublayer(LayerNorm(x))这个改动虽然简单,但对深层模型有明显帮助。GPT、BERT 后续版本以及很多开源大模型,都采用了 Pre-Norm 结构。理解这一点的价值在于:你在复现别人代码时,会看到两种不同的写法,不要觉得是错误,只是设计选择不同。
8.2 长序列优化
Transformer 平方复杂度的短板催生了一系列优化方法:
- sparse attention:只让每个 token 关注部分位置,而不是全部。
- FlashAttention:从访存优化的角度减少显存占用,不改变数学结果。
- Longformer、BigBird:针对超长文本设计稀疏注意力模式。
- 上下文扩展:在大模型中通过调整位置编码的方式支持更长上下文。
实际应用里,如果你只是处理几千 token 的文本,标准注意力完全够用。如果文本动辄几万甚至几十万 token,就需要考虑这些优化手段。
8.3 工程落地建议
在实际业务中,很少有人真的从随机初始化开始训练一个 Transformer。最稳妥的路径是:
- 使用预训练模型,比如 BERT、RoBERTa、ViT、Swin Transformer。
- 在自己的领域数据上做微调(fine-tuning)。
- 评估效果时,不仅看准确率,还要看推理延迟、显存占用、模型体积。
这里要特别提醒一点:如果你用预训练模型,上游模型使用的分词器和你的文本处理方式必须一致。很多乱码和效果差的问题,源头其实是分词器配置不对,而不是模型结构改错了。
8.4 关于“涨点”这件事
热搜里经常能看到“Transformer 涨点”的说法。所谓涨点,是指通过调整模型结构或训练策略,在某个 benchmark 上提升指标。常见涨点手段包括:
- 改位置编码:从绝对位置编码换成 RoPE。
- 调整初始化:某些初始化策略对深层网络的收敛速度影响很大。
- 用更好的激活函数:比如把 ReLU 换成 GELU 或 SwiGLU。
- 调整 dropout 位置:在 attention 计算后的 dropout 和 FFN 后的 dropout,效果差异较明显。
但涨点往往依赖具体数据和任务。别人在论文里涨点,不意味着你的业务也一定涨。务实的做法是,把改进当作实验变量,每次只改一个因素,记录效果,而不是一股脑堆叠所有技巧。
9. 常见问题与排查方法
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| Loss 不下降 | 学习率过大或过小;数据归一化不一致 | 打印梯度统计;尝试多个学习率 | 使用 warmup + 合适的学习率(如 3e-4) |
| 训练时报 NaN | 注意力分数过大;分母为 0 | 检查输入是否包含 NaN;检查位置编码 | 确认 mask 正确;增加 eps;降低学习率 |
| 序列长度变化时报错 | 位置编码 max_len 设置过小 | 查看报错堆栈中的 reshape 行 | 增大 max_len,或改用相对位置编码 |
| 多头注意力维度不匹配 | d_model 无法被 n_head 整除 | 检查 assert 条件 | 调整 d_model 或 n_head |
| 预测结果总是同一个类别 | 类别不平衡;模型未收敛 | 查看验证集 loss;打印 logits 分布 | 先训练足够轮数;必要时调整类别权重 |
| 位置编码无效 | 位置编码加在错误的维度上 | 打印位置编码 shape 和输入 shape | 对齐 max_len 与输入序列长度 |
| GPU 显存不足 | 序列过长,注意力矩阵太大 | 逐步缩小 batch_size 或 max_len | 使用梯度累积;采用 sparse attention |
第一排查优先级永远是数据。模型不会凭空产生错误输出,绝大多数问题都能追溯到输入数据、mask 或者词表构造阶段。代码里加上断言和日志,能省下大量调试时间。
10. 下一步实践路径
如果你读完这篇文章,最好的实践方式不是再去背概念,而是按下面两条路选一条走。
第一条路径:从零改代码。把我给出的实现继续完善,比如加上解码器、实现因果掩码,然后训练一个简单的中文文本生成模型。这个过程会逼迫你理解每个矩阵的维度变化。
第二条路径:用开源库做项目。如果你关心的是应用层,可以先把 Hugging Face Transformers 这类库用熟,用预训练 BERT 做文本分类,用 ViT 做图像分类。通过微调任务反向理解模型内部机制,也是很多工程师的实际学习路径。
我个人更推荐这两条路结合。先在开源库上跑通一个任务,再回来看代码实现,很多之前看不懂的概念会瞬间串起来。
Transformer 的核心其实不复杂,它只是把“根据上下文加权地更新每个元素”这件事做到了极致。真正复杂的是它衍生出来的工程实践:数据处理、训练技巧、推理优化、多模态融合。这些都需要你在具体项目中一点点积累。把这个最小实现跑通,是理解整个生态的第一步。
建议把文章里这几段代码保存成模板,下次遇到 Transformer 相关项目时,你会感谢当初愿意从零写一遍的自己。