KV Cache 优化与注意力机制加速:FlashAttention / PagedAttention

大型语言模型(LLM)的推理性能已经成为生产环境部署中的核心挑战。在 Transformer 架构中,自注意力机制(Self-Attention)虽然在建模长距离依赖方面表现出色,但其计算和内存开销随序列长度呈平方增长。

大型语言模型(LLM)的推理性能已经成为生产环境部署中的核心挑战。在 Transformer 架构中,自注意力机制(Self-Attention)虽然在建模长距离依赖方面表现出色,但其计算和内存开销随序列长度呈平方增长。当上下文窗口扩展到 128K 甚至 1M token 时,传统的注意力实现会迅速触及 GPU 硬件瓶颈。本文将深入探讨三项关键技术——FlashAttention、PagedAttention 与 KV Cache 优化——它们分别从计算效率和内存管理两个维度重塑了 LLM 推理的工程实践。

一、标准注意力机制及其瓶颈

1.1 注意力公式回顾

标准的缩放点积注意力(Scaled Dot-Product Attention)定义为:

Attention(Q, K, V) = Softmax(QK^T / sqrt(d_k))V

其中 Q、K、V 分别代表查询(Query)、键(Key)和值(Value)矩阵,维度均为 [batch_size, num_heads, seq_len, head_dim]d_k 为每个注意力头的维度。

1.2 O(n^2) 的复杂度困境

从矩阵乘法的维度分析可知:

  • QK^T 的计算量为 O(seq_len^2 * num_heads * head_dim)
  • Softmax 的输出矩阵维度为 [batch_size, num_heads, seq_len, seq_len]
  • 最终与 V 相乘的计算量同样为 O(seq_len^2 * num_heads * head_dim)

seq_len 从 1K 增加到 32K 时,计算量增长超过一千倍。更关键的是,这个 seq_len x seq_len 的注意力矩阵(Attention Score Matrix)必须在 GPU 的高带宽存储器(HBM, High Bandwidth Memory)中持久化——它以 FP16 精度存储时,32K 上下文长度下仅一个头就需要约 2GB 内存。对于拥有 32 到 64 个注意力头的标准模型,显存需求迅速膨胀至不可接受的程度。

1.3 真正的瓶颈是内存带宽

现代 GPU(如 A100/H100)的算力增长远超内存带宽的增长。标准注意力实现的流程如下:

  1. 从 HBM 读取 Q, K, V
  2. 在 SRAM(片上静态随机存储器)中计算 QK^T
  3. 将中间结果 S = QK^T 写回 HBM
  4. 从 HBM 读取 S,计算 Softmax 得到 P
  5. 将 P 写回 HBM
  6. 从 HBM 读取 P 和 V
  7. 计算 O = PV
  8. 将输出 O 写回 HBM

问题在于步骤 2 到 7 之间产生了大量对 HBM 的读写操作。HBM 的带宽虽然较高(A100 约 1.935 TB/s),但与 SRAM(A100 的 L2/SRAM 带宽可达数十 TB/s 级别)相比仍然慢一到两个数量级。因此,标准注意力是被内存带宽「卡脖子」的,而非单纯的算力不足

二、FlashAttention:用分块计算突破内存墙

2.1 核心洞察

FlashAttention 的核心思想来自一个简单但深刻的观察:能否避免将巨大的注意力矩阵写回 HBM?

GPU 的 SRAM(亦称共享内存或 Shared Memory,A100 每个 Streaming Multiprocessor 约 164KB)虽然容量极小,但访问速度极快。FlashAttention 将注意力计算拆分为足够小的块(tiling),使得所有中间计算都能驻留在 SRAM 中完成,只需将最终输出写回 HBM

2.2 分块计算与在线 Softmax

标准的 Softmax 计算需要全局的归一化因子(所有元素的和):

Softmax(x_i) = exp(x_i) / sum_j(exp(x_j))

在分块计算中,我们无法一次性看到所有元素。FlashAttention 使用**在线 Softmax(online softmax)**技巧解决了这个问题。其核心是利用一个不变式:通过维护两个统计量——当前块的最大值 m 和指数和 l——可以逐步修正之前的部分结果。

# 伪代码:在线 Softmax 的增量更新
def online_softmax_update(m_prev, l_prev, x_curr):
    m_curr = max(m_prev, max(x_curr))
    # 用新的最大值重新缩放旧的指数和
    l_curr = exp(m_prev - m_curr) * l_prev + sum(exp(x_curr - m_curr))
    return m_curr, l_curr

