拓冰建站拓冰建站
首页 / 资讯中心 / 正文

Transformer模型结构详解:从自注意力到PyTorch实现

做文本建模或者大模型方向的朋友早晚会绕不开一个东西——Transformer。我第一次对着那篇原始论文里的模型结构图看的时候光是把 Query、Key、Value 三个矩阵在脑子里对齐就花了大半个下午更别说后面多头拆分、位置编码、残差和归一化到底放在哪一层这种细节。后来自己从零手写了一遍跑通了小规模机器翻译任务才算真正把 Transformer 模型结构吃透。这篇就是我整理的一份超详细解读从整体骨架到每一个子模块再到能直接抄的 PyTorch 实现和踩过的坑尽量说人话。不管你是刚入门想搞懂注意力机制的新手还是已经能调库但说不清结构细节的熟手应该都能从里面捞到点东西。1. 从整体架构看Transformer的骨架设计很多人一上来就钻进注意力公式里其实先把整体骨架啃清楚后面很多细节是顺理成章的。Transformer 的原始设计是一个编码器-解码器结构专门为序列到序列任务准备的。但今天大家嘴上说的 Transformer往往泛指这一整类基于自注意力的结构包括只保留编码器的 BERT 系、只保留解码器的 GPT 系。1.1 编码器与解码器的分工到底差在哪编码器的任务是读懂输入它由 N 个完全相同的层堆叠而成每一层里有两个子模块一个是多头自注意力一个是前馈网络。注意这里只有自注意力没有掩码因为编码器处理的是完整的输入句子每个位置都能看到句子里所有其他位置这叫双向可见。解码器的任务则是生成输出它的每一层里有三个子模块带掩码的多头自注意力、编码器-解码器注意力、前馈网络。第一个自注意力加了因果掩码保证第 t 个位置只能看到前 t 个位置不能偷看未来。第二个注意力比较特殊它的 Query 来自解码器当前状态Key 和 Value 来自编码器的输出这一步是把读懂的输入和正在生成的输出对齐起来。我个人的理解是编码器像是一个把整段话压缩成向量表示的过程解码器像是一个一边看压缩表示一边逐字往外吐的过程。搞清这个分工后面看到 mask 为什么只加在解码器第一层、为什么编码器-解码器注意力不需要因果掩码就不会迷糊了。1.2 为什么选择堆叠而不是一味加宽原始论文里编码器和解码器都堆了 6 层模型维度 d_model 是 512。这里有个很关键的设计哲学深度优先于宽度。堆叠层数带来的是抽象层级的提升底层关注局部词形和短距离依赖中层开始捕捉句法结构高层才能表征语义和长距离关系。而单纯加宽只是增加单层容量很难形成这种逐级抽象。当然深度也不是越多越好。层数一多梯度消失和训练不稳定的问题就来了这也是后面 Post-LN 被 Pre-LN 逐渐取代的重要原因之一。我在实际做小任务时4 层到 6 层通常够用做大一点的任务12 层是个常见起点。层数选择本质上是在表达能力和训练难度之间找平衡。1.3 结构总览与全文的维度约定为了后面推导不混乱这里先把符号约定死全文都按这套来Bbatch size一次喂多少条样本L序列长度也就是 token 数量d_model模型隐藏维度比如 512h注意力头数比如 8d_k每个头的维度满足 d_k d_model / hd_ff前馈网络中间层维度通常是 4 * d_model举个例子输入张量形状是 (B, L, d_model) (2, 10, 512)h8那么 d_k 64。整套结构从头到尾就是在这几个维度之间来回搬运、拆分、合并。把这张维度地图记在脑子里看任何 Transformer 变体都会快很多。2. 核心组件逐个拆解从多头注意力到前馈网络骨架理清了接下来一个组件一个组件地拆。这一部分我会尽量把每一步的形状变化写出来因为形状对不上是新手写代码时最高频的报错来源。2.1 自注意力机制的计算全过程与维度推导自注意力的本质是一句话让序列里每个位置去综合其他所有位置的信息综合的权重由相关性决定。具体分三步。第一步把输入 X 分别乘上三个可学习矩阵 W_q、W_k、W_v得到 Q、K、V。这三个矩阵形状都是 (d_model, d_model)。所以 Q X W_q形状仍是 (B, L, d_model)。第二步算注意力分数。用 Q 乘 K 的转置得到形状 (B, L, L) 的分数矩阵再除以根号 d_k。为什么要除以根号 d_k因为当维度较大时Q 和 K 的点积数值会随维度增长而变大softmax 之后会变得非常尖锐几乎只有一个位置接近 1其余接近 0梯度会变得极小。除以根号 d_k 相当于把方差拉回到 1 附近这是数值稳定性的保障不是可有可无的装饰。第三步对分数做 softmax 得到注意力权重再乘 V输出形状 (B, L, d_k)。整个过程用公式写就是Attention(Q,K,V) softmax(QK^T / sqrt(d_k)) V。这里有个容易忽略的点softmax 是沿最后一维做的也就是每个 query 位置对全体 key 位置归一化方向反了结果就完全错了我在调试时被这个坑过不止一次。2.2 多头注意力的拆分与合并操作单头注意力只能捕捉一种相关性模式多头则是让模型同时从多个视角去看输入。实现上并不是真的跑了 h 次独立的注意力而是把 d_model 维度切成 h 份每份 d_k 维h 个头并行计算最后再拼回来。具体流程是这样的线性投影后得到 (B, L, d_model)用 view 重塑成 (B, L, h, d_k)再用 transpose 把 h 维换到前面变成 (B, h, L, d_k)。这样每个头就是独立的一层二维注意力。算完以后输出 (B, h, L, d_k)transpose 回 (B, L, h, d_k)用 contiguous 保证内存连续再 view 回 (B, L, d_model)最后乘输出矩阵 W_o。这里有个技术细节值得强调transpose 之后内存布局变了直接 view 会报错必须先 .contiguous()。这个报错信息往往很长很吓人其实原因就一句话——不连续的内存没法直接拉平。还有一点h 必须整除 d_model否则切不匀实现里通常加一句 assert 兜底。提示多头的意义不只是多个视角它还降低了单头的计算复杂度总开销因为每个头只在 d_k 维度上算总计算量和单头差不多。这也是它比堆 h 个完整注意力再拼更划算的原因。2.3 位置编码正弦编码与可学习编码的取舍自注意力有个先天缺陷它对输入顺序是无感的。你把句子里两个词调换位置只要它们携带的内容不变输出几乎不变。这对语言来说是灾难因为猫追狗和狗追猫完全不同。解决办法就是位置编码把位置信息注入进去。原始论文用的是正弦位置编码公式是偶数维用 sin(pos / 10000^(2i/d_model))奇数维用 cos(pos / 10000^(2i/d_model))。其中 pos 是位置i 是维度索引。这套设计的巧妙之处在于任意位置 posk 的编码可以表示成 pos 编码的线性变换这让模型有能力外推到更长的序列。另一种常见做法是可学习位置编码就是给每个位置准备一个可训练向量BERT 用的就是这个。它实现简单、效果稳定但最大长度在训练时就固定死了想扩展到更长序列得重新处理。我在短文本任务里通常用可学习编码处理变长或需要外推的场景更倾向正弦编码或旋转位置编码。选哪个没有绝对答案看任务。2.4 残差连接与LayerNorm的位置之争每个子模块外面都套了一层残差连接 层归一化写作 LayerNorm(x Sublayer(x))。残差连接的作用是给梯度开一条高速公路让深层网络也能训得动LayerNorm 则是稳定每层的激活分布加速收敛。但这里有个被反复讨论的细节归一化到底放在残差之前还是之后。原始论文是 Post-LN也就是先做子层再相加再归一化。后来研究发现 Post-LN 在层数多时训练很不稳定需要精细的学习率预热于是 Pre-LN 流行起来变成先归一化再做子层。Pre-LN 训练更稳、对学习率不那么敏感代价是最终效果在同等层数下可能略逊需要配合更深的堆叠来弥补。现在主流大模型几乎清一色 Pre-LN 或其变体如果你自己训练遇到深层不收敛先试试把归一化提到前面。3. 前馈网络、激活函数与归一化细节注意力负责在位置之间通信前馈网络负责在每个位置上独立加工。这两个部分交替出现构成了 Transformer 层的基本节奏。这一部分聊几个容易被忽视但影响很大的细节。3.1 前馈网络为什么是4倍扩张前馈网络结构很朴素两层线性变换中间夹一个激活函数写作 FFN(x) W_2 · activation(W_1 · x)。关键在于中间层的维度 d_ff 通常是 d_model 的 4 倍。原始论文里 d_model512d_ff2048正好是4倍。为什么是4倍而不是2倍或8倍我理解这是一个经验性的容量平衡点。注意力层已经负责了跨位置的信息聚合前馈层需要足够的宽度来对每个位置做非线性变换把特征投影到更高维再压回来形成一个瓶颈-扩张-瓶颈的结构类似自编码器的思路能增强表达能力。倍数太小非线性表达受限倍数太大参数量和计算量陡增收益递减。4 倍是大量实验下来的常用值T5 用过 8 倍也有用 2.67 倍的变体实际可以调。顺带说一句这部分参数量其实很可观。对 d_model512、d_ff2048 的单层来说FFN 参数量约 2 × 512 × 2048 ≈ 200 万而一层多头注意力四个投影矩阵加起来才约 4 × 512 × 512 ≈ 100 万。也就是说前馈网络占了单层参数的大头这一点很多人没意识到。3.2 GELU与ReLU的选用逻辑激活函数早期用 ReLU简单高效。后来 GELU 逐渐成为 Transformer 的标配BERT、GPT 系列都用它。GELU 的形式是对输入做高斯分布的累积概率加权可以粗略理解为平滑版的 ReLU在零点附近是光滑过渡的负值区域也不是直接截断为零而是保留一小部分。这个平滑性的好处是梯度更连续训练更稳定尤其深层网络里更明显。代价是计算比 ReLU 稍贵。如果你在做资源紧张的部署ReLU 或它的变体 Swish、SiLU 都是可以的替代实测差异在小任务上未必看得出来。我的经验是先用 GELU 跑通真到了要抠性能再换别一开始就在激活函数上纠结。3.3 归一化层的两种主流实现LayerNorm 和 BatchNorm 的区别经常被问。BatchNorm 是在 batch 维度上统计均值和方差依赖 batch 大小序列任务里 batch 内的长度还不一致用它很别扭。LayerNorm 则是针对每个样本、每个位置在自己这一条特征向量上做归一化和 batch 大小无关非常适合变长序列。LayerNorm 的计算是(x - mean) / sqrt(var eps) * gamma beta。其中 mean 和 var 是沿最后一维特征维算的gamma 和 beta 是可学习的缩放和平移参数初始化为 1 和 0。eps 是个很小的数防止除零一般取 1e-5 或 1e-6。再进阶一点现在不少模型把 LayerNorm 换成了 RMSNorm去掉了减均值和 beta 那一步只保留缩放计算更快效果基本持平。这是工程优化里的常见取舍属于知道就行、不必强求的细节。4. 从零手写一个Transformer光看懂结构不够自己写一遍才能发现所有藏起来的坑。这一部分给一份可以跑的 PyTorch 实现顺便把关键位置的形状标注清楚。我用的是自己教学时反复改过的版本去掉了花哨的东西突出主干。4.1 张量维度约定与整体代码骨架先定好超参d_model512n_heads8d_ff2048层数 6dropout 0.1。整体代码分成四块多头注意力、位置编码、编码器层、解码器层然后拼成完整模型。先看多头注意力这是最核心、也最容易写错的部分。import torch import torch.nn as nn import math class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_heads, dropout0.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.w_o nn.Linear(d_model, d_model) self.dropout nn.Dropout(dropout) def forward(self, q, k, v, maskNone): B q.size(0) # 投影 拆多头: (B, L, d_model) - (B, h, L, d_k) Q self.w_q(q).view(B, -1, self.n_heads, self.d_k).transpose(1, 2) K self.w_k(k).view(B, -1, self.n_heads, self.d_k).transpose(1, 2) V self.w_v(v).view(B, -1, self.n_heads, self.d_k).transpose(1, 2) # 缩放点积注意力: (B, h, L, L) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores scores.masked_fill(mask 0, -1e9) attn torch.softmax(scores, dim-1) attn self.dropout(attn) # 加权求和后合并多头: (B, L, d_model) out torch.matmul(attn, V) out out.transpose(1, 2).contiguous().view(B, -1, self.d_model) return self.w_o(out)这段里有三个点必须盯住。第一view 之前要能整除assert 是保险。第二transpose 换轴后必须 contiguous 再 view否则报错。第三masked_fill 用的值要足够小-1e9 是常用做法别用 0因为 softmax 后 0 还会分到权重。4.2 位置编码与前馈网络的实现位置编码可以写成一个预计算的矩阵训练时按序列长度取前 L 行加到词嵌入上。注意是加不是拼接。class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len5000): super().__init__() pe torch.zeros(max_len, d_model) position torch.arange(0, max_len).unsqueeze(1).float() 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) self.register_buffer(pe, pe.unsqueeze(0)) def forward(self, x): return x self.pe[:, :x.size(1)]这里 div_term 用的就是 10000^(2i/d_model) 的倒数形式写成 exp 能避免数值溢出。前馈网络更简单就是两个线性层加激活class FeedForward(nn.Module): def __init__(self, d_model, d_ff, dropout0.1): super().__init__() self.net nn.Sequential( nn.Linear(d_model, d_ff), nn.GELU(), nn.Dropout(dropout), nn.Linear(d_ff, d_model) ) def forward(self, x): return self.net(x)然后是编码器层把注意力和前馈串起来每个外面套 Pre-LN 残差class EncoderLayer(nn.Module): def __init__(self, d_model, n_heads, d_ff, dropout0.1): super().__init__() self.attn MultiHeadAttention(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, maskNone): x x self.dropout(self.attn(self.norm1(x), self.norm1(x), self.norm1(x), mask)) x x self.dropout(self.ffn(self.norm2(x))) return x注意这里三次传入的都是 norm1(x)这是自注意力的标志——Q、K、V 同源。解码器层会多一个交叉注意力其中 Q 来自解码器自身K 和 V 来自编码器输出这个区分是写对解码器的关键。4.3 掩码生成与训练调试要点因果掩码是个下三角矩阵形状 (L, L)下三角含对角线为 1其余为 0。表示第 t 个位置只能看到不超过 t 的位置。生成方式很简单def causal_mask(size): mask torch.tril(torch.ones(size, size)).bool() return mask.unsqueeze(0).unsqueeze(0) # (1, 1, L, L)同时别忘了处理 padding 掩码把补齐位置的注意力屏蔽掉否则模型会去关注无意义的填充符。训练时的几个调试心得学习率要配预热前几千步线性升到峰值再衰减这是原始论文的做法没有预热深层很容易发散dropout 在小数据集上别省能显著缓解过拟合如果 loss 一直不降先检查 mask 方向有没有反再看位置编码有没有真的加进去这两处是最隐蔽的错。5. 常见问题与排查技巧实录这一部分是实打实用时间换来的经验。很多问题看起来是训练技巧根子其实在结构实现上。我整理成速查表方便你对症下药。5.1 形状不匹配类问题最常见的一类报错就是形状对不上。我按照出现频率列一下对应的原因和修法都写清楚。报错现象可能原因修复办法view 处报 size 不匹配transpose 后内存不连续加 contiguous() 再 viewreshape 维度错误d_model 不能被 n_heads 整除调整头数或维度加 assert矩阵乘维度冲突注意力里 K 忘了转置对 K 用 transpose(-2, -1)广播失败mask 形状与 scores 不一致mask 补成 (B,1,L,L) 或 (1,1,L,L)拼接后维度翻倍concat 时算错头维度确认拼回的是 h*d_k d_model这类问题九成是维度换算没记清。我的习惯是在每个模块入口和出口都打印一次 shape跑一遍小数据确认整条链路都对再上大规模训练能省下大量瞎猜的时间。5.2 训练不收敛与效果异常排查形状对了但训练不收敛或者效果诡异也有套路可循。第一个要怀疑的是掩码方向。如果因果掩码打反了模型看到的就是未来信息训练 loss 可能很低但推理时完全崩表现是训练集无敌、生成一塌糊涂。第二个是学习率没有预热、峰值过大深层模型会在前几百步就崩掉。还有就是数值精度问题。用半精度训练时注意力里的分数容易溢出建议保留缩放那一步并且对 softmax 输入做一下 clamp。位置编码加错了也不会立刻报错只是模型学不到顺序表现是对词序不敏感你可以在一个小样本上故意打乱词序看输出是否变化变化不大就说明位置信息没生效。注意如果输出退化成所有位置生成同一个 token先检查 softmax 维度是不是用错再确认 logits 有没有被 mask 全屏蔽导致全是 -1e9。全屏蔽时 softmax 会输出均匀分布表现为随机或重复很迷惑人。5.3 分模块定位问题的思路面对一个不工作的 Transformer别一上来就调超参先做结构自检。我的顺序是先用一条长度为 2 的假数据跑前向确认不报错再把模型过拟合一个极小的批次比如 8 条样本如果连这点数据都过拟合不了说明结构或损失肯定有 bug最后才加正则、调学习率。这个先能过拟合、再谈泛化的思路非常实用。因为过拟合小数据是模型具备基本表达能力的必要条件做不到就说明前向或反向有问题跟超参无关。我见过太多人一开始就猛调学习率结果 bug 在结构里调多久都没用。6. 结构变体与现代工程取舍原版 Transformer 是 2017 年的设计这些年结构上做了不少演进。理解变体不是为了追新而是为了知道每个设计选择背后的权衡面试和实际选型都用得上。6.1 主流变体改了哪些结构点按改动部位来梳理会更清楚。归一化方面Post-LN 基本被 Pre-LN 取代部分模型换成 RMSNorm。位置编码方面可学习编码和正弦编码之外旋转位置编码RoPE在长文本场景流行起来它通过旋转矩阵把相对位置信息编码进注意力外推性更好。注意力计算方面为了省显存和算力出现了分组查询注意力和多查询注意力让多个 Query 头共享少量 Key、Value 头大幅降低推理时的 KV 缓存开销。前馈网络这块也有变化出现了用门控机制的变体比如把 FFN 拆成两路一路做门控能提升效果但增加参数。视觉领域的 Swin Transformer 则把注意力限制在局部窗口内降低计算量的同时引入层级结构。这些都是同一个自注意力内核在不同约束下的工程取舍理解原版结构是看懂它们的前提。6.2 被问最多的几个结构问题最后列几个我在交流里被问得最多的问题答案其实都藏在前面。第一个为什么注意力要缩放答案是高维点积方差过大导致 softmax 饱和、梯度消失。第二个多头和单头的计算量谁大答案是差不多因为每个头维度被切小了总运算量基本持平收益主要来自表达多样性。第三个为什么解码器第一层要加因果掩码交叉注意力不用因为解码器生成时必须保证不看未来而交叉注意力看的是完整编码结果本来就没有未来可言。第四个残差连接到底解决什么最直接的是让梯度能顺畅回传深层网络才训得动同时让每层只需学习增量优化目标更简单。把这些为什么能讲清楚比背公式有用得多。我自己也是写了、调了、错了才慢慢把这些点连成一片的。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门