大模型推理之所以慢,很大程度是因为自回归解码过程中存在大量重复计算。以 GPT 类模型为例,每生成一个新 token,模型都要把整条历史序列重新送入网络,重新计算所有位置的 Key 和 Value,再计算注意力分数。序列长度每增加一,计算量就会成比例增长,最终形成平方级的开销。针对这个问题,工程上有两条非常有效的优化路线:一是用 KV Cache 把已经计算过的键值对缓存起来,避免每一步都从头算;二是用 Flash Attention 减少注意力计算中的显存读写,降低 IO 瓶颈。二者并不冲突,反而可以叠加使用。

一、KV Cache:缓存历史键值,消除重复计算
在没有 KV Cache 的情况下,模型生成第 t 个 token 时,会把前面 t-1 个 token 连同当前 token 一起作为输入,所有历史 token 的 Key 和 Value 向量都会被重新投影计算。如果序列长度为 N,第 t 步的注意力计算量大约是 O(t·d),总计算量接近 O(N³·d),其中 d 是每个头的维度。这个开销里有一大部分是在重复计算已经算过的内容。
KV Cache 的做法非常直接:每一层、每一个注意力头都把历史 token 的 Key 和 Value 张量保存下来。生成下一个 token 时,只需要对新增的那个 token 计算 Q、K、V,然后把新得到的 K、V 拼接到缓存后面。注意力计算仍然需要新 token 与所有历史 token 交互,但省去了对历史 token 做线性投影的重复工作,单步计算量从 O(t·d) 降到 O(t·d) 的注意力部分加上 O(d) 的投影部分,整体复杂度从 O(N³·d) 下降到 O(N²·d)。在序列很长时,这个差距会非常惊人。
下面是一段简化的 PyTorch 代码,展示带 KV Cache 的前向过程:
# 假设 key_cache 和 value_cache 形状为 [batch, heads, seq_len, head_dim]
# query 形状为 [batch, heads, 1, head_dim]
# key、value 是当前 token 的投影结果
def forward_with_kv_cache(query, key, value, key_cache, value_cache):
# 将当前 token 的 key、value 拼接到历史缓存
key = torch.cat([key_cache, key], dim=2)
value = torch.cat([value_cache, value], dim=2)
# 更新缓存
key_cache = key
value_cache = value
# 计算注意力分数
scores = torch.matmul(query, key.transpose(-2, -1)) / (head_dim ** 0.5)
attn = torch.softmax(scores, dim=-1)
output = torch.matmul(attn, value)
return output, key_cache, value_cache
不过 KV Cache 并不是没有代价的。缓存会持续占用显存,其大小可以估算为:2 × B × L × N × H × d × 字节数,其中 B 是 batch 大小,L 是层数,N 是序列长度,H 是注意力头数,d 是每个头的维度。以 FP16 精度为例,一个 70B 模型在 batch=1、序列长度 4096 时,KV Cache 可能占用十几 GB 显存。因此,长文本和多轮对话场景下,KV Cache 的显存压力会迅速上升,这也催生了 KV Cache 量化、PagedAttention 等后续优化。
二、Flash Attention:分块计算,把数据留在高速缓存
标准注意力计算的另一个隐性瓶颈在于显存带宽。传统实现会先计算一个 N×N 的注意力分数矩阵,把它写入 HBM 显存,然后再从显存读出来做 softmax 和加权求和。当 N 较大时,这个矩阵本身就非常大,比如 N=4096 时,单个头就需要 4096×4096 个浮点数,多层多头加起来会产生大量的显存读写。Flash Attention 的核心思想是避免把完整的 N×N 矩阵写入 HBM,而是把 Q、K、V 分成小块,在 GPU 的 SRAM 高速缓存中逐块计算,利用 online softmax 技术维护全局归一化因子,最终只把结果矩阵 O 写回显存。
具体来说,Flash Attention 将查询矩阵 Q 分成多个块,对每一个 Q 块,遍历所有 K 块和 V 块,在 SRAM 中计算局部注意力分数,并通过不断更新最大值和归一化系数来合并结果。由于不需要存储中间注意力矩阵,显存复杂度从 O(N²) 降到 O(N),同时大幅减少了对 HBM 的读写次数。值得强调的是,Flash Attention 是一种精确算法,它得到的结果与标准 softmax 注意力在数学上是一致的,只是利用数值稳定的方式重新组织了计算顺序,因此不会损失模型精度。
下面是一段简化逻辑,展示分块计算和 online softmax 的大致过程:
# 简化版 Flash Attention 分块计算示意
# 实际实现依赖 CUDA kernel,这里用 Python 表达逻辑
def flash_attention(Q, K, V, block_size):
B, H, N, D = Q.shape
O = torch.zeros_like(Q)
L = torch.zeros(B, H, N, 1)
Q_blocks = Q.split(block_size, dim=2)
K_blocks = K.split(block_size, dim=2)
V_blocks = V.split(block_size, dim=2)
for i, Qi in enumerate(Q_blocks):
Oi = torch.zeros_like(Qi)
Li = torch.zeros(B, H, Qi.size(2), 1)
for j, (Kj, Vj) in enumerate(zip(K_blocks, V_blocks)):
Sij = torch.matmul(Qi, Kj.transpose(-2, -1)) / (D ** 0.5)
mij = Sij.max(dim=-1, keepdim=True).values
Pij = torch.exp(Sij - mij)
lij = Pij.sum(dim=-1, keepdim=True)
new_Li = Li + lij
Oi = Oi * (Li / new_Li) + torch.matmul(Pij, Vj)
Li = new_Li
O[:, :, i*block_size:(i+1)*block_size, :] = Oi / Li
return O
在实际工程中,PyTorch 从 2.0 开始提供的 scaled_dot_product_attention 函数已经集成了 Flash Attention 后端。只要 GPU 支持并且输入形状满足要求,它会自动选择 Flash Attention 或 Memory-Efficient Attention,开发者基本不需要手动调用底层 CUDA 接口。对于长序列推理,开启 Flash Attention 通常能带来数倍的加速,同时把显存占用削减到原来的几分之一。
三、实际落地:组合使用与工程调优
KV Cache 与 Flash Attention 解决的是不同层面的问题。KV Cache 减少了历史 token 投影的重复计算,Flash Attention 则降低了单次注意力计算中的显存读写。在实际推理框架中,两者往往同时存在。例如 vLLM 在 PagedAttention 管理 KV Cache 的同时,也会启用 Flash Attention 后端;TensorRT-LLM 和 Hugging Face 的新版本推理代码也采取了类似策略。
除了直接启用这两项技术,还可以进一步叠加其他优化。例如 GQA(分组查询注意力)和 MQA(多查询注意力)通过减少 KV 头的数量来降低 KV Cache 体积;KV Cache 量化可以在 INT8 甚至 INT4 精度下保存键值,进一步压缩显存;PagedAttention 则把 KV Cache 划分成固定大小的物理块,解决显存碎片问题并提高并发吞吐。对于长文档摘要、多轮对话、代码生成等长序列任务,这些优化带来的收益尤其明显;而短序列分类或单步推理场景下,Flash Attention 的优势可能不那么突出,因为矩阵较小,IO 开销本身不高。
下面是一个在 PyTorch 中启用 Flash Attention 的简单例子:
import torch import torch.nn.functional as F query = torch.randn(1, 8, 1024, 64, device="cuda", dtype=torch.float16) key = torch.randn(1, 8, 1024, 64, device="cuda", dtype=torch.float16) value = torch.randn(1, 8, 1024, 64, device="cuda", dtype=torch.float16) # PyTorch 会根据 GPU 和输入形状自动选择 Flash Attention output = F.scaled_dot_product_attention(query, key, value)
选型时可以先评估自己的业务负载:如果序列长度经常超过 1024,建议同时开启 KV Cache 和 Flash Attention;如果显存紧张,优先考虑 GQA 或 KV Cache 量化;如果需要高并发服务,可以引入 PagedAttention 提升显存利用率。工程上并不存在一个万能配置,但把 KV Cache 和 Flash Attention 作为基础优化,再根据瓶颈逐步叠加其他手段,是目前加速大模型推理最稳妥的路径。
KV CacheFlash Attention推理加速修改时间:2026-09-25 06:38:22