一次注意力计算到底长什么样?——从 QKV 投影到 KV Cache 的完整拆解

0 阅读9分钟

一次注意力计算到底长什么样?——从 QKV 投影到 KV Cache 的完整拆解

摘要:本文从单个 Token 的输入出发,逐段拆解 Transformer 中一次注意力计算的完整链路:QKV 投影、缩放点积打分、因果掩码、Softmax 归一、加权求和,以及 FFN 的升维—激活—降维。随后以「Prefill 8192 Token + Decode 1024 Token」为算例,给出注意力计算量与 KV Cache 显存占用的闭式公式,并指出原始推导中容易算错的两处细节。

关键词:Transformer;注意力机制;QKV 投影;KV Cache;Prefill / Decode;大模型推理;显存优化


一、为什么要拆开看「一次注意力」

大模型推理被天然切成两个阶段,二者的计算形态完全不同:

  • Prefill(预填充):一次性吃进整段 Prompt,做的是矩阵 × 矩阵,算力密集,可以打满 Tensor Core。
  • Decode(解码):每步只吐出一个 Token,做的是矩阵 × 向量,算术强度极低,瓶颈在显存带宽。

KV Cache、PagedAttention、FlashAttention、Prefix Caching、GQA 这些优化,本质上都是在这两个阶段的不同瓶颈上做文章。想真正看懂它们,前提是先能手推一遍一次 Attention 的张量形状和开销账。本文就是这件事的最小完整版本。


二、单个 Token 的前向链路

2.1 投影:一个 Token 变成 Q、K、V

输入 Token 先经 Embedding(叠加位置编码)得到向量 x ∈ R^(d_model)。随后与三个权重矩阵相乘,投影到语义空间:

Q = x · W_Q      # "我在找什么"
K = x · W_K      # "我能被什么找到"
V = x · W_V      # "找到了我提供什么内容"
符号含义形状
x单个 Token 的隐状态[d_model]
W_Q / W_K / W_V查询 / 键 / 值的投影矩阵[d_model, d_model]
Q / K / V投影结果[d_model]

批量输入时形状升一维:X: [B, L, d_model] → Q, K, V: [B, L, d_model]。

2.2 打分:Q·Kᵀ / √d_k

对第 i 个 Query 与第 j 个 Key 做点积得到匹配分数,再除以 √d_k 缩放:

score(i, j) = (Q_i · K_j) / √d_k

为什么必须除以 √d_k:假设 q、k 的各分量独立、均值 0、方差 1,那么点积的方差为 d_k、标准差为 √d_k。d_k 越大,logits 的绝对值越容易被放大,Softmax 就越容易推进饱和区——输出退化成近似 one-hot,梯度趋近 0。除以 √d_k 正是把方差拉回 1,让 Softmax 始终工作在梯度良好的区间。这一步是必需项,不是可选的 trick。

2.3 因果掩码与 Softmax

对分数矩阵沿最后一维做 Softmax,得到行和为 1 的注意力权重:

Attn = Softmax(Q·Kᵀ / √d_k)

在自回归场景下必须先施加因果掩码(Causal Mask):第 i 个 Query 只能看到位置 0..i 的 Key,未来的位置屏蔽为 -∞。掩码方式通常是加性掩码(屏蔽位加 -1e9 或直接置 -inf)后再做 Softmax,而不是 Softmax 之后再置零——后者会破坏归一化。

需要注意:Prefill 阶段序列长度大于 1,必须显式加掩码;而 Decode 每步只有一个新 Query,它能看到的是全部历史 Key,天然满足因果性,无需额外掩码。

2.4 加权求和

x_out = Attn · V

这一步才是真正的信息提取:注意力权重只是"配比",内容全部来自 V。因此业界常说"KV Cache 缓存的是内容,而不是分数"。

2.5 输出投影、残差与 LayerNorm

