大型语言模型(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)的算力增长远超内存带宽的增长。标准注意力实现的流程如下:
- 从 HBM 读取 Q, K, V
- 在 SRAM(片上静态随机存储器)中计算
QK^T - 将中间结果
S = QK^T写回 HBM - 从 HBM 读取 S,计算 Softmax 得到 P
- 将 P 写回 HBM
- 从 HBM 读取 P 和 V
- 计算
O = PV - 将输出 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 效果总结
| 指标 | 标准 Attention | FlashAttention |
|---|---|---|
| HBM 读写量 | O(N^2) | O(N) |
| 显存占用 | O(N^2) | O(N) |
| 典型加速比 | 1x | 2~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 = 80num_heads = 64(注意:key/value 头的数量可能因 GQA 而异,这里假设全头)head_dim = 128seq_len = 8192batch_size = 8dtype = 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 本身却在持续膨胀。这带来了两个新问题:
- 显存碎片化:不同序列长度、不同请求的 KV Cache 大小不一,导致存储不连续。
- 过度预留:为了支持最大上下文,系统通常按最坏情况(
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 中:
- 初始时,3 个请求共享 prompt 对应的物理 block。
- 每个请求维护独立的 block table,但指向相同的物理 block。
- 当某个请求生成新 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 的调度器在请求层面实现了一个精细的显存管理系统:
- Waiting Queue:新进入的请求在这里排队。
- Running Batch:当前正在 GPU 上执行的请求批次。PagedAttention 允许动态添加/移除请求,只要 block table 能管理。
- Swapping:当显存不足时,vLLM 可以将某些请求的 KV Cache block “换出”(swap out)到 CPU 内存;待显存释放后再 “换入”(swap in)。这类似于操作系统的内存交换。
- 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 + FlashAttention | KV 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 系统工程中最活跃的研究和开发方向之一。
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。