Transformer 模型在自然语言处理和多模态任务中的核心是自注意力机制。给定查询 Q、键 K 和值 V,标准计算会先得到注意力分数矩阵 S = QK^T,再经过 Softmax 后与 V 相乘。这个过程在序列长度为 N 时需要维护 N×N 的中间矩阵,其显存和访存开销成为长序列推理的瓶颈。以 N=4096、FP16 精度为例,单个注意力分数矩阵约占 32MB,多层堆叠后仅中间结果就可能占用数 GB。为了突破这一限制,业界出现了从不同层面切入的优化方案,其中 Flash Attention 聚焦计算过程中的 IO 效率,Paged Attention 聚焦推理服务中的 KV 缓存内存管理。下面从这两类机制展开分析。

一、标准注意力的显存与访存问题
标准自注意力的前向过程可以拆成三步:先计算 Q 和 K 的点积得到分数矩阵 S,再对每一行做 Softmax 得到权重矩阵 P,最后用 P 对 V 加权求和得到输出 O。当序列长度为 N 时,S 和 P 的形状都是 N×N。这个平方级的显存增长会迅速吃掉 GPU 的可用空间。以 N=2048 为例,一个注意力头在 FP16 下的 S 矩阵需要 2048×2048×2 字节,约 8MB;N 增加到 4096 时,矩阵大小直接变为约 32MB。多头注意力会进一步放大这个数值,而大型语言模型的层数通常有几十层,因此原始注意力在长序列场景下几乎不可用。
除了显存占用,访存模式也是性能瓶颈。GPU 的高带宽显存 HBM 容量大但带宽相对片上 SRAM 慢得多。标准实现会把完整的 S 和 P 矩阵写入 HBM,再从 HBM 读取出来参与后续计算。这种反复搬运并没有增加计算量,却消耗了大量时间。Flash Attention 的核心观察是:注意力计算并不需要在 HBM 中保存完整的 S 和 P,只需要最终输出 O 和用于反向传播的某些统计量。因此可以通过分块和重算来绕过完整矩阵的物化。
二、Flash Attention的分块计算与在线Softmax
Flash Attention 将 Q、K、V 切分成较小块,每次只把当前需要的块从 HBM 加载到 SRAM 中,在片上完成该块对应的注意力计算,再更新输出。普通 Softmax 需要先知道整行的最大值与所有元素的和,而分块之后某一行的数据被拆到多个块里,无法一次性获得全局最大值与求和。Flash Attention 使用在线 Softmax 技巧维护两个运行状态:当前所见的最大值 m 和累加和 l。当处理一个新的 K、V 块时,先计算局部分数,再用新的最大值修正之前累积的分母,并把旧输出按新旧最大值之差进行缩放。
这种方法在数学上等价于原始 Softmax,但避免了完整 N×N 矩阵的写入和读取。为了支持反向传播,Flash Attention 不缓存中间注意力矩阵,而是选择在前向时保存 Softmax 的归一化统计量,在反向时重算 S 和 P。虽然增加了一些重算代价,但整体显存和带宽收益远大于额外计算,因此训练和推理都能受益。PyTorch 从 2.0 开始集成 Flash Attention,调用方式如下:
import torch import torch.nn.functional as F # q, k, v 的形状为 (batch, heads, seq_len, head_dim) # scaled_dot_product_attention 会根据硬件和输入自动选择 Flash Attention 内核 out = F.scaled_dot_product_attention(q, k, v, dropout_p=0.0, is_causal=True)
如果想了解内部逻辑,下面是一个简化后的分块注意力前向伪代码:
def flash_attention_forward(Q, K, V, scale, block_size=128):
# Q, K, V: [N, d]
N = Q.shape[0]
O = torch.zeros_like(Q)
m = torch.full((N,), -float('inf'), device=Q.device)
l = torch.zeros(N, device=Q.device)
for start in range(0, N, block_size):
end = min(start + block_size, N)
Kj = K[start:end] # [B, d]
Vj = V[start:end] # [B, d]
S = Q @ Kj.T * scale # [N, B]
m_new = torch.maximum(m, S.max(dim=-1).values)
p = torch.exp(S - m_new[:, None])
l_new = torch.exp(m - m_new) * l + p.sum(dim=-1)
O = O * torch.exp(m - m_new)[:, None] + p @ Vj
m = m_new
l = l_new
O = O / l[:, None]
return O
这个实现本质上就是 Flash Attention 想要达到的数值等价效果。它没有产生 N×N 的完整中间矩阵,而是通过循环逐个块更新输出。实际工程实现还会结合张量核心、共享内存分配和寄存器优化,使计算效率接近理论峰值。
三、Paged Attention的KV缓存分页管理
在自回归解码阶段,模型每生成一个新 token 都会依赖之前所有 token 的 Key 和 Value。为了避免重复计算,推理框架会把每一层的 KV 向量保存为缓存,后续步骤直接从缓存读取。问题在于,不同请求的生成长度不同,有的请求可能先结束,有的请求持续生成很长内容。如果为每个请求预分配一段连续的 KV 缓存,短请求结束后会留下无法利用的空洞,新请求又需要新的连续空间。随着并发请求增加,显存碎片会越来越严重,导致可用空间远低于实际剩余空间。
Paged Attention 借鉴操作系统的分页机制,将 KV 缓存划分成固定大小的块,例如每块包含 16 或 32 个 token。每个请求只保留一个页表,记录自己的逻辑位置对应哪些物理块。物理块可以离散地分布在显存中,请求之间不再要求连续分配。当请求长度增加时,只需要分配新的物理块并更新页表。请求结束后,整页可以立即回收给新请求使用。这种管理方式几乎消除了显存碎片,提升了显存利用率和系统吞吐。
下面是一个简化的 KV 缓存块分配器伪代码,体现页表管理的基本思路:
class PagedKVCache:
def __init__(self, num_blocks, block_size, num_heads, head_dim):
self.num_blocks = num_blocks
self.block_size = block_size
self.num_heads = num_heads
self.head_dim = head_dim
# 每个物理块同时保存 K 和 V,形状: [num_blocks, 2, num_heads, block_size, head_dim]
self.physical_blocks = torch.zeros(num_blocks, 2, num_heads, block_size, head_dim)
self.free_blocks = list(range(num_blocks))
def allocate(self, num_tokens):
needed = (num_tokens + self.block_size - 1) // self.block_size
if len(self.free_blocks) < needed:
raise RuntimeError('No enough KV cache blocks')
blocks = [self.free_blocks.pop() for _ in range(needed)]
return blocks
def release(self, blocks):
self.free_blocks.extend(blocks)
实际系统中的页表还会记录每个块的有效 token 数,并在 attention 计算时根据页表把物理块组织成逻辑序列,再送给 Flash Attention 内核。vLLM 是较早实现 Paged Attention 的推理框架,它可以在接近零浪费的情况下管理动态增长的 KV 缓存,使单一 GPU 能同时服务更多请求。
四、两者协同与工程落地
Flash Attention 与 Paged Attention 的优化目标不同,但可以在推理系统中协同工作。Flash Attention 解决的是单次注意力计算中的显存与带宽压力,让长序列推理的内存峰值显著下降;Paged Attention 解决的是多请求并发时 KV 缓存的组织方式,减少碎片对显存总量的浪费。前者更偏算子级优化,后者更偏系统级内存管理。两者结合后,推理框架可以在有限的显存下容纳更大的 batch size 和更长的上下文窗口。
从工程选择角度看,是否需要两者同时启用取决于部署场景。单用户长文本生成通常首先受益于 Flash Attention,因为它直接降低单次计算的高峰显存;高并发在线服务则更依赖 Paged Attention 的碎片治理。如果一个框架使用 Paged Attention 但底层 attention 仍采用传统算子,虽然缓存分配更紧凑,单个请求的长序列计算依然可能触发显存高峰。因此主流推理方案通常同时集成两种机制。下表简要对比三者的定位:
| 方案 | 主要目标 | 核心手段 | 适用阶段 |
|---|---|---|---|
| 标准注意力 | 正确计算输出 | 完整物化 N×N 矩阵 | 训练与推理 |
| Flash Attention | 降低单次计算 IO | 分块与在线 Softmax | 训练与推理 |
| Paged Attention | 减少 KV 缓存碎片 | 分页映射管理缓存块 | 高并发推理 |
在实践中,开发人员可以为 PyTorch 训练和推理直接启用 Flash Attention,并在使用 vLLM、TensorRT-LLM 等推理框架时关注 Paged Attention 的块大小与显存占用。理解这两种优化手段的边界,有助于在模型超长上下文、批处理服务和低显存部署之间找到合适的平衡。
Flash AttentionPaged Attention注意力机制优化修改时间:2026-10-01 19:06:53