一次 Attention 到这里还没结束,还差三步(原始推导常漏掉这一段):

  1. 输出投影:AttnOut · W_O,把多头拼接结果映射回 d_model;(这里的w_o也是直接由Token和W_o权重矩阵计算直接可以得到的)。
  2. 残差连接:h = x + AttnOut · W_O;
  3. LayerNorm:对 h 做归一化,稳定后续 FFN 的输入分布。

2.6 FFN:升维 → 激活 → 降维

FFN(h) = W_2 · act(W_1 · h)

W_1 把维度从 d_model 升到 d_ff(经典设置 d_ff = 4·d_model;SwiGLU 结构约为 8/3·d_model),经激活函数(GELU / SwiGLU)后再由 W_2 降回 d_model。

这里就是"激活值"产生的地方。 激活张量形状为 [B, L, d_ff],是训练期显存的主要占用者之一。工程上常用激活重计算(Activation Checkpointing)在前向时丢弃它、反向时重算,用算力换显存。推理阶段不需要反向,激活值用完即可释放,但峰值仍要预留——这也是长序列下 Prefill 容易 OOM 的原因之一,而且它与 L 成正比,与 KV Cache 是两笔不同的账。

FFN 之后再接一次残差与 LayerNorm,得到一个完整 Transformer Block 的输出,送入下一层。


三、多头注意力与 KV Cache 的由来

实践中不会用单个 d_model 维的大头,而是拆成 h 个并行头,每头维度 d_head = d_model / h,最后拼接、经 W_O 投影。为了让 KV Cache 更小,衍生出三种形态:

形态Query 头数KV 头数每 Token KV Cache(相对量)代表模型
MHAhh1×Llama-2-7B
GQAhh / g(分组共享)1/g ×Llama-3-8B、Qwen2
MQAh11/h ×Falcon、PaLM 部分层

Decode 阶段每生成一个 Token,都要把新 Token 的 K、V 追加进缓存,并在下一步让新 Query 与全部历史 K/V 做注意力。这就是 KV Cache 的全部由来——它把 O(L²) 的重复计算压成 O(L) 的增量计算,代价是线性增长的显存。


四、算例:Prefill 8192 + Decode 1024

4.1 约定与口径

符号含义取值
PPrefill 阶段 Token 数(Prompt 长度)8192
DDecode 阶段新生成的 Token 数1024
H_kvKV 头数见下方模型表
d_head单头维度128

一处需要澄清的口径:原始推导中出现了 "Decode 1025 个 Token" 与公式里的 1024 并存。两者相差 1,通常源于"是否把首个生成 Token 单独计数"或"是否多算了一个结束符"。本文统一采用 D = 1024,即新生成 1024 个 Token,序列总长 8192 + 1024 = 9216。若按 1025 计算,存储结果只会多出 1 个 Token 的量(约 0.13 MiB),不影响任何结论。

4.2 计算量:Decode 阶段的 Q·Kᵀ

Decode 生成第 i 个 Token(i 从 1 开始)时,序列长度为 P + i,需要完成 P + i 次 Query-Key 点积。对 i = 1..D 求和:

总点积次数 = Σ(i=1..D) (P + i)
           = D·P + D·(D+1)/2
           = 1024 × 8192 + 1024 × 1025 / 2
           = 8,388,608 + 524,800
           = 8,913,408   (单头、单样本)

原始推导的一处偏差:原文写作 8192 × 1024 + 1024 × 1023 / 2,即 Σ(i=1..D) (P + i - 1),相当于漏掉了当前 Token 自身的 K/V。因果注意力是包含对角线的——当前 Token 必须能看到自己——正确项应为 D(D+1)/2 而非 D(D-1)/2。两者相差恰好 D = 1024 次,占比约 0.01%,量级上影响不大,但口径必须是自洽的,否则在推导更复杂的分块公式时会连锁出错。

换算成真实 FLOPs 还需乘三个系数:每次长度为 d_head 的点积约 2·d_head 次浮点运算(一次乘法一次加法),再乘头数 H_q、层数 N_layers 和批大小 B。

