Day45|大模型推理加速:KV Cache 和 FlashAttention,根本不是一回事
苦猿的大模型日记 · Day45 · 大模型推理加速-帮普通人把AI学进简历系列
前言:先给结论——KV Cache 和 FlashAttention 不能互相替代
大模型推理优化里,KV Cache 和 FlashAttention 经常被放在一起讲。
这种讲法容易让人误以为,它们都是给 Attention 提速的实现,选一个就够了。实际不是。
KV Cache 解决重复计算,FlashAttention 解决显存 IO。
- KV Cache 保存历史 token 的 Key 和 Value,避免 decode 时重复计算历史投影。
- FlashAttention 改写注意力的执行方式,减少中间矩阵在 HBM 与片上存储之间的读写。
两者作用于不同阶段,收益也依赖不同条件。KV Cache 对自回归 decode 几乎是基础配置;FlashAttention 在长序列 prefill 中通常更容易体现优势。到了 query length 等于 1 的 decode 阶段,瓶颈可能已经变成 KV 读取、权重带宽和调度,FlashAttention 未必还能带来同等幅度的收益。
所以,“开启 FlashAttention 能快多少”不是一个完整问题。至少还要补充四个条件:
- 慢的是 prefill 还是 decode?
- 输入长度、输出长度和 batch 分别是多少?
- 当前瓶颈是计算、显存带宽、KV 容量还是排队?
- 对比的是 attention kernel,还是包含调度与网络的端到端服务?
本文不提供一组所谓最佳参数,而是从计算路径、显存公式和可复现实验三条线,把这两个优化的边界讲清楚。

PART 01:先把推理拆成两段——prefill 和 decode
大模型生成一句话,不是一口气把整句吐出来。
它先读完输入,再一个 token、一个 token 往后生成。对应到系统里,就是两个性质完全不同的阶段。
第一段:prefill,一次读完整个 prompt
假设用户输入有 2048 个 token。模型会并行处理这 2048 个 token,算出每一层的隐藏状态以及 Key、Value。
prefill 的矩阵通常比较大,GPU 有机会把 Tensor Core 吃满,所以它更偏计算密集型。它直接影响 TTFT,也就是 Time To First Token。
第二段:decode,一次只生成一个 token
首个 token 出来之后,模型把它拼回上下文,再预测下一个 token。如此循环,直到遇到停止符或达到输出上限。
decode 每一步只有一个或少量 query,却要读取模型权重和越来越长的 KV Cache。计算块变小、循环次数变多,往往更偏显存带宽与调度受限。它直接影响 TPOT,也就是 Time Per Output Token。
这两个阶段慢起来,用户的感受完全不同:
- TTFT 很差、出字以后正常:多半是排队或 prefill 慢。
- TTFT 正常、打字机越吐越慢:重点看 decode、KV Cache 带宽和并发批次。
- 两者都差:可能已经超过服务容量,不能只盯注意力 kernel。
没有 KV Cache,重复计算到底有多夸张
先看最朴素的自回归生成伪代码:
tokens = prompt_tokens
for _ in range(max_new_tokens):
logits = model(tokens) # 每一步把完整历史重新送进模型
next_token = sample(logits[:, -1])
tokens = torch.cat([tokens, next_token], dim=-1)
第一步处理长度 2048,第二步处理 2049,第三步处理 2050。
历史 token 的 Q、K、V 投影和注意力关系,会被一遍遍重算。你明明只想知道“下一个 token 是什么”,却每次都把前面两千多个 token 的作业重新写一遍。
这里有个经常被文章说糊涂的复杂度问题。
对一段长度为 n 的完整序列,标准注意力本身是 O(n²d)。如果每生成一个 token 都对越来越长的完整序列重新做注意力,整个生成过程还要把每一步的成本累加。
而有了 KV Cache 后,decode 的新 token 只需要计算自己的 Q、K、V,再让一个 query 去读取全部历史 K、V。单步注意力从“重新计算整张 n×n 关系图”,降成“新增 query 与 n 个历史位置做匹配”。
KV Cache 不会减少历史长度,但能避免历史 token 的 K、V 投影在每个 decode 步骤重复执行。

