Chapter 04

训练与推理:并行与自回归的不对称

第 03 章给了静态网络结构——本章讲它如何"跑起来":训练时全位置并行、推理时一次一个 token 自回归,这套不对称是最大的认知陷阱,也是 KV-cache 与 O(n²) 成本的来源。

本章交付:训练并行机制 + 推理自回归循环 + KV-cache 原理与成本公式 + O(n²) 的显存本质 | 带走的能力:能解释为什么训练和推理的前向传播逻辑不同,能估算给定配置下 KV-cache 的显存占用

读完本章,你的脑子里会多出这几条

  • 训练用 teacher forcing + 因果掩码一次算完所有位置(并行)——因果掩码让每个位置只看左边,但所有位置的 loss 通过一次矩阵乘同时得到。
  • 推理无 ground truth,必须把上一步生成的 token 追加回输入,再跑一次前向,得到下一个——这是串行自回归,不可并行。
  • KV-cache 缓存过去 token 每层每头的 K、V 投影(不是权重),省去每步重算整片;代价是显存随序列长度线性增长。
  • attention 的 O(n²) 源于 QKᵀ 产生 n×n 分数矩阵,长上下文下 attention 是显存瓶颈,不只是算力瓶颈。

4.1训练:teacher forcing + 全位置并行

训练时将真实 token 序列整体喂入,因果掩码(第 02 章)防止每个位置偷看未来,所有位置的 next-token loss 经一次矩阵乘同时计算完毕。

训练的输入是一段真实文本序列,例如 [BOS, the, cat, sat, on, the, mat]。模型不需要知道"正确答案是什么"——答案就是序列本身向右移一位:输入位置 0 预测位置 1,位置 1 预测位置 2,以此类推。这种做法叫 teacher forcing(教师强制):每个位置的输入永远是真实 token,而不是模型上一步的预测结果,因此不存在误差累积。

能并行的关键在于:第 02 章介绍的因果掩码把注意力矩阵的上三角(未来位置)在 softmax 前置为 −∞,softmax 后权重归零——位置 t 只能看到位置 0…t 的信息,而这个约束对所有位置同时成立。于是一整个序列的 Q、K、V 可以同时投影、同时做矩阵乘,所有位置的注意力分数和 loss 通过单次前向传播同时得到,梯度也通过单次反向传播同时更新。

训练前向 · 伪代码 Python
# tokens: (batch, seq_len)  — 整段真实序列
logits = model(tokens[:, :-1])      # 输入:去掉最后一个 token
targets = tokens[:, 1:]             # 标签:去掉第一个 token(右移一位)
loss = cross_entropy(logits, targets)  # 所有位置一次算完
loss.backward()                     # 单次反向传播,全部梯度

交叉熵 loss 的形式:对每个位置 t,取 logits[t] 经 softmax 后在真实 token 上的负对数概率,再对整个序列平均:

next-token 交叉熵 数学
L = -1/T · Σₜ log P(xₜ₊₁ | x₁..xₜ)

  T       : 序列长度
  P(·|·)  : 模型 softmax 输出的概率
  xₜ      : 位置 t 处的真实 token
想一想

训练时所有位置能并行计算,但推理时却只能一步一步生成。两者用的是同一套模型权重——为什么同一套权重会产生这种差异?

展开答案(先停 10 秒再点)

差异不在权重,在输入从哪里来。训练时每个位置的输入都是真实 token(teacher forcing),整段序列已知,因此可以批量矩阵乘一次算完。推理时位置 t 的输入是模型刚才生成的 token——它在生成位置 t 之前根本不存在,必须先生成 t,才能把它追加进去,再生成 t+1。这个"生成 → 追加 → 再生成"的依赖链决定了推理必须串行。

训练(并行) the cat sat on 输入序列(teacher forcing) ▼ 因果掩码(下三角可见) ✓ ✗ ✗ ✗ ✓ ✓ ✗ ✗ ✓ ✓ ✓ ✗ ✓ ✓ ✓ ✓ ↓ 所有位置同时 L₁ L₂ L₃ L₄ 各位置 loss,求均值 mean(L) → backward 推理(串行自回归) prompt: "the cat sat" 前向传播 → logits 生成 token: "on" 追加 输入: "the cat sat on" 前向传播 → logits 生成 token: "the" 循环直到 EOS
图 4.1 左:训练并行——因果掩码构成下三角可见矩阵(✓=可见,✗=被掩盖),所有位置的 loss 经一次前向传播同时得到。右:推理串行——每步只生成一个 token,追加进输入后再跑一次前向,循环至 EOS 或达到长度上限。注意:两侧使用完全相同的权重,差异只在"输入从哪来"——训练时来自真实序列,推理时来自模型自己上一步的输出。