FlashAttention 的外层循环遍历 Q 的块,内层循环遍历 K、V 的块,逐步累积注意力的输出和归一化因子。

2.3 详细计算流程

# FlashAttention 核心逻辑伪代码
# 假设 SRAM 可容纳大小为 Br x d 和 Bc x d 的块

def flash_attention(Q, K, V):
    # Q, K, V shape: [N, d], N = seq_len
    # 分块大小
    Br = 64   # Q 的行块大小
    Bc = 64   # K/V 的列块大小
    
    # 初始化输出矩阵 O,以及每行的统计量
    O = zeros(N, d)
    L = zeros(N)      # 存储行间指数和(用于反向传播)
    m = full(N, -inf) # 每行当前最大值
    l = zeros(N)      # 每行当前缩放后的指数和
    
    # 外层循环:按行遍历 Q(分成 Tr 块)
    for i in range(0, N, Br):
        Qi = Q[i:i+Br]  # 加载 Qi 到 SRAM
        mi = m[i:i+Br]
        li = l[i:i+Br]
        Oi = zeros(Br, d)
        
        # 内层循环:按列遍历 K, V(分成 Tc 块)
        for j in range(0, N, Bc):
            Kj = K[j:j+Bc]  # 加载 Kj, Vj 到 SRAM
            Vj = V[j:j+Bc]
            
            # 在 SRAM 中计算 Sij = Qi * Kj^T
            Sij = Qi @ Kj.T  # shape: [Br, Bc]
            
            # 在线 Softmax:更新当前块的行最大值和指数和
            mij_local = max(Sij, axis=1)  # [Br]
            m_new = max(mi, mij_local)
            
            # 重新缩放旧的输出和指数和
            alpha = exp(mi - m_new)
            beta = exp(mij_local - m_new)
            
            # 计算当前块的指数权重
            Pij = exp(Sij - m_new[:, None])
            
            # 增量更新输出:Oi = alpha * Oi + Pij @ Vj
            Oi = alpha[:, None] * Oi + Pij @ Vj
            
            # 更新全局统计量
            li = alpha * li + beta * sum(Pij, axis=1)
            mi = m_new
        
        # 归一化并写回 HBM
        O[i:i+Br] = Oi / li[:, None]
        L[i:i+Br] = mi + log(li)  # 存储 log-sum-exp 用于反向传播
        m[i:i+Br] = mi
        l[i:i+Br] = li
    
    return O, L

2.4 重计算策略

标准注意力在反向传播时可以直接读取前向传播缓存的注意力矩阵 P。但 FlashAttention 为了节省内存,不保存巨大的 P 矩阵。取而代之的是,它在反向传播时重新计算 P——由于可以再次利用分块策略,重计算的额外开销很小,却能换来数量级的内存节省。在纯推理场景(无反向传播)中,这一 trade-off 更加有利。

2.5 FlashAttention-2 与 FlashAttention-3

**FlashAttention-2(2023)**的主要改进包括:

  • 减少非矩阵乘法的 FLOPs:通过更精细的并行调度,减少 warp 之间的同步开销。
  • 更好的工作划分:不再按批次和注意力头循环,而是让 warp groups 专注于不同的注意力头,减少空闲线程。
  • 序列并行:在序列维度上并行化,对于长序列尤为有效。
  • 实际测试中,FlashAttention-2 相比初代可达到 1.5~2 倍加速

**FlashAttention-3(2024)**则面向新一代 GPU(H100/H200 的 Hopper 架构):

  • 异步数据传输:利用 Tensor Memory Accelerator(TMA)实现异步的块加载/存储,与计算重叠。
  • FP8 低精度支持:在 Hopper 的 FP8 Tensor Core 上实现低精度注意力,进一步加速。
  • Warp Specialization:将不同的 warps 专职用于数据加载与计算,实现流水线并行。

2.6 效果总结

指标标准 AttentionFlashAttention
HBM 读写量O(N^2)O(N)
显存占用O(N^2)O(N)
典型加速比1x2~4x
最大支持序列~8K-16K>128K

三、LLM 推理中的 KV Cache

3.1 什么是 KV Cache

在 Transformer 的解码阶段,生成过程是自回归的:每个新 token 的预测都需要用到之前所有 token 的上下文。如果不做优化,每次生成新 token 时都会重新计算之前所有位置的 K 和 V,造成大量重复计算。