PART 02:KV Cache——用显存换取 decode 计算
Transformer 每层注意力都会做三次投影:
Q = XWq
K = XWk
V = XWv
Attention(Q, K, V) = softmax(QKᵀ / √d) V
自回归 decode 时,过去 token 的 K 和 V 不会再变化。既然它们已经算过,就可以缓存在显存里。
下一步只计算新 token 的 K、V,再追加到缓存:
new_k = key_proj(new_hidden)
new_v = value_proj(new_hidden)
k_cache = torch.cat([k_cache, new_k], dim=-2)
v_cache = torch.cat([v_cache, new_v], dim=-2)
output = attention(new_q, k_cache, v_cache)
真实推理引擎不会在每一步用 torch.cat 重新复制整块缓存。它们会预分配、分页或用专门的数据结构管理 cache。这里的代码只是把逻辑展示清楚。
一行公式,算出单个 token 吃多少显存
KV Cache 的估算公式是:
KV bytes per token
= 2 × num_layers × num_kv_heads × head_dim × dtype_bytes
其中:
- 2 代表 K 和 V 两份缓存。
num_kv_heads要看模型是 MHA、MQA 还是 GQA,不能直接拿 attention heads 代替。dtype_bytes对 FP16/BF16 通常是 2,对 FP8 通常是 1。
假设一个模型有 36 层、8 个 KV heads、head dimension 是 128,KV Cache 使用 FP16:
2 × 36 × 8 × 128 × 2
= 147456 bytes
≈ 144 KiB / token
一条总长度 8192 token 的请求,仅 KV Cache 理论值就接近:
8192 × 144 KiB ≈ 1.125 GiB
如果同时驻留 16 条这样的请求,KV Cache 就可能奔着 18 GiB 去。模型权重、激活值、CUDA context 和运行时 buffer 还没算。
这就是为什么模型明明能塞进卡里,并发一高却还是 OOM。
GQA 为什么对推理这么重要
传统 Multi-Head Attention 中,Query、Key、Value 的 head 数相同。GQA 让多个 Query heads 共享更少的 KV heads。
如果 Query 有 32 个 heads,KV 只有 8 个 heads,那么 KV Cache 体积大致只剩传统 MHA 的四分之一。
注意,它不是把所有注意力计算都砍成四分之一,而是显著减少需要存储和读取的 K、V。对长上下文 decode 来说,这个收益非常实在。
所以比较两个模型的部署成本,别只看“都是 7B”。参数量接近,不代表 KV Cache 成本接近。模型层数、KV heads、head dimension 和上下文长度,都会把账单改掉。
PagedAttention 解决的不是每 token 字节数
很多人会把 KV Cache 和 PagedAttention 混成一个概念。
KV Cache 是“把历史 K、V 存下来”。PagedAttention 是“怎么管理这块动态增长、长短不一的存储”。
如果为每条请求一次性预留最大上下文,浪费会非常严重。PagedAttention 把 cache 切成 block,需要多少就分多少,并允许不连续的物理块映射到一条逻辑序列。
它降低的是碎片和预留浪费,提升的是并发管理效率。它不会神奇地改变 144 KiB/token 这笔物理账。

KV Cache 的三个进阶方向
第一项是 FP8 KV Cache。
从 2 字节降到 1 字节,理论容量接近翻倍。但这不是免检开关:硬件、模型和引擎必须支持,标定方式也会影响质量。上线前至少要用自己的长上下文评测集检查困惑度、事实问答和多轮一致性。
第二项是 Prefix Caching。
KV Cache 复用的是同一条序列内部的历史;Prefix Caching 进一步复用不同请求之间相同的前缀,比如固定 system prompt、工具定义和 few-shot 示例。
它主要省重复前缀的 prefill,不会让新增 token 的 decode 免费。动态时间戳、请求 ID、随机排序如果放在 prompt 前部,也会把命中前缀截断。
第三项是 滑动窗口或 KV 淘汰。
它牺牲可见历史,换取固定 cache 上限。适合模型结构和业务允许只保留最近上下文的场景;如果任务需要跨 20K token 精确引用早期事实,随便淘汰只会把性能问题改造成质量问题。
KV Cache 是明确的空间换时间:部署前要计算容量,空间不足时要明确牺牲的是上下文、并发、精度还是硬件成本。
PART 03:FlashAttention——它没有少算注意力,只是少搬了很多数据
KV Cache 解决重复投影之后,注意力仍然可能慢。
原因不一定是 GPU 算不过来,而是数据搬得太多。
GPU 大致可以分成两类地方:
- HBM,也就是大家在
nvidia-smi里看到的显存,容量大、离计算核心远一些。 - SRAM,包括 shared memory 和寄存器,容量小,但离计算单元更近、带宽更高。
标准注意力如果按教科书实现,会经历几步:先算 S = QKᵀ,把巨大的分数矩阵写回 HBM;再读出来做 softmax;再写回;最后又读出来与 V 相乘。
序列一长,n×n 的注意力矩阵非常大。GPU 很多时间不是在乘加,而是在 HBM 和片上存储之间运数据。
FlashAttention 的核心是 IO-aware
FlashAttention 把 Q、K、V 分块搬进更快的片上存储,在块内完成矩阵乘、在线 softmax 和结果累加,不再把完整的 n×n 注意力矩阵物化到 HBM。
这里最关键的不是“分块”两个字,而是三个动作一起发生:
- Tiling:一次只处理能放进 SRAM 的小块。
- Kernel fusion:减少中间结果落回 HBM 的次数。
- Online softmax:分块时维护全局最大值和归一化因子,保证结果与标准 softmax 等价,而不是做近似注意力。
因此,FlashAttention 仍然是精确注意力。它没有把 O(n²) 的数学计算复杂度变成 O(n),而是把额外存储从显式 n² 中间矩阵压下来,并大幅减少 HBM IO。
它优化的是数据路径,不是修改注意力定义。

