苦猿的大模型日记 · 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 能快多少”不是一个完整问题。至少还要补充四个条件:

  1. 慢的是 prefill 还是 decode?
  2. 输入长度、输出长度和 batch 分别是多少?
  3. 当前瓶颈是计算、显存带宽、KV 容量还是排队?
  4. 对比的是 attention kernel,还是包含调度与网络的端到端服务?

本文不提供一组所谓最佳参数,而是从计算路径、显存公式和可复现实验三条线,把这两个优化的边界讲清楚。

KV Cache 与 FlashAttention 优化目标对照


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 步骤重复执行。

Prefill、Decode 与 KV Cache 时间轴


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 与 PagedAttention 显存块管理

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。

这里最关键的不是“分块”两个字,而是三个动作一起发生:

  1. Tiling:一次只处理能放进 SRAM 的小块。
  2. Kernel fusion:减少中间结果落回 HBM 的次数。
  3. Online softmax:分块时维护全局最大值和归一化因子,保证结果与标准 softmax 等价,而不是做近似注意力。

因此,FlashAttention 仍然是精确注意力。它没有把 O(n²) 的数学计算复杂度变成 O(n),而是把额外存储从显式 n² 中间矩阵压下来,并大幅减少 HBM IO。

它优化的是数据路径,不是修改注意力定义。

标准 Attention 与 FlashAttention 的 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 / 128prefill、FlashAttention、TTFT
短输入、长输出256 / 1024decode、KV 带宽、TPOT
长输入、长输出4096 / 1024KV 容量、抢占、尾延迟和 OOM

每组至少记录:

  • P50 / P95 TTFT。
  • P50 / P95 TPOT。
  • input tokens/s 与 output tokens/s,二者不要混成一个数。
  • 峰值显存、KV Cache 使用率、等待请求和抢占次数。
  • 满足延迟目标的请求 goodput,而不只是极限吞吐。

三种组合,结论应该怎么读

方案PrefillDecode显存特点适合回答的问题
无 KV Cache、普通 Attention重复处理历史重复计算严重cache 少,但计算浪费大只适合教学基线
KV Cache、普通 Attention基本不变大幅减少重复投影cache 随驻留 token 线性增长KV Cache 到底省了多少 decode 时间
KV Cache、FlashAttention / 高效 SDPAIO 更高效收益依 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 都可能影响输出行为或适用边界。

性能验收和质量验收必须一起跑:

  1. 先用固定模型与数据验证输出差异。
  2. 再做长上下文、多轮对话、事实引用和结构化输出回归。
  3. 最后才看吞吐收益是否值得这份风险。

只验证接口能够返回,不足以证明性能优化可以上线。

生产推理上线检查清单


结尾:先定位阶段,再选择优化

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 学进简历

更多推荐