我先把丑话说在前头:这篇不打算给你那种“概念几十条、代码跑不通”的科普文。今天讲的Transformer,我会从它到底解决了什么问题、注意力为什么这么设计、到用PyTorch手写一个能跑的最小实现、再到训练和推理阶段的避坑经验,一次性讲到底。我看过太多人把《Attention Is All You Need》背得滚瓜烂熟,但一上手写代码就崩,一训就loss发散,一到推理就被速度卡死。所以这篇不是翻译论文,而是把论文、代码、训练、部署之间那层窗户纸捅破。
如果你是第一次接触Transformer,看完这篇你会知道每一步为什么这么设计;如果你已经写过Transformer,我也建议你重点看第5节和第8节,那些训练和排查的坑,是我实打实踩过的,不是教科书里抄来的。
1. 先搞清楚Transformer到底解决了什么问题
1.1 老架构的痛点:RNN为什么慢
在2017年之前,序列建模的主流是RNN(循环神经网络),以及它的改进版LSTM、GRU。它们的核心思路很朴素:把序列一个token一个token地吃进去,每吃一个就更新一次隐藏状态,把“到目前为止看到的信息”压缩在一个向量h_t里。
h_t = f(h_{t-1}, x_t)
问题就出在这个f上。第t个时间步的输出必须等第t-1步算完才能开始,这是天然的串行依赖。你在GPU上跑一个长度为512的句子,就算显卡再强,也得老老实实按顺序跑512步。那时候训练一个翻译模型,动不动要几周,其中很大一部分时间浪费在了这种“排队式”的计算上。
更麻烦的是长距离依赖。信息要从前面的词传到后面,中间要经过很多步,每经过一步就经过一次非线性变换,梯度要么指数级消失,要么指数级爆炸。LSTM用门控机制缓解了梯度消失,但本质上还是在一个窄窄的“信息管道”里传递,距离一远,前面的信息早就被冲淡了。我印象很深,早期做机器翻译时,长句子的翻译质量会断崖式下跌,就是因为模型记不住句子的开头。
1.2 Self-Attention凭什么赢
《Attention Is All You Need》给出的答案是:不要一步步传递信息了,让每个位置直接看到全部位置。这句话听起来简单,背后是三个深层优势:
第一是并行。Self-Attention不是逐时间步计算的。它一次性接收整个序列,通过矩阵乘法直接计算所有位置两两之间的关系,整个计算过程在GPU上可以高度并行。同样跑一个512长度的句子,RNN要串行512步,Transformer只需要几次大矩阵乘法。
第二是全局感受野。RNN里第i个位置要看到第1个位置,路径长度是i-1;而在Self-Attention里,任意两个位置之间都只隔一步(一次注意力计算),也就是说信息传递的路径长度恒等于1。长距离依赖问题在结构上被彻底绕开了。
第三是计算效率的取舍。Self-Attention的时间复杂度是O(n²·d),n是序列长度,d是维度;RNN是O(n·d²)。当序列长度中等(比如几百)时,Transformer的复杂度是完全可接受的,换来的是远超RNN的并行能力和信息容量。两相对比,RNN几乎没有胜算。这也是为什么从2018年开始,BERT、GPT、T5这些基于Transformer的模型统治了NLP,后来连视觉、语音也都把它抢走了。
2. 注意力机制全景拆解:从QKV到缩放点积
2.1 Q、K、V到底是什么
注意力机制在Transformer之前就已经存在,最早用在Seq2Seq翻译模型里,让解码器在生成每个词时“回看”编码器的所有隐藏状态,自己决定更关注哪些位置。Transformer把它升级成了Self-Attention:序列自己跟自己算注意力。
理解Q、K、V最舒服的类比是查资料:
- Query(查询):你脑子里的问题,比如“这段里谁在做什么”。
- Key(键):资料库里的索引标签,比如文章的标题、标签。
- Value(值):索引对应的正文内容。
查资料时,你拿着Query和所有Key做相似度匹配,匹配度高的,对应的Value就多看几眼。注意力机制就是把这件事数学化:先算Query和所有Key的点积,得到相似度分数,过softmax变成权重,再对所有Value做加权平均。
在Self-Attention里,Q、K、V都来自同一个输入X,但各自乘了不同的权重矩阵W_Q、W_K、W_V做了线性变换。为什么要变换而不直接用X本身?因为直接拿原始向量互相比对,表达能力太弱,打架也严重;经过不同的投影矩阵后,模型可以让Query去“表达关注什么”,让Key去“表达我有什么”,各司其职。
2.2 缩放点积注意力的公式拆解
论文里的核心公式长这样:
Attention(Q, K, V) = softmax(QK^T / sqrt(d_k)) V
很多新手第一次看到这个公式直接懵了,其实拆开看只有四步:
- Q乘K的转置,得到两个token之间的匹配分数,得到一个形状为[n, n]的矩阵,第i行第j列表示第i个token对第j个token的注意力得分。
- 除以sqrt(d_k),也就是Key向量维度的平方根,做一次缩放。
- 对每一行做softmax,把所有位置的分数归一化成权重,和为1。
- 用这个权重矩阵去乘V,得到加权求和后的输出。
注意一个细节:对每一行做softmax,意思是在固定“当前要看第i个token”的前提下,决定“我把注意力分给序列里的其他token各多少”。所以注意力矩阵的第i行,就是token i对整个序列的注意力分布。
2.3 为什么要除以根号d_k
这是面试必问题。原因可以讲得很数学:假设Q和K里每个元素都是独立同分布的,均值0、方差1。那么两个d_k维向量的点积,结果的均值是0,方差却是d_k。也就是说,维度越大,点积的结果方差越大,数值分布越分散。
分散的数值一进softmax就麻烦了。因为softmax对大的输入值非常敏感,一旦某个分数特别大,softmax出来后就会非常接近1,其他位置接近0,梯度趋近于0——模型学不动了。除以sqrt(d_k)相当于把点积的方差重新拉回到1,让softmax的输入保持在一个“温和”的范围内,梯度可以正常回传。
我自己第一次实现Transformer时,忘了加这个缩放,训小模型还能勉强跑,稍微把d_model调大一点就直接不收敛了。所以这不是论文里写得好玩的,是真有实际意义的。
2.4 Mask到底在干什么
注意力机制里有两类Mask,初学者几乎必搞混。第一类是Padding Mask。一个batch里的句子长短不齐,短的句子要补padding到统一长度。但padding位不是真实内容,不能让模型注意到它们,否则就相当于让模型去关注一堆没意义的0。做法是在softmax之前,把padding位置的注意力分数设成一个非常大的负数(常见是-1e9而不是-float('inf')),这样softmax算出来的权重就无限接近0。
为什么要用-1e9而不是负无穷?因为-inf在参与softmax内部的指数运算时,可能在某些框架里出现NaN,-1e9足够“骗过”softmax,又不会引起数值问题。
第二类是Causal Mask,也叫因果掩码或下三角掩码。Decoder生成第i个token时,不能看到未来位置,否则训练时就“作弊”了。做法是把注意力矩阵的上三角全部置为-1e9,这样第i行只能attention到前i个位置。这个掩码和Padding Mask可以叠加使用:先把padding位置遮掉,再遮掉未来位置。
3. 完整Transformer架构逐层拆解
3.1 输入嵌入与位置编码
Transformer本身对顺序没有概念,如果把一句话拆成一堆token直接丢进注意力层,模型看到的只是一个无序集合。打个比方,你拿到一叠卡片,每张卡片上写了一个词,但卡片没有编号,你根本不知道哪个词先哪个词后。所以必须显式地把位置信息加进去。
论文用了一组正弦函数来编码位置:
PE(pos, 2i) = sin(pos / 10000^(2i/d_model)) PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
其中pos是位置下标,i是维度下标。直接把位置向量和词向量相加,作为Transformer的输入。
为什么要选正弦函数?两个原因。第一,正弦函数的值域在[-1, 1],不会像可学习位置编码那样训练时出现数值不稳定的风险;第二,借助正弦和余弦的线性组合性质,模型可以比较容易地学到“相对位置”关系——比如要表达“当前位置偏移k”时,可以表示成当前位置编码的线性组合。这意味着即使序列比训练时更长,正弦位置编码也有一定外推能力。当然,后来的BERT、GPT们大多用了可学习的位置编码,效果也不错,因为外推问题可以通过训练更长序列来解决,但理解论文里的原始设计,对理解位置编码的本质很有帮助。
3.2 多头注意力:一组注意力还是多组注意力
多头注意力(Multi-Head Attention)是Transformer抓住我眼球的设计之一。一个注意力头只能学到一种“相关性”,比如它可能只关注词序关系,或只关注语义相似性,但真实语言里的关系是复杂的:语法依存、指代消解、语义关联、局部共现……这些模式不是一种注意力能覆盖的。
多头机制的做法是:把Q、K、V分别线性投影到h个低维子空间,每个子空间各算一次注意力,得到h个输出结果,把这h个结果拼接起来,再经过一次线性变换得到最终输出。论文用的h=8,每个头的维度d_k = d_model / h = 64,这样总计算量与单头注意力基本持平,却让模型有了h个“视角”。
打个比方:单头注意力像你请了一个专家看问题,专家水平再高也有盲区;多头注意力像你请了一个委员会,每个专家各自关注不同维度,最后综合意见。委员会的整体判断往往比单一专家更稳定、更丰富。实际实现时,多头不会真的h次循环,而是通过reshape操作一次矩阵乘法同时完成所有头的计算,这也是第4节代码里要重点演示的地方。
3.3 前馈网络、残差与LayerNorm
注意力层之后,每个位置还要过一个前馈网络(FFN)。FFN是两层全连接,中间有ReLU或GELU激活函数:
FFN(x) = max(0, xW_1 + b_1)W_2 + b_2
中间层的维度一般是d_model的4倍。很多初学者不理解:注意力已经那么强了,为什么还要FFN?我的理解是,注意力机制本质是“信息路由”,它决定哪些信息需要被传递、加权、组合,但它本身是一种线性加权(虽然QKV投影里有非线性,但注意力分布的计算是线性的),表达能力有限。FFN是对每个位置独立做一次非线性特征变换和存储,模型的知识很大程度是存在FFN里的。有研究发现FFN的第一层类似记忆检索,第二层类似特征整合,所以别小看这两层全连接。
为了让网络能堆到很深,Transformer每个子层(包括注意力和FFN)外面都包了残差连接和LayerNorm:
output = LayerNorm(x + Sublayer(x))
这里的残差连接解决了深层网络的梯度传播问题;LayerNorm对每个token的向量做归一化,让训练更稳定。
这里有个值得知道的细节:论文原始结构是Post-LN,也就是先加残差再LayerNorm。但这种结构在层数很深时容易不稳定。后来GPT、BERT等现代实现大多采用Pre-LN,也就是先LayerNorm再进入子层,最后加残差。Pre-LN训练更稳,可以用更大的学习率,虽然早期收敛稍慢,但从实践角度看更省心。
3.4 从全局看Encoder-Decoder结构
整个Transformer由Encoder和Decoder两大部分组成:
- Encoder:由多层相同的编码器层堆叠而成,每层包含一个多头自注意力、一个FFN,加上残差和LayerNorm。Encoder的每一个位置都能看到整个输入序列,所以叫双向编码,适合理解类任务。
- Decoder:结构稍复杂,每层包含两个注意力子层。第一个是Masked多头自注意力,用因果掩码保证解码时只能看已经生成的内容;第二个是Cross-Attention(交叉注意力),这个注意力子层的Query来自Decoder上一层的输出,而Key和Value来自Encoder的输出。也就是说,解码器在生成每个词时,都会主动去“查阅”编码器编码好的源文本信息,这就是机器翻译里“注意”机制的直接体现。
三种注意力的分工:Encoder的自注意力负责理解输入序列内部关系;Decoder的Masked自注意力负责建模已生成内容之间的关系;Cross-Attention负责连接输入和输出。把这三个搞清楚,Transformer的骨架就算拿下了。
4. 手写一个极简Transformer(PyTorch实战)
4.1 场景设定与超参数选择
理论讲完,代码必须跟上。我写一个能给“译前训练”用的极简Transformer,不追求SOTA,只求逻辑清晰、能跑通。场景是英译中,数据就用十几条日常用语,比如“hello -> 你好”、“how are you -> 你好吗”这种,主要用来验证模型能否正常收敛、注意力是否在正常工作。
超参数设计遵循论文的比例,但为了小数据集能快速运行,做了缩小:
- d_model=64(嵌入和模型维度)
- num_heads=4(d_model必须能被num_heads整除)
- num_encoder_layers=2、num_decoder_layers=2
- d_ff=128(FFN中间层)
- dropout=0.1
- batch_size=32,序列最大长度=16
4.2 核心代码实现:从多头注意力到Transformer块
先写多头注意力,这是最核心的部分。我推荐新手在初始阶段就这么写,逻辑最直白,不要一上来就优化成全矩阵合并版本:
import torch import torch.nn as nn import torch.nn.functional as F import math class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads, dropout=0.1): super().__init__() assert d_model % num_heads == 0 self.d_model = d_model self.num_heads = num_heads self.head_dim = d_model // num_heads self.wq = nn.Linear(d_model, d_model) self.wk = nn.Linear(d_model, d_model) self.wv = nn.Linear(d_model, d_model) self.out_proj = nn.Linear(d_model, d_model) self.dropout = nn.Dropout(dropout) def forward(self, query, key, value, mask=None): batch_size = query.size(0) # 1. 线性投影到 Q、K、V Q = self.wq(query) # [batch, seq_len, d_model] K = self.wk(key) V = self.wv(value) # 2. 拆成多个头: [batch, seq_len, num_heads, head_dim] # 转置后变成 [batch, num_heads, seq_len, head_dim] Q = Q.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) K = K.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) V = V.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) # 3. 缩放点积注意力 scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.head_dim) if mask is not None: scores = scores.masked_fill(mask == 0, -1e9) attn_weights = F.softmax(scores, dim=-1) attn_weights = self.dropout(attn_weights) context = torch.matmul(attn_weights, V) # [batch, heads, seq_len, head_dim] # 4. 合并多头 context = context.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model) output = self.out_proj(context) return output这段代码里有几个关键点要解释。view和transpose是拆多头的核心操作:view把最后一个维度切成num_heads份,transpose把head维度换到第2维,这样后续矩阵乘法就能把所有头一次性算完,不用写for循环遍历每个头。masked_fill时用-1e9而不用-float('inf'),原因前面已经说过。softmax是在最后一维(每个位置对其他位置的注意力分布)上做的,方向千万别搞错。
然后是位置编码和Transformer子层。位置编码实现时注意要算好“偶数是sin、奇数是cos”的规律:
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) # [1, max_len, d_model] self.register_buffer('pe', pe) def forward(self, x): x = x + self.pe[:, :x.size(1)] return self.dropout(x)Transformer Block就是把多头注意力、FFN、残差和LayerNorm组装起来。这里我采用Pre-LN的结构(先归一化再进子层),训练稳定性更好:
class EncoderBlock(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout=0.1): super().__init__() self.self_attn = MultiHeadAttention(d_model, num_heads, dropout) self.ffn = nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model) ) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.dropout = nn.Dropout(dropout) def forward(self, x, src_mask=None): x = x + self.dropout(self.self_attn(self.norm1(x), self.norm1(x), self.norm1(x), src_mask)) x = x + self.dropout(self.ffn(self.norm2(x))) return xDecoder的区别在于多了Cross-Attention,并且第一层自注意力要加因果掩码。完整的Transformer就是把Encoder和Decoder串起来,这里不再贴全部代码,核心的注意力实现你已经看到了,其余部分按照“先Masked自注意力,再交叉注意力,再FFN”的顺序搭积木就行。
4.3 训练循环要点
训练Transformer常见的一个坑是学习率策略。论文用的是Noam Scheduler:先做warmup让学习率从非常小线性升到峰值,之后按步数的倒数平方根衰减。这一步很重要,我见过太多次因为直接用固定学习率导致loss震荡甚至发散的案例。
class NoamScheduler: def __init__(self, optimizer, d_model, warmup_steps=4000): self.optimizer = optimizer self.d_model = d_model self.warmup_steps = warmup_steps self.step_num = 0 def step(self): self.step_num += 1 lr = self.d_model ** (-0.5) * min(self.step_num ** (-0.5), self.step_num * (self.warmup_steps ** (-1.5))) for param_group in self.optimizer.param_groups: param_group['lr'] = lr训练循环就是一个标准的外层epoch、内层batch,loss用交叉熵,同时把padding位置设成ignore_index=-1,避免模型去预测padding位。小数据集上通常几十个epoch就能看到loss明显下降。
4.4 全程张量形状追踪
我把每个阶段的张量形状列成一张表,初学时照着这张表检查自己的代码,能少走很多弯路:
| 阶段 | 形状变化 |
|---|---|
| 输入token序列 | [batch, src_len] |
| Embedding后 | [batch, src_len, 64] |
| 加位置编码后 | [batch, src_len, 64] |
| Q/K/V投影后 | [batch, src_len, 64] |
| view+transpose拆多头 | [batch, 4, src_len, 16] |
| QK^T后 | [batch, 4, src_len, src_len] |
| 乘V后 | [batch, 4, src_len, 16] |
| 合并多头后 | [batch, src_len, 64] |
| FFN输出 | [batch, src_len, 64] |
| 最终输出的logits | [batch, tgt_len, vocab_size] |
注意head_dim=16,是因为d_model=64、heads=4,二者必须整除。一旦除不尽,view就会报错,这也是新手最常见的报错之一。
5. 训练Transformer的5个保命技巧
5.1 学习率策略:不要忽略warmup
Transformer对学习率极其敏感。熟悉RNN的人刚切过来时,习惯性地用固定学习率,结果loss在前几百步就飞了,还以为是自己代码写错了,其实只是缺少warmup。
warmup的原理是:训练初始阶段,模型参数是随机初始化的,梯度噪声非常大,如果直接上大学习率,极易把参数推到不收敛的“坏区域”。warmup充当一个“热身期”,让模型先用很小的学习率慢慢探索,等梯度走向稳定了再加速。论文里warmup_steps=4000,但在小数据集上可以适当减小到1000或500,不然光热身就要耗掉很多步数。
5.2 Dropout与标签平滑
Dropout在Transformer里至少有三个位置需要加:输入嵌入后、每个子层的输出后、注意力权重的计算后(就是4.2代码里softmax之后那一次dropout)。最后这个注意力Dropout很容易被忽略,但它的作用很大,可以强制模型不要过分依赖某一条注意力路径,提升泛化能力。
标签平滑也是一个朴素的技巧:把one-hot的硬标签换成软标签,比如0.9给正确答案,剩下的0.1平均分给其他词。直观理解是不要逼模型“自信过头”,留一点容错空间。机器翻译任务里标签平滑对BLEU提升很明显,尤其是训练数据量不大的时候。
5.3 梯度裁剪与混合精度
梯度裁剪(gradient clipping)在Transformer训练里几乎属于必需品。虽然Self-Attention不像RNN那样容易梯度爆炸,但深层堆叠+长序列仍然可能让梯度范数变得很大。我习惯把max_norm设为1.0,加了之后训练稳定性肉眼可见地上升。
混合精度训练(AMP,Automatic Mixed Precision)则是省显存和加速训练的利器。原理很直白:用FP16存储和运算,关键环节(如梯度累加、Loss Scaling)保持FP32,既能把显存占用几乎砍半,又能利用GPU的Tensor Core加速。PyTorch里用torch.cuda.amp就够,不需要额外安装任何东西。
5.4 Batch Size和序列长度的经验值
小batch size(比如8、16)在Transformer上训练极不稳定,因为目标函数对每个batch的波动太敏感。我实测下来,batch_size至少32起步会比较舒服;如果显存紧张,优先减小序列长度而不是batch size,或者用梯度累积(accumulate)模拟大步长。
序列长度也不是越长越好。注意力复杂度是序列长度的平方,长度翻倍计算量翻4倍。训练初期可以用较短序列先让模型学会基本语法,再逐步增加长度,这种“课程学习”思路在Transformer训练里非常实用。
6. 推理阶段的性能优化:KV Cache与Flash Attention
6.1 KV Cache:为什么生成时能省数倍计算
很多人训练完Transformer就开始用,结果发现生成一句话慢得离谱。原因在于推理阶段是自回归的,也就是每生成一个token,都要把这句话重新在模型里算一遍。不带缓存的话,生成长度为n的序列,总计算量是O(n³)级别的,效率极低。
KV Cache的核心思想是:每一步生成时,之前的token已经计算过各自的Key和Value,而这些Key和Value其实不会随着后续token的生成而改变,完全可以缓存下来。每步生成新token时,只需要计算当前token的Q、K、V,其中Q只和当前token有关,而K、V则把缓存里之前所有token的K、V拼接起来用。这样每步的计算量从O(n²)降到了O(n),生成速度提升非常大。注意,Q是不能缓存的,因为每个新token的Query都不一样,它要去“检索”之前所有Token。
6.2 Flash Attention:把计算吃进SRAM
KV Cache解决的是重复计算问题,而Flash Attention解决的是显存带宽瓶颈。标准注意力实现把QK^T这个中间矩阵(形状是[n, n])写到显存(HBM),再读出来做softmax,最后乘V。序列一长,这个中间矩阵极其占用带宽,计算单元反而在“等数据”。
Flash Attention的思路是分块(tiling):把Q、K、V切成小块,让每一步计算和softmax都在GPU的SRAM(高速缓存)里完成,不把中间矩阵写回HBM。它对输出结果没有精度损失,却能把速度提升好几倍,还省显存。实际使用中,我用Flash Attention跑长序列推理,最高有4-5倍的速度提升,显存占用更是从O(n²)降到了O(n)。这已经成了当前大模型框架里的标配优化。
6.3 注意力机制的现代变体
Push O(n²)的极限,研究界提出不少注意力变体:Sparse Attention用固定稀疏模式只计算部分位置两两之间的注意力;Linear Attention把softmax中的指数核替换成可分解的核函数,把复杂度压到O(n);滑动窗口注意力(如Swin里的实现)只关注相邻窗口内的token。这些变体各有适用场景,比如超长文本、高分辨率图像,但核心的“Query匹配Key,加权Value”逻辑没有变。理解了基础版,看这些变体基本就是一层窗户纸的事。
7. Transformer家族谱:从BERT/GPT到ViT/Swin
7.1 Encoder-only与Decoder-only的路线分歧
Transformer论文给的是一个Encoder和Decoder都有的整体架构,但后来者并没有死守所有组件,而是根据自己的任务各取所需,形成了三大流派:
- Encoder-only(如BERT):只保留编码器,用双向注意力,句子里的每个词都能看到前后所有词。适合需要“理解整句话”的任务,比如文本分类、命名实体识别、语义匹配。
- Decoder-only(如GPT系列):只保留解码器,用因果注意力,每个词只能看到它之前的词。适合生成任务,比如文本续写、对话、代码生成。现在的大语言模型几乎都是这条路线。
- Encoder-Decoder(如T5、BART):保留完整结构,适合翻译、摘要这类“输入一段,输出一段”的任务。
这三者用的是同一套Self-Attention骨架,区别主要在Mask的形态:双向注意力没有Mask,因果注意力需要上三角Mask。理解了这个,你就理解了BERT和GPT的本质区别。
7.2 视觉Transformer的困境与解法
视觉Transformer(ViT)把图像切成固定大小的patch(比如16x16像素一块),每块展平成向量,像token一样送入Transformer。这个方法在足够大的数据集上效果超过了CNN,但它早期有个致命弱点:图像是二维的,直接展平成一维序列,序列长度随像素数呈平方增长,一张224x224的图切成16x16的patch有196个token,还能接受;到了高清图,序列长度就爆炸了。
而且,图像中的“局部结构”是极为重要的,相邻像素之间的关系远比隔得很远的像素更紧密。ViT初始版本没有这个先验,需要海量数据才能学会局部模式。
Swin Transformer给出的解法很有启发:把注意力限制在窗口内(比如7x7的patch),并且模仿CNN的金字塔结构,逐层合并相邻patch,让网络先看局部、再看全局,同时窗口设计使计算复杂度从平方级降到了线性级。这告诉我们一个道理:Attention的全局性是好东西,但硬套到超高分辨率输入时,需要通过结构设计来约束它。
8. 常见问题与排查技巧实录
8.1 训练Loss不降或直接NaN
这是最多新手遇到的问题,我按概率排序给出排查清单:
- 检查缩放因子sqrt(d_k)有没有加。忘加缩放,大维度下softmax会饱和,loss必然不稳。
- 检查Mask有没有用-1e9,以及有没有在softmax之前加。mask写错位置或值设成-inf,都会带来NaN。
- 检查d_model能不能被num_heads整除。除不尽时view直接报错,这还算好的;更隐蔽的是某些实现里整除后head_dim=0,模型完全没学习能力。
- 检查学习率是不是一开始就太大。没有warmup或者warmup步数太小,loss发散的概率很高。
- 检查标签平滑和dropout是否在推理时被错误保留(eval模式没切)。这条我曾经排查了一下午。
8.2 显存OOM
OOM时我建议按这个顺序尝试:先减小batch_size,再减小序列长度,这两个改动最简单直接。如果还不满足,可以使用梯度累积模拟大batch_size。再不够,用混合精度训练(AMP),显存占用几乎减半。最后,如果你训练的是深层Transformer,可以开启gradient checkpointing,用计算换显存,虽然在长序列上训练时间会变长,但至少能让模型跑起来。
8.3 训练速度慢得离谱
如果是你自己写的注意力实现,先检查是不是用for循环遍历head了,或者对每个序列片断串行算注意力。把Q、K、V reshape成多头后一次性矩阵相乘,是最基本的加速手段。再看看是否在GPU上运行了CUDA(很多人本地用CPU跑还抱怨慢)。推理阶段的话,先确认有没有开KV Cache,没有缓存的情况下,长文本生成速度慢是理所当然的。
8.4 Mask用错的典型症状
Cross-Attention里忘了用Encoder的padding mask,会让Decoder注意到Encoder padding的位置,症状是生成的译文在结束token后莫名多出几个重复的“嗯”之类的内容。训练时因果mask写反了(mask了下三角而不是上三角),大多loss不降还伴随乱码输出。这些小坑都是看loss曲线看不出名堂、必须盯输出才能发现的。
8.5 给新手的路线图建议
如果让我给一个第一次接触Transformer的人安排一周时间,我会建议这样分配:第一天精读论文的注意力公式和三张结构图,第二天照着这篇博文的代码亲手敲一遍,第三到第五天换到真实任务上(用HuggingFace跑一个翻译或分类模型),第六天把KV Cache和Flash Attention的代码读一遍,第七天尝试消融实验——删掉缩放、删掉多头、删掉位置编码,看看效果都怎么变。最后这步非常重要,只有亲手验证了每个组件的价值,你才算真正理解了Attention Is All You Need这句话,而不只是背住了论文标题。