为什么 decode 阶段不一定“快两倍”
这是第二个容易写错的地方。
在长 prompt 的 prefill 中,Q 和 K 都很长,完整注意力矩阵巨大,FlashAttention 的分块与融合收益通常更明显。
但普通自回归 decode 每一步的 query length 往往只有 1。此时根本没有一张巨大的 q_len×kv_len 矩阵需要物化,瓶颈更可能是读取全部 KV Cache、读取模型权重以及小 kernel 调度。
所以你可能看到:prefill 明显变快,长上下文显存峰值下降,但 TPOT 只改善一点;某些短序列、小 batch 场景甚至差异很小。
这不是 FlashAttention 失效了,而是你的瓶颈换了。decode 更依赖 paged KV 管理、continuous batching、量化 kernel、Flash-Decoding 一类针对性优化。
现在不一定非要手装 flash-attn
如果你使用较新的 PyTorch,scaled_dot_product_attention 会根据 GPU、dtype、shape 和 mask 条件,自动在 FlashAttention、memory-efficient attention 或 math 实现之间选择。
import torch
import torch.nn.functional as F
q = torch.randn(2, 16, 1024, 64, device="cuda", dtype=torch.float16)
k = torch.randn(2, 16, 1024, 64, device="cuda", dtype=torch.float16)
v = torch.randn(2, 16, 1024, 64, device="cuda", dtype=torch.float16)
with torch.inference_mode():
out = F.scaled_dot_product_attention(
q, k, v,
dropout_p=0.0,
is_causal=True,
)
print(out.shape)
Hugging Face Transformers 也可以通过模型支持的 attention implementation 选择实现:
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained(
"Qwen/Qwen2.5-7B-Instruct",
torch_dtype="auto",
device_map="auto",
attn_implementation="flash_attention_2",
)
具体支持情况跟模型、Transformers、PyTorch 和硬件版本有关。生产环境不要只因为参数写进去了,就默认内核一定命中。
真要安装 flash-attn,先固定四个版本
常见安装方式是:
pip install flash-attn --no-build-isolation
但在执行之前,先记录:
python -c "import torch; print(torch.__version__); print(torch.version.cuda)"
nvcc --version
nvidia-smi
python --version
这里最容易踩三类坑。
第一类:PyTorch 使用的 CUDA runtime 与本机 nvcc 工具链不兼容。
nvidia-smi 显示的 CUDA Version 是驱动能支持的最高版本,不等于你编译扩展时实际使用的 toolkit 版本。判断环境时,至少把 torch.version.cuda 和 nvcc --version 放在一起看。
第二类:编译把内存和 CPU 打满。
可以限制并行任务数:
MAX_JOBS=4 pip install flash-attn --no-build-isolation
Windows PowerShell 写法是:
$env:MAX_JOBS = "4"
pip install flash-attn --no-build-isolation
第三类:硬件架构并不在当前版本的高效支持范围内。
不同 FlashAttention 大版本、CUDA backend 和 GPU 架构的支持矩阵会变。Ampere、Ada、Hopper 与更老的 Turing、Volta 不能用同一条经验判断。先看项目当前版本说明,再决定走 FlashAttention、PyTorch SDPA、xFormers 还是推理引擎自带 kernel。
为了安装一个扩展,连续修改 PyTorch、CUDA、编译器和驱动,会显著扩大变更范围。应先在隔离环境完成兼容性验证,不要直接修改正在提供服务的运行环境。
PART 04:建立可复现的 prefill 与 decode 基准
“FlashAttention 提速 2 倍,显存下降 40%”这类结论,如果没有实验条件,基本无法复用。
GPU、batch、输入长度、dtype、mask 类型和执行阶段,任一项变化都可能改变结果。尤其不能把 attention kernel 的加速比直接写成端到端模型吞吐提升。
一套能用的 benchmark,至少要固定:
- GPU 型号与数量。
- PyTorch、CUDA、Transformers、推理引擎版本。
- 模型 revision、权重 dtype、KV dtype、量化方式。
- batch size、输入长度、输出长度和长度分布。
- 是否预热、是否同步 CUDA、是否包含 tokenizer 和网络耗时。
一个最小注意力 kernel 对比脚本
下面的脚本对比朴素注意力与 PyTorch SDPA。它测的是 attention kernel,不是完整模型吞吐,但能帮你观察不同序列长度下,显式物化注意力矩阵的代价。
import argparse
import math
import statistics
import time
import torch
import torch.nn.functional as F
def naive_attention(q, k, v, causal):
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(q.size(-1))
if causal:
q_len, kv_len = q.size(-2), k.size(-2)
mask = torch.ones(q_len, kv_len, device=q.device, dtype=torch.bool).tril(
diagonal=kv_len - q_len
)
scores = scores.masked_fill(~mask, float("-inf"))
probs = torch.softmax(scores, dim=-1)
return torch.matmul(probs, v)
def benchmark(fn, warmup=10, repeat=50):
for _ in range(warmup):
fn()
torch.cuda.synchronize()
samples = []
for _ in range(repeat):
start = time.perf_counter()
fn()
torch.cuda.synchronize()
samples.append((time.perf_counter() - start) * 1000)
return statistics.median(samples)
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--seq-len", type=int, default=2048)
parser.add_argument("--batch", type=int, default=1)
parser.add_argument("--heads", type=int, default=16)
parser.add_argument("--head-dim", type=int, default=64)
args = parser.parse_args()
shape = (args.batch, args.heads, args.seq_len, args.head_dim)
q = torch.randn(shape, device="cuda", dtype=torch.float16)
k = torch.randn(shape, device="cuda", dtype=torch.float16)
v = torch.randn(shape, device="cuda", dtype=torch.float16)
with torch.inference_mode():
naive_ms = benchmark(lambda: naive_attention(q, k, v, causal=True))
sdpa_ms = benchmark(
lambda: F.scaled_dot_product_attention(
q, k, v, dropout_p=0.0, is_causal=True
)
)
print(f"shape={shape}")
print(f"naive median: {naive_ms:.3f} ms")
print(f"sdpa median: {sdpa_ms:.3f} ms")
print(f"speedup: {naive_ms / sdpa_ms:.2f}x")
if __name__ == "__main__":
main()
运行三档长度:
python bench_attention.py --seq-len 512
python bench_attention.py --seq-len 1024
python bench_attention.py --seq-len 2048
如果显存较小,朴素实现可能在长序列直接 OOM。这本身就是结果:显式 n×n 注意力矩阵的峰值显存,正在变成瓶颈。
但别拿这个 kernel benchmark 宣称“模型推理提速多少”。完整服务还包含 QKV 投影、MLP、归一化、采样、KV 管理、调度、tokenizer 和网络传输。
生产压测要分四组看
| 组别 | 输入 / 输出 | 主要观察 |
|---|---|---|
| 短输入、短输出 | 256 / 128 | 调度开销、小 batch kernel 效率 |
| 长输入、短输出 | 4096 / 128 | prefill、FlashAttention、TTFT |
| 短输入、长输出 | 256 / 1024 | decode、KV 带宽、TPOT |
| 长输入、长输出 | 4096 / 1024 | KV 容量、抢占、尾延迟和 OOM |
每组至少记录:
- P50 / P95 TTFT。
- P50 / P95 TPOT。
- input tokens/s 与 output tokens/s,二者不要混成一个数。
- 峰值显存、KV Cache 使用率、等待请求和抢占次数。
- 满足延迟目标的请求 goodput,而不只是极限吞吐。
三种组合,结论应该怎么读
| 方案 | Prefill | Decode | 显存特点 | 适合回答的问题 |
|---|---|---|---|---|
| 无 KV Cache、普通 Attention | 重复处理历史 | 重复计算严重 | cache 少,但计算浪费大 | 只适合教学基线 |
| KV Cache、普通 Attention | 基本不变 | 大幅减少重复投影 | cache 随驻留 token 线性增长 | KV Cache 到底省了多少 decode 时间 |
| KV Cache、FlashAttention / 高效 SDPA | IO 更高效 | 收益依 shape 而定 | 避免大中间矩阵,仍需存 KV | 长 prompt 与端到端组合收益 |
不要预设第三组一定“全面碾压”。如果输入只有 64 token、batch 很小,kernel 启动开销可能掩盖优化;如果 decode 已经卡在权重读取,prefill kernel 再快也救不了 TPOT。