KV Cache 的解决方案是:在第一次前向传播(Prefill 阶段)时,计算并存储所有历史 token 的 K、V 张量;在后续的解码步骤(Decode 阶段)中,只需将新生成 token 的 K、V 追加到缓存中,然后复用历史 K、V 进行注意力计算。

3.2 KV Cache 的内存开销

KV Cache 的显存占用可以用以下公式估算:

KV_cache_size = 2 * num_layers * num_heads * head_dim * seq_len * batch_size * sizeof(dtype)

以 Llama-2-70B 为例:

  • num_layers = 80
  • num_heads = 64(注意:key/value 头的数量可能因 GQA 而异,这里假设全头)
  • head_dim = 128
  • seq_len = 8192
  • batch_size = 8
  • dtype = fp16 (2 bytes)
KV_cache = 2 * 80 * 64 * 128 * 8192 * 8 * 2 bytes
         ≈ 163.8 GB

这已经超过一张 A100-80GB 的显存容量。随着 128K、甚至 1M 上下文窗口的模型出现,KV Cache 成为了推理部署中最主要的显存消耗来源。

3.3 解码阶段的新挑战

在解码阶段,每个新 token 的注意力计算是逐 token 进行的(单个查询向量对全部历史的 K、V),但 KV Cache 本身却在持续膨胀。这带来了两个新问题:

  1. 显存碎片化:不同序列长度、不同请求的 KV Cache 大小不一,导致存储不连续。
  2. 过度预留:为了支持最大上下文,系统通常按最坏情况(max_seq_len)预先分配显存,造成大量浪费。

四、PagedAttention(vLLM):虚拟内存思想管理 KV Cache

4.1 问题定义

在 vLLM 提出 PagedAttention 之前,主流推理框架(如 FasterTransformer、Hugging Face TGI)管理 KV Cache 的方式相当粗糙:

  • 为每个请求预留连续的大块显存:大小为 max_seq_len,无论实际生成了多少 token。
  • 无法共享 KV Cache:当多个并行采样请求共享同一个输入前缀时,各自的 KV Cache 完全独立存储。
  • 外部内存碎片:不同长度的序列释放后,残留的 “空洞” 难以被有效复用。

4.2 分页存储:从操作系统借来的灵感

PagedAttention 的核心创新是将操作系统的虚拟内存分页机制引入 KV Cache 管理:

  • Block:将 KV Cache 划分为固定大小的块(通常是 16 个 token 为一页),每页内存储 K、V 向量。
  • Block Table:维护一个从「逻辑 KV Cache」(连续的 token 序列)到「物理块」(可能不连续的 GPU 显存地址)的映射表。
  • 按需分配:只在实际需要时分配新的 block,而非预先分配最大长度。
逻辑视角(用户看到的):
Token [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, ...]
        |-- Block 0 --|-- Block 1 --|-- Block 2 --|

物理视角(GPU 显存中的实际布局):
Block 0 → GPU Mem Page 7
Block 1 → GPU Mem Page 3
Block 2 → GPU Mem Page 12
        (不要求物理连续)

4.3 Copy-on-Write 与前缀共享

PagedAttention 引入了类似操作系统 COW(Copy-on-Write)的机制来高效处理前缀共享场景:

场景示例:用户发送了一个 1000 token 的 prompt,要求模型生成 3 个不同的续写结果(并行采样 or 束搜索)。

在传统方案中,3 个请求的 KV Cache 各自独立,prompt 部分被存储了 3 份。

在 PagedAttention 中:

  1. 初始时,3 个请求共享 prompt 对应的物理 block。
  2. 每个请求维护独立的 block table,但指向相同的物理 block。
  3. 当某个请求生成新 token 需要写入时,系统复制该 block(COW),新请求获得写权限,其他请求继续共享旧版本。
# 简化的 PagedAttention Block Table 操作伪代码
class BlockTable:
    def __init__(self, block_size=16):
        self.block_size = block_size
        self.logical_to_physical = []  # 逻辑 block_id -> 物理 block_id
        self.physical_blocks = {}      # 物理 block_id -> 实际张量
    
    def allocate(self, num_tokens):
        """为新 token 分配物理 block"""
        needed_blocks = ceil(num_tokens / self.block_size)
        while len(self.logical_to_physical) < needed_blocks:
            phys_id = memory_allocator.alloc_block()
            self.logical_to_physical.append(phys_id)
    
    def get_kv(self, token_position):
        """根据 token 位置找到对应的物理 block 和偏移"""
        block_id = token_position // self.block_size
        offset = token_position % self.block_size
        phys_id = self.logical_to_physical[block_id]
        return self.physical_blocks[phys_id][offset]
    
    def fork(self):
        """Copy-on-Write:创建新 BlockTable 共享物理块"""
        new_table = BlockTable(self.block_size)
        new_table.logical_to_physical = self.logical_to_physical.copy()
        new_table.physical_blocks = self.physical_blocks
        # 标记为共享,写时复制
        block_ref_counts.increment_all(new_table.logical_to_physical)
        return new_table