对照一下 Prefill:因果掩码下只需算下三角,点积次数为 P²/2 = 8192²/2 = 33,554,432,是 Decode 的约 3.8 倍——但 Prefill 是高度并行的矩阵乘,实际耗时远低于 Decode。这就是"Prefill 算得多、Decode 跑得慢"的直观来源。

4.3 存储量:KV Cache 占多少显存

单个 Token 的 KV Cache 字节数:

per_token = 2 × N_layers × H_kv × d_head × dtype_bytes
            └ K 和 V 两份

总占用 = per_token × (P + D)。以 FP16(dtype_bytes = 2)为例:

模型层数KV 头数每 Token KV Cache9216 Token 总占用
Llama-3-8B328(GQA)128 KiB1.125 GiB
Llama-2-7B3232(MHA)512 KiB4.5 GiB
Qwen2-7B284(GQA)56 KiB0.49 GiB

三个模型参数量相近,KV Cache 却相差近 10 倍——决定 KV Cache 的是 N_layers × H_kv × d_head,而不是参数量。这也是 GQA 能以极低精度代价换来巨大显存收益的原因。

4.4 注意事项

  1. KV Cache 与激活值是两笔账:前者随序列长度线性增长且全程驻留;后者只在 Prefill 峰值出现。估算显存时不能混算。
  2. 上述只是单层单头的相对口径:落地到具体模型时务必乘上 N_layers、H_kv、dtype_bytes;换成 FP8 / INT8 KV Cache 可直接减半或减到 1/4。
  3. 多用户并发时 KV Cache 才是主瓶颈:单条 9216 Token 请求占 1.125 GiB,若并发 64 路就是 72 GiB,远超模型权重本身。这也是 PagedAttention、Prefix Caching 存在的理由。
  4. 计算量公式只统计了 Q·Kᵀ:完整的 Attention 还有 Attn·V(同量级)以及 FFN(约 2 × d_model × d_ff per Token,通常比 Attention 更大)。本文口径与原始推导一致,仅用于横向对比 Decode 内部的 KV 增长。

五、总结

  1. 一次注意力的完整链路是:QKV 投影 → Q·Kᵀ/√d_k → 因果掩码 → Softmax → Attn·V → 输出投影 → 残差与 LayerNorm → FFN(升维—激活—降维)→ 残差与 LayerNorm。其中 /√d_k 用于把 logits 方差拉回 1、避免 Softmax 饱和;残差与 LayerNorm 是最容易被漏掉但结构必需的环节。
  2. Decode 阶段 Q·Kᵀ 点积次数的闭式解为 D·P + D·(D+1)/2,代入 P=8192, D=1024 得 8,913,408(单头单样本)。需要注意原推导的 D(D-1)/2 漏算了当前 Token 自身的 K/V。
  3. KV Cache 容量由 2 × N_layers × H_kv × d_head × dtype_bytes 决定,与模型参数量无直接关系。Llama-3-8B 在 9216 Token 下约占 1.125 GiB,而同为 7B 量级的 Llama-2-7B 因使用 MHA 需要 4.5 GiB。
  4. 优化方向因此非常明确:降 H_kv(GQA/MQA)、降 dtype_bytes(FP8/INT8 量化)、降重复前缀(Prefix Caching)、降碎片(PagedAttention)。

参考资料

  1. Vaswani A, et al. Attention Is All You Need. NeurIPS 2017.(缩放点积注意力与多头机制的原始定义)
  2. Ainslie J, et al. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. EMNLP 2023.(GQA 的提出与显存收益分析)
  3. Kwon W, et al. Efficient Memory Management for LLM Serving with PagedAttention. SOSP 2023.(KV Cache 显存管理与碎片问题)
  4. Dao T, et al. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. NeurIPS 2022.(IO 感知的注意力实现)
  5. Meta. Llama 3 Model Card, 2024.(32 层、GQA 8 头、d_head=128 的结构参数)