1. 为什么绕不开Transformer:从一个序列建模难题说起
做深度学习这几年,我越来越觉得Transformer已经从一个"可选的模型架构"变成了"默认的基线方案"。无论你做什么方向——NLP、CV、语音、时间序列预测,最后几乎都会遇到同一个问题:该不该上Transformer?怎么用Transformer?用了之后为什么效果不稳定?
几年前序列建模的主流还是RNN和LSTM。它们的问题很本质:必须一个词一个词按顺序处理。这意味着训练慢、难以并行,而且长距离依赖捕捉能力受限于链式传递带来的梯度衰减。后来有人用CNN替代RNN,比如TCN(时间卷积网络),靠膨胀卷积扩大感受野,虽然能并行,但感受野始终是有上限的。2017年"Attention Is All You Need"这篇论文出来后,Transformer直接换了一条路:不依赖递归,也不依赖卷积,而是纯靠注意力机制去建模序列中任意两个位置之间的关系。
为什么最后是Transformer?我的理解是,它同时解决了三个问题:并行训练、长距离依赖、统一架构。这三个优势放在一起,让它在几乎所有模态上都能碾过传统结构。而且Transformer的"上手门槛"其实不高,几十行PyTorch代码就能实现一个版本跑起来。这篇文章我不打算复述论文,而是从一个实际使用的角度,把我踩过的坑、调参的经验、以及在不同场景下的适配方案完整梳理一遍,希望你看完能直接拿来用。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 核心原理与实际拆解:自注意力、多头机制与位置编码
2.1 自注意力是怎么解决"一对多"依赖的
在RNN里,你要计算第t个词的表征,必须经历从第1个词到第t个词的链式传播。Transformer的做法完全不一样:它一次性把所有词放入矩阵,再通过三个不同的线性投影计算出Q、K、V。Q和K做点积得到注意力分数,softmax归一化后对V加权求和。这个过程没有任何递归依赖,任意两个位置的距离在计算层面都是一步。
举个例子,假设句子是"那个苹果我昨天没吃,今天想吃"。在Transformer里,第6个词"吃"可以直接通过注意力分数和"苹果"建立高权重关联,这个距离是概念上的"一步"。换成RNN,这个关联可能要跨越多个隐状态传播,很容易在传播过程中被稀释。实际项目中,我经常用注意力权重可视化来检查模型是否学到了合理的语义关联,如果发现权重分布近乎均匀,那多半是训练不充分或者学习率设置有问题。
自注意力的代价是计算复杂度O(n²),其中n是序列长度。序列一长,显存和耗时都会暴涨。这一点后面在工程优化部分我会详细讲怎么处理。
2.2 多头注意力:让模型拥有"多个视角"
单独做一个注意力机制,等价于每个token只看一个全局分布下的聚合结果。问题在于"相关性"这个概念本身是复杂的:两个词可能因为语法关系相关,也可能因为语义相似相关,还可能因为共现模式相关。如果把所有这些相关性混在一个注意力分布里,模型最后只能取一个折中,表达能力受限。
多头注意力做的事情,是把Q、K、V切成h份,每一份在不同的子空间里做注意力计算,最后把结果拼接起来再过一层线性投影。直观理解就是:同一句话,8个头分别从不同角度去理解"谁和谁有关系",最后汇总出综合判断。
实践中的一个经验是:头的数量不是越多越好。我试过在中等规模数据集上把head从8加到16,效果反而下滑,因为头数增多的同时,每个头的维度变小了(总维度固定时),模型反而损失了每个头内的表达能力。视觉类任务(ViT)里有个有趣的现象:不同层不同头的注意力模式差异很大,底层偏局部、高层偏全局,这说明多头机制确实在和任务语义做匹配。
2.3 位置编码:给无序集合补上"顺序感"
注意力机制本身对位置是不敏感的。打乱词序,"苹果吃我"和"我吃苹果"算出来的自注意力结果如果只按集合来看,完全一样。但语言和图像里顺序/空间信息显然是有意义的,所以必须把位置信息显式注入输入。
Transformer原始做法是用固定频率的正弦余弦函数生成位置编码,不需要学习。后来很多工作发现,可学习的位置编码在部分任务上效果更好,于是出现了可学习位置嵌入(learned positional embedding)。我在做BERT类模型时更倾向于可学习位置编码,因为它在训练数据量充足时能自适应学到更符合数据分布的位置关系;但如果数据量级不大,固定位置编码的泛化性更稳,不容易过拟合到训练集特定的位置组合上。
还有一个细节很容易忽略:位置编码加在输入Embedding之后,是逐元素相加而不是拼接。这个设计的考量是,如果拼接,embedding维度会翻倍,计算量增大;而直接相加,模型通过在早期层里把位置信号和语义信号"叠加编码",后续层有能力自己把二者分离或重组,参数效率更高。
3. 用PyTorch从零搭建Transformer:核心代码与关键参数
3.1 最小可用版本:实现Encoder
接下来我直接给出一个我在项目中常用的、最精简的Transformer Encoder实现。它是理解完整Transformer的最佳切入点,也适合作为baseline做各类改造。
python复制import torch
import torch.nn as nn
import math
class MultiHeadSelfAttention(nn.Module):
def __init__(self, d_model, n_heads, dropout=0.1):
super().__init__()
assert d_model % n_heads == 0, "d_model必须能被n_heads整除"
self.d_model = d_model
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.out_proj = nn.Linear(d_model, d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x, mask=None):
batch_size, seq_len, _ = x.size()
Q = self.w_q(x).view(batch_size, seq_len, self.n_heads, self.d_k)
K = self.w_k(x).view(batch_size, seq_len, self.n_heads, self.d_k)
V = self.w_v(x).view(batch_size, seq_len, self.n_heads, self.d_k)
Q = Q.transpose(1, 2)
K = K.transpose(1, 2)
V = V.transpose(1, 2)
attn_scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
if mask is not None:
attn_scores = attn_scores.masked_fill(mask == 0, float('-1e9'))
attn_weights = torch.softmax(attn_scores, dim=-1)
attn_weights = self.dropout(attn_weights)
attn_output = torch.matmul(attn_weights, V)
attn_output = attn_output.transpose(1, 2).contiguous()
attn_output = attn_output.view(batch_size, seq_len, self.d_model)
return self.out_proj(attn_output)
这里有个关键点:d_model必须能被n_heads整除。我经常在代码里加assert,因为一旦整除不了,运行时报错信息会很绕,提前断言可以省去不少调试时间。另外,attn_scores除以math.sqrt(self.d_k)这一步叫缩放,是论文里非常关键的设计。为什么要缩放?因为当d_k增大时,点积的方差会跟着增大,softmax的梯度会趋近于消失,除以缩放因子后分数范围被压回稳定区,训练才能稳。
3.2 搭建完整Encoder层
单有自注意力还不够,一个标准的Encoder层还包含前馈网络(FFN)、残差连接和层归一化(LayerNorm)。这部分我在工程里按"Pre-LN"结构实现,因为实践下来它比原论文的"Post-LN"更容易训练,尤其在深层的Transformer里。
python复制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):
return self.linear2(self.dropout(torch.relu(self.linear1(x))))
class TransformerEncoderLayer(nn.Module):
def __init__(self, d_model, n_heads, d_ff, dropout=0.1):
super().__init__()
self.self_attn = MultiHeadSelfAttention(d_model, n_heads, 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(self.norm1(x), mask))
x = x + self.dropout(self.ffn(self.norm2(x)))
return x
注意这里我把d_ff单独作为参数传进去,没有写死在代码里。一般经验是d_ff取d_model的4倍左右,但这个比例不是固定不变的。比如在Swin Transformer里,不同stage的FFN比例会有差异;在一些轻量级模型里,d_ff会适当缩小以控制参数量。
3.3 位置编码的实现细节
python复制class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=512):
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) # (1, max_len, d_model)
self.register_buffer('pe', pe)
def forward(self, x):
return x + self.pe[:, :x.size(1)]
register_buffer是个实用的细节,它会把位置编码注册成模型的buffer,这样调用.to(device)时它会跟着移动,而且不会出现在优化器的参数列表里。我见过不少新手直接把它写成一个普通tensor,结果训练时忘了搬设备直接报错。
3.4 模型参数怎么选:d_model、depth、dropout
参数选型是整个使用过程中最容易让人迷茫的部分。我给出一个在通用任务上能直接用的推荐配置,以文本分类为例:
d_model = 256:序列特征维度,太小表达能力不够,太大会拖慢训练n_heads = 8:配合d_model=256,每个head维度32num_layers = 4:小数据集4层够用,大数据集可以到6层或12层d_ff = 1024:4倍于d_modeldropout = 0.1:默认值,过拟合时可以提升到0.3max_len = 512:序列长度上限
这些参数需要配合数据规模调整。数据只有几万条,把depth堆到12层,很容易过拟合;反过来数据巨大但模型太小,又欠拟合。一个比较实用的判断标准是:先跑一个小规模子集,模型看训练集的loss能否降到接近0,如果降不下去,说明模型容量或结构有bug;如果瞬间降到0,再换进全量数据。
4. 训练Transformer的工程要点:优化器、学习率与损失设置的坑
4.1 优化器选择:为什么AdamW是标配
Transformer训练和标准CNN训练最大的不同,就是对优化器更"敏感"。SGD在Transformer上表现很一般,原因在于注意力矩阵的梯度分布变化剧烈,自适应学习率算子的好处能很好应对这种情况。但Adam有个问题:权重衰减实现方式不对,导致某些参数被错误地衰减。AdamW从解耦角度修正了这一点,实践效果几乎总是优于Adam。
我自己的经验是:除非你有明确的理由,否则直接用AdamW(torch.optim.AdamW),初始学习率设置在1e-4到5e-4之间。这个范围对绝大多数CV和NLP任务都适用。
4.2 带Warmup的学习率调度
Transformer原论文用了一个特殊的学习率调度:先用一段warmup把学习率从0逐步升到峰值,之后按倒数平方根规律衰减。很多人第一次跑Transformer,直接固定学习率,结果loss震荡下不去。本质上是因为Transformer的优化曲面初期很"陡峭",一上来就用大学习率容易一步跨到坏区域。
我常用的实践是:warmup_steps设为总训练steps的10%左右,峰值学习率设为3e-4,之后用余弦退火或线性衰减。PyTorch的get_cosine_schedule_with_warmup可以直接用,省去手写调度器的麻烦。
注意:warmup阶段不等于让模型学习到全面特征,而是让LayerNorm里的均值和方差统计先稳定下来。过早进入高学习率阶段,注意力分布还可能集中在初始随机状态,结果后面的训练容易一直在"修正早期错误"。
4.3 损失函数:标签平滑不能忘
分类任务上,很多人用交叉熵就用默认的one-hot标签。但在Transformer这种"容量大、拟合快"的模型上,one-hot标签会让模型过度自信,输出概率分布过于尖锐,泛化性变差。标签平滑(label smoothing)把one-hot变成soft target,比如分类数C=10时,把目标概率设为[0.9, 0.01, 0.01, ...]而不是[1, 0, 0, ...]。
标签平滑的合理性可以这样理解:它不让模型为训练样本的每个细节"死磕",相当于告诉模型"正确答案附近的小概率分布也允许存在",降低过拟合风险。我在自己的项目中,smoothing系数通常取0.1,上下浮动0.05都试过,不建议超过0.2,否则模型会变得过度保守,loss明明降不下去,但准确率也上不来。
4.4 梯度裁剪与混合精度
Transformer的梯度范数波动很大,经常出现某个层的梯度过大,其他层梯度正常的情况。训练时加上clip_grad_norm_(model.parameters(), max_norm=1.0)几乎成了标配,可以显著减少训练崩溃的概率。
另一个工程优化是混合精度训练(AMP)。PyTorch里一行with torch.autocast(device_type='cuda', dtype=torch.float16):就能启用。对Transformer来说,混合精度不仅能省显存,还能加速训练。但要注意:如果用的不是较新版本的PyTorch,AMP在LayerNorm和Softmax这两个操作上有可能出现数值不稳定。实测中发现,把这两类操作强制切回FP32能避免NaN问题。
5. 优化与调参:让Transformer在目标任务上真正"生效"
5.1 序列长度压缩与截断策略
Transformer的O(n²)复杂度让长序列成为天然瓶颈。一个我在做长文本分类时的方案:不是把整篇文本一股脑塞进去,而是用滑窗切段,每个窗口独立编码,最后做池化或再加一层融合。这样既控制了计算量,又能捕捉局部语义。实测在BERT类任务里,用512长度的滑窗加全局池化,在长文档分类上能接近直接跑4096长度效果,而显存占用低了一个数量级。
对于机器翻译或生成类任务,如果序列确实过长,另一个办法是使用分层的Transformer:先用低层Transformer对局部窗口建模,再用高层Transformer对窗口级token建模。这个概念对应到CV里,就是Swin Transformer的窗口注意力与窗口移位思路,本质都是"先局部后全局"。
5.2 Padding Mask 与 Attention Mask 的正确用法
训练时一个batch内序列长短不一,需要padding到统一长度。但padding位置没有实际语义,把它也纳入注意力计算会让模型学到"padding位置和真实token相关"的错误概念。所以必须传入mask,让注意力分数在padding位置变成负无穷,softmax之后权重为0。
我在代码中遇到的一个常见错误是:mask的维度搞错。注意力分数的shape是(batch, n_heads, seq_len, seq_len),所以mask也需要是(batch, 1, seq_len, seq_len)或能广播到这个shape。很多新手直接把(batch, seq_len)的mask传进去,触发广播错误或者根本没起屏蔽作用。写mask时建议多一步:
python复制# key_padding_mask: (batch, seq_len),True表示该位置是pad
mask = key_padding_mask.unsqueeze(1).unsqueeze(2) # (batch, 1, 1, seq_len)
这个mask会被广播到所有attention头上,正确屏蔽掉pad位置。另外,如果是Decoder里的自注意力,还需要加上因果mask(causal mask),保证第i步只能看到第i步之前的信息,避免未来信息泄漏。
5.3 从NLP扩展到CV:Vision Transformer的适配要点
Vision Transformer(ViT)的核心改动是把图像切成patch,每个patch拉平后经过线性投影得到token,再加上位置编码。虽然结构上和NLP Transformer几乎一样,但使用时有几个特殊的地方:
- patch size是超参,常用的是16x16或32x32。patch越小,token数量越多,计算量越大,但能保留更多空间细节。我实际测试中,小数据集上patch size 8的效果比16明显好,但也更容易过拟合。
- ViT需要比CNN更多的训练数据。没有大规模预训练时,ViT在中小数据集上往往打不过ResNet。这是结构特性:注意力机制的自由度更高,需要更多样本约束。
- 位置编码在ViT里通常采用可学习式,因为图像patch的排列比较固定,学习式编码更灵活。Swin Transformer则是通过窗口机制,让位置编码变得局部化,这也是它在目标检测和语义分割上比ViT实用的重要原因。
5.4 时间序列预测:TCN+Transformer的组合实战
把Transformer用于时间序列预测,是最近特别火的方向。热词里有"TCN时间卷积网络+transformer实战股票预测",这个组合思路很直接:TCN负责提取局部时序模式(比如短期波动趋势),Transformer负责捕捉长程依赖(比如一段周期性的整体走势)。两者串联,往往比单独使用效果更稳。
我在做类似任务时的实践是:
- 先用TCN对原始序列做下采样或特征提取,把序列长度缩短,同时保留关键局部信息。
- 把TCN输出的特征序列输入一个层数较少的Transformer Encoder(2-3层足够)。
- 最后一层输出的所有token做全局池化,再接回归层。
注意时间序列预测和NLP有一个本质不同:标签不是离散类别,而是连续值。所以损失函数用MSE或Huber Loss。而且时间序列数据很容易出现不同时间段的分布偏移,训练时建议做窗口化采样,而不是简单随机打乱。我还试过把股票数据做对数收益变换后输入模型,相比直接用原始价格序列,训练稳定性和预测精度都更高。
6. 常见问题与排查技巧实录
6.1 训练Loss出现NaN怎么办
这是Transformer训练最头疼的问题。我的排查顺序是这样的:
先检查输入数据里有没有NaN或Inf。很多时间序列数据里有空值或极端值,直接在Batch里带入NaN,模型再怎么调也救不回来。数据层面干净之后再看结构层面。结构上最典型的NaN来源是注意力分数在FP16溢出。如果用了混合精度,试试把注意力计算切到FP32,或者给注意力分数加一个缩放因子。另外,标签平滑和梯度裁剪对NaN也有一定缓解,因为它们都能限制模型的"极端值输出"。
如果以上都没解决,检查学习率。Transformer对学习率比较敏感,过大的学习率很容易让loss在前几步就冲出合理区间,然后梯度爆炸产生NaN。把warmup加上、把峰值学习率降一个数量级,大概率能解决。
6.2 训练不收敛或loss一直高位震荡
多数情况下是数据问题或学习率分配不当。我先建议做"过拟合单样本测试":取一个batch,反复在这个batch上训练几十步,看看loss能不能降到很低。如果连单batch都降不到接近0,说明模型有结构性问题(比如mask写错、维度不对、层与层之间信息传递断裂)。如果单batch能过拟合,在全部数据上loss高位震荡,那就要检查学习率调度、batch size、数据是否分类均衡。
还有一点容易被忽略:LayerNorm的epsilon参数。在FP16训练时,epsilon默认1e-5可能偏小,导致归一化分母出现数值不稳定;改成1e-6或1e-7可以缓解。这个参数很小,但确实影响训练稳定性。
6.3 显存不足怎么办
显存不足是Transformer使用的高频问题。模型太大、序列太长、batch size太大,总有一个是罪魁祸首。推荐的优化方案按优先级排列:
- 把
batch size降低,这一步最直接。 - 使用梯度累积:把多个小batch的梯度累积后统一更新,等价于大batch效果,显存占用却小得多。
- 开启混合精度,约能节省40%-50%显存。
- 检查是否用了gradient checkpointing,它用计算换显存,对深层Transformer效果显著。PyTorch里可以用
torch.utils.checkpoint.checkpoint_sequential,或者直接用HuggingFace里封装好的gradient_checkpointing_enable()。
我个人的配置经验是:一张24GB显存的卡,8层Transformer、d_model=768的模型,batch size开到32、序列长度512,在启用AMP和gradient checkpointing后能顺利训练。
6.4 推理速度慢的优化思路
如果模型训练完成但部署推理太慢,有几个成熟路径:
- 把注意力替换成线性注意力或FlashAttention,把复杂度从O(n²)降到O(n)或大幅降低常数因子。
- 使用ONNX Runtime或TensorRT做模型导出与优化。
- 对模型做知识蒸馏,用小模型逼近大模型效果。
- 在部署时做动态长度裁剪,避免为最长序列预留所有token计算。
实践中,FlashAttention的收益最直观,尤其在长序列场景下,推理加速可以达到数倍且几乎无损。前提是硬件较好且框架支持。
7. 变形与应用扩展:从Swin Transformer到多模态统一架构
7.1 Swin Transformer:窗口注意力与移动窗口设计
Swin Transformer是ViT之后最值得关注的视觉Transformer架构之一。它在标准Transformer基础上引入了两个核心改动:窗口注意力(window attention)和移动窗口(shifted window)。
窗口注意力把特征图划分成不重叠的小窗口,每个窗口内独立计算注意力。这能大幅降低计算量——假设特征图尺寸是H×W,窗口尺寸是M×M,全局注意力的复杂度是O(H²W²),窗口注意力则变成O(HW×M²)。移动窗口的关键在于,相邻层之间窗口位置错开,让不同窗口的信息有机会跨窗口交互。
我实际使用Swin Transformer做图像分类和目标检测时,最直接的感受是它比ViT更适合中小规模数据。原因是窗口注意力天然带有局部先验,相当于"结构化的归纳偏置",模型不需要从头学出局部性这个常识,这在数据量不足时是巨大优势。Swin Transformer还专门做了相对位置编码,比ViT的绝对位置编码在平移等变性上表现更好,这也是它在检测任务中表现突出的原因之一。
7.2 从单模态到多模态:Transformer的统一潜力
Transformer还有一个重要属性:模态无关。它不关心输入是词、图像patch、音频帧还是时间序列窗口,只要能表示成token序列就能统一处理。这带来了多模态模型的天然土壤。
在多模态场景里,标准做法是不同模态分别用不同的encoder提取特征,然后作为token序列拼接送入共同Transformer层。图像一侧常用ViT或ResNet提取patch特征,文本一侧用tokenizer加embedding,两条分支的序列在同一个注意力空间里互相交互。这个设计让模型能够在跨模态场景(图文匹配、视觉问答、图文生成)中学习到模态间的对齐关系。
跨模态的场景我踩过不少坑:一个是不同模态的特征尺度差异巨大,直接拼接后注意力可能只关注高范数的模态。解决办法是给各模态的特征分别做LayerNorm,或者引入模态专用embedding(模态type embedding)。另一个是模态间不平衡:文本信息稠密、图像信息相对稀疏,训练时经常出现视觉部分学不动的情况。这时候可以分别调节各模态encoder的学习率,视觉部分用更大的学习率来追赶。
7.3 模型剪枝与轻量化:把Transformer搬到边上
Transformer虽然效果好,但参数量大一直是部署痛点。现在做轻量化Transformer有两个方向:一是结构重设计,比如Swin和MobileViT;二是训练后压缩,比如剪枝、量化和蒸馏。
我在实际部署中优先试量化:把FP32权重转为INT8,模型体积减少75%,推理速度提升2-4倍,精度损失一般控制在1-2个点以内。做量化时注意,注意力计算中的softmax对低精度非常敏感,最好保持softmax部分为FP32运算。如果再配合模型蒸馏——用大模型蒸馏小模型,部署时的精度损失还能进一步收窄。
剪枝在Transformer里的策略也和CNN不太一样:CNN剪的通常是通道,而Transformer里更适合做注意力头剪枝。实践中发现不少注意力头是冗余的,去掉它们对最终结果影响很小,甚至因为减少噪声而略微提升性能。有大厂工作经验的朋友提到过,LLM推理时注意力头稀疏化能带来可观的加速,这和我的观察一致。
8. 一个完整的实践项目:用Transformer做文本分类的端到端流程
为了让前面讲的内容落地,我整理一个完整的小项目流程:用中文新闻标题做情感分类(正面/负面/中性),使用Transformer Encoder作为特征抽取器。整个流程包含数据准备、模型定义、训练配置、评估和日志记录。
8.1 数据处理与Tokenizer
中文文本处理比英文多一步分词。我这次用简单的中文分词工具,把每个词映射成token id,加上[CLS]和[SEP]标记。注意:中文里的词表大小通常比英文大很多,控制在3万-5万比较合理,太小了OOV问题严重,太大了模型embedding层涨参数量。
python复制import jieba
def encode_text(text, vocab, max_len=128):
tokens = list(jieba.cut(text))[:max_len - 2]
ids = [vocab.get('[CLS]')] + [vocab.get(w, vocab.get('[UNK]')) for w in tokens] + [vocab.get('[SEP]')]
ids = ids + [vocab.get('[PAD]')] * (max_len - len(ids))
mask = [1 if i < len(ids) else 0 for i in range(max_len)] # 注意这里要重新计算
return ids, mask
这里有个小坑:计算mask时,如果先pad再算,会把pad位置也算成1。正确做法是在pad之前先记录真实长度,再生成mask。上面代码只是一个示意,实际实现时建议先处理出原ids长度,再统一padding和mask。
8.2 模型定义:Encoder + 分类头
python复制class TransformerClassifier(nn.Module):
def __init__(self, vocab_size, d_model=256, n_heads=8, num_layers=4,
d_ff=1024, num_classes=3, max_len=128, dropout=0.1):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model, padding_idx=0)
self.pos_encoding = PositionalEncoding(d_model, max_len)
self.encoder_layers = nn.ModuleList([
TransformerEncoderLayer(d_model, n_heads, d_ff, dropout)
for _ in range(num_layers)
])
self.norm = nn.LayerNorm(d_model)
self.classifier = nn.Linear(d_model, num_classes)
self.dropout = nn.Dropout(dropout)
def forward(self, input_ids, mask=None):
x = self.embedding(input_ids)
x = self.pos_encoding(x)
for layer in self.encoder_layers:
x = layer(x, mask)
x = self.norm(x)
cls_repr = x[:, 0, :] # 取[CLS]位置的输出
return self.classifier(self.dropout(cls_repr))
标准的NLP分类方案是取[CLS]位置的输出作为整句的语义表示。这个设计受BERT影响很深:预训练时[CLS]位置被显式训练为聚合全句信息,微调时直接接分类头即可。如果不想加[CLS],也可以对所有token输出做全局池化,效果通常差不多,但我个人倾向于保留[CLS]方案,因为它在多任务场景(比如同时做分类和NER)扩展性更好。
8.3 训练循环:完整的工程化写法
python复制def train_epoch(model, dataloader, optimizer, criterion, scaler, device, clip_grad=1.0):
model.train()
total_loss = 0.0
for batch in dataloader:
input_ids = batch['input_ids'].to(device)
attention_mask = batch['attention_mask'].to(device)
labels = batch['labels'].to(device)
optimizer.zero_grad()
with torch.autocast(device_type='cuda', dtype=torch.float16):
outputs = model(input_ids, attention_mask)
loss = criterion(outputs, labels)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), clip_grad)
scaler.step(optimizer)
scaler.update()
total_loss += loss.item() * input_ids.size(0)
return total_loss / len(dataloader.dataset)
scaler.unscale_(optimizer)这一步很关键。如果直接scaler.step(),梯度裁剪时使用的是缩放后的梯度,阈值就完全失去了意义。所以标准流程是:先unscale,再裁剪,再step。这个顺序写错的话,梯度裁剪形同虚设。
8.4 评估与模型保存
评估阶段的注意事项和训练阶段不同:需要手动指定torch.no_grad(),同时关闭dropout。model.eval()会切换LayerNorm和Dropout的状态。注意eval()不会禁用梯度记录,所以必须配合no_grad()。
python复制def evaluate(model, dataloader, device):
model.eval()
preds, labels_all = [], []
with torch.no_grad():
for batch in dataloader:
input_ids = batch['input_ids'].to(device)
attention_mask = batch['attention_mask'].to(device)
outputs = model(input_ids, attention_mask)
preds.extend(torch.argmax(outputs, dim=-1).cpu().tolist())
labels_all.extend(batch['labels'].tolist())
return accuracy_score(labels_all, preds)
保存模型时,我习惯只保存state_dict而不保存整个模型,这样在版本切换时更灵活。另外强烈建议同时把模型配置(d_model、num_layers等)也保存为一个json文件,否则过几天你自己都记不清当时用了什么结构。这个教训我吃过不止一次。
9. 常见问题速查表与避坑指南
| 问题现象 | 可能原因 | 解决方案 |
|---|---|---|
| 训练loss为NaN | 输入含NaN/Inf | 检查数据,处理缺失值和异常值 |
| 训练loss为NaN | FP16下注意力分数溢出 | 注意力计算切到FP32,或降学习率 |
| 训练loss为NaN | 学习率过大导致梯度爆炸 | 加warmup,降峰值学习率,开梯度裁剪 |
| loss不下降 | mask维度/逻辑错误 | 检查mask是否传对,单batch过拟合测试 |
| loss不下降 | 学习率过低或调度器错误 | 调大学习率至1e-4级别,检查warmup |
| 过拟合严重 | 模型容量过大 | 增加dropout、降层数、加数据增强/正则 |
| 显存不足 | batch size过大 | 调小batch,开启AMP或gradient checkpointing |
| 推理太慢 | 序列长、模型大 | FlashAttention、量化、蒸馏、ONNX导出 |
| 多模态效果差 | 模态特征尺度不平衡 | 各模态独立LayerNorm,模态type embedding |
| ViT在中小数据上效果差 | 数据量不足 | 用Swin/CNN预训练特征,或降低patch尺寸 |
这个表是我在多个项目里总结出来的高频问题集合。如果你碰到的场景不在此列,最推荐的调试策略还是那一步:单batch过拟合测试。它能用几秒钟的时间把"模型结构问题"和"训练策略问题"分隔开,是Transformer调试里性价比最高的一步。
10. 一些个人经验与后续想法
做了这几年Transformer相关的工作,我最深的体会是:Transformer本身不难用,难的是理解每个模块在特定任务里是怎么协作的。很多人一上来就调大模型、堆数据,忽略了mask写没写对、学习率调度合不合理这些基础问题。我反而建议,如果你是新接触Transformer,先用一个小数据集、一个浅层小模型,完整跑通整个链路,再逐步做规模扩展。这样排查问题的时候,任何一层出故障,你都清楚是哪里的原因。
另一个很实用的习惯是:每次改模型结构或训练策略,都坚持做对比实验并记录日志。我把每次实验的d_model、层数、dropout、学习率、最终指标统一记录在一个表格里,时间久了,哪些配置在哪些任务上效果好,就形成了自己的经验库,很多时候不用重新搜索,直接查表就能选定初始配置。
Transformer到目前为止还在快速演进,从标准架构到FlashAttention、Linear Attention、Sparse Attention,再到基于Transformer的多模态大模型,本质上都是在"效果好"和"算得快"之间做平衡。如果你已经熟练掌握了标准Transformer的用法,接下来值得尝试的方向有三个:一是理解FlashAttention的底层IO优化,这对长序列项目帮助巨大;二是学习Swin Transformer的局部注意力设计思路,对CV类任务非常实用;三是多模态token融合方法,这是当前最活跃也最有前景的领域。把这些点吃透,你就能在Transformer这条路上走得比别人更远。
