共计 2878 个字符,预计需要花费 8 分钟才能阅读完成。
背景与痛点
刚接触 Transformer 架构时,很多同学会被其复杂的结构吓退。特别是实现 Encoder 部分时,最常见的三大拦路虎是:

-
位置编码的理解:为什么不能直接用传统 RNN 的序列处理方式?正弦 / 余弦位置编码的数学意义是什么?
-
注意力机制实现:QKV 矩阵究竟如何计算?为什么需要缩放点积注意力?
-
训练不稳定问题:模型初期 loss 震荡剧烈,甚至出现 NaN 值该怎么办?
这些痛点在实际做 AG News 分类任务时会被放大——我们需要处理可变长度的新闻文本,同时要保证模型对关键词的捕捉能力。
技术选型:库 vs 原生实现
面对两种主流方案,我的建议是:
- HuggingFace 优势
- 三行代码调用预训练模型
- 内置优化过的 Attention 计算
-
适合快速原型验证
-
原生 PyTorch 优势
- 彻底掌握模型细节
- 方便自定义修改(比如调整 Encoder 层数)
- 更轻量无依赖
考虑到本文的教学目的,我们选择从零实现。放心,我会带你避开所有深坑!
核心实现四步走
第一步:数据预处理
AG News 数据集包含 4 类新闻标题和描述,我们需要:
from torchtext.datasets import AG_NEWS
from torchtext.data.utils import get_tokenizer
tokenizer = get_tokenizer('basic_english')
train_iter = AG_NEWS(split='train')
# 构建词汇表
vocab = build_vocab_from_iterator(map(tokenizer, [text for label, text in train_iter]),
specials=['<unk>', '<pad>']
)
vocab.set_default_index(vocab['<unk>'])
# 文本向量化函数
def text_pipeline(text):
return vocab(tokenizer(text))
关键点说明:
- 使用基础英文分词器处理标点
- 预留
<unk>和<pad>两个特殊 token - 最终生成形如
[23, 156, 792]的数值序列
第二步:实现 Encoder 层
精简版 SingleHeadAttention 实现(完整版见后续代码):
import torch
import torch.nn as nn
class SelfAttention(nn.Module):
def __init__(self, embed_size):
super().__init__()
self.query = nn.Linear(embed_size, embed_size)
self.key = nn.Linear(embed_size, embed_size)
self.value = nn.Linear(embed_size, embed_size)
self.softmax = nn.Softmax(dim=-1)
def forward(self, x):
Q = self.query(x)
K = self.key(x)
V = self.value(x)
# 缩放点积注意力
scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(Q.size(-1)))
attention = self.softmax(scores)
return torch.matmul(attention, V)
第三步:组装完整模型
class TransformerClassifier(nn.Module):
def __init__(self, vocab_size, embed_size=128, num_classes=4):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed_size)
self.position_encoding = PositionalEncoding(embed_size) # 需自行实现
self.encoder_layer = nn.TransformerEncoderLayer(
d_model=embed_size,
nhead=8,
dim_feedforward=512
)
self.transformer_encoder = nn.TransformerEncoder(self.encoder_layer, num_layers=3)
self.fc = nn.Linear(embed_size, num_classes)
def forward(self, x):
x = self.embedding(x)
x = self.position_encoding(x)
x = self.transformer_encoder(x)
x = x.mean(dim=1) # 全局平均池化
return self.fc(x)
第四步:训练技巧
三个关键超参数设置:
-
学习率:使用带 warmup 的 AdamW 优化器
optimizer = AdamW(model.parameters(), lr=5e-5) scheduler = get_linear_schedule_with_warmup(optimizer, num_warmup_steps=100, num_training_steps=1000) -
Batch Size:根据 GPU 显存选择(建议 32-64)
-
梯度裁剪:防止梯度爆炸
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
五大避坑指南
- 梯度消失问题
- 解决方案:每层添加残差连接
-
代码实现:
x = x + self.attention(x) # 残差连接 -
过拟合现象
- 对策:在 Embedding 后添加 Dropout 层
-
推荐参数:
nn.Dropout(p=0.1) -
位置编码失效
- 关键检查:确保 PE 值范围与 Embedding 匹配
-
调试方法:可视化前几个位置的编码向量
-
内存溢出(OOM)
-
应急方案:
torch.cuda.empty_cache() reduce_batch_size() -
预测时结果随机
- 根本原因:忘记
model.eval()模式 - 完整预测流程:
with torch.no_grad(): model.eval() output = model(input)
进阶实战建议
当模型准确率稳定在 90%+ 后,可以考虑:
-
模型压缩:使用知识蒸馏技术,将大模型的能力迁移到小模型
student_model = TinyTransformer() distil_loss = KLDivLoss(teacher_logits, student_logits) -
部署优化:转换为 ONNX 格式提升推理速度
torch.onnx.export(model, dummy_input, "ag_news.onnx")
思考题
- 如果新闻文本特别长(如超过 512 个 token),应该如何修改当前架构?
- 如何修改注意力机制使其能捕捉局部特征(类似 CNN 的效果)?
- 在多语言场景下,位置编码需要做哪些特殊处理?
希望这篇笔记能帮你打通 Transformer 的任督二脉!遇到问题欢迎在评论区交流~
正文完
发表至: 未分类
近两天内
