Transformer架构深度解析:从Self-Attention原理到PyTorch代码实战
1. 为什么Transformer值得你花时间啃透如果你最近在学深度学习大概率已经被Transformer这四个字反复轰炸过了。不管是做NLP、CV、语音还是多模态Transformer几乎成了绕不开的基础设施。但很多人第一次翻开那篇论文的时候看到Encoder-Decoder、Multi-Head Attention、Positional Encoding这些词脑子里基本是一团浆糊。我当初也是这么过来的第一次看论文原文光Self-Attention的公式就盯了半小时没缓过神。这篇内容就是把我自己从“看不懂”到“能给别人讲明白”这个过程里踩过的坑、总结的方法完整地梳理一遍。核心目标很明确让你真正理解Transformer架构到底长什么样、每个模块为什么这么设计、代码层面怎么对应、以及实际用的时候有哪些容易忽略的细节。不管你是刚入门的新手还是已经用过BERT、GPT但没深究过底层原理的开发者都能从这里拿到对自己有用的东西。需要提前说明的是这篇内容不会只停留在“画个图讲概念”的层面。我会把架构拆到每个张量维度、每个矩阵乘法的层面同时给出PyTorch的可运行代码片段。你看完之后应该能做到拿一张白纸自己把Transformer的结构画出来并且说清楚每一层在干什么、为什么这么干。2. Transformer整体架构拆解与设计逻辑2.1 从RNN的痛点说起为什么需要Transformer要理解一个东西为什么被设计出来最好的方式是先看它要解决什么问题。在Transformer出现之前处理序列数据的主流方案是RNN和LSTM。RNN的核心思路是“按顺序读”每读一个词就把前面的信息压缩到一个隐藏状态里然后传给下一步。这个思路很直观但有两个致命问题。第一个问题是串行计算。你没法并行处理一个句子里的所有词因为第t步的计算必须等第t-1步完成。这意味着在GPU上RNN的训练效率极低序列越长越慢。第二个问题是长距离依赖。当序列很长的时候前面信息传到后面已经被稀释得差不多了梯度消失让模型很难学到远距离的关联。虽然LSTM用门控机制缓解了一部分但本质上没有根治。Transformer的解法非常干脆彻底放弃循环结构用注意力机制直接建模任意两个位置之间的关系。这样一来所有位置的计算可以同时进行并行度拉满同时任意两个词之间的距离都是1步不存在信息衰减的问题。这个设计思路的转变是Transformer最核心的贡献。2.2 编码器-解码器结构总览Transformer的整体架构分为两大部分编码器Encoder和解码器Decoder。如果你拿的是原始论文《Attention Is All You Need》里的图左边是编码器右边是解码器。编码器由N个相同的层堆叠而成论文里N6。每一层包含两个子层一个是多头自注意力机制一个是前馈神经网络。每个子层后面都跟着残差连接和层归一化。解码器同样由N6层堆叠但每层有三个子层带掩码的多头自注意力、交叉注意力Query来自解码器Key和Value来自编码器输出、以及前馈网络。这里有一个很多人一开始会困惑的点编码器和解码器到底分别负责什么用一个翻译任务来类比编码器负责“读懂”源语言句子把它压缩成一组富含语义信息的向量表示解码器负责“生成”目标语言句子生成的时候会参考编码器的输出同时只能看到已经生成的部分。这个分工在后续的BERT只用编码器和GPT只用解码器里被进一步分化但根源都在这里。2.3 各模块职责与数据流走向把数据流完整走一遍你就能把整个架构串起来。假设输入是一个句子“I love NLP”处理流程是这样的第一步词嵌入。每个词被映射成一个d_model维的向量论文里d_model512。这一步就是把离散的词变成连续向量。第二步位置编码。因为Transformer没有循环结构本身不知道词的顺序所以需要额外注入位置信息。论文用的是正弦余弦函数直接加到词嵌入上。第三步进入编码器层。数据依次经过多头自注意力、残差归一化、前馈网络、残差归一化。这个流程重复6次。第四步编码器的最终输出是一组向量每个向量对应输入序列的一个位置。这组向量会被送到解码器的交叉注意力层。第五步解码器层。解码器的输入是已经生成的目标序列训练时用teacher forcing先经过带掩码的自注意力再经过交叉注意力最后经过前馈网络。同样重复6次。第六步输出层。解码器最后一层的输出经过一个线性层映射到词表大小再经过Softmax得到每个位置的概率分布。整个流程里残差连接保证梯度能顺畅回传层归一化稳定训练过程掩码机制确保解码器不会偷看未来的词。这些设计每一个都有明确的工程动机不是拍脑袋加的。3. 核心机制深度剖析注意力到底在算什么3.1 Self-Attention的计算过程与维度变化Self-Attention是Transformer的灵魂也是最多人卡住的地方。我用最直白的方式讲一遍。假设输入序列长度为n每个词的向量维度是d_model。我们通过三个线性变换矩阵W_Q、W_K、W_V把输入分别映射成Query、Key、Value三个矩阵。维度分别是d_k、d_k、d_v。论文里d_kd_v64多头情况下h8所以8×64512d_model。计算过程用公式表示就是Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) V拆开看每一步QK^T计算每个Query和每个Key的点积得到一个n×n的矩阵。这个矩阵的第(i,j)个元素表示第i个位置对第j个位置的“关注程度”。除以sqrt(d_k)这是缩放操作。为什么要除因为当d_k很大的时候点积结果会很大Softmax之后会变得非常尖锐梯度会变得很小训练不稳定。除以sqrt(d_k)可以把方差拉回到1附近。Softmax按行做归一化让每一行的注意力权重加起来等于1。乘以V用注意力权重对Value加权求和得到每个位置的输出。维度变化是这样的输入(n, d_model) → Q(n, d_k), K(n, d_k), V(n, d_v) → QK^T(n, n) → softmax(n, n) → 输出(n, d_v)。多头的情况下把d_model拆成h份每份独立做Attention最后拼接再过一个线性层。3.2 多头注意力为什么要“多”很多人会问一个注意力不够吗为什么要搞多头这个问题我当初也想了很久。核心原因在于不同的头可以关注不同的关系模式。比如在翻译“The animal didnt cross the street because it was too tired”这句话时一个头可能关注“it”和“animal”的指代关系另一个头可能关注“cross”和“street”的动宾关系。如果只有一个头Softmax之后注意力会被平均化多种关系模式会互相干扰。从实现角度看多头并不是简单地复制多份计算。而是把d_model维度的向量切成h份每份维度是d_model/h分别做Attention最后把h个输出拼接起来。这样总计算量和单头全维度差不多但表达能力更强。论文里h8每个头64维。这里有个实操细节多头注意力的实现通常用一个大矩阵一次性算完而不是写循环。因为把h个头的W_Q拼在一起就是(d_model, d_model)的矩阵一次矩阵乘法就能得到所有头的Q。这也是Transformer高效的原因之一。3.3 位置编码没有循环怎么知道顺序Transformer最大的特点是没有循环也没有卷积这带来了并行化的优势但也带来了一个问题它本身对词的顺序完全不敏感。“我打你”和“你打我”在纯注意力机制下得到的表示是一样的这显然不行。位置编码的解法是给每个位置生成一个固定维度的向量直接加到词嵌入上。论文用的是正弦余弦函数PE(pos, 2i) sin(pos / 10000^(2i/d_model)) PE(pos, 2i1) cos(pos / 10000^(2i/d_model))这个设计有几个巧妙之处。第一每个位置都有唯一的编码。第二任意两个位置之间的编码差异可以通过线性变换表示这让模型容易学到相对位置关系。第三它可以外推到比训练时更长的序列。后来也有很多变体比如可学习的位置编码BERT用的就是这种、相对位置编码T5、Transformer-XL用的。但原始的正弦编码是最经典的理解它有助于理解后续的改进。3.4 残差连接与层归一化的工程意义残差连接和层归一化在Transformer里看起来不起眼但缺了它们训练根本跑不起来。残差连接解决的是深层网络的梯度消失问题。公式是LayerNorm(x Sublayer(x))。注意这里x是子层的输入Sublayer(x)是子层的输出。残差连接让梯度可以直接跳过子层回传即使子层学得很差至少还有恒等映射保底。这也是为什么Transformer能堆到几十层甚至上百层的原因。层归一化和BatchNorm不同它是在每个样本内部做归一化不依赖batch大小。这在序列任务里很重要因为不同序列长度不同BatchNorm统计量不稳定。LayerNorm对每个位置的d_model维向量做归一化均值为0方差为1然后再用可学习的参数缩放平移。有一个细节值得注意原始论文用的是Post-LN也就是先做子层再做归一化。但后来很多工作发现Pre-LN先归一化再进子层训练更稳定不需要warmup。现在主流实现大多用Pre-LN。这个细节在复现论文的时候很容易踩坑。4. 从零手写Transformer核心代码4.1 环境准备与依赖说明动手写代码之前先把环境理清楚。我用的是PyTorch版本建议1.10以上因为后面用到的nn.MultiheadAttention和nn.Transformer在旧版本里接口有差异。不过为了讲清楚原理我会从最基础的矩阵运算开始写不直接调高级API。需要的依赖很简单pip install torch numpy如果你要用GPU加速确保CUDA版本和PyTorch匹配。我实测下来即使是CPU版本跑一个小规模的Transformer做demo也完全够用不用一上来就折腾环境。4.2 多头注意力模块的完整实现先写最核心的多头注意力。我会把每一步的维度变化都注释清楚。import torch import torch.nn as nn import math class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() assert d_model % num_heads 0, d_model必须能被num_heads整除 self.d_model d_model self.num_heads num_heads self.d_k d_model // num_heads # 三个线性变换把输入映射到Q、K、V 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) def forward(self, query, key, value, maskNone): batch_size query.size(0) # 1. 线性变换并拆分成多头 # (batch, seq_len, d_model) - (batch, seq_len, num_heads, d_k) Q self.W_q(query).view(batch_size, -1, self.num_heads, self.d_k) K self.W_k(key).view(batch_size, -1, self.num_heads, self.d_k) V self.W_v(value).view(batch_size, -1, self.num_heads, self.d_k) # 2. 转置成 (batch, num_heads, seq_len, d_k) Q Q.transpose(1, 2) K K.transpose(1, 2) V V.transpose(1, 2) # 3. 计算注意力分数 # (batch, num_heads, seq_len, d_k) x (batch, num_heads, d_k, seq_len) # - (batch, num_heads, seq_len, seq_len) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) # 4. 应用掩码如果有 if mask is not None: scores scores.masked_fill(mask 0, -1e9) # 5. Softmax归一化 attn_weights torch.softmax(scores, dim-1) # 6. 加权求和 # (batch, num_heads, seq_len, seq_len) x (batch, num_heads, seq_len, d_k) # - (batch, num_heads, seq_len, d_k) context torch.matmul(attn_weights, V) # 7. 拼接多头 # (batch, num_heads, seq_len, d_k) - (batch, seq_len, num_heads, d_k) context context.transpose(1, 2).contiguous().view( batch_size, -1, self.d_model ) # 8. 输出线性变换 output self.W_o(context) return output这段代码有几个容易出错的地方。第一view之前必须保证张量是连续的所以用了.contiguous()。第二掩码的填充值用-1e9而不是-inf因为-inf在某些情况下会导致NaN。第三缩放因子是sqrt(d_k)不是sqrt(d_model)这个细节很多人搞混。4.3 前馈网络与编码器层的组装前馈网络看起来简单但它是Transformer里参数量最大的部分。结构就是两个线性层加一个ReLUclass PositionwiseFeedForward(nn.Module): def __init__(self, d_model, d_ff, dropout0.1): super().__init__() self.linear1 nn.Linear(d_model, d_ff) self.linear2 nn.Linear(d_ff, d_model) self.dropout nn.Dropout(dropout) self.relu nn.ReLU() def forward(self, x): # (batch, seq_len, d_model) - (batch, seq_len, d_ff) - (batch, seq_len, d_model) return self.linear2(self.dropout(self.relu(self.linear1(x))))论文里d_ff2048是d_model的4倍。这个比例在后续模型里基本被沿用比如BERT-base也是4倍。为什么是4倍经验上这个比例在表达能力和计算量之间取得了比较好的平衡太小了模型容量不够太大了计算开销吃不消。接下来组装编码器层class EncoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, num_heads) self.feed_forward PositionwiseFeedForward(d_model, d_ff, dropout) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) def forward(self, x, maskNone): # 子层1多头自注意力 残差 归一化 attn_output self.self_attn(x, x, x, mask) x self.norm1(x self.dropout1(attn_output)) # 子层2前馈网络 残差 归一化 ff_output self.feed_forward(x) x self.norm2(x self.dropout2(ff_output)) return x这里用的是Post-LN结构和原始论文一致。如果你想用Pre-LN把norm放到子层前面就行。我实测下来Pre-LN在小数据集上收敛更快但最终效果差异不大。4.4 位置编码的实现与可视化验证位置编码的实现有很多种写法我用的是最直观的版本class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len5000, dropout0.1): super().__init__() self.dropout nn.Dropout(dropout) # 创建一个足够长的位置编码矩阵 pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.float).unsqueeze(1) # 计算分母项 div_term torch.exp( torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model) ) # 偶数维度用sin奇数维度用cos pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) # 增加batch维度 pe pe.unsqueeze(0) # (1, max_len, d_model) self.register_buffer(pe, pe) def forward(self, x): # x: (batch, seq_len, d_model) x x self.pe[:, :x.size(1), :] return self.dropout(x)写完可以做个简单的验证把pe矩阵画出来你会看到不同维度上有不同频率的波形低频维度变化慢高频维度变化快。这种多频率的设计让模型能同时捕捉粗粒度和细粒度的位置信息。4.5 完整模型串联与维度检查把上面所有模块串起来就是一个完整的Transformer编码器class TransformerEncoder(nn.Module): def __init__(self, vocab_size, d_model512, num_heads8, num_layers6, d_ff2048, max_len5000, dropout0.1): super().__init__() self.embedding nn.Embedding(vocab_size, d_model) self.pos_encoding PositionalEncoding(d_model, max_len, dropout) self.layers nn.ModuleList([ EncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) self.d_model d_model def forward(self, x, maskNone): # 词嵌入并缩放 x self.embedding(x) * math.sqrt(self.d_model) x self.pos_encoding(x) for layer in self.layers: x layer(x, mask) return x注意词嵌入之后乘了sqrt(d_model)这是论文里的做法目的是让词嵌入的尺度和位置编码匹配。如果不乘位置编码的值会相对过大影响训练。写完之后一定要做维度检查。我习惯在forward里加print确认每一步的shape符合预期。比如输入(batch2, seq_len10)经过嵌入变成(2, 10, 512)经过位置编码还是(2, 10, 512)经过6层编码器输出还是(2, 10, 512)。如果哪一步维度对不上大概率是view或者transpose用错了。5. 训练调优与常见问题排查5.1 学习率调度与warmup策略Transformer的训练对学习率非常敏感。原始论文用的是warmup加逆平方根衰减lr d_model^(-0.5) * min(step^(-0.5), step * warmup_steps^(-1.5))这个公式的意思是前warmup_steps步线性增加学习率之后按步数的平方根倒数衰减。warmup_steps通常设4000。为什么要warmup因为训练初期模型参数是随机的梯度方向不稳定大学习率容易导致训练发散。先用小学习率让模型“热身”等梯度稳定了再加大。我踩过的一个坑是用Adam优化器但不加warmuploss经常在前几百步就炸了。后来加上warmup训练曲线明显平滑很多。如果你用的是Pre-LN结构warmup可以短一些甚至不用但Post-LN基本是必须的。5.2 掩码机制的正确使用方式掩码在Transformer里有两个用途padding mask和causal mask。Padding mask用于处理不同长度的序列。一个batch里句子长度不同短的用0填充。但注意力计算时不能让模型关注到这些填充位置所以要把对应位置的注意力分数设成负无穷。实现上就是构造一个(batch, 1, 1, seq_len)的布尔矩阵填充位置为False。Causal mask用于解码器的自注意力确保位置i只能看到位置0到i不能看到后面的。实现上是一个下三角矩阵。这个在训练时特别重要如果忘了加模型会“偷看”答案训练loss很低但推理时效果极差。def create_causal_mask(seq_len): # 下三角矩阵对角线及以下为1 mask torch.tril(torch.ones(seq_len, seq_len)).bool() return mask注意PyTorch的masked_fill里mask为True的位置是保留的False的位置被填充。这个逻辑和有些框架相反用之前一定要确认清楚。5.3 常见报错与排查速查表报错信息可能原因解决方法RuntimeError: mat1 and mat2 shapes cannot be multiplied线性层输入维度不匹配检查d_model和输入最后一维是否一致RuntimeError: The size of tensor a must match tensor b残差连接时维度不一致确认子层输出维度和输入维度相同CUDA out of memorybatch太大或模型太大减小batch size或用梯度累积loss变成NaN学习率太大或没有warmup加warmup降低学习率检查是否有除零训练loss正常但推理效果差忘了加causal mask检查解码器自注意力是否加了掩码模型完全不收敛位置编码没加或加错确认位置编码在词嵌入之后加上这个表是我自己调试时积累的基本上覆盖了80%的常见问题。遇到报错先查维度维度对了再查掩码掩码对了再查学习率这个顺序能帮你快速定位问题。5.4 实操心得从demo到可用的关键细节最后分享几个我在实际项目里总结的经验。第一dropout的位置很重要。原始论文在三个地方加了dropout注意力权重、子层输出、位置编码。我试过只在子层输出加效果会差一些。注意力权重的dropout能防止模型过度依赖某几个位置。第二权重初始化不能忽视。Transformer默认用Xavier初始化但有些实现会用更小的初始化方差。如果训练初期loss震荡厉害可以试试把初始化方差调小。第三batch size和序列长度的权衡。Transformer的计算复杂度是O(n^2 * d)序列长度翻倍计算量翻四倍。如果显存不够优先减小序列长度而不是batch size因为batch size太小训练不稳定。第四验证集loss和训练集loss差距大不一定是过拟合。Transformer在小数据集上很容易过拟合但更大的可能是数据预处理有问题比如padding太多导致有效信息被稀释。我遇到过一次后来发现是padding token的embedding没有固定为0改完之后效果提升明显。第五不要迷信论文的超参数。论文里的d_model512、num_layers6是针对特定任务的。实际用的时候要根据数据量和任务复杂度调整。数据量小的时候层数减到2-3层、d_model降到128-256效果反而更好训练也快得多。6. 从理解架构到灵活运用把Transformer的架构吃透之后你会发现后续很多模型都是在这个基础上做加减法。BERT把编码器拿出来做预训练GPT把解码器拿出来做生成T5把两者结合做统一框架。Vision Transformer把图像切成patch当成序列处理Swin Transformer加了窗口注意力降低计算量。这些变体看起来五花八门但底层逻辑都是你上面看到的那些东西。我自己在学的时候最大的体会是不要一上来就追求看懂所有细节先把数据流走通再逐个模块深入。第一遍看论文知道有个编码器有个解码器有个注意力机制就够了。第二遍看搞清楚QKV是怎么算的。第三遍看理解为什么要多头、为什么要位置编码。每看一遍都会有新的收获。代码层面也是一样先跑通一个最小的demo哪怕只是复制粘贴看到loss在下降你就有信心继续往下挖。然后再逐行改代码把某个模块替换成自己的实现对比效果。这种“先跑通再优化”的路径比一开始就追求完美实现要高效得多。如果你现在还在纠结某个公式看不懂我的建议是先放一放去跑一遍代码打印中间张量的shape和数值很多困惑会在看到实际数据的那一刻自然消解。Transformer没有那么神秘它本质上就是一堆矩阵乘法和归一化的组合只是组合的方式比较巧妙而已。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →