注意力之后为何还要逐位置变换?Transformer Block 的 FFN、残差与 Pre-LN
从注意力只负责位置混合的局限出发,拆解逐位置前馈网络、两条残差路径,并比较 Pre-LN 与 Post-LN 的数据流和梯度路径。
上一篇把缩放点积自注意力(Scaled Dot-Product Self-Attention)拆成了 、mask、softmax 与 的加权和。它让一个 token 能读取其他位置,却没有回答另一个问题:读回来的信息怎样在每个位置内部完成非线性加工?
Transformer Block 用两个互补子层回答它:多头自注意力负责沿序列位置混合信息,逐位置前馈网络(Position-wise Feed-Forward Network, FFN)负责沿特征维度变换信息。每个子层外再放残差连接、LayerNorm 与 dropout,才组成可堆叠的块。
本文只讲透这一个块:先追踪注意力与 FFN 的张量轴,再比较归一化放在残差分支之前还是之后。Encoder–Decoder 的交叉注意力、KV Cache 与完整语言模型生成留到后续文章。
01 只有自注意力,为什么仍然不够?#
设输入为:
是 batch 大小, 是序列长度, 是模型宽度。自注意力对第 个位置输出:
它擅长决定“第 个位置该从哪些位置读取什么”,但输出仍以 Value 的加权混合为核心。若一个位置读回了“主语”“否定词”和“动作”三类线索,还需要一个共享的非线性函数把这些特征组合成新的表示。
FFN 对每个位置独立应用同一组参数:
其中:
| 变量 | 形状 | 作用 |
|---|---|---|
[D] | 第 个 token 当前表示 | |
[F,D] | 把特征从 扩张到中间宽度 | |
[F] | 第一层偏置 | |
[D,F] | 把中间特征投影回 | |
[D] | 第二层偏置 | |
| 逐元素 | ReLU、GELU 等非线性 |
对整个 batch,nn.Linear 只改变最后一维:
X [N,L,D]
│
├─ Linear(D → F) ─► H [N,L,F]
├─ activation ─► H' [N,L,F]
├─ dropout ─► H''[N,L,F]
└─ Linear(F → D) ─► Z [N,L,D]
位置 1、2、...、L 使用相同的 W1/W2,但彼此不做求和。text所以两类子层分工非常清楚:
self-attention:沿 L 轴交换信息,关系矩阵是 [L,L]
FFN: 沿 D/F 轴加工信息,每个位置独立text没有位置混合,FFN 看不到别的 token;没有特征变换,自注意力读回的信息缺少逐位置的非线性加工。
02 一个 Transformer Block 的完整数据流#
先看 Pre-LN(Pre-Layer Normalization)版本。LN 位于每个残差分支的输入端:
完整数据流是:
X [N,L,D]
│
├─────────────────────────────────────────────────────┐
└─ LN1 ─► Multi-Head Self-Attention ─► Dropout ─► (+) ─► U [N,L,D]
│
┌──────────────────────────┘
├─────────────────────────────────────────────┐
└─ LN2 ─► Linear D→F ─► GELU ─► Linear F→D │
└─► Dropout ─► (+) ─► Y [N,L,D]
▲
U ──────────────────────────────────────────────┘text为避免图中紧凑标注造成误读,FFN 的精确宽度变化是 D → F → D。两次残差相加都要求分支输入与输出为 [N,L,D]; 只存在于 FFN 内部。
若 ,仅 FFN 两个权重矩阵就约有:
个参数;标准 Q/K/V 与输出投影合计约 。忽略偏置时,FFN 常比注意力投影拥有更多参数。注意力矩阵可能主导长序列的激活显存,FFN 则常主导块内参数与逐 token 计算量;不能只优化其中一边。
03 用两个特征手算一次 FFN 与残差#
暂时只看一个 token,令 :
第一层得到:
使用 ReLU 后:
再令:
则 FFN 修正量为:
残差相加后:
FFN 没有重新生成整个 token 表示,而是把第二个特征向下修正了 1.5。若第二个序列位置输入不同,它会独立经过完全相同的 ;两位置在这一步不会相互读取。
现在忽略 ,对 做不带仿射参数的 LayerNorm。均值为 ,方差为:
所以标准化结果为 [1,-1]。LayerNorm 改变的是单个 token 内特征的中心与尺度;它不沿 batch 或序列长度统计,也不会让两个 token 互相通信。
04 Post-LN 与 Pre-LN 究竟差在哪里?#
原始 Transformer 使用 Post-LN(Post-Layer Normalization)。对任一子层 :
Pre-LN 改为:
Post-LN:
x ───────────────┐
└─ Sublayer ─────┴─ (+) ─► LayerNorm ─► y
Pre-LN:
x ─────────────────────────────┐
└─ LayerNorm ─► Sublayer ──────┴─ (+) ─► ytext前向形状完全相同,差别却不只是“代码顺序”。设 和 分别为子层与 LayerNorm 对输入的雅可比矩阵,则局部梯度路径可写为:
Post-LN 中,残差相加后的所有信号还要经过 LayerNorm 的雅可比。Pre-LN 则有:
这里出现一条显式恒等项 :即使残差分支的局部梯度很小,仍有一条不经过本块 LayerNorm 与子层的直接路径。深层网络中,Pre-LN 往往更容易在训练初期维持梯度传播;这也是许多现代 Transformer 采用它的原因。
Pre-LN 堆叠后通常还会在整个栈末尾加一次最终 LayerNorm:
tokens ─► embedding + position ─► Block₁ ─► ... ─► Block_K ─► final LN ─► task headtext漏掉最终归一化,数值范围和已有实现的输出契约都会改变。
05 LayerNorm 到底沿哪条轴计算?#
对 token 向量 ,LayerNorm(Layer Normalization)计算:
对 [N,L,D] 调用 nn.LayerNorm(D) 时,每个 batch、每个位置分别沿最后一维 统计,输出仍是 [N,L,D]。 是逐特征可学习参数;统计量来自当前 token,在训练和推理时都这样计算,不维护 BatchNorm 式 running mean/variance。
这里有三个常见混淆:
LayerNorm(D)不会跨 batch 统计,batch size 从 32 改成 1 不会切换统计公式。- 它也不会沿 统计,因此 padding 位置不会直接污染真实 token 的 LayerNorm 统计;padding 仍需在注意力和损失端处理。
eps在平方根内用于数值稳定;混合精度下若出现非有限值,既要检查eps,也要检查进入归一化前的激活范围。
06 不调用 Transformer 封装,写出透明的 Pre-LN Block#
下面只用基础模块组装一个 Encoder block。它使用非因果自注意力;valid_tokens=True 表示真实 token,而 nn.MultiheadAttention 的布尔 key_padding_mask=True 表示应忽略,所以传入时要取反。
import torch
from torch import nn
class PreNormEncoderBlock(nn.Module):
def __init__(
self,
model_dim: int,
num_heads: int,
ffn_dim: int,
dropout: float = 0.1,
) -> None:
super().__init__()
assert model_dim % num_heads == 0
assert ffn_dim >= model_dim
self.norm1 = nn.LayerNorm(model_dim)
self.self_attn = nn.MultiheadAttention(
embed_dim=model_dim,
num_heads=num_heads,
dropout=dropout, # attention 权重上的 dropout
batch_first=True,
)
self.attn_output_dropout = nn.Dropout(dropout)
self.norm2 = nn.LayerNorm(model_dim)
self.ffn = nn.Sequential(
nn.Linear(model_dim, ffn_dim),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(ffn_dim, model_dim),
)
self.ffn_output_dropout = nn.Dropout(dropout)
def forward(
self,
x: torch.Tensor, # [N,L,D]
valid_tokens: torch.Tensor, # [N,L], bool, True = 真实 token
) -> torch.Tensor:
n, length, width = x.shape
assert valid_tokens.shape == (n, length)
assert valid_tokens.dtype == torch.bool
assert valid_tokens.any(dim=1).all()
qkv = self.norm1(x) # [N,L,D]
attn_out, _ = self.self_attn(
qkv, qkv, qkv,
key_padding_mask=~valid_tokens, # MHA 中 True = 忽略 key
need_weights=False,
) # [N,L,D], None
x = x + self.attn_output_dropout(attn_out)
ffn_out = self.ffn(self.norm2(x)) # [N,L,D]
x = x + self.ffn_output_dropout(ffn_out)
return x
block = PreNormEncoderBlock(
model_dim=32, num_heads=4, ffn_dim=128, dropout=0.1
)
tokens = torch.randn(2, 6, 32)
valid = torch.tensor([
[True, True, True, True, False, False],
[True, True, True, True, True, True],
])
output = block(tokens, valid)
assert output.shape == (2, 6, 32)pythonself_attn 内部的 dropout 作用于注意力权重;残差分支输出处的 dropout 是另一处随机化,不能因为都叫 dropout 就合并。nn.Dropout 在训练时把保留元素按 缩放,在 eval() 时成为恒等映射。
key_padding_mask 只阻止真实查询读取 padding 键列。padding 查询行仍可能产生非零输出,残差也会继续携带它们。若任务头做平均池化,应显式用 valid_tokens 做 masked mean;若做 token 级损失,应使用 ignore_index 或等价 mask。
07 怎样改成 Post-LN?#
模块参数可以不变,只改前向顺序:
def post_norm_forward(self, x, valid_tokens):
attn_out, _ = self.self_attn(
x, x, x,
key_padding_mask=~valid_tokens,
need_weights=False,
)
x = self.norm1(x + self.attn_output_dropout(attn_out))
x = self.norm2(x + self.ffn_output_dropout(self.ffn(x)))
return xpython不要在 Pre-LN 代码上“顺手”保留相加后的第二次归一化,否则会变成第三种结构。架构实验必须把每个 LayerNorm 的输入、输出和残差相加位置画出来,而不是只在配置里记录一个含糊的 pre_norm=True。
若要让 Pre-LN 块在初始化时接近恒等映射,可把残差分支的最后输出投影初始化得很小或为零;但这会改变默认初始化,必须记录并单独验证,不能默默加入“透明实现”。
08 与 PyTorch 2.13 当前官方层对齐#
PyTorch 2.13 的 nn.TransformerEncoderLayer ↗ 是用于理解基础架构的参考实现。关键参数为:
from torch import nn
official = nn.TransformerEncoderLayer(
d_model=32,
nhead=4,
dim_feedforward=128,
dropout=0.1,
activation="gelu",
batch_first=True,
norm_first=True, # True = Pre-LN;默认 False = Post-LN
layer_norm_eps=1e-5,
bias=True,
)
src = torch.randn(2, 6, 32)
src_key_padding_mask = ~valid
out = official(src, src_key_padding_mask=src_key_padding_mask)
assert out.shape == src.shapepython当前 API 有几项值得写进契约测试:
batch_first=True才使用[N,L,D];默认仍是[L,N,D]。norm_first=True表示注意力和 FFN 之前做 LayerNorm;默认False对应 Post-LN。dim_feedforward是 ,不会改变最终输出宽度 。activation当前可用字符串"relu"、"gelu"或一元 callable;默认是 ReLU。src_key_padding_mask=True表示忽略该键;is_causal是因果 mask 的提示,错误提示可能导致不正确执行。- 该层是基础参考实现,只提供有限的现代 Transformer 特性;不能把“官方类”误解为所有场景下最快或最完整的生产实现。
官方文档还列出了推理优化路径的条件,例如 .eval()、关闭 autograd、三维 batch-first 输入、受支持激活,以及 mask 组合限制。优化是否命中应以 profiler 和所用版本为准,不要靠类名猜测。
09 从 token 到分类结果,一次训练怎样流动?#
以文本分类为例, 是词表大小, 是类别数:
token_ids [N,L]
├─ Embedding(V,D) ─► token vectors [N,L,D]
├─ position vectors [L,D](广播到 batch)
▼
X [N,L,D]
└─ K 个 Pre-LN Block ─► final LayerNorm ─► H [N,L,D]
│
valid_tokens [N,L] ─► masked mean ──────┘
▼
pooled [N,D]
▼
Linear(D,C)
▼
logits [N,C]text关键部分可以写成:
def masked_mean(x, valid_tokens):
weights = valid_tokens.unsqueeze(-1).to(x.dtype) # [N,L,1]
summed = (x * weights).sum(dim=1) # [N,D]
counts = weights.sum(dim=1).clamp_min(1.0) # [N,1]
return summed / counts
model.train()
logits = model(token_ids, valid_tokens) # [N,C],未经 softmax
loss = nn.functional.cross_entropy(logits, labels) # labels [N]
optimizer.zero_grad(set_to_none=True)
loss.backward()
grad_norm = nn.utils.clip_grad_norm_(
model.parameters(), max_norm=1.0, error_if_nonfinite=True
)
optimizer.step()python推理时需要同时切换模块行为与关闭梯度记录:
model.eval()
with torch.inference_mode():
logits = model(token_ids, valid_tokens) # [N,C]
probabilities = logits.softmax(dim=-1) # [N,C]
predictions = probabilities.argmax(dim=-1) # [N]pythonmodel.eval() 会关闭模块式 Dropout,但它本身不关闭 autograd;torch.inference_mode() 关闭梯度记录,却不会替你把模型切到 eval。两者职责不同。
10 一条可执行的调试路径#
- 先过拟合一个极小 batch。 用 4 条固定长度样本,关闭 dropout,确认损失能快速接近 0;否则先查标签、mask 和残差顺序。
- 逐点打印形状。 注意力、两次残差相加、块输出都应为
[N,L,D];FFN 中间才是[N,L,F]。 - 做 padding 不变性测试。 固定真实前缀,只替换 padding token 的 embedding;真实位置输出和 pooled logits 应保持不变。
- 单独关掉子层。 把注意力输出投影或 FFN 第二个 Linear 置零,Pre-LN 块应分别退化为另一子层加恒等路径。
- 记录逐层残差比例。 监控
||S(LN(x))|| / ||x||;突然从小量级跳到数十倍常预示学习率、初始化或数值问题。 - 记录深度方向梯度。 对每层输入保留梯度,比较浅层到深层的范数;只看全模型总梯度会掩盖 Post-LN 的局部衰减。
- 固定随机性比较 train/eval。 Dropout 开启时两次训练前向可以不同;eval 前向应一致。若验证仍抖动,检查是否调用了函数式 dropout 且忘传
training=self.training。 - 对齐官方层。 在小维度、dropout=0 下复制参数,逐子层比较输出;不要只比较最终 loss。
padding 不变性测试应同时覆盖“真实位置表示”和“任务头输出”。只检查注意力权重的 padding 列为零,仍可能在无 mask 的平均池化处把 padding 查询混入结果。
11 最常见的“形状正确,结构却错了”#
- 把 FFN 写成跨序列卷积或先展平
[L,D]。 标准逐位置 FFN 共享参数但不混合位置。 - 第一层扩到 后忘记投回 。 残差相加因此失败,或被迫引入未经设计的投影。
- 把两次 LayerNorm 复用成同一个实例。 两个位置通常各有自己的 ;共享会改变参数化。
- 在 Pre-LN 相加后又做 LayerNorm。 这不再是本文公式中的 Pre-LN。
- 只给 FFN 加残差,漏掉注意力残差。 信息与梯度都必须穿过注意力分支。
- 同一个 Dropout 实例并非错误,但把不同 dropout 位置当成一次操作是错误。 注意力权重、FFN 隐层和残差分支输出的作用点不同。
key_padding_mask真值方向反了。 在 MHA 中True=忽略,与上一篇 SDPA 布尔 mask 的语义相反。- 只 mask 注意力,不 mask 池化或损失。 padding 查询仍可进入任务头。
- Pre-LN 栈漏掉 final LayerNorm。 最终表示尺度与常见实现不一致。
- 用
eval()代替关闭梯度。 仍会构建 autograd 图并占用内存。 - 用
inference_mode()代替eval()。 Dropout 仍可能保持训练行为。 - 比较 Pre/Post-LN 时沿用同一最优学习率就下结论。 两者优化条件不同,应分别调参并报告初始化与 warmup。
12 这个块会在哪些场景失败?#
- 超长序列。 标准自注意力的 关系矩阵仍是主要瓶颈;FFN 不会修复它。
- FFN 宽度过大。 参数、激活显存和逐 token 计算迅速增长,尤其在大词表或长 batch 下。
- 残差分支尺度失控。 恒等路径不能阻止非有限激活、过大学习率或错误 mask 注入异常值。
- 数据很少。 宽 FFN 与多头注意力提供高容量,也更容易记忆训练集;需要学习曲线和任务匹配的正则化。
- 精确算法任务。 连续向量与有限深度未必可靠执行长位数算术、栈操作或长度外推。
- padding 占比极高。 逻辑 mask 保证语义正确,却不自动省掉所有密集计算;要另行评估打包、Nested Tensor 或变长内核。
- 分布外长度。 位置表示、残差尺度和训练上下文共同限制长度外推,不能只替换 FFN 激活就解决。
13 与相近结构怎样区分?#
| 结构 | 跨位置混合 | 每位置特征变换 | 归一化位置 | 主要用途 |
|---|---|---|---|---|
| 仅自注意力层 | 是 | 主要是 Q/K/V 与输出线性投影 | 未规定 | 建立 token 关系 |
| Transformer Encoder Block | 自注意力 | FFN | Pre-LN 或 Post-LN | 双向上下文编码 |
| Transformer Decoder Block | causal 自注意力,可再加交叉注意力 | FFN | 依架构而定 | 自回归生成/条件生成 |
| 卷积残差块 | 局部卷积 | 通道与空间共同变换 | BN/LN 等 | 图像或局部序列建模 |
| MLP-Mixer 类块 | 显式 token-mixing MLP | channel-mixing MLP | 通常有 | 不用注意力的全局混合 |
| MoE Transformer | 注意力不变 | 只路由到部分专家 FFN | 通常沿用主干 | 增大参数容量而控制单 token 计算 |
门控 FFN(如 GLU/SwiGLU)改变的是逐位置非线性分支;稀疏注意力改变的是位置混合图;FlashAttention 优化的是精确注意力的内存访问。它们解决不同层面的问题,不能都笼统称为“更快的 Transformer”。
14 今天真正需要记住什么?#
- 自注意力沿序列轴混合 token,FFN 用共享的
D → F → D非线性网络独立加工每个位置,两者缺一不可。 - 一个标准块包含两次残差相加:注意力子层一次、FFN 子层一次;每次分支输出都必须回到
[N,L,D]。 - Post-LN 是
LN(x+S(x)),Pre-LN 是x+S(LN(x));后者的局部梯度含显式恒等项,深层训练常更稳定。 LayerNorm(D)对[N,L,D]的最后一维逐 token 统计,训练与推理都使用当前输入统计。- mask、池化 mask 与损失 mask 负责不同边界;
eval()与inference_mode()也不能互相替代。
15 思考题与小练习#
- 延续手算例,把第二个 token 设为
[-1,1],用同一组 计算其 FFN 输出。说明两个 token 为何共享函数却没有在 FFN 中互相影响。 - 为
PreNormEncoderBlock写 padding 不变性测试:保持valid_tokens不变,随机替换 padding 位置输入,验证所有真实位置输出不变;再故意去掉key_padding_mask,观察测试失败。 - 令 ,忽略偏置,计算 FFN 与 Q/K/V+输出投影各自的参数量。再让序列长度从 512 翻倍到 1024,解释参数量为何不变,而注意力分数元素数为何约增至 4 倍。
相关工作#
- Vaswani et al. (2017), Attention Is All You Need ↗:提出由多头注意力、逐位置 FFN、残差与归一化组成的原始 Post-LN Transformer。
- Xiong et al. (2020), On Layer Normalization in the Transformer Architecture ↗:分析 Pre-LN 与 Post-LN 在初始化时的梯度行为及 warmup 需求。
- Shazeer (2020), GLU Variants Improve Transformer ↗:研究门控逐位置前馈网络及其激活变体。
- Wang et al. (2022), DeepNet: Scaling Transformers to 1,000 Layers ↗:通过残差与初始化缩放研究极深 Transformer 的稳定训练。
- Dao et al. (2022), FlashAttention ↗:优化精确注意力的 IO 路径,帮助区分架构数学与内核实现问题。
16 下一篇预告#
一个 Encoder block 已能让所有 token 交换并加工信息,但自回归生成还要求“只能看过去”,条件生成还要求解码器读取另一条源序列。下一篇将组装 Transformer Decoder,逐层区分 causal self-attention、cross-attention 与 FFN 的 Query/Key/Value 来源,并追踪训练和逐 token 推理的数据流。