在 16.2 节中我们提到,大模型训练与推理对 GPU 的需求存在本质差异:训练阶段是计算密集型,受限于前向/反向传播的海量矩阵运算;推理阶段则是显存带宽密集型,受限于逐 Token 生成时频繁的显存读写。而自注意力机制(Self-Attention)作为 Transformer 的绝对核心,恰好同时命中了这两个瓶颈——它的计算复杂度与显存占用都是序列长度的平方级(O(N²))。
传统 PyTorch 的 scaled_dot_product_attention 在标准实现下,需要将巨大的 Query-Key 相似度矩阵(N×N)和注意力权重矩阵完整地写入高带宽显存(HBM)。当序列长度达到 4K、8K 甚至 128K 时,HBM 的读写带宽迅速成为天花板,导致 GPU 的计算单元(Tensor Core)大量空转等待数据。FlashAttention 和 PagedAttention 分别从计算内核优化和显存管理策略两个维度,击穿了这个瓶颈。理解它们的原理,是你在千卡集群上调通长上下文训练、或在单卡上部署高并发推理服务的必修课。
一、FlashAttention:用“分块+重计算”打破内存墙
1.1 问题根源:HBM 读写远大于计算
在标准 Attention 的计算流程中(Q, K, V 均为 N×d 矩阵):
- 从 HBM 读取 Q, K,计算 S = QKᵀ,将 S(N×N)写回 HBM;
- 从 HBM 读取 S,计算 Softmax(P),将 P(N×N)写回 HBM;
- 从 HBM 读取 P 和 V,计算输出 O = PV,写回 HBM。
这里的 N×N 中间矩阵是致命伤。以 FP16、序列长度 8192 为例,仅一个 Attention Head 的 S 矩阵就占 128 MB;若模型有 32 个 Head、多层网络,中间激活值的显存峰值将轻松吞掉数十 GB HBM,而且大量时间耗费在“搬运”而非“计算”上。
1.2 核心思想:IO-Aware 的分块计算
FlashAttention 的核心洞察来自 2010 年代的循环分块(Tiling)思想,但它针对性地解决了 Softmax 的“全局归一化”难题。
传统矩阵乘法可以分块,因为每个输出块只依赖对应的输入块。但 Softmax 的分母需要整行的指数和,这似乎是全局操作。FlashAttention 通过 Online Softmax 技巧,将全局 Softmax 拆解为可增量更新的局部统计量:
- 每次只将一小块 Q 和一小块 K 加载到 GPU 片上高速缓存(SRAM,如 A100 每 SM 192 KB Shared Memory);
- 在 SRAM 内完成局部矩阵乘法,维护当前已见的局部最大值
m和局部指数和l; - 当新的分块加入时,利用代数变换更新全局
m和l,无需一次性看到整行; - 最终输出的块直接写回 HBM,全程不在 HBM 中存储完整的 N×N 中间矩阵。
反向传播时,标准实现需要缓存前向的 N×N 注意力矩阵以计算梯度。FlashAttention 选择重计算(Recomputation)策略:不保存巨大的中间激活,而是在反向传播时重新分块计算前向的 Attention 矩阵。这以少量额外计算换取了巨大的显存节省,而 GPU 上计算单元往往闲置,因此整体 wall-clock 时间反而更短。
1.3 版本演进与工程收益
| 版本 | 关键改进 | 工程意义 |
|------|---------|----------|
| FlashAttention | 提出分块 Tiling + Online Softmax + 重计算 | 将 Attention 显存从 O(N²) 降至 O(N),首次实现 2-4× 端到端加速 |
| FlashAttention-2 | 减少非矩阵乘法的 CPU 开销,优化 Warp 级并行,更好的序列并行 | 在 A100 上实现约 2× 于 FlashAttention-1 的吞吐;反向传播更稳定 |
| FlashAttention-3 | 针对 Hopper 架构(H100)利用 Tensor Core 的异步拷贝(TMA)、FP8 低精度、Warp Group Cluster | 进一步逼近硬件峰值,支持更长序列的极低延迟计算 |
实用落地信息:
- PyTorch 2.0+ 已将 FlashAttention 内核封装进
torch.nn.functional.scaled_dot_product_attention,默认在 CUDA 后端自动调用,无需手动安装。 - 在千卡训练集群中,DeepSpeed、Megatron-LM、Colossal-AI 均已原生集成 FlashAttention。开启后,长上下文(32K+)预训练不再因显存溢出(OOM)而被迫缩小批次。
- 局限:FlashAttention 对 head dimension 有要求(通常需 ≤ 128 或 256),且对变长序列(Padding 场景)需要额外的掩码处理;在极短序列(< 1K)下收益不明显,因为 HBM 带宽尚未成为瓶颈。
二、PagedAttention:用“虚拟内存”治理 KV Cache
如果说 FlashAttention 解决的是“计算时如何少搬数据”,PagedAttention 解决的是“推理时如何少占显存”。它并非训练加速库,而是推理服务层的显存管理革命。
2.1 问题根源:KV Cache 的显存浪费与碎片化
在自回归推理中,为了避免重复计算,模型会将每一层的 Key 和 Value 向量缓存起来(KV Cache)。对于长度为 N 的序列,KV Cache 的显存占用约为:
2 × 层数 × 头数 × 头维度 × 序列长度 × 批次大小 × 精度字节数
以 Llama-2-70B(80 层、8 KV Head、128 维)为例,单个 4K 序列的 FP16 KV Cache 约需 2 × 80 × 8 × 128 × 4096 × 2 ≈ 1.3 GB。
更致命的是分配策略问题:
- 静态预分配:传统框架(如早期 HuggingFace)为每个请求按最大可能长度(如 2048)预先分配一块连续显存。如果用户只输入 50 个 Token,剩余 1998 个位置就空占着。
- 碎片问题:不同请求长度不一,释放后产生大量不连续的小块空闲显存,导致后续大请求无法分配,尽管总空闲显存足够。
- 并行采样冗余:当使用 Beam Search 或并行生成多个候选时,同一个前缀的 KV Cache 被重复存储多份。
2.2 核心思想:操作系统式的页式内存管理
PagedAttention 借鉴了操作系统的虚拟内存与页表机制:
- 分块(Blocking):将 KV Cache 划分为固定大小的连续块(Block),例如每个 Block 存储 16 个 Token 的 K/V 向量。
- 按需分配:不再为整个序列预分配连续显存,而是生成多少个 Token,就动态申请多少个 Block。Block 之间在物理显存上不必连续。
- 块表(Block Table):为每个请求维护一个类似页表的映射结构,记录逻辑 Token 块到物理显存块的映射关系。Attention 计算时通过 Block Table 间接寻址。
- Copy-on-Write(COW):在 Beam Search 或并行采样时,多个候选序列共享相同的物理前缀 Block。只有当某个候选需要写入新的 KV 时才复制该 Block,极大减少冗余显存占用。
2.3 工程收益与落地
- 显存利用率提升:从传统方案的 40%–60% 提升到 90%+。这意味着在同等 GPU 显存下,推理服务的并发批次(Batch Size)可提升 2–4 倍,直接转化为更高的吞吐量(Throughput)。
- 消除长度预设:无需在启动时猜测
max_seq_len并预分配,支持真正的动态长度伸缩。 - 与 Continuous Batching 配合:PagedAttention 是 vLLM(见 23.1 节)的核心底座,两者结合使得推理引擎可以在一个批次内动态加入新请求、驱逐已完成请求,而无需等待整批结束。
实用落地信息:
- 如果你使用 vLLM、TensorRT-LLM 或 HuggingFace 的
Text Generation Inference(TGI),PagedAttention 已被内置。作为使用者,你只需关注 Block Size 的调参(默认 16,在长上下文场景可适当调大以减少块表开销)。 - 注意:PagedAttention 本身不改变 Attention 的计算复杂度,它解决的是显存分配策略。因此,若你的瓶颈是计算而非显存(如短序列高并发),PagedAttention 的收益会相对有限。
三、两者的关系与选型建议
| 维度 | FlashAttention | PagedAttention |
|------|----------------|----------------|
| 主攻阶段 | 训练 + 推理 | 推理(含 serving) |
| 优化目标 | 减少 HBM 读写,加速 Attention 计算 | 减少 KV Cache 显存占用,提升并发 |
| 技术层级 | CUDA Kernel / 计算优化 | 显存管理 / 系统层优化 |
| 直接收益 | 更快的长序列训练,更低显存峰值 | 更高的推理吞吐,支持更长上下文服务 |
| 典型集成 | PyTorch SDPA、DeepSpeed、Megatron | vLLM、TensorRT-LLM、SGLang |
互补而非替代:在生产级的推理服务中,最先进的系统(如 vLLM、SGLang)往往同时采用两者——用 FlashAttention 计算注意力,用 PagedAttention 管理 KV Cache 的显存生命周期。
给开发者的 checklist:
- 训练长上下文(>8K):检查你的训练框架是否启用了 FlashAttention。若未启用,长序列训练很可能因为激活值显存爆炸而无法运行。
- 部署高并发推理服务:确认你的推理引擎是否基于 PagedAttention。如果没有,你的 GPU 显存可能大量浪费在预分配和碎片上,导致实际吞吐远低于理论值。
- 硬件适配:FlashAttention-3 需要 Hopper(H100)才能发挥极致;Ampere(A100)和 Ada(RTX 4090)通常运行 FlashAttention-2 即可。PagedAttention 对架构代际要求较宽松,但显存越大收益越明显。
四、小结
FlashAttention 用分块 Tiling 与重计算绕开了 Attention 的 HBM 内存墙,让计算单元不再“饿着肚子等数据”;PagedAttention 用页式显存管理根治了推理 KV Cache 的碎片化与冗余,让宝贵的 HBM 不再被空洞的预分配吞掉。两者分别从 Kernel 计算层和系统管理层,定义了现代大模型 Attention 优化的工程基线。
在 17.4 节中,我们将继续深入 CUDA 生态,审视 TensorRT、cuBLAS、Cutlass 这些张量加速库如何在更广泛的矩阵运算层面,进一步压榨 GPU 的理论算力峰值。