一次注意力计算到底长什么样?------从 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)。随后与三个权重矩阵相乘,投影到语义空间:
ini
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 缩放:
scss
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 的注意力权重:
ini
Attn = Softmax(Q·Kᵀ / √d_k)
在自回归场景下必须先施加因果掩码(Causal Mask) :第 i 个 Query 只能看到位置 0..i 的 Key,未来的位置屏蔽为 -∞。掩码方式通常是加性掩码(屏蔽位加 -1e9 或直接置 -inf)后再做 Softmax,而不是 Softmax 之后再置零------后者会破坏归一化。
需要注意:Prefill 阶段序列长度大于 1,必须显式加掩码;而 Decode 每步只有一个新 Query,它能看到的是全部历史 Key,天然满足因果性,无需额外掩码。
2.4 加权求和
ini
x_out = Attn · V
这一步才是真正的信息提取:注意力权重只是"配比",内容全部来自 V。因此业界常说"KV Cache 缓存的是内容,而不是分数"。
2.5 输出投影、残差与 LayerNorm
一次 Attention 到这里还没结束,还差三步(原始推导常漏掉这一段):
- 输出投影 :
AttnOut · W_O,把多头拼接结果映射回d_model;(这里的w_o也是直接由Token和W_o权重矩阵计算直接可以得到的)。 - 残差连接 :
h = x + AttnOut · W_O; - LayerNorm :对
h做归一化,稳定后续 FFN 的输入分布。
2.6 FFN:升维 → 激活 → 降维
scss
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(相对量) | 代表模型 |
|---|---|---|---|---|
| MHA | h |
h |
1× |
Llama-2-7B |
| GQA | h |
h / g(分组共享) |
1/g × |
Llama-3-8B、Qwen2 |
| MQA | h |
1 |
1/h × |
Falcon、PaLM 部分层 |
Decode 阶段每生成一个 Token,都要把新 Token 的 K、V 追加进缓存,并在下一步让新 Query 与全部历史 K/V 做注意力。这就是 KV Cache 的全部由来------它把 O(L²) 的重复计算压成 O(L) 的增量计算,代价是线性增长的显存。
四、算例:Prefill 8192 + Decode 1024
4.1 约定与口径
| 符号 | 含义 | 取值 |
|---|---|---|
P |
Prefill 阶段 Token 数(Prompt 长度) | 8192 |
D |
Decode 阶段新生成的 Token 数 | 1024 |
H_kv |
KV 头数 | 见下方模型表 |
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 求和:
ini
总点积次数 = Σ(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 字节数:
ini
per_token = 2 × N_layers × H_kv × d_head × dtype_bytes
└ K 和 V 两份
总占用 = per_token × (P + D)。以 FP16(dtype_bytes = 2)为例:
| 模型 | 层数 | KV 头数 | 每 Token KV Cache | 9216 Token 总占用 |
|---|---|---|---|---|
| Llama-3-8B | 32 | 8(GQA) | 128 KiB | 1.125 GiB |
| Llama-2-7B | 32 | 32(MHA) | 512 KiB | 4.5 GiB |
| Qwen2-7B | 28 | 4(GQA) | 56 KiB | 0.49 GiB |
三个模型参数量相近,KV Cache 却相差近 10 倍------决定 KV Cache 的是 N_layers × H_kv × d_head,而不是参数量。这也是 GQA 能以极低精度代价换来巨大显存收益的原因。
4.4 注意事项
- KV Cache 与激活值是两笔账:前者随序列长度线性增长且全程驻留;后者只在 Prefill 峰值出现。估算显存时不能混算。
- 上述只是单层单头的相对口径 :落地到具体模型时务必乘上
N_layers、H_kv、dtype_bytes;换成 FP8 / INT8 KV Cache 可直接减半或减到 1/4。 - 多用户并发时 KV Cache 才是主瓶颈:单条 9216 Token 请求占 1.125 GiB,若并发 64 路就是 72 GiB,远超模型权重本身。这也是 PagedAttention、Prefix Caching 存在的理由。
- 计算量公式只统计了 Q·Kᵀ :完整的 Attention 还有
Attn·V(同量级)以及 FFN(约2 × d_model × d_ffper Token,通常比 Attention 更大)。本文口径与原始推导一致,仅用于横向对比 Decode 内部的 KV 增长。
五、总结
- 一次注意力的完整链路 是:QKV 投影 →
Q·Kᵀ/√d_k→ 因果掩码 → Softmax →Attn·V→ 输出投影 → 残差与 LayerNorm → FFN(升维---激活---降维)→ 残差与 LayerNorm。其中/√d_k用于把 logits 方差拉回 1、避免 Softmax 饱和;残差与 LayerNorm 是最容易被漏掉但结构必需的环节。 - Decode 阶段 Q·Kᵀ 点积次数 的闭式解为
D·P + D·(D+1)/2,代入P=8192, D=1024得8,913,408(单头单样本)。需要注意原推导的D(D-1)/2漏算了当前 Token 自身的 K/V。 - 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。 - 优化方向 因此非常明确:降
H_kv(GQA/MQA)、降dtype_bytes(FP8/INT8 量化)、降重复前缀(Prefix Caching)、降碎片(PagedAttention)。
参考资料
- Vaswani A, et al. Attention Is All You Need. NeurIPS 2017.(缩放点积注意力与多头机制的原始定义)
- Ainslie J, et al. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. EMNLP 2023.(GQA 的提出与显存收益分析)
- Kwon W, et al. Efficient Memory Management for LLM Serving with PagedAttention. SOSP 2023.(KV Cache 显存管理与碎片问题)
- Dao T, et al. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. NeurIPS 2022.(IO 感知的注意力实现)
- Meta. Llama 3 Model Card , 2024.(32 层、GQA 8 头、
d_head=128的结构参数)