4.4 vLLM 调度器与 Block 管理

vLLM 的调度器在请求层面实现了一个精细的显存管理系统:

  1. Waiting Queue:新进入的请求在这里排队。
  2. Running Batch:当前正在 GPU 上执行的请求批次。PagedAttention 允许动态添加/移除请求,只要 block table 能管理。
  3. Swapping:当显存不足时,vLLM 可以将某些请求的 KV Cache block “换出”(swap out)到 CPU 内存;待显存释放后再 “换入”(swap in)。这类似于操作系统的内存交换。
  4. Block Allocator:负责维护空闲 block 池,分配和回收物理 block。

这种设计使得 vLLM 可以:

  • 消除内部碎片:固定大小的 block 避免了变长分配。
  • 消除外部碎片:block 可以从空闲池中复用,无需连续内存。
  • 提高 batch size:不再为每个请求预留最大长度,显存利用率提升 2~4 倍,batch size 可提升同等量级。
  • 支持动态 batching:新请求可以随时加入 running batch,只要显存允许。

4.5 PagedAttention 的效果

根据 vLLM 论文(Kwon et al., 2023)的实验数据:

  • 相比 Orca(当时 SOTA),在相同延迟约束下,吞吐量提升 2~4 倍
  • 在并行采样场景中(共享前缀),KV Cache 显存占用降低为原来的 1/n(n 为并行路径数)。
  • 支持 连续批处理(Continuous Batching) 与动态增删请求,GPU 利用率显著提升。

五、其他关键优化技术

5.1 Multi-Query Attention(MQA)与 Grouped-Query Attention(GQA)

MQA(Shazeer, 2019):所有注意力头共享同一组 K、V 投影,只有 Q 维持多头。KV Cache 显存需求降至 1 / num_heads

GQA(Ainslie et al., 2023):介于 MQA 和全多头注意力之间的折中方案。K、V 分为少量组(如 8 组),每组被多个 Q 头共享。

以 Llama-2-70B 为例,它使用 GQA,将 key/value 头数从 64 减少到 8,KV Cache 额外节省 8 倍空间,同时保留大部分多头注意力的表达能力。

# MQA vs GQA vs MHA 的维度示意
# MHA: num_kv_heads = num_q_heads (e.g., 64)
# GQA: num_kv_heads = num_q_heads // group_size (e.g., 8)
# MQA: num_kv_heads = 1

# KV Cache 大小比例:MQA : GQA : MHA = 1 : 8 : 64(对于 Llama-2-70B)

5.2 滑动窗口注意力(Sliding Window Attention)

灵感来自 Longformer 和 Mistral 模型,滑动窗口注意力将每个 token 的注意力限制在局部窗口内(如左侧 4K token),实现了 O(w * n) 的线性复杂度。

Mistral-7B 证明了配合 FlashAttention,滑动窗口机制可以在不牺牲太多质量的前提下支持极长上下文。

5.3 投机解码(Speculative Decoding)

标准自回归解码每步只能生成一个 token,而每次前向传播的 GPU kernel 启动开销很大。投机解码的核心思想是:

  • 用一个小型「草稿模型」(draft model)快速生成若干个候选 token。
  • 用「目标模型」(target model,即实际部署的大模型)并行验证这些候选 token。
  • 接受所有匹配的 token,拒绝首个错误 token并重新采样。

Medusa:在目标模型上增加多个解码头,每个头负责预测未来 n 步的 token,无需额外草稿模型。

EAGLE:基于自回归特征的轻量级外推模型,生成质量更高的候选序列。

Lookahead Decoding:无需草稿模型,利用 n-gram 局部重复模式进行自投机,实现 < 2x 加速。

投机解码的理论上限是将解码步骤减少 1 + gamma 倍(gamma 为候选 token 数),前提是草稿模型足够快且准确。

六、工程实践与选型指南

6.1 性能测量

在实际部署中,以下指标值得关注:

# 关键性能指标测量示例
import torch
import time

