第 04 章 · 核心
Transformer 主干:注意力怎样工作
顺着一个 Token 的数据流,拆开自注意力、残差、归一化和前馈网络。
学习准备与本章目标
- 开始前
- 矩阵乘法直觉 · Embedding · Softmax
- 学完后
- 画出 Decoder Block · 推导注意力形状 · 解释残差和归一化的作用
本章路线图
Transformer 是全套课程的核心。本章按一条隐藏状态的数据流拆解:归一化 → 自注意力 → 残差 → 归一化 → FFN → 残差。读完后要能写出每一步形状,并解释信息在序列位置之间何时发生交换。
Query、Key、Value 可以怎样理解
每个位置向量分别乘三个可学习矩阵,得到 Q、K、V。Q 表示“我在寻找什么”,K 表示“我能被怎样匹配”,Q 与所有 K 的点积给出相关分数;Softmax 后的权重再对 V 加权求和,形成该位置汇总到的信息。
公式是 Attention(Q,K,V)=softmax(QKᵀ/√d)V。除以 √d 是为了控制点积随维度增大而变得过大,避免 Softmax 过早饱和。
因果注意力为什么不能偷看未来
生成模型在预测第 t 个位置时,只能看它之前的 Token。训练虽然会并行计算所有位置,但通过上三角因果 mask,把未来位置的注意力分数设为负无穷,使其 Softmax 概率为零。
如果 mask 方向写反或标签错位,训练 loss 可能看起来很好,生成却完全失败。检查时可改动未来 Token,确认当前输出不受影响。
多头、残差与前馈网络各做什么
多头注意力把隐藏维度拆成多个 head,不同 head 可以学习不同关系,再拼接回原维度。注意力负责跨位置通信;FFN 则对每个位置独立做非线性变换,通常先升维再降维,是模型参数和计算的重要来源。
残差连接让每个子层学习“在原表示上增加什么”,为深层网络保留梯度通道;LayerNorm 或 RMSNorm 控制表示尺度。现代 Decoder-only 模型多采用 Pre-Norm,即先归一化再进入子层,以提高深层训练稳定性。
从一个 Block 到完整 Decoder-only 模型
输入经过 Token Embedding 和位置信息,依次通过许多 Decoder Block,最后归一化并投影到词表 logits。训练时所有位置一起算;生成时重复拿最后一个位置的 logits 采样。
参数量并不只由层数决定,还与隐藏维度、FFN 宽度、词表和是否采用 MoE 有关。理解架构时始终同时追踪三个问题:张量形状是什么、位置之间在哪里通信、参数和计算主要花在哪里。
用形状推导一次多头注意力
输入 X 为 [B, T, D],投影后将隐藏维拆成 H 个头,每头维度 Dh = D / H:
Q, K, V [B, H, T, Dh]
Q @ Kᵀ [B, H, T, T]
attention prob [B, H, T, T]
prob @ V [B, H, T, Dh]
concat + Wo [B, T, D]
缩放点积注意力为:
除以 是为了避免维度增大时点积方差过大,导致 Softmax 过早饱和。遮罩 M 把不可见位置设成足够小的值。
因果遮罩、Padding 遮罩与数值问题
因果遮罩形成下三角可见区域,保证位置 t 不能访问未来标签;padding 遮罩排除补齐位置。两者可以组合,但广播维度必须与 [B, H, T, T] 对齐。
若某一整行都被遮罩,Softmax 可能面对全负无穷并产生 NaN。低精度实现通常在更高精度中计算归约,或使用经过验证的融合内核。FlashAttention 改变的是 IO 和计算分块方式,不改变数学定义。
MHA、MQA 与 GQA
标准多头注意力为每个头保留独立 Q/K/V。Multi-Query Attention 让所有 Query 头共享一组 K/V;Grouped-Query Attention 则让若干 Query 头共享一组 K/V。
训练阶段三者计算差异有限,生成阶段却很重要:KV Cache 大小与 KV 头数线性相关。减少 KV 头可显著降低长上下文和高并发的缓存,但共享过度可能影响质量,因此 GQA 常作为折中。
FFN、激活函数与 SwiGLU
注意力负责位置之间通信,FFN 对每个位置独立进行非线性变换。传统形式是 D → 4D → D;现代模型常使用门控结构:
其中上投影与门控投影产生两个分支,逐元素相乘后再降回隐藏维。FFN 通常占据大量参数和 FLOPs;MoE 正是把这一部分替换为按 Token 选择的专家网络。
残差、Pre-Norm 与 RMSNorm
残差让子层学习相对原表示的增量,也为深层反向传播保留接近恒等的通路。Pre-Norm 形式可以写成:
x = x + Attention(Norm(x))
x = x + FFN(Norm(x))
LayerNorm 会减均值再按方差缩放;RMSNorm 只按均方根缩放,计算更简单。归一化的维度通常是单个 Token 的隐藏维,不依赖 batch 内其他样本。
RoPE 如何把距离带进注意力
RoPE 对 Q/K 的二维分量施加随位置变化的旋转。两个位置旋转后的点积只依赖相对位移,因此注意力分数自然包含距离信息。
长上下文扩展方法会调整旋转频率或插值策略,但“窗口可设为更长”不等于模型能稳定使用远处信息。必须同时验证短文本退化、长文检索、多跳信息和真实任务,并计入注意力与 KV Cache 成本。
计算量和显存花在哪里
全注意力分数矩阵是 [T, T],因此朴素注意力随序列长度呈平方增长。线性投影和 FFN 主要随 T × D² 增长。短序列、大隐藏维时线性层可能占主导;超长序列时注意力矩阵与 KV Cache 更突出。
训练要保存反向所需激活,推理 Decode 则保存历史 K/V。FlashAttention 通过分块避免把完整分数矩阵频繁写回显存,activation checkpointing 则在反向时重算部分前向,二者解决不同问题。
一个最小自注意力实现
import math, torch
def causal_attention(q, k, v):
# q/k/v: [B, H, T, Dh]
scores = q @ k.transpose(-2, -1) / math.sqrt(q.size(-1))
T = q.size(-2)
mask = torch.triu(torch.ones(T, T, dtype=torch.bool), diagonal=1)
scores = scores.masked_fill(mask.to(scores.device), float("-inf"))
probs = torch.softmax(scores.float(), dim=-1).to(q.dtype)
return probs @ v
这个版本用于理解形状,不适合生产。实际训练应使用框架提供的 scaled dot-product attention 或经过验证的融合内核。
阅读任意模型配置的方法
拿到配置文件时依次找:词表大小、隐藏维、层数、Q 头数、KV 头数、每头维度、FFN 宽度、归一化类型、位置编码、上下文长度和是否使用 MoE。然后估算:
- Attention 投影参数约与
D²同阶。 - FFN 参数常是每层的大头。
- KV Cache 由层数、KV 头、每头维、序列长度、batch 和 dtype 决定。
- 总参数与激活参数在 MoE 中必须分开。
常见误区
- 注意力权重高不等于可直接解释模型因果决策。
- 多头并不是人工指定“语法头”“事实头”,功能由训练形成。
- FlashAttention 不是稀疏注意力,也不会修改输出定义。
- RoPE 扩展窗口不保证远距离推理能力。
- Decoder-only 的训练可并行计算所有位置,只有生成必须逐 Token 进行。
章末检查
- 从
[B,T,D]推导 Q/K/V、分数矩阵与输出形状。 - 解释因果遮罩为何不会妨碍训练并行。
- MQA/GQA 为什么主要改善推理缓存?
- 注意力与 FFN 分别承担哪类信息变换?
- FlashAttention、KV Cache 和 activation checkpointing 分别优化哪个阶段?
延伸资料
本章进阶内容
原理专题与代码实践
完成主教材后,按顺序阅读原理专题、完成代码实践,并用掌握标准复查本章内容。
- 01
先追踪单层 Transformer 的张量流
- 02
深入注意力、位置编码、归一化与长上下文
- 03
从 Embedding 开始实现可运行 Transformer
- 04
对每一步打印形状并做因果遮罩检查
本章术语
六个关键词
- Self-Attention
- 同一序列内各位置按相关度汇总信息的机制。
- Q / K / V
- 注意力中的查询、键和值三组投影。
- 因果遮罩
- 阻止当前位置访问未来 Token 的下三角约束。
- GQA
- 多个 Query 头分组共享较少 K/V 头的注意力结构。
- SwiGLU
- 带 SiLU 门控分支的前馈网络结构。
- Pre-Norm
- 先归一化再进入注意力或 FFN 的残差结构。
章节记录