4.2推理:自回归循环

推理没有 ground truth 可供 teacher forcing,每步把刚生成的 token 追加到输入序列,重新前向传播,取最后一位的 logits 采样或 argmax,循环至 EOS 或达到长度上限。

推理的全过程:

  1. 以 prompt(可以是 [BOS] 或用户的问题 token 序列)作为初始输入。
  2. 前向传播,取最后一个位置的 logits,经 softmax 得到 next-token 概率分布。
  3. 从分布中采样(temperature 采样、top-k、top-p 等)或取 argmax,得到 token t。
  4. 把 token t 追加到输入末尾,回到步骤 2。
  5. 直到生成 <EOS> 或序列长度达到上限,停止。

关键:步骤 3 生成的 token t 在步骤 2 之前根本不存在,所以步骤 2→3→4→2 的循环无法被展开成并行矩阵操作——这是数据依赖带来的串行性,不是工程局限。

贪心 vs 采样

argmax(每步选概率最大的 token,即贪心解码)速度最快、结果确定,但容易陷入重复或过于保守的输出。temperature 采样引入随机性,temperature > 1 使分布更平坦(更多样)、< 1 使分布更尖锐(更保守)。top-k 和 top-p(nucleus sampling)截断低概率的长尾,在质量与多样性之间取平衡。这些策略都在步骤 3 改变如何从 logits 取 token,不影响模型权重本身。

4.3不对称是最大的认知陷阱

这套并行训练 / 串行推理的不对称,是初学者最高频的卡点:用训练式前向实现"推理",输出永远是垃圾——因为少了把输出喂回输入的循环。

按上下文的 §5 所记:很多人实现 transformer 时,先写出 logits = model(tokens) 的训练前向,然后直接拿 logits[-1] 当"生成结果",却发现无论怎么调参,输出都像随机噪声。症结在这里:

  • 训练前向:一次喂入完整 token 序列,同时得到所有位置的 logits。这里的 logits[t] 预测 token t+1,是拿 t 之前的真实 token 算的。
  • 推理前向:每次只取最后位置的 logits 采样一个新 token,把这个 token 追加进去,再跑整个序列——少了这个追加步骤,模型永远在对同一段 prompt 重复预测,不是在"生成"。
失败模式

nanoGPT(Karpathy)的 generate() 函数之所以是经典参考,正是因为它把这个循环写得清清楚楚:for _ in range(max_new_tokens): logits = model(idx); next_token = sample(logits[:, -1, :]); idx = cat(idx, next_token)。去掉 cat(idx, next_token) 这一行,"推理"就退化成反复对 prompt 打分,产不出任何新 token 序列。

这也是 aha #3 的核心:两个阶段的前向传播在数学上调用的是同一套 attention + FFN 权重(第 03 章介绍的完整 block),差异纯粹在输入构造和调用方式上。理解了这一点,KV-cache 的必要性就立刻显现——既然推理每步都要重跑整个序列,所有过去 token 的 K、V 投影每次都要重算,这是巨大的浪费。

4.4KV-cache:别重算过去

推理时每一步新生成的 token 只需要计算自己的 Q,过去 token 的 K 和 V 已经算过且不会改变——KV-cache 把它们存下来,把每步复杂度从 O(n) 次矩阵乘降到 O(1) 次。

理解 KV-cache 需要先回忆 第 02 章的 scaled dot-product attention:

self-attention 前向(单头) 数学
Q = X · Wq    (n, d_k)
K = X · Wk    (n, d_k)   ← 过去 token 的 K 每步都一样
V = X · Wv    (n, d_v)   ← 过去 token 的 V 每步都一样

scores = Q · Kᵀ / √d_k   (n, n)
out    = softmax(scores) · V   (n, d_v)

生成 token t+1 时,新输入只有 token t+1 本身。位置 0…t 的 X 没变,所以它们的 K 和 V 投影与上一步完全相同。KV-cache 做的事:在每层每头把已算过的 K、V 矩阵拼接存入缓存,新 token 只算自己的 Q(以及自身新的 K、V 追加入缓存),再和缓存里的 K、V 做 attention。

