导读:本期聚焦于赵六创作的《推理速度慢怎么办?KV Cache与Flash Attention如何加速大模型推理》,敬请观看详情。自回归语言模型每生成一个 token 都要重新计算前面所有 token 的注意力,这种重复计算是推理延迟的主要来源。KV Cache 的思路是把每一层已经算好的 Key 和 Value 缓存起来,下一步只计算新 token 的表示,再与历史缓存拼接,从而把单步计算量从平方级降为线性级。不过缓存会占用大量显存,长序列场景下尤为明显。Flash Attention 则从另一个维度入手,它不改变注意力计算的数学结果,而是通过分块与在线 softmax 把中间矩阵留在高速缓存中,大幅减少对高带宽显存的读写次数。两者一个优化计算量,一个优化 IO 开销,组合使用可以在不损失精度的情况下明显提升推理速度。本文会拆解它们的原理、实现要点以及落地时的配置建议。

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

推理速度慢怎么办?KV Cache与Flash Attention如何加速大模型推理

一、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

免责声明:已尽一切努力确保本网站所含信息的准确性。网站作品多为原创整理与精心创作,观点力求客观中立。本站旨在免费分享,内容仅供个人学习、研究或参考使用。若引用了第三方作品,版权归原作者所有。如内容涉及您的权益,请联系我们进行处理Email:chomcom@qq.com。
引用或转载本作品时,请注明当前出处:https://www.ipipp.com/html/0925/61613.html,基于非商业用途的前提下,欢迎转载或二创本作品。
内容垂直聚焦
专注技术核心技术栏目,确保每篇文章深度聚焦于实用技能。从代码技巧到架构设计,为用户提供无干扰的纯技术知识沉淀,精准满足专业提升需求。
知识结构清晰
覆盖从开发到部署的全链路。AI、前端、编程、数据库、服务器、建站、系统层层递进,构建清晰学习路径,帮助用户系统化掌握开发与运维所需的核心技术。
深度技术解析
拒绝泛泛而谈,深入技术细节与实践难点。无论是数据库优化还是服务器配置,均结合真实场景与代码示例进行剖析,致力于提供可直接应用于工作的解决方案。
专业领域覆盖
精准对应开发生命周期。从前端界面到后端编程,从数据库操作到服务器运维,形成完整闭环,一站式满足全栈工程师和运维人员的技术需求。
即学即用高效
内容强调实操性,步骤清晰、代码完整。用户可根据教程直接复现和应用于自身项目,显著缩短从学习到实践的距离,快速解决开发中的具体问题。
持续更新保障
专注既定技术方向进行长期、稳定的内容输出。确保各栏目技术文章持续更新迭代,紧跟主流技术发展趋势,为用户提供经久不衰的学习价值。