ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

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

一次注意力计算到底长什么样?——从 QKV 投影到 KV Cache 的完整拆解 一次注意力计算到底长什么样——从 QKV 投影到 KV Cache 的完整拆解摘要本文从单个 Token 的输入出发逐段拆解 Transformer 中一次注意力计算的完整链路QKV 投影、缩放点积打分、因果掩码、Softmax 归一、加权求和以及 FFN 的升维—激活—降维。随后以「Prefill 8192 Token Decode 1024 Token」为算例给出注意力计算量与 KV Cache 显存占用的闭式公式并指出原始推导中容易算错的两处细节。关键词Transformer注意力机制QKV 投影KV CachePrefill / 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 到这里还没结束还差三步原始推导常漏掉这一段输出投影AttnOut · W_O把多头拼接结果映射回d_model这里的w_o也是直接由Token和W_o权重矩阵计算直接可以得到的。残差连接h x AttnOut · W_OLayerNorm对h做归一化稳定后续 FFN 的输入分布。2.6 FFN升维 → 激活 → 降维FFN(h) W_2 · act(W_1 · h)W_1把维度从d_model升到d_ff经典设置d_ff 4·d_modelSwiGLU 结构约为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-7BGQAhh / g分组共享1/g ×Llama-3-8B、Qwen2MQAh11/h ×Falcon、PaLM 部分层Decode 阶段每生成一个 Token都要把新 Token 的K、V追加进缓存并在下一步让新 Query 与全部历史 K/V做注意力。这就是 KV Cache 的全部由来——它把O(L²)的重复计算压成O(L)的增量计算代价是线性增长的显存。四、算例Prefill 8192 Decode 10244.1 约定与口径符号含义取值PPrefill 阶段 Token 数Prompt 长度8192DDecode 阶段新生成的 Token 数1024H_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个 Tokeni从 1 开始时序列长度为P i需要完成P i次 Query-Key 点积。对i 1..D求和总点积次数 Σ(i1..D) (P i) D·P D·(D1)/2 1024 × 8192 1024 × 1025 / 2 8,388,608 524,800 8,913,408 单头、单样本原始推导的一处偏差原文写作8192 × 1024 1024 × 1023 / 2即Σ(i1..D) (P i - 1)相当于漏掉了当前 Token 自身的 K/V。因果注意力是包含对角线的——当前 Token 必须能看到自己——正确项应为D(D1)/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)。以 FP16dtype_bytes 2为例模型层数KV 头数每 Token KV Cache9216 Token 总占用Llama-3-8B328GQA128 KiB1.125 GiBLlama-2-7B3232MHA512 KiB4.5 GiBQwen2-7B284GQA56 KiB0.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·(D1)/2代入P8192, D1024得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_kvGQA/MQA、降dtype_bytesFP8/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_head128的结构参数
返回列表