重要区分:缓存的是 K 和 V 的投影结果(形状 (past_len, d_k) 和 (past_len, d_v)),不是权重矩阵 Wk、Wv(权重在整个推理过程中始终不变,本来就无需重算)。

显存成本公式

KV-cache 的显存占用可以精确估算:

KV-cache 显存 数学
KV_cache_bytes = seq_len × layers × heads × head_dim × 2 × dtype_bytes

  seq_len    : 上下文窗口长度(已生成 + prompt 的 token 数)
  layers     : Transformer 层数
  heads      : 每层注意力头数(KV 头,GQA 下小于 Q 头数)
  head_dim   : 每头的维度 = d_model / heads
  × 2        : K 和 V 各一份
  dtype_bytes: fp16 = 2 bytes, bf16 = 2 bytes, int8 = 1 byte

以 7B 规模的典型配置(Llama-2-7B 同规格)具体计算:

7B 模型 · 8K 上下文 · fp16 的 KV-cache 数学
seq_len    = 8 192
layers     = 32
heads      = 32        (MHA,KV 头 = Q 头)
head_dim   = 128       (d_model=4096 / 32头)
dtype      = fp16 = 2 bytes

KV = 8192 × 32 × 32 × 128 × 2 × 2
   = 8192 × 32 × 32 × 512
   ≈ 4 294 967 296 bytes
   ≈ 4 GB  ×  2  (K 和 V)
   ≈ 8 GB

结论:7B / 32层 / 32头 / d128 / 8K上下文 / fp16
     ≈ 8 GB KV-cache(不含模型权重本身的 ~14 GB)

这 8 GB 随序列长度线性增长:上下文翻倍到 16K,KV-cache 也翻倍到 16 GB。GQA(Grouped Query Attention,第 05 章)通过让多个 Q 头共享一组 KV 头,直接按比例压缩这个数字——32 头 GQA 分成 8 组,则 KV 头数从 32 降到 8,缓存降至约 2 GB。

无 KV-cache(每步重算) 生成 token t+1 时 t₀ t₁ t₂ … tₜ new tₜ₊₁ ▼ 重算所有 K, V X · Wk, X · Wv 全部 t+1 个 token 重算 Attention → 新 token 每步计算量 O(t²) — 随长度平方增长 t=1000 步:重算 ~10⁶ K/V 对 有 KV-cache(只算新 token) 生成 token t+1 时 t₀…tₜ(K,V 已在缓存) new tₜ₊₁ ▼ 只算新 token 的 Q(及 K, V 追加缓存) 新 Q = tₜ₊₁ · Wq 只算 1 个 token 读缓存 K[0…t], V[0…t] 无重算,直接读显存 Attention → 新 token 显存随长度 线性增长 7B/32L/8K/fp16 ≈ 8 GB
图 4.2 左:无 KV-cache——每步生成新 token 时,整段历史序列的 K、V 投影全部重算,计算量随 token 数平方增长。右:有 KV-cache——过去 token 的 K、V 只算一次后存入缓存,新步只算当前 token 的 Q,并把新 K、V 追加入缓存。注意:节省的是计算,代价是显存——缓存大小随序列长度线性增长,这直接催生了第 05 章的 GQA/MQA/MLA 优化。

4.5O(n²):长上下文为什么贵

attention 的 O(n²) 根源在于 QKᵀ 产生 n×n 分数矩阵,这个矩阵必须落地(或用 FlashAttention 分块重算避免);在长上下文下,显存瓶颈先于算力瓶颈到来。

回到 第 02 章的公式:scores = Q·Kᵀ / √d_k,其中 Q 形状 (n, d_k),Kᵀ 形状 (d_k, n),相乘结果是 (n, n)——n 个 query 对 n 个 key 各算一个相似度分数。

这个 n×n 矩阵的问题有两个维度:

  • 算力:矩阵乘的浮点操作数是 O(n²·d_k),随序列长度平方增长。
  • 显存:n×n 矩阵要暂存在 GPU SRAM 或 HBM(高带宽存储)中以便后续 softmax 和乘 V。n=8192 时,每头分数矩阵有 8192×8192 ≈ 6700 万个 float16 元素 ≈ 128 MB;32 头 × 32 层 = 超过 100 GB,远超单卡 HBM 容量。
