大模型基础知识与 KV Cache 详解

目录
- [Activation 与 Weight 的关系](#Activation 与 Weight 的关系)
- [Transformer 架构基础](#Transformer 架构基础)
- [Self-Attention 机制](#Self-Attention 机制)
- 大模型训练流程
- 大模型推理流程
- [KV Cache 核心概念](#KV Cache 核心概念)
- [KV Cache 显存计算](#KV Cache 显存计算)
- [KV Cache 优化技术](#KV Cache 优化技术)
- 8.1 [GQA / MQA:减少 KV 头数量](#GQA / MQA:减少 KV 头数量)
- 8.2 PagedAttention (vLLM)
- 8.3 [KV Cache 量化](#KV Cache 量化)
- 8.4 [Sliding Window Attention](#Sliding Window Attention)
- 8.5 [KV Cache 卸载与压缩](#KV Cache 卸载与压缩)
- 8.6 其他优化方向
- 推理加速技术
- 关键论文与参考资料
1. Activation 与 Weight 的关系
在深度学习中,Weight(权重)是模型学习到的计算规则,Activation(激活值)是输入经过这些规则计算后产生的动态中间结果。基本关系为:
z = W x + b , a = f ( z ) z = Wx + b, \qquad a = f(z) z=Wx+b,a=f(z)
其中, x x x 是上一层的 Activation, W W W 和 b b b 是当前层的 Weight 与偏置, f f f 是激活函数, a a a 是当前层输出的 Activation。因此,上一层的 Activation 会作为下一层的输入,逐层向前传播。
| 对比项 | Weight | Activation |
|---|---|---|
| 本质 | 模型训练得到的参数 | 输入在网络中产生的中间数据 |
| 是否随输入变化 | 通常不变 | 随输入变化 |
| 生命周期 | 模型加载期间长期存在 | 通常仅在当前计算期间存在 |
| 规模主要取决于 | 参数量、数据类型 | Batch Size、序列长度、隐藏维度、数据类型 |
| Transformer 示例 | W Q W_Q WQ、 W K W_K WK、 W V W_V WV、MLP 权重 | Hidden States、Q、K、V、Logits |
在 Transformer 中,输入 Activation X X X 与权重矩阵共同生成新的 Activation:
Q = X W Q , K = X W K , V = X W V Q = XW_Q, \qquad K = XW_K, \qquad V = XW_V Q=XWQ,K=XWK,V=XWV
- 前向传播:Weight 决定如何把输入 Activation 转换为输出 Activation。
- 反向传播:训练时需要利用前向传播保留的 Activation 计算梯度,再由优化器更新 Weight。这是训练显存占用通常显著高于推理的重要原因。
- 量化表示:例如 W8A8 表示 Weight 和 Activation 均使用 8 bit;W4A16 表示 Weight 使用 4 bit、Activation 使用 16 bit。
- 与 KV Cache 的关系:KV Cache 是缓存下来的 K、V Activation,不属于 Weight。它由当前输入和模型权重共同计算产生,并随会话内容及序列长度增长。
简而言之:Weight 是模型学到的规则,Activation 是数据按照这些规则流经网络时产生的状态;前向传播由 Weight 生成 Activation,反向传播再利用 Activation 更新 Weight。
2. Transformer 架构基础
当前主流大模型(GPT 系列、LLaMA、Qwen、Mistral 等)均采用 Decoder-Only Transformer 架构。
核心组件
| 组件 | 作用 | 说明 |
|---|---|---|
| Token Embedding | 将 token ID 映射为稠密向量 | 查表操作,维度 d_model |
| Positional Encoding | 注入位置信息 | RoPE(旋转位置编码)是当前主流 |
| Self-Attention | 捕捉 token 间依赖关系 | 模型的核心,计算量主要来源 |
| FFN (Feed-Forward Network) | 非线性变换,增加模型容量 | 通常占模型参数量的 2/3 |
| Layer Norm | 稳定训练 | Pre-Norm 架构(先归一化再计算) |
| Residual Connection | 缓解梯度消失 | 每个子层输出 = 子层输入 + 子层(子层输入) |
结构概览
输入 Token IDs
↓
[Token Embedding + Positional Encoding] ← 位置编码注入
↓
┌─────────────────────────────────┐
│ Transformer Block × N layers │
│ ┌─────────────────────────┐ │
│ │ LayerNorm │ │
│ │ Masked Self-Attention │ ← 核心子层 1
│ │ + Residual │ │
│ ├─────────────────────────┤ │
│ │ LayerNorm │ │
│ │ Feed-Forward Network │ ← 核心子层 2
│ │ + Residual │ │
│ └─────────────────────────┘ │
└─────────────────────────────────┘
↓
[LayerNorm → Linear (vocab projection)]
↓
Logits → Softmax → Next Token Probability
关键超参数
| 参数 | 含义 | 典型值 (LLaMA-3 8B) |
|---|---|---|
| d_model | 隐藏层维度 | 4096 |
| n_heads | 注意力头数 | 32 |
| d_head | 每个头的维度 | 128 |
| n_layers | Transformer 层数 | 32 |
| d_ff | FFN 中间层维度 | 14336 |
| vocab_size | 词表大小 | 128256 |
3. Self-Attention 机制
核心公式
Attention(Q, K, V) = softmax(Q · K^T / √d_k) · V
Q / K / V 的含义
| 矩阵 | 含义 | 直觉类比 |
|---|---|---|
| Q (Query) | 当前 token 在"提问" | "我需要关注什么信息?" |
| K (Key) | 每个 token 在"应答" | "我有什么信息可以提供?" |
| V (Value) | 每个 token 的实际内容 | "这是我的具体内容。" |
计算流程
-
输入 X 经过三个线性投影得到 Q、K、V:
Q = X · W_Q,形状 (seq_len, d_k)K = X · W_K,形状 (seq_len, d_k)V = X · W_V,形状 (seq_len, d_v)
-
计算注意力分数:
scores = Q · K^T / √d_k,形状 (seq_len, seq_len) -
应用因果掩码(Causal Mask):将上三角设为 -∞,确保每个位置只能看到前面的 token
-
Softmax 归一化:
weights = softmax(scores),每行求和为 1 -
加权求和:
output = weights · V,形状 (seq_len, d_v)
多头注意力 (Multi-Head Attention)
将 d_model 维度拆分为 n_heads 个子空间,每个头独立计算 Attention,最后拼接:
MultiHead(Q, K, V) = Concat(head_1, ..., head_h) · W_O
其中 head_i = Attention(Q·W_Q_i, K·W_K_i, V·W_V_i)
- 每个头的维度:
d_head = d_model / n_heads - 不同头学习不同子空间的注意力模式
RoPE(旋转位置编码)
当前主流模型(LLaMA、Qwen 等)使用 RoPE 替代绝对位置编码:
- 通过对 Q、K 施加旋转矩阵来注入相对位置信息
- 优势:天然支持外推(扩展上下文长度)、相对位置感知
4. 大模型训练流程
三阶段训练范式
阶段 1: 预训练 (Pre-training)
↓ 大规模无监督语料,Next Token Prediction
阶段 2: 指令微调 (SFT - Supervised Fine-Tuning)
↓ 高质量指令-回答对,学习对话格式
阶段 3: 对齐训练 (RLHF / DPO)
↓ 人类偏好对齐,提升安全性和有用性
各阶段详解
| 阶段 | 目标 | 数据 | 方法 | 数据量 |
|---|---|---|---|---|
| 预训练 | 学习语言知识和世界知识 | 网页、书籍、代码等 | 自回归 Next Token Prediction | 万亿 token 级 |
| SFT | 学习遵循指令 | 人工编写的指令-回答对 | 监督学习 | 10K-1M 条 |
| RLHF/DPO | 对齐人类偏好 | 人类偏好标注数据 | 强化学习 / 直接偏好优化 | 100K-1M 条 |
关键概念
- Loss 函数:Cross-Entropy Loss(预训练和 SFT)
- RLHF (Reinforcement Learning from Human Feedback) :
- 训练奖励模型 (Reward Model)
- 用 PPO 等算法优化策略模型
- DPO (Direct Preference Optimization):无需训练奖励模型,直接从偏好数据优化,更简单稳定
5. 大模型推理流程
两个关键阶段
大模型推理分为 Prefill 和 Decode 两个阶段,这是理解 KV Cache 的前提:
Prefill 阶段(预填充)
- 输入:用户完整的 prompt(如 "解释什么是 KV Cache")
- 操作:一次性处理所有 prompt token,计算并缓存它们的 K、V
- 特点:计算密集型(compute-bound),可并行处理所有 token
- 输出:第一个生成 token + 完整的 KV Cache
Decode 阶段(解码)
-
输入:上一步生成的单个 token
-
操作:只计算新 token 的 Q、K、V,将新 K、V 追加到 Cache,用新 Q 与全部 Cache K 计算注意力
-
特点:内存密集型(memory-bound),每次只处理 1 个 token,GPU 利用率低
-
循环:重复直到生成结束符 (EOS) 或达到最大长度
Prefill: [t1, t2, t3, t4] → 计算所有 K,V → 生成 t5
↓ 缓存 K1-K4, V1-V4
Decode: [t5] → 计算 Q5,K5,V5 → 追加 Cache → 生成 t6
↓ 缓存 K5, V5
Decode: [t6] → 计算 Q6,K6,V6 → 追加 Cache → 生成 t7
↓ ...
推理的性能瓶颈
| 阶段 | 瓶颈类型 | 原因 | 优化方向 |
|---|---|---|---|
| Prefill | 计算密集 (compute-bound) | 大量矩阵乘法 | FlashAttention、Tensor Core 优化 |
| Decode | 内存密集 (memory-bound) | 每步加载全部 KV Cache 但只算 1 个 token | KV Cache 优化、 batching、Speculative Decoding |
关键观察 :Decode 阶段的瓶颈在于------每生成一个 token,都要将整个 KV Cache 从 GPU HBM 读到 SRAM,但只做极少量计算。这就是 KV Cache 优化如此重要 的根本原因。
6. KV Cache 核心概念
什么是 KV Cache
KV Cache 是大模型推理时的一项基础优化:在自回归生成过程中,缓存已计算 token 的 Key 和 Value 矩阵,避免重复计算。
为什么需要 KV Cache
在自回归生成中,每生成一个新 token,Self-Attention 需要计算新 token 与所有历史 token 的注意力。如果不缓存:
- 生成第 n 个 token 时,需要重新计算前 n-1 个 token 的 K 和 V
- 总计算复杂度:O(n³)
- 序列越长,浪费的计算越多
使用 KV Cache 后:
- 历史 token 的 K、V 已缓存,只需计算新 token 的 Q、K、V
- 总计算复杂度:O(n²)
- 降低了整整一个量级
KV Cache 的工作流程
Step 1: 输入 prompt [A, B, C]
→ 计算 A, B, C 的 K, V
→ 存入 KV Cache: [K_A, K_B, K_C] [V_A, V_B, V_C]
→ 用 Q_C 与 Cache K 计算注意力 → 生成 D
Step 2: 输入新 token [D]
→ 只计算 D 的 K_D, V_D
→ 追加到 Cache: [K_A..K_D] [V_A..V_D]
→ 用 Q_D 与全部 Cache K 计算注意力 → 生成 E
Step 3: 输入新 token [E]
→ 只计算 E 的 K_E, V_E
→ 追加到 Cache: [K_A..K_E] [V_A..V_E]
→ 用 Q_E 与全部 Cache K 计算注意力 → 生成 F
...
本质:用显存换计算
KV Cache 是典型的 space-time tradeoff:
- 代价:消耗 GPU 显存存储历史 K、V
- 收益:避免重复计算,大幅提升推理速度
7. KV Cache 显存计算
计算公式
KV Cache 的显存占用取决于模型配置和序列长度:
KV Cache 大小 = 2 × n_layers × seq_len × n_kv_heads × d_head × precision_bytes
其中:
2 → K 和 V 两个矩阵
n_layers → Transformer 层数
seq_len → 序列长度(prompt + 已生成 token)
n_kv_heads → KV 头的数量(MHA = n_heads,GQA < n_heads)
d_head → 每个头的维度
precision → FP16=2, FP32=4, INT8=1, FP8=1
具体示例
LLaMA-2 7B(MHA,32 头)
参数: n_layers=32, n_heads=32, d_head=128, seq_len=4096
KV Cache = 2 × 32 × 4096 × 32 × 128 × 2 bytes
= 2,147,483,648 bytes
≈ 2.0 GB (FP16, batch_size=1)
LLaMA-2 70B(GQA,8 KV 头)
参数: n_layers=80, n_kv_heads=8, d_head=128, seq_len=4096
MHA Cache = 2 × 80 × 4096 × 64 × 128 × 2 ≈ 10.7 GB (如果用 MHA)
GQA Cache = 2 × 80 × 4096 × 8 × 128 × 2 ≈ 1.3 GB (实际使用 GQA)
KV Cache vs 模型权重的显存占比
| 模型 | 权重大小 (FP16) | KV Cache (4K ctx) | KV Cache (32K ctx) | KV Cache (128K ctx) |
|---|---|---|---|---|
| LLaMA-7B (MHA) | 14 GB | 1.0 GB | 8.0 GB | 32 GB |
| LLaMA-70B (MHA) | 140 GB | 5.0 GB | 40 GB | 160 GB |
| LLaMA-70B (GQA-8) | 140 GB | 0.6 GB | 5.0 GB | 20 GB |
关键结论:
- 长上下文场景下,KV Cache 的显存可远超模型权重
- GQA 可将 KV Cache 压缩 8-16 倍,是长上下文模型的基础
- batch_size > 1 时,KV Cache 按倍数线性增长
8. KV Cache 优化技术
8.1 GQA / MQA:减少 KV 头数量
这是从模型架构层面减少 KV Cache 大小的方案,在训练时就确定了。
三种方案对比
| 方案 | KV 头数量 | Cache 大小 | 模型质量 | 代表模型 |
|---|---|---|---|---|
| MHA (Multi-Head Attention) | = Q 头数 | 1.0x (基准) | 最优 | GPT-2, BERT, LLaMA-2 7B |
| MQA (Multi-Query Attention) | 1 | 1/h x (最小) | 有损 | PaLM, Falcon, StarCoder |
| GQA (Grouped-Query Attention) | 1 < g < h | ~g/h x (折中) | 接近 MHA | LLaMA-2 70B, LLaMA-3, Qwen-2 |
原理
- MHA:每个 Q 头都有独立的 K 头和 V 头,KV Cache 最大
- MQA:所有 Q 头共享 1 个 K 头和 1 个 V 头,KV Cache 最小,但质量损失较大
- GQA:将 Q 头分成 g 组,每组共享 1 个 K/V 头。当 g = h 时退化为 MHA,当 g = 1 时退化为 MQA
为什么 GQA 成为主流
GQA 论文 (Ainslie et al., 2023) 表明:
- GQA 在质量上接近 MHA
- 在推理速度和显存占用上接近 MQA
- 是质量和效率的最佳平衡点
当前主流模型几乎全部采用 GQA:LLaMA-3、Qwen-2、Mistral、DeepSeek-V2 等。
8.2 PagedAttention (vLLM)
PagedAttention 是 vLLM 框架提出的 KV Cache 内存管理优化方案,灵感来自操作系统的虚拟内存分页机制。
解决的问题
传统 KV Cache 分配方式的痛点:
- 内部碎片:预分配最大序列长度的连续内存,但大多数请求用不满 → 浪费 60-80% 显存
- 外部碎片:请求结束后释放内存,留下不连续的空闲块 → 无法用于新请求
- 无法共享:不同请求有相同前缀(如系统提示词)时,KV Cache 无法复用
PagedAttention 方案
传统方式:
┌──────────────────────────────────┐
│ Request 1: [████████░░░░░░░░░░░░] │ ← 预分配 max_len,大量浪费
│ Request 2: [████░░░░░░░░░░░░░░░░] │
└──────────────────────────────────┘
PagedAttention:
┌─────┐ ┌─────┐ ┌─────┐ ┌─────┐
│Blk 0│ │Blk 1│ │Blk 2│ │Blk 3│ ← 固定大小的 Block(如 16 token/block)
│R1: │ │R1: │ │R2: │ │R1: │
│tok0 │ │tok1 │ │tok0 │ │tok2 │
│tok15│ │tok16│ │tok15│ │tok17│
└─────┘ └─────┘ └─────┘ └─────┘
R1 → [Block 0] → [Block 1] → [Block 3] ← 链表组织,按需分配
R2 → [Block 2] ← 不浪费空间
核心优势
| 优势 | 说明 |
|---|---|
| 消除内部碎片 | Block 是最小分配单元,最多浪费不到 1 个 Block |
| 消除外部碎片 | Block 大小固定,可任意分配给任何请求 |
| 支持共享前缀 | 相同前缀的请求共享 Block,通过引用计数管理 |
| 提升吞吐量 | 显存利用率从 ~20% 提升到 ~96%,吞吐量提升 2-4 倍 |
Block 大小选择
- 通常 16 个 token/block
- 太小:Block 表(Block Table)开销大
- 太大:内部碎片增加
8.3 KV Cache 量化
将 KV Cache 从 FP16 量化到低精度,直接减少显存占用。
量化方案
| 方案 | 精度 | 压缩比 | 质量影响 | 适用场景 |
|---|---|---|---|---|
| FP16 | 16-bit | 1.0x | 无 | 基准 |
| FP8 | 8-bit | 2.0x | 几乎无损 | H100+ GPU |
| INT8 | 8-bit | 2.0x | 轻微下降 | 通用 |
| INT4 | 4-bit | 4.0x | 有一定下降 | 极限压缩 |
| KV Cache FP8 (H100) | 8-bit | 2.0x | 几乎无损 | H100 原生支持 |
关键技术
-
KVCache Quantization (论文: KIVI, Hooper et al., 2024):
- Key 按通道量化 (per-channel),Value 按 token 量化 (per-token)
- INT4 量化下仍能保持接近 FP16 的质量
-
FP8 KV Cache:NVIDIA H100 原生支持 FP8,几乎无损压缩 2 倍
8.4 Sliding Window Attention
限制每个 token 只关注最近的 W 个 token,KV Cache 只保留最近 W 个位置的 K/V。
原理
标准 Attention: token_i 关注 token_0, token_1, ..., token_i (全部历史)
Sliding Window: token_i 只关注 token_{i-W+1}, ..., token_i (最近 W 个)
效果
- KV Cache 大小固定为 W,不随序列长度增长
- 适合长文本场景
代表模型
- Mistral 7B:W = 4096
- Gemma 2:使用 Sliding Window + Full Attention 交替
局限
- 丢失长距离依赖信息
- 通常与其他注意力模式交替使用(如每 N 层用一次 Full Attention)
8.5 KV Cache 卸载与压缩
KV Cache Offloading(CPU 卸载)
- 将不活跃的 KV Cache 从 GPU 显存卸载到 CPU 内存
- 需要时再按需加载回 GPU
- 代价:增加 PCIe 传输延迟
- 适用场景:超长上下文、显存不足时
KV Cache 驱逐 (Eviction)
- 基于注意力分数判断哪些 token 的 KV Cache "不重要",予以丢弃
- 策略:保留注意力分数高的 token,丢弃分数低的
- 代表方法:H2O (Heavy-Hitter Oracle)、StreamingLLM
StreamingLLM
- 保留 attention sink(开头的几个 token)+ 最近的滑动窗口
- 发现:丢弃开头的 token 会导致注意力分数爆炸,必须保留
- 可以让模型在百万级 token 上稳定生成
8.6 其他优化方向
Sparse Attention(稀疏注意力)
- 不是所有 token 都需要关注所有历史 token
- 模式:局部窗口 + 全局 token + 随机连接
- 代表:Longformer, BigBird, MoBA (DeepSeek)
Multi-Token Attention 变体
- MLA (Multi-head Latent Attention) :DeepSeek-V2 提出
- 将 KV 压缩到低秩潜空间
- KV Cache 可压缩 93.3%(相比 MHA)
- 质量优于 MHA
Prefix Caching(前缀缓存)
- 多个请求共享相同前缀(如系统提示词)时,复用前缀的 KV Cache
- vLLM、SGLang 等框架已支持
- 大幅降低重复计算
9. 推理加速技术
FlashAttention
不是 KV Cache 优化,而是 Attention 计算优化,但与 KV Cache 密切相关。
核心思想
标准 Attention 的瓶颈在于:Q×K^T 产生 (seq_len × seq_len) 的中间矩阵,需要反复在 HBM 和 SRAM 之间搬运。
FlashAttention 通过 tiling(分块) 和 online softmax 技术:
- 将 Q、K、V 分块加载到 SRAM
- 在 SRAM 内完成计算,避免中间矩阵写回 HBM
- 减少 HBM 读写量,提升计算效率
版本演进
| 版本 | 关键改进 | 加速比 |
|---|---|---|
| FlashAttention v1 | 分块 + online softmax | 2-4x |
| FlashAttention v2 | 更好的并行度、减少非矩阵乘法 | 2x over v1 |
| FlashAttention v3 | FP8 支持、异步化 | 1.5-2x over v2 (H100) |
Speculative Decoding(投机解码)
核心思想
用一个小模型(draft model)快速生成多个候选 token,再用大模型一次性验证,减少大模型的前向传播次数。
流程
1. 小模型快速生成 k 个候选 token: [t1, t2, t3]
2. 大模型一次前向传播验证这 k 个 token
3. 接受正确的 token,从第一个错误处重新生成
4. 如果全部正确,一次前向传播就生成了 k 个 token
优势
- Decode 阶段从每次生成 1 个 token → 可生成多个 token
- 在不损失质量的情况下提升 2-3 倍速度
- 代表实现:Medusa, EAGLE, Lookahead Decoding
Continuous Batching(连续批处理)
解决的问题
静态批处理中,同一 batch 内不同请求长度不同,短请求完成后要等长请求,GPU 空转。
方案
-
每完成一个请求就立即从等待队列中插入新请求
-
batch 大小动态变化,GPU 始终保持满载
-
配合 PagedAttention 实现高效内存管理
静态批处理:
Time →
Req1: [████████████████] 完成
Req2: [████████] 等待... 等待...
Req3: [████████████████████████] 完成Continuous Batching:
Req1: [████████████████] 完成
Req2: [████████] 完成
Req3: [████████████████████████] 完成
Req4: [████████████████] 完成 ← Req2完成后立即插入
Req5: [████████████████] ← Req1完成后立即插入
10. 关键论文与参考资料
基础架构
| 论文 | 年份 | 贡献 |
|---|---|---|
| Attention Is All You Need | 2017 | Transformer 架构 |
| RoFormer: Enhanced Transformer with Rotary Position Embedding | 2021 | RoPE 旋转位置编码 |
| RMSNorm | 2019 | 替代 LayerNorm,更高效 |
| SwiGLU | 2020 | GLU 激活函数变体,LLaMA 采用 |
KV Cache 优化
| 论文 | 年份 | 贡献 |
|---|---|---|
| Fast Transformer Decoding (Noam Shazeer) | 2019 | Multi-Query Attention (MQA) |
| GQA: Training Generalized Multi-Query Transformer Models | 2023 | Grouped-Query Attention (GQA) |
| Efficient Memory Management for LLMs with PagedAttention (vLLM) | 2023 | PagedAttention,分页式 KV Cache 管理 |
| DeepSeek-V2 Technical Report | 2024 | MLA (Multi-head Latent Attention) |
| KIVI: A Tuning-Free Asymmetric 2bit Quantization for KV Cache | 2024 | INT4 KV Cache 量化 |
| StreamingLLM | 2023 | Attention Sink + 滑动窗口,无限长度生成 |
| H2O: Heavy-Hitter Oracle for Efficient Generative Inference | 2023 | KV Cache 驱逐策略 |
推理加速
| 论文 | 年份 | 贡献 |
|---|---|---|
| FlashAttention | 2022 | IO 感知的注意力计算 |
| FlashAttention-2 | 2023 | 更好的并行度 |
| FlashAttention-3 | 2024 | FP8 + 异步化 (H100) |
| Medusa: Simple LLM Inference Acceleration | 2024 | Speculative Decoding |
| EAGLE: Speculative Sampling Requires Rethinking Feature Uncertainty | 2024 | 改进的投机解码 |
| Orca: A Distributed Serving System for Transformer-Based Generative Models | 2022 | Continuous Batching |
长上下文
| 论文 | 年份 | 贡献 |
|---|---|---|
| Longformer | 2020 | 局部+全局注意力模式 |
| Ring Attention | 2023 | 跨设备分布式长上下文注意力 |
| YaRN | 2023 | RoPE 外推方法,扩展上下文长度 |
| MoBA: Mixture of Block Attention | 2025 | DeepSeek 提出的混合块注意力 |
推理框架
| 框架 | 特点 |
|---|---|
| vLLM | PagedAttention + Continuous Batching,工业界最广泛使用 |
| SGLang | RadixAttention 前缀缓存,结构化生成优化 |
| TensorRT-LLM | NVIDIA 官方,深度优化,支持 FP8 |
| DeepSpeed-FastGen | 微软出品,Dynamic Splitfuse |
| MLC-LLM | 跨平台部署,支持移动端 |
附:KV Cache 优化技术速查表
┌────────────────────────────────────────────────────────────────┐
│ KV Cache 优化技术全景图 │
├────────────────────┬───────────────────────────────────────────┤
│ 优化维度 │ 具体技术 │
├────────────────────┼───────────────────────────────────────────┤
│ 模型架构 (训练时) │ GQA / MQA / MLA │
│ │ Sliding Window Attention │
│ │ Sparse Attention (Longformer, MoBA) │
├────────────────────┼───────────────────────────────────────────┤
│ 内存管理 (推理时) │ PagedAttention (vLLM) │
│ │ Prefix Caching │
│ │ Continuous Batching │
├────────────────────┼───────────────────────────────────────────┤
│ 精度压缩 │ FP8 / INT8 / INT4 量化 │
│ │ KIVI (非对称量化) │
├────────────────────┼───────────────────────────────────────────┤
│ Cache 驱逐/卸载 │ StreamingLLM (attention sink + window) │
│ │ H2O (heavy-hitter eviction) │
│ │ CPU Offloading │
├────────────────────┼───────────────────────────────────────────┤
│ 计算加速 │ FlashAttention v1/v2/v3 │
│ │ Speculative Decoding (Medusa, EAGLE) │
└────────────────────┴───────────────────────────────────────────┘
KV-Cache 详解
来源 :CSDN 博客 - 卡洛驰(2024-09-18 发布,2025-10-23 修改)
核心问题:什么是 KV Cache?KV Cache 在哪里使用?KV Cache 节省了 Self-Attention 层中哪部分的计算?
目录
- [什么是 KV Cache](#什么是 KV Cache)
- [KV Cache 在哪里使用](#KV Cache 在哪里使用)
- [KV Cache 节省了哪部分计算](#KV Cache 节省了哪部分计算)
- 总结
一、什么是 KV Cache
1.1 基本概念
在 自回归模型(autoregressive models) 中,模型逐个生成文本的每个 token,每次新的预测都依赖于之前的上下文:
- 预测第 1000 个 token → 需要前 999 个 token 的信息
- 预测第 1001 个 token → 需要前 999 + 第 1000 个 token 的信息
KV Cache 的作用:通过存储之前 K、V 的计算结果,在后续 token 生成时复用,从而避免重复计算。
1.2 工作机制
KV Cache 在自回归生成模型中充当一个 内存库 的角色:
| 步骤 | 操作 | 说明 |
|---|---|---|
| 1 | 计算当前 token 的 Q、K、V | 对输入进行线性变换得到三个向量 |
| 2 | 缓存 K、V | 将当前步的 K、V 追加到缓存中 |
| 3 | 复用历史 K、V | 从缓存中检索之前所有 token 的 K 和 V |
| 4 | 计算 Attention | 用当前 Q 与缓存中所有 K 做 attention,加权求和 V |
1.3 适用范围
| 模型类型 | 是否使用 KV Cache | 原因 |
|---|---|---|
| Decoder-Only(如 GPT) | ✅ 使用 | 自回归生成,causal mask |
| Encoder-Decoder(如 T5) | ✅ 解码部分使用 | decoder 是 causal 的 |
| Encoder-Only(如 BERT) | ❌ 不使用 | 非生成型,无自回归过程 |
关键 :KV Cache 只在 decoder 中使用 ,因为 decoder 是 causal 的(一个 token 的注意力只依赖于它前面的 token)。
1.4 为什么只存储 K 和 V,而不存储 Q?
核心原因:Decoder 模型是 causal 的,即一个 token 的注意力只依赖于它前面的 token。
Transformer 是 自回归模型 ,参照序列中所有先前输入来预测下一个输出。但由于它不像 RNN 那样逐个接收序列,而是 一次性接受整个序列 ,因此需要 look-ahead mask 来限制注意力范围:
| Mask 行为 | 说明 |
|---|---|
| 遮蔽右侧 token | 查询词右侧的所有 token 的注意力得分被置为极小值 |
| 仅关注左侧 + 自身 | 查询词只关注自身及序列中位于其左侧的所有 token |
| 仅用于 decoder 第一层 | look-ahead mask 只应用于每个解码器层的第一个注意力子层 |
为什么不缓存 Q? 因为 Q 只与当前 token 相关、每步都会重新计算,而且历史 Q 在自回归注意力中不会再用(已经计算过了),缓存它没有任何收益。
1.5 数学推导示例
原始 Q、K、V 矩阵
Q = 0.212 0.04 0.63 0.36 0.1 0.14 0.86 0.77 0.31 0.36 0.19 0.72 Q = \begin{bmatrix} 0.212 & 0.04 & 0.63 & 0.36 \\ 0.1 & 0.14 & 0.86 & 0.77 \\ 0.31 & 0.36 & 0.19 & 0.72 \end{bmatrix} Q= 0.2120.10.310.040.140.360.630.860.190.360.770.72
K = 0.31 0.84 0.963 0.57 0.45 0.94 0.73 0.58 0.36 0.83 0.1 0.38 K = \begin{bmatrix} 0.31 & 0.84 & 0.963 & 0.57 \\ 0.45 & 0.94 & 0.73 & 0.58 \\ 0.36 & 0.83 & 0.1 & 0.38 \end{bmatrix} K= 0.310.450.360.840.940.830.9630.730.10.570.580.38
V = 0.36 0.83 0.1 0.38 0.31 0.36 0.19 0.72 0.31 0.84 0.963 0.57 V = \begin{bmatrix} 0.36 & 0.83 & 0.1 & 0.38 \\ 0.31 & 0.36 & 0.19 & 0.72 \\ 0.31 & 0.84 & 0.963 & 0.57 \end{bmatrix} V= 0.360.310.310.830.360.840.10.190.9630.380.720.57
Step 1:计算注意力分数并应用 Mask
Q K T d k = 0.4556 0.4009 0.1547 0.7078 0.6255 0.2654 0.4959 0.5171 0.3515 \frac{QK^T}{\sqrt{d_k}} = \begin{bmatrix} 0.4556 & 0.4009 & 0.1547 \\ 0.7078 & 0.6255 & 0.2654 \\ 0.4959 & 0.5171 & 0.3515 \end{bmatrix} dk QKT= 0.45560.70780.49590.40090.62550.51710.15470.26540.3515
Mask 矩阵(上三角置为 − ∞ -\infty −∞):
M = 0 − 1 e 9 − 1 e 9 0 0 − 1 e 9 0 0 0 M = \begin{bmatrix} 0 & -1\text{e}9 & -1\text{e}9 \\ 0 & 0 & -1\text{e}9 \\ 0 & 0 & 0 \end{bmatrix} M= 000−1e900−1e9−1e90
相加后:
Q K T d k + M = 0.4556 − 1 e 9 − 1 e 9 0.7078 0.6255 − 1 e 9 0.4959 0.5171 0.3515 \frac{QK^T}{\sqrt{d_k}} + M = \begin{bmatrix} 0.4556 & -1\text{e}9 & -1\text{e}9 \\ 0.7078 & 0.6255 & -1\text{e}9 \\ 0.4959 & 0.5171 & 0.3515 \end{bmatrix} dk QKT+M= 0.45560.70780.4959−1e90.62550.5171−1e9−1e90.3515
Step 2:应用 Softmax
沿行应用 softmax,极小值( − 1 e 9 -1\text{e}9 −1e9)变为 0:
softmax ( Q K T d k + M ) = 1.0 0 0 0.5206 0.4794 0 0.3464 0.3538 0.2998 \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}} + M\right) = \begin{bmatrix} 1.0 & 0 & 0 \\ 0.5206 & 0.4794 & 0 \\ 0.3464 & 0.3538 & 0.2998 \end{bmatrix} softmax(dk QKT+M)= 1.00.52060.346400.47940.3538000.2998
Step 3:不存储 Q 的情况(仅存储 K、V)
分步的 Q 矩阵(每步只有一个 token 的 Q 参与计算):
Q 1 = 0.212 0.04 0.63 0.36 --- --- --- --- --- --- --- --- , Q 2 = --- --- --- --- 0.1 0.14 0.86 0.77 --- --- --- --- , Q 3 = --- --- --- --- --- --- --- --- 0.31 0.36 0.19 0.72 Q_1 = \begin{bmatrix} 0.212 & 0.04 & 0.63 & 0.36 \\ --- & --- & --- & --- \\ --- & --- & --- & --- \end{bmatrix}, \quad Q_2 = \begin{bmatrix} --- & --- & --- & --- \\ 0.1 & 0.14 & 0.86 & 0.77 \\ --- & --- & --- & --- \end{bmatrix}, \quad Q_3 = \begin{bmatrix} --- & --- & --- & --- \\ --- & --- & --- & --- \\ 0.31 & 0.36 & 0.19 & 0.72 \end{bmatrix} Q1= 0.212------0.04------0.63------0.36------ ,Q2= ---0.1------0.14------0.86------0.77--- ,Q3= ------0.31------0.36------0.19------0.72
逐步增长的 K 矩阵(缓存累积):
K 1 = 0.31 0.84 0.963 0.57 --- --- --- --- --- --- --- --- , K 2 = 0.31 0.84 0.963 0.57 0.45 0.94 0.73 0.58 --- --- --- --- , K 3 = 0.31 0.84 0.963 0.57 0.45 0.94 0.73 0.58 0.36 0.83 0.1 0.38 K_1 = \begin{bmatrix} 0.31 & 0.84 & 0.963 & 0.57 \\ --- & --- & --- & --- \\ --- & --- & --- & --- \end{bmatrix}, \quad K_2 = \begin{bmatrix} 0.31 & 0.84 & 0.963 & 0.57 \\ 0.45 & 0.94 & 0.73 & 0.58 \\ --- & --- & --- & --- \end{bmatrix}, \quad K_3 = \begin{bmatrix} 0.31 & 0.84 & 0.963 & 0.57 \\ 0.45 & 0.94 & 0.73 & 0.58 \\ 0.36 & 0.83 & 0.1 & 0.38 \end{bmatrix} K1= 0.31------0.84------0.963------0.57------ ,K2= 0.310.45---0.840.94---0.9630.73---0.570.58--- ,K3= 0.310.450.360.840.940.830.9630.730.10.570.580.38
分别计算注意力分数(只有下三角有值):
Q 1 K 1 T d k = 0.4556 --- --- --- --- --- --- --- --- \frac{Q_1 K_1^T}{\sqrt{d_k}} = \begin{bmatrix} 0.4556 & --- & --- \\ --- & --- & --- \\ --- & --- & --- \end{bmatrix} dk Q1K1T= 0.4556------------------------
Q 2 K 2 T d k = --- --- --- 0.7078 0.6255 --- --- --- --- \frac{Q_2 K_2^T}{\sqrt{d_k}} = \begin{bmatrix} --- & --- & --- \\ 0.7078 & 0.6255 & --- \\ --- & --- & --- \end{bmatrix} dk Q2K2T= ---0.7078------0.6255------------
Q 3 K 3 T d k = --- --- --- --- --- --- 0.4959 0.5171 0.3515 \frac{Q_3 K_3^T}{\sqrt{d_k}} = \begin{bmatrix} --- & --- & --- \\ --- & --- & --- \\ 0.4959 & 0.5171 & 0.3515 \end{bmatrix} dk Q3K3T= ------0.4959------0.5171------0.3515
相加后的结果:
Q 1 K 1 T d k + Q 2 K 2 T d k + Q 3 K 3 T d k = 0.4556 --- --- 0.7078 0.6255 --- 0.4959 0.5171 0.3515 \frac{Q_1 K_1^T}{\sqrt{d_k}} + \frac{Q_2 K_2^T}{\sqrt{d_k}} + \frac{Q_3 K_3^T}{\sqrt{d_k}} = \begin{bmatrix} 0.4556 & --- & --- \\ 0.7078 & 0.6255 & --- \\ 0.4959 & 0.5171 & 0.3515 \end{bmatrix} dk Q1K1T+dk Q2K2T+dk Q3K3T= 0.45560.70780.4959---0.62550.5171------0.3515
应用 softmax:
softmax ( ∑ i = 1 3 Q i K i T d k ) = 1.0 --- --- 0.5206 0.4794 --- 0.3464 0.3538 0.2998 \text{softmax}\left(\sum_{i=1}^{3} \frac{Q_i K_i^T}{\sqrt{d_k}}\right) = \begin{bmatrix} 1.0 & --- & --- \\ 0.5206 & 0.4794 & --- \\ 0.3464 & 0.3538 & 0.2998 \end{bmatrix} softmax(i=1∑3dk QiKiT)= 1.00.52060.3464---0.47940.3538------0.2998
Step 4:结论
✅ 分步计算与一次性计算的结果完全一致!
这证明了在 KV Cache 中只存 K 和 V 是正确的------因为 Q 每步只参与当前 token 的计算,历史 Q 不再需要。
二、KV Cache 在哪里使用
2.1 自回归生成过程
循环直到 eos_token:
1. 将新生成的 token append 到序列末尾
2. 将新序列作为输入传入模型
3. 模型计算 Q、K、V → Attention → FFN → 输出下一个 token
问题 :每次新序列输入时,都需要 重复计算 前面 n-1 个 token 的 q、k、v,浪费资源。
2.2 KV Cache 的使用方式
| 操作 | 不使用 KV Cache | 使用 KV Cache |
|---|---|---|
| 历史 token 的 K、V | 每步重新计算 | 从缓存中读取 |
| 当前 token 的 Q、K、V | 每步计算 | 每步计算(仅当前 token) |
| Attention 计算 | Q × 全部 K^T | 当前 Q × 缓存 K^T |
| 计算复杂度 | O(n²) 每步 | O(n) 每步(仅当前行) |
三、KV Cache 节省了 Self-Attention 层中哪部分的计算
3.1 Self-Attention 机制回顾
| 向量 | 来源 | 作用 |
|---|---|---|
| Query(Q) | 输入线性变换 | 当前 token 的"查询"表示 |
| Key(K) | 输入线性变换 | 用于计算与 Q 的相似度 |
| Value(V) | 输入线性变换 | 注意力加权求和的目标 |
计算流程:
Attention(Q, K, V) = softmax(QK^T / √d_k) · V
- Q × K^T → 注意力分数(相似度矩阵)
- / √d_k → 缩放(防止梯度消失)
- softmax → 归一化为概率分布
- × V → 加权求和得到输出
3.2 KV Cache 节省的部分
| 计算环节 | 是否被 KV Cache 节省 | 说明 |
|---|---|---|
| 历史 token 的 K 计算 | ✅ 节省 | 从缓存读取,无需重新线性变换 |
| 历史 token 的 V 计算 | ✅ 节省 | 从缓存读取,无需重新线性变换 |
| 历史 token 的 Q 计算 | 不涉及 | Q 每步只算当前 token |
| Scaled Dot-Product(QK^T) | ❌ 不节省 | 仍需用当前 Q 与所有 K 做点积 |
| Softmax | ❌ 不节省 | 仍需对完整注意力行做归一化 |
| 与 V 的加权求和 | ❌ 不节省 | 仍需用注意力权重加权所有 V |
⚠️ 重要 :KV Cache 节省的是 历史 token 的 K 和 V 的线性变换计算 ,但 不节省 Scaled Dot-Product Attention 本身的计算(QK^T、softmax、×V 仍需每步执行)。
四、总结
| 问题 | 答案 |
|---|---|
| 什么是 KV Cache? | 在自回归模型中缓存之前计算过的 K、V 的优化技术,避免重复计算历史 token 的键值对 |
| 在哪里使用? | 仅在 decoder 中使用(GPT 等 decoder-only 模型,或 T5 等 encoder-decoder 模型的解码部分),BERT 等 encoder-only 模型不涉及 |
| 节省了哪部分计算? | 节省历史 token 的 K 和 V 的重复计算(线性变换),但 不节省 Scaled Dot-Product Attention 本身的计算 |
| 为什么只存 K、V 不存 Q? | Q 只与当前 token 相关、每步重新计算,历史 Q 在自回归注意力中不再使用,缓存无收益 |
核心要点
- KV Cache = 用显存换计算:缓存历史 K、V,避免每步重复线性变换
- 仅限 Decoder:encoder-only 模型(如 BERT)不涉及
- 不缓存 Q:Q 每步只参与当前 token,历史 Q 不再需要
- 不节省 Attention 本身:QK^T、softmax、×V 仍需每步完整执行
五、参考资料
- Transformers KV Caching Explained --- João Lages
- transformer之KV Cache --- Takoony
- Understanding Attention In Transformers Models --- Alvaro Henriquez
- The Illustrated GPT-2 --- Jay Alammar
- Scaled Dot-Product Attention详解 --- 卡洛驰
本文档持续更新,如有遗漏或错误请指正。