基准测试的目的不是证明某项技术有效,而是验证它是否命中了当前系统的主要瓶颈。
PART 05:生产落地必须处理的五个边界
算法收益成立,不代表生产服务可以直接使用。下面五个边界会直接影响稳定性。
坑一:自己用 padding 管动态 batch
不同请求长度不同,如果把它们 padding 到同一长度,短请求会替最长请求一起交显存税和计算税。
有人会说,那就按长度分桶。方向没错,但在线请求持续到达,旧请求还在 decode,新请求又来 prefill,batch 每一步都在变化。你很快会遇到 cache 重排、请求完成后的空洞、取消请求回收和多轮 prefix 复用。
这也是为什么生产环境通常交给 vLLM、TGI、TensorRT-LLM 等成熟引擎管理 paged cache 和 continuous batching,而不是在业务代码里手写一个 torch.cat。
手写 KV Cache 适合理解原理;在线调度与缓存管理则应优先采用经过大规模请求验证的推理引擎。
坑二:只限制输入长度,不限制总长度
KV Cache 取决于“输入 token + 已生成 token”。
如果 API 只限制 prompt 不能超过 8K,却允许 max_tokens=8K,一条合法请求的最坏驻留长度仍然可能接近 16K。
上线前至少明确三个边界:
max_input_tokens
max_output_tokens
max_total_tokens
而且不能只看平均值。平均输入 800 token 的系统,完全可能被 1% 的 20K 长请求拖出 cache 抢占和长尾延迟。
坑三:把 OOM 当成唯一失败信号
KV Cache 不够时,成熟引擎未必立刻 OOM。它可能让请求等待、抢占正在运行的序列、丢掉 cache 后重新计算,或者把部分数据换到 CPU。
服务还活着,但有效吞吐下降,P95 延迟越来越差。
所以监控至少要有:运行中请求、等待请求、KV Cache 使用率、抢占次数、TTFT、TPOT 和失败率。具体指标名会随引擎版本变化,应该从当前服务的 /metrics 实际输出确认,而不是复制博客里的 Prometheus 字段。
坑四:FlashAttention 装上了,实际没有走
mask 类型不支持、dtype 不匹配、head dimension 不合适、GPU 架构不支持,或者框架回退到 math kernel,都可能发生。
最简单的验证不是看安装日志,而是做 profiler 和 A/B benchmark。对相同输入,固定随机种子、预热、CUDA 同步,观察实际 kernel、耗时和峰值显存。
配置成功不等于路径命中,路径命中也不等于端到端变快。
坑五:为了提速,悄悄改坏了质量
FP8 KV、滑动窗口、权重量化、speculative decoding 都可能影响输出行为或适用边界。
性能验收和质量验收必须一起跑:
- 先用固定模型与数据验证输出差异。
- 再做长上下文、多轮对话、事实引用和结构化输出回归。
- 最后才看吞吐收益是否值得这份风险。
只验证接口能够返回,不足以证明性能优化可以上线。