n=8192 时每头 scores 矩阵大小 数学
n = 8 192
每头 scores = n × n = 8192 × 8192 ≈ 67 108 864 个元素
fp16: 67 108 864 × 2 bytes ≈ 128 MB / 头

32 层 × 32 头 × 128 MB ≈ 128 GB(scores 矩阵总量)
→ 这是为什么朴素 attention 在长上下文下首先触达显存天花板
  而不是算力天花板

朴素实现将整个 n×n scores 矩阵写入 HBM,再读回来做 softmax,再写一次,再读回来乘 V——内存带宽是真正的瓶颈(HBM 带宽约 2 TB/s,但矩阵反复读写造成大量浪费)。FlashAttention(第 05 章)的核心思路正是把这个矩阵"分块留在 SRAM 里",避免落地 HBM,从根本上规避这一瓶颈——且数学输出完全等价,不是近似。

洞察 · 显存先于算力成为瓶颈

GPT-4 / Claude 3 等的长上下文支持(128K+ token)之所以需要 FlashAttention,不是因为算力不够,而是因为标准 attention 的 n×n 矩阵根本放不进 GPU 显存。这一约束催生了第 05 章的 FlashAttention IO-aware 设计以及 sliding-window attention 等长上下文方案。

4.6这些成本催生了下一章的优化

KV-cache 的线性显存增长和 O(n²) 的平方瓶颈,是第 05 章所有现代优化的直接动机——不理解成本的来源,就无法理解那些方案在解决什么。

本章揭示的两条成本线,在第 05 章各对应一簇优化方案:

表 4.1 · 本章成本 → 第 05 章对应优化
本章揭示的成本量级第 05 章对应方案
KV-cache 显存随 seq_len 线性增长 7B/8K/fp16 ≈ 8 GB GQA / MQA(减少 KV 头数)/ MLA(DeepSeek 低秩压缩 KV)
attention O(n²) 算力 + 显存 n=8K 每头 128 MB scores FlashAttention 1/2/3(IO-aware 分块)/ sliding-window / attention sinks

第 05 章在这两条线上都会给出机制级解释:GQA 如何把 KV 头数从 32 压到 8;FlashAttention 如何用 SRAM 分块消除 HBM 落地;MLA 如何把 K/V 压成低秩潜向量再按需解压(DeepSeek-V2,2024-05)。把本章的成本公式和第 05 章的优化方案对照起来看,才能真正理解那些方案在省什么、代价是什么。

§本章 self-check

先合上教程,把能想到的答案写在纸上或编辑器里,写完再展开对照。

  1. 训练时为什么所有位置可以并行计算 loss?第 02 章的哪个机制保证了并行不等于"偷看答案"?
  2. 推理时"自回归"的具体步骤是什么?少了哪一步会导致生成永远是垃圾?
  3. KV-cache 缓存的是什么(精确到矩阵名称)?为什么不用也不能缓存权重矩阵 Wk、Wv?
  4. 对于一个 32 层、32 头、head_dim=128、上下文 4K 的 fp16 模型,KV-cache 大约占多少显存?
答案(先做完再展开)
  1. teacher forcing 让每个位置的输入都是真实 token,整段序列已知,因此 Q、K、V 可以整批矩阵乘。因果掩码(第 02 章)把注意力矩阵的上三角在 softmax 前置 −∞,保证位置 t 只能看到 0…t,不泄露未来——在此约束下所有位置一次算完。
  2. 步骤:① 前向传播取最后位置 logits;② 采样或 argmax 得到新 token t;③ 把 token t 追加进输入序列;④ 回到步骤 ①,循环至 EOS。少了步骤 ③,模型一直对同一段 prompt 打分,输出永远是对初始 prompt 末尾的预测,不是真正的生成。
  3. 缓存的是每层每头的 K 投影结果(X·Wk)和 V 投影结果(X·Wv)。权重矩阵 Wk、Wv 在整个推理过程中从不改变,始终在显存里,根本不需要缓存——缓存是为了避免对同一段历史 token 重复做 X·Wk 和 X·Wv 的乘法。
  4. 4096 × 32 × 32 × 128 × 2 × 2 bytes = 4 096 × 32 × 32 × 512 ≈ 2 147 483 648 bytes ≈ 2 GB(约为 8K 情形的一半)。