def benchmark_attention(func, Q, K, V, warmup=10, repeats=50):
    # 预热
    for _ in range(warmup):
        _ = func(Q, K, V)
    torch.cuda.synchronize()
    
    # 正式测试
    start = time.perf_counter()
    for _ in range(repeats):
        _ = func(Q, K, V)
    torch.cuda.synchronize()
    elapsed = time.perf_counter() - start
    
    avg_ms = elapsed * 1000 / repeats
    # 计算 FLOPs: 2 * seq_len^2 * head_dim (QK^T + PV)
    flops = 2 * Q.size(0) * Q.size(1) * Q.size(2) * K.size(2)
    tflops = flops / (avg_ms / 1000) / 1e12
    
    print(f"平均耗时: {avg_ms:.3f} ms")
    print(f"有效算力: {tflops:.2f} TFLOPS")
    return avg_ms

# 显存分析
print(f"分配的显存: {torch.cuda.memory_allocated() / 1e9:.2f} GB")
print(f"预留的显存: {torch.cuda.memory_reserved() / 1e9:.2f} GB")

6.2 优化技术选型矩阵

场景推荐优化理由
长上下文 Prefill(>8K)FlashAttention-2/3降低 O(N^2) 显存,实际加速 2~4x
高并发推理服务vLLM + PagedAttention消除显存碎片,提升 batch size 2~4x
显存极度受限(边缘设备)MQA/GQA + FlashAttentionKV Cache 缩小 4~8 倍
超长文档(>100K)滑动窗口 + FlashAttention线性复杂度,适合局部相关性强的任务
低延迟交互场景投机解码(Medusa/EAGLE)减少解码步数,延迟降低 1.5~3x
多轮对话、共享前缀PagedAttention COW前缀 KV 共享,显存随对话轮数线性增长

6.3 与 TensorRT-LLM 和 vLLM 的集成

**TensorRT-LLM(NVIDIA)**已经在其内核实现中原生集成了 FlashAttention 和 PagedAttention:

  • trtllm-build 阶段启用 --gpt_attention_plugin 即可自动使用优化的多头注意力内核。
  • KV Cache 由 TensorRT-LLM 的 KVCacheManager 管理,支持 PagedAllocation 策略。
  • 适合部署在 NVIDIA GPU 上的生产环境,与 Triton Inference Server 无缝集成。

vLLM则是一个开源、框架无关的推理服务引擎:

  • 默认启用 PagedAttention,无需额外配置。
  • 通过 vllm.LLM 或 OpenAI-compatible API 提供服务。
  • 社区维护活跃,对多种模型架构(Llama、Qwen、Baichuan、Mixtral 等)支持良好。
# vLLM 使用示例
from vllm import LLM, SamplingParams

llm = LLM(model="meta-llama/Llama-2-7b-hf", 
          tensor_parallel_size=1,
          gpu_memory_utilization=0.9)

# PagedAttention 自动管理 KV Cache,支持长上下文
outputs = llm.generate("FlashAttention 的核心优化原理是", 
                       SamplingParams(max_tokens=512))

七、总结

LLM 推理中的注意力优化已经从单一的算法改进演化为全栈的系统工程

  • FlashAttention 通过 SRAM 分块计算和在线 Softmax,从根本上消除了巨大的 HBM 读写量,使长上下文注意力的内存复杂度从 O(N^2) 降至 O(N)。
  • PagedAttention 借鉴操作系统虚拟内存思想,以固定大小的 block 和页表机制管理 KV Cache,解决了显存碎片化、过度预留和无法共享的问题。
  • MQA/GQA、滑动窗口、投机解码等技术从不同角度进一步压缩了 KV Cache 体积或减少了解码步数。

在生产环境中,这些优化往往不是孤立使用的:vLLM 将 FlashAttention 内核与 PagedAttention 内存管理结合,TensorRT-LLM 在编译期融合多层优化。对于工程师而言,理解这些技术的原理与 trade-offs,才能根据实际场景(延迟敏感 vs 吞吐量优先、短 prompt vs 超长上下文、单用户 vs 高并发)做出正确的架构选择。

随着上下文窗口继续向百万级扩展,以及多模态模型(视觉-语言)的注意力维度进一步膨胀,注意力优化仍将是 LLM 系统工程中最活跃的研究和开发方向之一。

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「ai」更多文章

  1. 模型量化技术详解:INT8、FP16 与混合精度推理
  2. 模型剪枝与知识蒸馏:从压缩到加速全链路
  3. 推理引擎终极对比:TensorRT vs ONNX Runtime vs OpenVINO