结尾:先定位阶段,再选择优化
KV Cache 和 FlashAttention 可以同时使用,但不能混为一个问题。
KV Cache 主要减少自回归 decode 中的重复投影计算,代价是随驻留 token 数线性增长的显存占用。FlashAttention 主要减少注意力计算的 HBM IO 和中间矩阵存储,不改变标准注意力的 O(n²) 计算复杂度。
工程上可以按下面的顺序判断:
- 历史 token 被反复计算:看 KV Cache。
- 长 prompt 的注意力中间矩阵与 HBM IO 太重:看 FlashAttention 或高效 SDPA。
- KV 空间碎片和动态并发难管理:看 PagedAttention 与成熟推理引擎。
- 队列增长、到达率超过服务率:做容量规划和扩容,别再假装是 kernel 参数问题。
完成瓶颈定位后,再决定启用 KV Cache、FlashAttention、PagedAttention、continuous batching、量化或扩容。没有阶段拆分和基线数据,优化项叠得越多,越难解释收益来自哪里。
推理优化的基本纪律,是先测清时间和显存花在哪里,再修改对应的计算路径。
互动时间:你们线上最难守的是 TTFT,还是 TPOT?有没有遇到“GPU 没满、接口却越来越慢”的情况?评论区把模型、显卡和输入输出长度留下来,我们一起拆瓶颈。
下一篇预告:推理速度解决之后,怎么用 continuous batching 把并发真正吃满?关注苦猿,Day46 不见不散。
— END —
苦猿 · 帮普通人把 AI 学进简历
更多推荐
所有评论(0)