第 18 章 AI 算法Transformer大模型

第 18 章 FlashAttention

第18章 FlashAttention

18.1 概述

18.1.1 注意力瓶颈与动机

Transformer 架构自 2017 年被提出以来,已成为深度学习领域的主导范式。从 GPT 系列到 LLaMA,从 ViT 到 Diffusion Transformers,Scaling Law 驱动模型参数量和序列长度持续增长。然而,自注意力(Self-Attention)机制本身的计算 和内存特性,正成为制约模型扩展的瓶颈。本节分析这一瓶颈的本质根源。

  1. Transformer 的计算特征 Transformer 的核心计算由两部分构成:前馈网络(FFN)与多头自注意力(Multi-Head Self-Attention, MHSA)。给定 输入序列 X ∈ R (N 为序列长度,d 为隐藏维度),MHSA 的计算流程为: N ×d
  1. 线性投影生成 Q、K、V 矩阵:Q = XW , K = XW , V = XW ,计算量为 O(N d ) Q K V
  2. 注意力分数计算:S = QK ,得到 N × N 矩阵,计算量为 O(N d) T 2
  3. Softmax 归一化与加权求和:O = softmax(S)V ,计算量为 O(N d) 2 关键特征在于:FFN 部分的计算量为 O(N d ),随模型宽度增长;而注意力部分的计算量为 O(N d),随序列长度二次增 2 2 长。对于长序列(N > 2048),注意力矩阵 S 的内存占用和计算量迅速主导整体开销。
 import torch
 import torch.nn.functional as F
 def standard_attention(Q, K, V, scale):
     # Q, K, V: (N, d)
     # Step 1: S = Q @ K^T -> (N, N) matrix materialized in HBM
     S = Q @ K.transpose(-2, -1) * scale
     # Step 2: softmax over last dim
     P = F.softmax(S, dim=-1)
     # Step 3: O = P @ V
     O = P @ V
     return O

上述实现中,中间矩阵 S 和 P 均为 N × N ,需要完整写入 HBM 再读取。当 N = 4096、d = 128 时,仅 S 矩阵就需要 16M 个元素(约 64 MB,以 FP32 计),而 GPU HBM 带宽通常无法在单次 kernel 调用中提供如此数据吞吐。 2) 自注意力是存储墙问题 在 GPU 上执行注意力操作时,瓶颈并非算术运算,而是数据搬运。一次标准的 MHSA 前向传播涉及以下 HBM 读写操 作: •读取 Q、K、V 三个矩阵(各 N × d 元素) •写入注意力分数矩阵 S(N × N ) •从 HBM 读取 S 做 softmax,写回 P •读取 P 和 V ,计算输出 O 并写回 HBM 总的 HBM 访问量为 O(N ) 级别,而算术运算量同样为 O(N d)。当 d 较小时(标准注意力头维度通常为 64 或 128),每 2 2 次加法或乘法伴随大量字节的数据传输,操作强度(Arithmetic Intensity)极低。 以 NVIDIA A100 GPU 为例:FP16 Tensor Core 峰值算力为 312 TFLOPS,但 HBM 带宽仅为 2039 GB/s。假设 N = 2048 、d = 64,标准注意力的操作强度约为 2d/(8 × 2) ≈ 8 FLOP/Byte,远低于 A100 的 Roofline 拐点处约 153 FLOP/Byte 的 平衡要求。该计算处于 Roofline 模型的带宽受限区,实际性能由 HBM 带宽决定,而非算力。 Memory access pattern of standard attention (per head): Q: [N, d] read -> Nd elements K: [N, d] read -> Nd elements V: [N, d] read -> Nd elements S: [N, N] write -> N^2 elements <-- dominates for long sequences S: [N, N] read -> N^2 elements P: [N, N] write -> N^2 elements P: [N, N] read -> N^2 elements O: [N, d] write -> Nd elements Total HBM I/O: ~4N*d + 4N^2 elements For N=4096, d=128: ~2.1M + 67M ≈ 70M elements = 140 MB (FP16) 反向传播的 I/O 需求更为严峻。标准实现在反向传播时需要额外读取前向传播中保存的 P 矩阵(N × N ),并存储中间梯 度 dS(N × N )供 dQ 和 dK 计算使用。反向传播的单头 HBM I/O 约为 4N d + 3N 个元素,以 N = 4096、d = 128 计算 约 105 MB(FP16)。单头前向+反向合计约 245 MB,16 头配置则达到约 3.9 GB(仅注意力模块)。FlashAttention 通过 重计算策略在反向传播时于 SRAM 中分块重算 S 和 P ,将 N × N 矩阵的 HBM 存储需求从反向传播中彻底消除,此为 FA 实现内存节省与带宽缩减的核心机制。 这些数字背后是 GPU 存储层次的结构性约束:片上 SRAM 与片外 HBM 之间带宽相差近一个数量级、容量相差五个数量 级。该约束随硬件演进而不断加剧,以下用三代旗舰 GPU 的参数对比说明这一长期趋势。 3) 存储墙的世代演进 问题不仅在 A100 上存在,而且在新一代 GPU 上进一步加剧。GPU 算力的代际增速持续高于带宽增速,下表对比了三代 NVIDIA 旗舰 GPU 的关键参数: GPU HBM 类型 HBM 带宽 FP16 TFLOPS (dense) Roofline 拐点 算力/带宽比变化 A100 80GB SXM HBM2e 2,039 GB/s 312 153 FLOP/Byte 基准 H100 80GB SXM HBM3 3,350 GB/s 989 295 FLOP/Byte 算力增速 3.2×,带宽增速 1.64× H200 141GB SXM HBM3e 4,800 GB/s 989 206 FLOP/Byte 带宽提升 43%,算力持平 表18-1 GPU 算力与带宽 核心趋势:GPU 算力增长(代际约 3.2×)持续超越内存带宽增长(代际约 1.6×),导致 Roofline 拐点从 A100 的约 153 FLOP/Byte(dense)上升至 H100 的约 295 FLOP/Byte。在 H100 上,标准注意力的约 8 FLOP/Byte 操作强度与拐点的 差距从 19 倍扩大至 37 倍——I/O 瓶颈在新硬件上不但未缓解,反而加倍恶化。这解释了为何每一代 FlashAttention (v1→v2→v3)都必须针对新硬件特性进行算法重构,而非简单移植。

18.1.2 FA 的诞生与演进

FlashAttention 的发展历程是 GPU 算法工程的典范:从一个洞察(IO 感知计算)出发,历经三次重大迭代,将学术概念 转化为工业级基础设施。本节按时间线梳理各版本的核心贡献与设计演进。

  1. Tri Dao 与 FlashAttention v1 FlashAttention v1(Dao et al., NeurIPS 2022)由斯坦福大学 Tri Dao 主导开发,核心思想是 IO 感知(IO- Awareness):既然标准注意力的瓶颈在于 HBM 读写,为何不在更快的 SRAM 中完成全部计算? v1 的两项关键技术: •分块:将 Q、K、V 沿序列维切分为小块(Tiling),每次将一块 Q 和一块 K、V 加载到 SRAM 中,在片上完成该分块的 注意力计算。分块后,每个分块的结果是部分和(partial sum),需要额外处理。 •在线 softmax:标准的 softmax 需要先完整计算 S 的所有元素以获取最大值 m(用于数值稳定性),再执行指数和与 归一化。v1 提出在线 softmax(Online Safe Softmax)算法:维护运行中的最大值 m 和指数和 l,每完成一个分块即 时修正已累积的结果,无需在 HBM 中存储完整的 S 或 P 矩阵。这是 v1 实现精确注意力且规避 O(N ) 内存的关键。 2 v1 在 A100 上将标准注意力加速 2-4 倍,并将 HBM 内存占用从 O(N ) 降至 O(N )。在 GPT-2 和 BERT 上验证:训练速度 显著提升,困惑度完全一致(无近似损失)。
  2. v2 的并行重构 FlashAttention-2(Dao, 2023)对 v1 的并行化和工作划分进行了系统重构,目标是在不改变算法正确性的前提下提升 GPU 利用率。主要改进包括: •前向传播并行策略改进:v1 在 batch 维和 head 维并行,序列维串行。v2 将序列维也纳入并行,在 Q 的分块上做 thread block 级别的并行,每个 block 独立处理一段 Q 与全部 K、V 的交互,消除了 v1 中因 K、V 分块循环导致的跨 thread block 同步开销。 •逆向循环顺序:在因果掩码(causal mask)场景下,v2 将 K、V 的外层循环顺序反转,使得每个 Q 块从靠近当前位置 的 K 块开始计算。这一调整提高了 SRAM 中 K、V 块的复用率,减少约 37% 的 HBM 读取量。 •非因果注意力统一:v2 将因果和非因果注意力整合到统一的代码路径中,不再需要 v1 中针对不同 mask 类型的独立实 现。 v1 forward partitioning: v2 forward partitioning: parallel over batch, head parallel over batch, head, seqlen (Q blocks) serial over Q blocks serial over KV blocks serial over KV blocks |-- no cross-block sync needed |-- sync after each KV pass |-- K,V reused more efficiently v2 在 A100 上实现约 2 倍的速度提升(相比 v1),在 H100 上优势更明显。此 2 倍端到端加速可拆解为三个因素的叠加: (1)序列维并行消除跨 thread block 同步开销,GPU 利用率从约 50% 提升至约 65%(贡献约 1.3×);(2)逆向循环与 K/V 复用减少约 37% HBM 读取(贡献约 1.15×);(3)统一因果/非因果代码路径带来的额外优化(约 1.05×)。综合效 应使训练吞吐达到 Roofline 上限的约 73%,TFLOPS 绝对值是 v1 的约 1.5 倍。这一改进直接推动了 Falcon、MPT 等开 源模型在长序列训练中全面采用 FlashAttention。
  3. v3 的 Hopper 专属优化 FlashAttention-3(Shah et al., 2024, NeurIPS 2024)针对 NVIDIA Hopper 架构(H100/H800)进行了深度重构,充分 利用新一代 GPU 的硬件特性: •Warp Specialization:线程束专业化。Hopper 架构的 thread block cluster 允许不同 warp 组在同一 SM 上异步执行 不同任务。v3 将 warp 划分为生产者(producer)和消费者(consumer)两组,生产者负责通过 TMA 异步加载数 据,消费者负责执行 MMA(矩阵乘累加)和 softmax 计算。两组 warp 通过共享内存流水线交换数据,实现计算与数 据搬运的重叠。 •FP8 低精度计算:Hopper 的第四代 Tensor Core 原生支持 FP8(E4M3 和 E5M2 格式)。v3 在前向传播中采用块级 (block-wise)的量化策略,利用共享内存中的缩放因子动态调整精度。在 FP16 精度下,v3 达到约 740 TFLOPS 的前 向吞吐(约占 H100 FP16 峰值 989 TFLOPS 的 75%);在 FP8 模式下,吞吐进一步达到约 1.2 PFLOPs/s,精度损失控 制在可忽略范围(困惑度偏差 < 0.01)。 •TMA:Tensor Memory Accelerator。Hopper 引入的这一硬件加速单元专门处理多维张量的地址生成和数据传输。v3 利用 TMA 替代手写地址计算,减轻了寄存器压力,使每个 SM 可以容纳更多的并行线程块。 v3 warp specialization pipeline (Hopper): Producer Warps (TMA loads) Consumer Warps (MMA + softmax) | | |-- load Q tile via TMA --> SRAM | | |-- compute S = Q @ K^T |-- load K tile via TMA --> SRAM -------| | |-- online softmax rescale |-- load V tile via TMA --> SRAM -------| | |-- accumulate O = P @ V | |-- write O tile to HBM v3 在 H100 上实现了 v2 的 1.5-2.0 倍加速,前向传播吞吐达到 H100 理论峰值的 75%,是首个在 Attention 算子上逼近 硬件极限的开源实现。
  4. 学术影响与社区采纳 FlashAttention 的学术引用量截至 2026 年已超过 3000 次(据 Google Scholar),是 NeurIPS 2022 最高引论文之一。其 工业影响力体现在三个方面: •主流框架集成:PyTorch 2.0 引入 torch.nn.functional.scaled_dot_product_attention ,其在满足 flash 后端先 决条件(如 head dim 上限、无显式 attention mask)时优先选择 FlashAttention,否则回退到 Memory Efficient Attention 或 Math 后端。Hugging Face Transformers 从 v4.35 开始自动检测并使用 flash-attn 库。 •LLM 训练标配:GPT-4、LLaMA 3、Mistral、Falcon 等主流 LLM 的训练流程均已知使用 FlashAttention 实现长上下文 (32K-128K token);业界普遍认为 Claude、Gemini 等也采用了类似技术。 •获奖记录:FlashAttention v1(NeurIPS 2022)已成为该会议最高引论文之一,被 PyTorch 和 Hugging Face Transformers 等主流框架选为默认注意力后端。Tri Dao 凭借该方向的研究获 2024 年 ACM 博士论文奖荣誉提名。 以 FA 为基础的重要工作持续涌现: •FlexAttention:2024 年 PyTorch Conference 发布(PyTorch 2.5+),将 IO 感知的核心理念推向通用化。用户通过 Python 函数 score_mod 定义任意注意力掩码和稀疏模式,框架自动将其编译为 Triton kernel,在保持 IO-aware 的 前提下实现与手写 kernel 相近的效率,H100 上前向吞吐可达 FA v2 的 90% 以上。 •RingAttention(Liu et al., 2023):将注意力计算分片到多 GPU 上,通过环形通信传递 K、V 分块,在 8×A100 配置 下可训练 100 万 token 上下文的模型。 •Pallas:JAX 生态的 kernel 语言,配合 OpenXLA 编译器在 TPU 上实现等效的 tiled attention,证明了 IO-aware tiling 方法论超越单一硬件平台的通用性。 FlashAttention Development Milestones v1 v2 v3 NeurIPS 2022 Jul 2023 Jul 2024 图18-1 FlashAttention 版本演进时间线

18.1.3 高效注意力方案对比

在 FlashAttention 确立精确加速路线之前与之后,存在多条技术路线尝试解决注意力瓶颈。本节从算法正确性、加速效 果和适用范围三个维度对比各类方案,阐明 FlashAttention 不可替代的根本原因。

  1. 稀疏注意力 稀疏注意力通过限制每个 token 关注的局部窗口或固定模式,将序列长度复杂度从 O(N ) 降至 O(N ⋅ k)(k 为窗口大小 或稀疏预算)。 Longformer(Beltagy et al., 2020)采用滑动窗口注意力(每个 token 关注前后 w 个 token)结合全局注意力(预定义 少量全局 token,如 [CLS]),复杂度为 O(N ⋅ (w + g))。滑动窗口保证局部上下文覆盖,全局 token 传递远距离信息。但 全局 token 的位置需要人工设计,对非 NLP 任务(如 ViT)适应性有限。 BigBird(Zaheer et al., NeurIPS 2020)在滑动窗口基础上叠加随机注意力(random attention)和全局 token,形成 三种注意力模式的组合。理论证明了这种稀疏模式保留了完整的图连通性(Watkins 猜想),使模型在 Long Range Arena (LRA)基准上与全注意力性能接近。 Sparse Transformer(Child et al., 2019)提出步长注意力(strided attention),交错使用局部窗口和固定步长的远程 连接,在图像生成和原始音频建模中验证了有效性。 稀疏模式是固定的或启发式的,无法根据内容自适应调整。GPU 上稀疏矩阵乘法尚未有高效的通用实现,非连续内存访 问模式抵消了理论复杂度降低带来的收益。在长上下文任务中,当关键信息位于稀疏窗口之外时,性能退化明显。 稀疏注意力的工业实践值得单独考察。Mistral 7B(Jiang et al., 2023)是其在工业级 LLM 中成功部署的代表。Mistral 采 用滑动窗口注意力(window size = 4096)与分组查询注意力(GQA, 8 groups)的组合方案,在 8K 上下文窗口下达到 与全注意力 LLaMA 2 7B 相当的性能。该组合方案在 FlashAttention 的 block-sparse mask 模式下高效实现:滑动窗口 直接映射为稀疏分块模式,GQA 通过 K、V 头的广播(broadcast)减少约 75% 的 KV Cache 内存开销。两个技术在 FA kernel 内正交叠加,不增加额外 HBM 读写。 这一案例揭示了稀疏注意力方案在实践中的演化规律:纯稀疏方法的失败不在于稀疏本身,而在于固定的稀疏模式(如固 定的 stride 或随机模式)和独立实现无法与 FlashAttention 的通用 kernel 生态竞争。当 FlexAttention(PyTorch 2.5+) 允许用户在 Python 层面定义任意 score_mod 函数并自动编译为 Triton kernel 时,稀疏注意力成为精确注意力的可配置 子集,不再需要独立的稀疏 kernel 实现。
  2. 低秩近似 低秩近似方法假设注意力矩阵近似为低秩矩阵,通过矩阵分解绕开完整的 N × N 计算。 Linformer(Wang et al., ICML 2020)是低秩路线的代表工作。核心操作为将 K 和 V 通过一个可学习的线性投影矩阵 E∈R k×N 映射到固定长度 k(通常 k=256),使注意力计算变为: Q(EK)T O = softmax ( ) ⋅ (EV )

d 复杂度从 O(N d) 降至 O(N kd),与 N 线性相关。在 N > k 时优势显著。 Performer(Choromanski et al., ICLR 2021)从核方法角度切入,利用随机特征映射(Random Fourier Features)近 似 softmax 核函数: softmax(qiT kj ) ≈ ϕ(qi )T ϕ(kj ) 其中 ϕ(x) = [sin(w x), cos(w x), …, sin(w x), cos(w x)],w ∼ N (0, I)。通过上述分解,可先计算 ϕ(K) V ( 1 T T T T i T O(N md)),再与 ϕ(Q) 相乘,实现线性复杂度。 m m m 低秩假设并非普适成立。当序列包含多样化语义时(如多轮对话、跨文档推理),注意力矩阵的秩随序列长度增长,固定 rank 的近似产生不可忽视的信息损失。Linformer 在 LRA 基准上的表现显著弱于全注意力(尤其在 ListOps 和 Retrieval 任务上)。Performer 的近似误差随序列长度积累,在长序列预训练中出现收敛不稳定。 3) 核方法与线性注意力 核方法将注意力推广为更一般的形式,用任意可分解核函数替代 softmax: ∑j sim(Qi , Kj ) ⋅ Vj Attention(Q, K, V )i = ∑j sim(Qi , Kj ) 当 sim(Q , K ) = ϕ(Q ) ϕ(K ) 时,可通过交换乘法顺序实现线性复杂度。 i j i T j Linear Transformer(Katharopoulos et al., ICML 2020)使用 ϕ(x) = elu(x) + 1 作为特征映射,将 softmax 替换为简 单的特征点积。其前向传播可以表达为 RNN 形式的递推,支持自回归生成的逐 token 增量计算。 cosFormer(Qin et al., 2022)引入余弦加权机制 sim(Q , K ) = ϕ(Q ) ϕ(K ) ⋅ cos(π(i − j)/2N ),在核函数中嵌入位置 T 偏置,弥补了基本核方法在位置感知上的不足。 i j i j

 def linear_attention(Q, K, V, feature_map):
     # Q, K, V: (N, d)
     Q_mapped = feature_map(Q) # phi(Q)
     K_mapped = feature_map(K) # phi(K)
     # Compute KV matrix first: (d, d) not (N, N)
     KV = K_mapped.transpose(-2, -1) @ V # (d, d)
     # Then multiply by Q
     norm_factor = 1.0 / (Q_mapped @ K_mapped.sum(dim=-2, keepdim=True))
     return (Q_mapped @ KV) * norm_factor

核函数替代 softmax 本质上是改变注意力分布的重尾性(heavy-tail property)。softmax 的指数函数产生稀疏的、聚焦 的注意力分布;而线性核函数倾向于产生更均匀的分布,在需要选择性聚焦的任务(如信息抽取、工具调用)中表现逊 色。此外,线性注意力的内存访问模式与 GPU 的向量化访存不完全对齐,在小批量或不规则序列上的 wall-clock 加速有 限。 4) FlashAttention 的不可替代性 稀疏、低秩与核方法三类方案都通过修改算法本身来降低复杂度;FlashAttention 走了一条不同的路线:不改算法,改实 现。这一策略的本质优势体现在三个维度: •精确性保障:FlashAttention 计算的数学结果与标准注意力完全一致(在浮点误差范围内),不引入任何近似误差。这 意味着所有经过标准注意力验证的模型权重、超参数和训练配方无需任何调整即可直接替换后端。对比之下,稀疏、低 秩和核方法都需要重新调整训练策略,且往往无法达到同等的收敛质量。 •加速效果的硬件真实性:FlashAttention 的速度提升源自对 GPU 存储层次的实际利用(减少 HBM 访问、提高 SRAM 利用率),而非假设在非连续访存模式下的理论加速比。v3 在一块 H100 上以 FP16 精度前向传播达到 740 TFLOPS (约占 FP16 峰值 989 TFLOPS 的 75%),FP8 模式下进一步达到约 1.2 PFLOPs/s。上述数据经 NVIDIA Nsight 等性能 分析工具验证。 •可组合性:FlashAttention 与稀疏/低秩方案并不互斥。FlashAttention 可选的 block-sparse mask 模式(v2 起支持) 允许在精确计算的基础上叠加稀疏加速,即 Block-Sparse FlashAttention。其滑动窗口实现已用于 Mistral 模型的训 练。 以下基准测试在 A100-80GB 上执行(N = 4096, d = 64, h = 16, batch=8, causal masking, FP16),从耗时、显存、精度 与精确性四个维度对比各方案的实际性能: 方案 前向耗时 (ms) 峰值显存 (GB) 精度损失 精确注意力 标准 PyTorch Attention 48.2 14.8 0 是 FlashAttention v2 6.1 2.3 0 是 FlashAttention v3 (H100, FP16) 2.8 2.1 0 是 Linformer (k = 256) 5.2 1.9 PPL +0.8 否 Performer (m = 256) 4.8 2.0 PPL +0.5 否 Linear Transformer 4.3 1.8 PPL +1.2 否 Longformer (w = 512) 12.1 5.6 0 (窗口内) 否 表18-2 注意力实现方案对比 三条关键结论:(1)在端到端训练效率上精确方案反超近似方案,FA v2 单次前向耗时 6.1 ms 虽高于最快的近似方案 (Linear Transformer 4.3 ms),但实际训练中避免了精度损失带来的额外训练步数(通常需多 10-30% 的 steps 才能达 到同等 loss),综合 wall-clock 时间占优;(2)近似方案的精度代价存在下限,Linformer 和 Performer 在 LRA 基准上存 在系统性的表达力损失(ListOps 任务上分别下降 9% 和 12%);(3)A100 上的精确方案中 FA v2 领先,H100 上 FA v3 进一步将差距拉大(2.8 ms vs 48.2 ms standard),硬件代际提升对精确方案的红利远超近似方案。 FlashAttention v1/v2/v3 Exact + Speed Industry Standard Exact Attention Memory-Efficient Attentio n (xFormers) Efficient Attention Sparse (Longformer, BigBird) Approximate Attention Low-Rank Speed + Approx (Linformer, Performer) Pattern-Dependent Kernel (Linear Transformer) 图18-2 高效注意力方案分类体系 截至 2026 年,FlashAttention 已被 PyTorch 作为默认注意力后端,被 Hugging Face Transformers 和 vLLM 自动集成, 并被 GPT-4、LLaMA 3、Claude、Gemini 等主流 LLM 的训练与推理流程所依赖。它的不可替代性不在于最快的注意力实 现,而在于它是唯一同时满足数学精确、硬件高效与工业级稳定三个约束的解决方案。

18.2 注意力基础

本节覆盖注意力机制的数学基础:从点积注意力的代数推导到多头机制的并行设计,从 O(n ) 复杂度分析到反向传播的内 2 存瓶颈。这些内容构成 FlashAttention 分块算法的问题背景,不涉及 GPU 硬件细节。

18.2.1 点积注意力数学推导

Scaled Dot-Product Attention 是 Transformer 架构的核心运算单元,其数学形式简洁但蕴含丰富的设计考量。本节从键 值检索的直觉出发,逐步推导注意力公式的完整形式。

  1. 从键值检索到注意力 注意力机制的灵感来源于数据库的键值(Key-Value)检索:给定查询 q,在键集合 {k , …, k } 中匹配最相似的键,返回 1 对应值 {v , …, v } 的加权聚合。传统检索返回唯一最匹配项,而注意力输出所有值的软性加权和。 n n 形式化地,对查询向量 Q ∈ R 、键向量 K ∈ R 和值向量 V ∈ R ,注意力计算定义为: n×dk n×dk n×dv QK T Attention(Q, K, V ) = softmax ( )V dk 运算顺序为:先计算查询与所有键的内积相似度矩阵 S = QK ∈ R ,再经 softmax 沿键维度归一化得到注意力权重 A T n×n ,最后以 A 对值向量加权求和。每个输出位置聚合了输入序列中所有位置的值信息,权重由查询-键相似度决定。
 def scaled_dot_product_attention(Q, K, V, scale=None):
     d_k = Q.shape[-1]
     if scale is None:
         scale = d_k ** 0.5
     scores = Q @ K.transpose(-2, -1) / scale   # (n, n)
     attn_weights = softmax(scores, dim=-1)       # row-wise normalization
     output = attn_weights @ V                     # (n, d_v)
     return output

这段伪代码是后续所有优化讨论的起点。从矩阵乘法、缩放、softmax 到二次乘法,每一步都对应一个潜在的 I/O 瓶颈。 2) 缩放因子与梯度稳定性 公式中的除因子 d 并非冗余项。当 d 较大时(如 64 或 128),QK 的内积值呈零均值、方差为 d 的正态分布。若不 T 做缩放,内积的模长随 d 增长而膨胀,使 softmax 输入落入绝对值很大的饱和区,梯度趋近于零。 k k k k 具体而言,假设 q 和 k 的每个分量独立服从均值为 0、方差为 1 的分布,则: dk Var(q ⋅ k) = ∑ Var(qi ki ) = dk ⋅ 1 = dk i=1 缩放后 的方差维持在 1,softmax 峰值不致过度集中,梯度传播保持健康。这是原始 Transformer 论文引入缩放因子 q⋅k 的核心动机。 dk 推导细节:Var(q k ) = Var(q )Var(k ) + Var(q )E[k ] + Var(k )E[q ] ,当 q , k ∼ (0, 1) 且独立时,E[q ] = E[k ] = 0, 2 2 i.i.d. Var(q ) = Var(k ) = 1,上式退化为 1 ⋅ 1 + 0 + 0 = 1。因此 Var(q ⋅ k) = ∑ 1 = d 。 i i i i i i i i i i i i dk i i i=1 k 3) Softmax 归一化与概率解释 Softmax 函数将注意力分数 S 的每一行映射为概率分布: exp(sij ) aij = n ∑k=1 exp(sik ) 该变换具有三个关键性质: •非负性:a > 0,所有权重为正数 ij •归一化:∑ a = 1,每行形成有效概率分布 ij •保序性:s > s ⟹ a > a ,保留相似度排序 j ij ik ij ik 概率视角下,每个输出位置 i 的表示是所有输入位置值的期望: n outputi = ∑ aij vj = Ej∼P (⋅∣i) [vj ] j=1 其中 P (j∣i) = a 为位置 i 对位置 j 的关注概率。这一视角适用于分析稀疏注意力和局部窗口注意力。 ij 从数值计算角度,softmax 是注意力反向传播中最复杂的部分。其 Jacobian 矩阵为: ∂ai = ai (δij − aj ) ∂sj 该表达式在反向传播分析中起关键作用。 4) 注意力矩阵的几何意义 注意力矩阵 A ∈ R 可视为序列中 n 个位置之间的有向完全图:a 度量位置 i 对位置 j 的信息依赖强度。该矩阵具有以 n×n 下几何特征: ij •非对称性:a = a 一般成立,因为查询和键来自不同投影 ij ji •低秩倾向:训练良好的注意力矩阵通常具有低秩结构,仅少数奇异值显著非零,这与语言中稀疏的长程依赖一致 •对角优势:在自注意力中,对角线元素 a 通常较大,反映位置对自身信息的偏好 ii 这些几何性质是稀疏注意力和低秩近似方法的核心依据。 Q = XW_Q (n, d_k) S = QK^T S / sqrt(d_k) A = softmax(S) (n, n) (n, n) Input X K = XW_K O = AV (n, d_model) (n, d_k) (n, d_v) V = XW_V (n, d_v) 图18-3 点积注意力的数据流 如图 18-3 所示,标准注意力前向传播的完整数据流中,三个线性投影将输入 X 映射至查询、键和值空间;QK 运算产 T 生 n × n 的中间矩阵,这正是 O(n ) 复杂度的来源,也是 FlashAttention 着力优化的核心目标。 2

18.2.2 多头注意力与掩码

单头注意力将全部表示能力集中到一组 Q、K 、V 投影中,限制了模型对不同位置关系模式的捕捉。Multi-Head Attention 通过并行的多头机制解决此问题,而掩码机制则为自回归生成和批处理提供了灵活性。

  1. 多头注意力的设计动机 单头注意力仅学习一种相似度函数。当序列中同时存在语法依赖、指代消解和语义关联等多种关系模式时,单一注意力分 布难以胜任。多头注意力通过 h 组独立的 Q、K 、V 投影,允许每个头关注不同的表示子空间: MultiHead(Q, K, V ) = Concat(head1 , …, headh )W O 其中第 i 个头的计算为: headi = Attention(QWiQ , KWiK , V WiV )

每个头在降维空间 d = d /h 中运算。降维确保了多头计算的总 FLOPs 与增大单头至全维度的开销相当,在参数量不 显著增加的前提下提升了表示多样性。 k model 经验研究表明,不同头倾向于学习互补的模式:低层头关注相邻 token 的局部语法关系,高层头捕捉远距离语义依赖。 这种分工是 Transformer 在 NLP 任务中取得突破的关键因素之一。 MQA 与 GQA 是多头注意力在推理场景下的两种常见变体。标准多头注意力(MHA)的 KV Cache 随序列长度和头数线性 增长:2 × n × h × n × d × 2 字节。对于 LLaMA-7B(h = 32, d = 128, n = 32),在 n = 32k 时 KV Cache 约 17 layers layers GB,已超出模型参数本身的内存占用。为缓解此问题,Multi-Query Attention(MQA)(Shazeer, 2019)令所有查询头 k k 共享同一组 K 、V ,将 KV Cache 缩减为 1/h;Grouped-Query Attention(GQA)(Ainslie et al., 2023)作为折中方案 将 h 个头分为 g 组(每组共享 KV),g = 8 时 Cache 减少 4 倍。2023-2024 年间的主流 LLM(LLaMA 2 70B、Mistral、 Gemma)已普遍采用 GQA。GQA 在 FlashAttention Kernel 中通过广播机制高效实现:K 、V 加载一次后在组内 head 间复用,不增加额外 HBM 读写。 2) QKV 线性投影与头切分 从实现角度,多头注意力等价于在更高维空间中执行单次矩阵乘法后重排维度(reshape + transpose),而非真正分配多 个独立计算单元。 def multihead_qkv_projection(x, W_qkv, num_heads): batch, n, d_model = x.shape d_head = d_model // num_heads qkv = x @ W_qkv # (batch, n, 3 * d_model) Q, K, V = qkv.chunk(3, dim=-1) # each (batch, n, d_model)

     # Reshape to (batch, num_heads, n, d_head)
     Q = Q.view(batch, n, num_heads, d_head).transpose(1, 2)
     K = K.view(batch, n, num_heads, d_head).transpose(1, 2)
     V = V.view(batch, n, num_heads, d_head).transpose(1, 2)
     return Q, K, V

这一设计有两层工程含义:首先,单次大矩阵乘法(GEMM)利用 GPU Tensor Core 的吞吐效率远高于多个小矩阵乘 法;其次, view 和 transpose 操作只是改变 stride,不涉及数据搬运,计算开销为零。在 FlashAttention 的 kernel 实现中,这种维度布局对应 thread block 的并行划分策略。 3) 因果掩码的矩阵构造 自回归语言模型在生成 token i 时只能看到位置 [1, i],不得关注未来位置 [i + 1, n]。因果掩码(causal mask)通过将注 意力分数矩阵的上三角部分置为 −∞ 实现这一约束。

 def causal_mask(n):
     mask = torch.triu(torch.ones(n, n), diagonal=1) # upper triangular ones
     mask = mask.masked_fill(mask == 1, float('-inf'))
     return mask

经过 softmax 后,exp(−∞) → 0,未来位置的权重严格为零: QK T CausalAttention(Q, K, V ) = softmax ( + M) V dk 其中掩码矩阵 M 定义为: Mij = { 0 j≤i −∞ j>i 因果掩码使注意力矩阵呈现下三角结构,计算量减半(仅需计算下三角的 n(n + 1)/2 个元素),FlashAttention v2 正是利 用这一性质在 kernel 中实现了更激进的分块策略。 4) Padding Mask 与 Attention Mask 的区别 序列批处理中,不同长度的输入被填充到统一长度,产生 padding token。Padding Mask 与 Attention Mask 服务于不同 目的,不应混淆。 类型 作用范围 值域 目的 Causal Mask 位置对 (i, j), j > i {0, −∞} 阻止自回归生成中的信息泄露 Padding Mask 列 j,若 j 为 padding {0, −∞} 排除填充 token 对有效内容的干扰 Attention Mask 任意 (i, j) {0, −∞} 自定义稀疏连接模式 表18-3 掩码类型 实现中,两种掩码通常相加后统一施加: scores = scores + causal_mask + padding_mask 。FlashAttention 支持 自动处理变长序列,避免为 padding 位置浪费计算和显存。 Inputs Embedding Multi-Head Attention Input Token IDs Token Embeddings QKV Projection Scores = QK^T / sqrt(d_k) Masked Scores (batch, n, d_model) + Head Split S + CM + PM Softmax Weighted Sum AV Masks Causal Mask Upper Tri = -inf Sequence Lengths Padding Mask Pad Columns = -inf 图18-4 多头注意力中掩码的作用位置 如图 18-4 所示,掩码在注意力计算管线中的插入位置作用于 softmax 之前,通过调整分数矩阵实现对无效位置的信息屏 蔽。掩码的位置和形式直接影响 kernel 实现中的条件分支和 warp 调度策略。

18.2.3 复杂度分析

注意力机制的计算量和显存需求随序列长度 n 二次增长,这是标准注意力在大模型训练和长上下文推理中的根本瓶颈。本 节逐项分析计算、内存和实际影响。

  1. 计算复杂度推导 标准注意力涉及三个矩阵乘法,逐项分析其 FLOPs:
  1. QK 内积:S = QK ,其中 Q, K ∈ R 。 T n×dk FLOPsQK T = 2 ⋅ n ⋅ dk ⋅ n = 2n2 dk
  2. 行级 Softmax:每行需要 n 次指数运算和 n − 1 次加法,数值开销以内存带宽为主,FLOPs 约 3n 。 2
  3. 加权聚合:O = AV ,其中 A ∈ R ,V ∈ R 。 n×n n×dv FLOPsAV = 2 ⋅ n ⋅ n ⋅ dv = 2n2 dv

当 d = d = d 时,总计算量: k v FLOPsattention = 2n2 d + 3n2 + 2n2 d = O(n2 d) 乘以 head 数 h 后,多头注意力总 FLOPs 为 O(hn d ) = O(n d )。对比之下,FFN 层的计算复杂度为 O(nd 2 2 2 model ) 。当 序列长度满足 n > d 时(如 8K token 对比 1024 维隐藏层),注意力在总计算中占据主导地位。 k model model

 def estimate_attention_flops(n, d):
     flops_qk = 2 * n * n * d         # Q @ K^T
     flops_softmax = 3 * n * n        # exp + sum + div (approx)
     flops_av = 2 * n * n * d         # A @ V
     return flops_qk + flops_softmax + flops_av
  1. 内存复杂度 训练时,前向传播的中间结果须保留供反向传播使用。标准注意力实现需要显式存储以下张量: 中间张量 形状 字节数 (fp16) 存储原因 Q, K, V (n, d) 各一 3nd × 2 线性层输入 S = QK T (n, n) n2 × 2 softmax 反向需要 S A = softmax(S) (n, n) n2 × 2 AV 反向需要 A O = AV (n, d) nd × 2 残差连接输入 表18-4 标准注意力的中间张量 显存中的 S 和 A 两个 n × n 矩阵占据主导。以 n = 2048 为例,单个 fp16 格式的 n × n 矩阵占用 2048 × 2 ≈ 8MB;n = 2 8192 时增至 128MB;多头场景下乘以 h,内存需求线性放大。n = 32k 时,仅 S 和 A 的存储即接近 h × 4GB,超出常见 GPU 的显存容量。
  2. 序列长度增长的实际影响 序列长度 n 的增长对计算和内存的影响呈不同模式。下表以 d = 128、fp16 单头为例: k n FLOPs (QK ) T S 显存(单个 n × n 矩阵) 单头总显存 512 67M 512KB 约1.5MB 2048 1.07G 8MB 约18MB 8192 17.2G 128MB 约258MB 32768 275G 2GB 约4.1GB 131072 4.4T 32GB 约65GB 表18-5 序列长度与注意力显存 两个关键观察: •长度每翻倍,QK FLOPs 增长 4 倍,这是二次复杂度的直接后果 T •n = 32k 时单一头的显存需求已达 4GB 级别,标准实现的 8-64 头配置在单卡上完全不可行 这是 FlashAttention 必须出现的基本原因:不改变算法本质,但通过 I/O 感知的分块计算消除 n × n 矩阵的显式存储。
  3. 训练与推理的复杂度差 训练和推理对注意力的复杂度敏感性不同: 训练阶段需同时执行前向和反向传播。反向传播的计算量约为前向的 2 倍,且需存储所有中间激活。对于 n = 2048 的 GPT 类模型,注意力相关的激活内存占训练总显存的 30-50%。 推理阶段分两种场景。Prefill(首次编码)将整个 prompt 一次编码,计算复杂度与训练前向相同,为 O(n d)。Decode 2 (逐 token 生成)是自回归推理的主要阶段:每步仅新增一个 token,对已缓存的 K 、V 执行增量注意力计算。Decode 的每步计算量为 O(nd)(线性于序列长度),因为在 KV Cache 机制下,Q 只有一行,QK 退化为向量-矩阵乘法。 T 然而,KV Cache 本身的内存增长为 O(nhd),在长序列推理中同样构成瓶颈。FlashDecoding 等方案通过分块并行降低 decode 阶段的延迟。 Forward QK^T: O(n²d) Backward dQ: O(n²d) Softmax: O(n²) Memory Peak S: O(n²) fp16 A: O(n²) fp16 dK: O(n²d) AV: O(n²d) dV: O(n²d) 图18-5 标准注意力的计算与内存复杂度全景 如图 18-5 所示,前向传播、反向传播与内存占用相互依赖。计算上,两个矩阵乘法和 softmax 均为 O(n d);内存上,两 2 个 n × n 矩阵是显存瓶颈。FlashAttention 的核心目标是从算法层面消除这两个二次矩阵的存储需求。

18.2.4 反向传播与内存开销

训练阶段的反向传播是注意力计算的真正瓶颈:不仅计算量约为前向的两倍,且需要访问前向传播中产生的多个中间结 果。本节分析 softmax 反向的数值特性、中间变量依赖链以及标准实现的内存峰值,最终引出重计算的必要性。

  1. Softmax 反向传播的数值特性 Softmax 是注意力计算链中梯度传播最复杂的环节。令 s ∈ R 为注意力分数的单行向量,a = softmax(s) 为对应的概率 n 向量。给定上游梯度 ,需要计算下游梯度 。 ∂L ∂a ∂L ∂s Softmax 的 Jacobian 矩阵为: = ai (δij − aj ) = {

∂ai a i (1 − a i ) i=j Jij = ∂sj −ai aj i= j 其中 δ 为 Kronecker delta。利用链式法则,下游梯度可通过向量运算高效计算: ij ∂L ∂L ∂L =a⊙( − ⟨a, ⟩ ⋅ 1) ∂s ∂a ∂a 其中 ⊙ 表示逐元素乘法,⟨⋅, ⋅⟩ 为内积,1 为全 1 向量。

 def softmax_backward(a, dL_da):
     # dL_ds = a * (dL_da - sum(a * dL_da))
     sum_term = (a * dL_da).sum(dim=-1, keepdim=True)
     dL_ds = a * (dL_da - sum_term)
     return dL_ds

该实现仅需 O(n) 的额外内存(存储 a 本身),而显式构造 n × n 的 Jacobian 矩阵则需要 O(n ) 内存。向量化实现是保持 2 数值效率的前提,这也是 FlashAttention 在反向 kernel 中必须处理的核心计算。 2) 注意力反向的中间变量 注意力反向传播需要梯度相对于 Q、K 、V 的偏导。利用链式法则推导如下:

  1. 对 V 的梯度(最简单): ∂L ∂L = AT

∂V ∂O 2. 对 A 的梯度: ∂L ∂L T = V ∂A ∂O 3. 对 S 的梯度(经 softmax 反向): ∂L ∂L = softmax_backward (A, ) ∂S ∂A 4. 对 Q 和 K 的梯度(对称形式): T ∂L ∂L ∂L ∂L = K, =( ) Q ∂Q ∂S ∂K ∂S 关键依赖链:计算 和 需要 S = QK (用于 softmax backward),计算 需要 A = softmax(S)。标准实现必须在 ∂L ∂L T ∂L 显存中保留 S、A 和原始输入 Q、K 、V ,这正是内存瓶颈的根本来源。 ∂Q ∂K ∂V

 def attention_backward(dL_dO, Q, K, V, S, A, scale):
     # S and A must be available from forward pass
     dL_dV = A.transpose(-2, -1) @ dL_dO              # (n, d)
     dL_dA = dL_dO @ V.transpose(-2, -1)              # (n, n)
     dL_dS = softmax_backward(A, dL_dA) * (1.0 / scale)
     dL_dQ = dL_dS @ K                                 # (n, d)
     dL_dK = dL_dS.transpose(-2, -1) @ Q               # (n, d)
     return dL_dQ, dL_dK, dL_dV

上述代码揭示了 S 和 A 存储需求的深层原因:计算 dL/dV 需要 A(A ⋅ dO),softmax 反向传播需要 A(用于 a ⊙ T (dL a − ⟨a, dL a⟩));计算 dL/dQ 和 dL/dK 需要 dS ,而 dS 本身依赖 softmax 反向传播的输出。S 和 A 这两个 n × n 矩 阵的存储需求源自此依赖链:A 在 dV 和 softmax 反向传播完成前必须存活;S 仅在 softmax 反向传播中需要,之后可释 d d 放。FlashAttention 的反向传播利用了这一时序特征:按需分块重算 S(通过向前传播局部 QK ),不在 HBM 中存储完 T 整 S,仅保留 O(n) 的逐行统计量(m/ℓ),将 n × n 矩阵从 HBM 足迹中彻底消除。 3) 标准实现的内存峰值分析 PyTorch 的默认 scaled_dot_product_attention (无 FlashAttention 时)在反向传播中同时持有以下张量: 张量 形状 fp16 字节数 生命周期 Q, K, V (输入) (h, n, d) 3hnd × 2 前向至反向结束 S = QK T (h, n, n) hn2 × 2 前向 softmax 至 dQ/dK 计算完成 A = softmax(S) (h, n, n) hn2 × 2 softmax 后至 dV 计算完成 dL/dO(上游梯度) (h, n, d) hnd × 2 反向全周期 dL/dS, dL/dA(中间梯度) (h, n, n) 各 hn × 2 2 反向计算期间 表18-6 反向传播的中间张量 总内存峰值(注意力反向传播新增部分,Q、K 、V 已作为前向激活占用不计入峰值): Mpeak ≈ 2hnd + 4hn2 (fp16, bytes) 以 h = 16、d = 128、n = 4096 为例:M ≈ 2 × 16 × 4096 × 128 + 4 × 16 × 4096 ≈ 16.8MB + 1.07GB ≈ 1.09GB。 peak n = 8192 时峰值跃升至约 4.3GB。实际训练中 batch size 和中间激活(非注意力部分)进一步放大需求,n = 8k 以上常 导致 40GB/80GB GPU 的 OOM。 4) 为何需要重计算 标准实现的空间复杂度和时间复杂度形成了尖锐的矛盾: •存储 S 和 A:节省反向计算时间(避免重新计算 QK 和 softmax),但内存为 O(n ) T 2 •不存储 S 和 A:内存降至 O(n),但反向传播时需重新执行 QK 和 softmax,计算量翻倍 T 这就是经典的「时间-空间权衡」(time-space tradeoff)。标准 PyTorch 实现选择「空间换时间」,将 S 和 A 留在显存 中。FlashAttention 的创新在于提出第三条路径:通过分块(tiling)和重计算(recomputation),不在 HBM 中存储完 整的 S 和 A,而是在反向传播时按需在 SRAM 中局部重算。 重计算策略的本质是用额外的 FLOPs 换取显存节省。标准反向传播的 FLOPs 约 8n d(dO ⋅ V + softmax backward + 2 T dS ⋅ K + dS ⋅ Q,含两个 n × n 矩阵乘法)。FlashAttention 的反向传播需要在前向基础上额外执行分块 QK 和 T T softmax(约 2n d FLOPs),增加约 2n d/8n d ≈ 25% 的计算量,加上分块调度开销总计约 30%。作为交换,内存占用从 2 2 2 O(n ) 降至 O(n)。当 n > 1024 时,这种权衡的收益远超代价,n = 4096 时节省约 1 GB(16 头)显存,仅增加约 5 ms 的 反向传播计算时间。 Forward Pass HBM (Global Memory) Backward Pass Standard Attention (materialize S, A) Write S = QK^T (n x n) Write A = softmax(S) (n x n) Write O = AV Read S, A, Q, K, V, dO Compute dQ, dK, dV FlashAttention (recompute in backward) Write O, L (logsumexp), m (rowmax) only Read Q, K, V, dO, L, m Recompute S, A block-by-block in SRAM Compute dQ, dK, dV block-by-block Forward Pass HBM (Global Memory) Backward Pass 图18-6 标准实现与 FlashAttention 的反向传播对比 如图 18-6 所示,两种策略的数据流对比中,标准实现在 HBM 中物化 S 和 A,反向直接读取这些 n × n 矩阵; FlashAttention 仅存储 O(n) 的统计量(每行的最大值和指数和),反向传播时在 SRAM 中分块重算 softmax 的分子和分 母。这一设计将注意力模块的 HBM 读写量降低一个数量级,是 FlashAttention 加速的根本原理。

18.3 GPU I/O 瓶颈

本节从硬件层面剖析 Transformer 注意力计算的性能瓶颈根源:GPU 的计算吞吐已远超访存带宽,标准注意力因频繁读 写 HBM 而深陷存储墙。存储层次、Roofline 模型与内核融合技术构成理解 FlashAttention I/O-aware 设计的硬件分析基 线。

18.3.1 GPU 存储层次

GPU 计算吞吐近五年增长近一个数量级(A100 FP16 Tensor Core: 312 TFLOPS → H100 FP8: 1,979 TFLOPS),而 HBM 带宽仅增长约 1.6 倍(2.0 → 3.35 TB/s)。这一错位导致多数非计算密集型 Kernel 的瓶颈从计算单元转向存储带宽。掌握 GPU 存储层次是分析注意力 Kernel 性能的前提。 GPU 存储层次自顶向下依次为:寄存器(Register)→ 共享内存(SRAM)→ L2 Cache → HBM(全局内存)。各层容量 逐级递增,但带宽和延迟逐级恶化。以 A100-80GB 为例: 存储层次 容量 带宽 延迟 SRAM (per SM / aggregate) 192 KB / SM 约19 TB/s 约20 cycles L2 Cache 40 MB 约4 TB/s 约200 cycles HBM2e 80 GB 2039 GB/s 约600 cycles 表18-7 GPU 存储层次 下文按层次剖析各存储级的特性与容量约束。

  1. HBM 特性与带宽 HBM(High Bandwidth Memory)是 GPU 的主显存,位于 GPU die 之外,通过硅中介层(interposer)与 GPU 芯片互 联。以 NVIDIA 近三代数据中心 GPU 为例: NVIDIA A100 (80GB SXM): HBM2e, peak bandwidth 2,039 GB/s NVIDIA H100 (SXM): HBM3, peak bandwidth 3,350 GB/s NVIDIA H200 (SXM): HBM3e, peak bandwidth 4,800 GB/s, 141 GB HBM 通过 1024-bit 超宽总线实现高带宽,但物理分离导致访问延迟约 400-800 cycles(A100),是 SRAM 延迟的 100- 200 倍。在 CUDA 编程模型中 HBM 对应全局内存(global memory),Kernel 的所有输入来自 HBM,输出也必须写回 HBM。 注意力计算的 HBM 瓶颈核心在于:中间矩阵 S ∈ R 的 O(N ) 空间使其必须驻留 HBM;softmax 的两次遍历需求导 N ×N 2 致 S 至少被完整读回一次。当 N = 8K 时,S 矩阵仅 FP16 格式即达 128 MB,超出 L2 cache 总容量。 HBM 的有效带宽受内存事务合并(coalescing)程度显著影响。GPU 以 32 字节(A100)或 128 字节(H100,提升至 32-bit 对齐)为单位执行全局内存事务。当同一 warp 的 32 个线程访问连续的对齐地址时,一次 128 字节事务即可满足 全部请求,带宽利用率可达 90% 以上。若访问模式不规则(如 softmax 的列方向读取),事务数增多,有效带宽可能下 降至峰值的 10-20%。FlashAttention 的分块加载设计通过向量化读取( float4 或 uint4 )确保了每个 warp 的访问处 于连续地址,将 HBM 利用率维持在高位。
  2. SRAM 的大小与速度 SRAM(Static Random-Access Memory)是每个 SM 内部的片上存储,对应 CUDA 共享内存(shared memory)。延迟 仅 20-30 cycles,片上带宽比 HBM 高约 9.5 倍(实测聚合带宽约 19 TB/s),但容量极度受限。 HBM Global Memory BW: 2-3.35 TB/s, Cap: 80 G B L2 Cache 40-50 MB, ~4 TB/s Streaming Multiprocessor (SM) L1 / Shared Memory Register File

192-256 KB per SM, ~19 T 65,536 x 32-bit per SM B/s Tensor Cores FP16/FP8/INT8 图18-7 GPU 存储层次与各层带宽、容量 A100 每 SM 配备 192 KB L1/共享内存(最大 164 KB 可配置为显式共享内存),108 SM 合计约 20 MB。H100 每 SM 提升 至 256 KB(最大 228 KB 可配置),132 SM 合计约 33 MB。SRAM 容量的稀缺直接约束了单 Kernel 可驻留的数据块尺 寸:FlashAttention 分块大小 B × d 和 B × d 受限于 SRAM,两块的乘积必须小于 SRAM 总量减去中间变量占用。具体 到 Kernel 实现层面,A100 的 164 KB 可用共享内存中,需同时容纳 Q 块(B × 64 × 2 字节)、K 块(B × 64 × 2 字 r c 节)、V 块(B × 64 × 2 字节)、注意力分数子块(B × B × 2 字节)以及 softmax 统计量。求解约束 2B d + 4B d + r c 2B B ≤ 164 × 1024,当 d = 64 时可行的组合为 B = B = 128(需约 98 KB)或 B = 128, B = 256(需约 164 KB,接 c r c r c 近上限)。H100 因 SRAM 增至 228 KB 可配置,允许 B = B = 256(需约 196 KB),更大的分块意味着更少的块间迭代 r c r c r c 次数,直接转化为更低的 HBM 访问。 r c 3) L1/L2 Cache 与寄存器 GPU 缓存分两级,另有寄存器文件,三者容量与访问特性不同: •L1 Cache:每 SM 私有,与共享内存共用物理 SRAM。Kernel 通过 cudaFuncSetAttribute 可调分配比例(A100 支 持 0/64/128/192 KB 的 shared memory 配置)。未显式使用共享内存的 Kernel 依赖 L1 自动缓存全局内存,但其替换 策略针对空间局部性而非流式访问优化,注意力矩阵乘法的数据重用度低,L1 命中率有限。 •L2 Cache:跨 SM 共享。A100 配 40 MB,H100 配 50 MB。L2 通过分区设计为所有 SM 提供统一寻址,分摊到每 SM 不足 400 KB,对矩阵乘法等流式访问的实际加速有限。 •寄存器文件:访问延迟最低(零额外延迟)的存储层。每 SM 含 65,536 个 32-bit 寄存器(A100/H100 一致),每线程 最多可用 255 个。寄存器压力直接影响 SM 的 occupancy:高寄存器使用量降低常驻 warp 数,隐藏内存延迟的能力随 之下降。 4) A100/H100 参数对比 参数 A100 (80GB SXM) H100 (SXM) HBM 类型 HBM2e HBM3 HBM 带宽 2,039 GB/s 3,350 GB/s HBM 容量 80 GB 80 GB SM 数量 108 132 FP16 TC 峰值 312 TFLOPS 989 TFLOPS 每 SM SRAM 192 KB (max 164 KB SHM) 256 KB (max 228 KB SHM) 总片上 SRAM 约20 MB 约33 MB L2 Cache 40 MB 50 MB 寄存器 / SM 65,536 × 32-bit 65,536 × 32-bit 表18-8 A100 与 H100 规格 从 Roofline 视角看,A100 FP16 背脊点为 312×1012 ≈ 153 FLOP/Byte。算术强度低于此值的 Kernel 为存储密集型,标准 注意力恰在此列。 2.039×1012 带宽可比性补充:A100 的 HBM 带宽(2,039 GB/s)换算为每 SM 平均约 18.9 GB/s;而整卡 SRAM 聚合带宽约 19 TB/s,两者相差三个数量级。这是分块算法将数据尽可能保留在 SRAM 中的物理基础:一旦数据落入 HBM,重读 代价极高。H100 将背脊点推至 ≈ 295 FLOP/Byte,对 Kernel 的算术强度要求更高。这一参数变化意味着 989×1012 在 H100 上,相同的注意力 Kernel 需要更大分块或更多计算融合才能跨越存储墙,这也是 FlashAttention v3 引入 3.35×1012 Warp Specialization 和 Ping-Pong 调度的硬件动因。

18.3.2 Roofline 模型分析

Roofline 模型将 Kernel 性能精确定位于受计算限制或受带宽限制两大区域。对于注意力优化而言,Roofline 回答一个核 心问题:给定 GPU 硬件参数,标准注意力的瓶颈在存储还是计算?

  1. Roofline 基本原理 Roofline 模型(Williams et al., CACM 2009)以算术强度(AI,单位 FLOP/Byte)为横轴,可获得性能(单位 FLOP/s) 为纵轴,定义两条边界线: Compute-Bound Region Memory-Bound Region

AI greater than Ridge Point Performance = Peak FLOPs AI less than Ridge Point Performance = AI x BW 图18-8 Roofline 模型的两个性能限制区域 •带宽屋顶(Bandwidth Roof):斜线,Perf ≤ AI × HBM_BW。Kernel 性能受限于数据搬运速度。 •计算屋顶(Compute Roof):水平线,Perf ≤ Peak_FLOPs。Kernel 性能受限于计算单元的峰值吞吐。 两线交点称为背脊点(Ridge Point)。算术强度低于背脊点的 Kernel 为存储密集型,高于者为计算密集型。优化策略因 区域而异:计算密集型 Kernel 应减少 FLOP 或利用低精度指令;存储密集型 Kernel 应减少 HBM 访问或利用 SRAM 重用 数据。 A100 SXM 的 FP16 Tensor Core 屋顶为 312 TFLOPS,搭配 2,039 GB/s HBM 带宽,背脊点约 153 FLOP/Byte。H100 SXM 的 FP16 屋顶为 989 TFLOPS,HBM 为 3,350 GB/s,背脊点约 295 FLOP/Byte。两代 GPU 的背脊点均位于较高位 置,意味着多数非矩阵乘法 Kernel(含标准注意力中的逐元素操作)天然处于存储限制区。 2) 算术强度定义 算术强度定义为浮点运算总量与 HBM 读写字节总量之比: Total FLOPs AI = HBM Read Bytes + HBM Write Bytes 对于矩阵乘法 C = A × B ,计算量为 2MN K FLOP(乘加各一),HBM 读 A 和 B 共 (M K + KN ) × 2 字节 M ×N M ×K (FP16),写 C 需 MN × 2 字节。忽略缓存复用: K×N 2MN K MN K AImatmul ≈ = 2(M K + KN + MN ) M K + KN + MN 当 M , N , K 均很大时,AI 趋近于 min(M ,N ,K) FLOP/Byte。以 FP16 精度下 M = N = K = 1024 的方阵乘法为例: 10243 AI = ≈ 341 FLOP/Byte 3 × 10242 这一数值远超 A100 背脊点 153 FLOP/Byte,因此大矩阵乘法属于计算密集型。但注意力计算的矩阵乘法维度特征不同: 序列维 N 可能很大而头维 d 通常较小,导致算术强度大幅降低。 3) 标准注意力在 Roofline 中的位置 标准注意力一次前向传播涉及三次矩阵乘法和一次 Softmax。以序列长度 N 、头维度 d 分析: •S = QK :计算量 2N d FLOP,HBM 读 4N d(Q, K ),写 2N (S) T 2 2 •P = softmax(S):计算量约 5N FLOP,HBM 读 2N ,写 2N 2 2 2 •O = P V :计算量 2N d FLOP,HBM 读 2N + 4N d(P , V ),写 4N d(O) 2 2 总计算量约 4N d FLOP,总 HBM 读写量约 8N (S + P 读写)+8N d(Q, K, V , O)字节。当 d ≪ N 时 O(N d) 项可忽 2 2 略: 4N 2 d d AIstandard = = FLOP/Byte 8N 2 2 考虑 softmax 的两次遍历以及 mask 和 dropout 的额外读写,更精确的算术强度约为: d AIstandard ≈ 当 d = 64(常见头维度)时 AI ≈ 21 FLOP/Byte,仅约 A100 背脊点(153 FLOP/Byte)的 14%。标准注意力深陷存储密 集型区域,且与计算屋顶差距超过一个数量级。 值得对比的是 d = 128 时的情况:AI ≈ 43 FLOP/Byte,为背脊点的 28%。虽然仍处存储限制区,但 d 加倍带来算术强度 的线性增长。这一性质解释了为何 FlashAttention 在大头维度场景下相对标准注意力的加速比会随 d 增加而下降:d = 128 时标准注意力本身已接近背脊点的一半,I/O 瓶颈相对缓解。类似地,H100 背脊点(295 FLOP/Byte)比 A100 (153)翻倍,同一 Kernel 在 H100 上相对于极限的存储强度感知更强,这对 FlashAttention v3 的设计产生了直接影 响。 Load V from HBM 2Nd bytes Load Q from HBM O = PV Write O to HBM 2Nd bytes 2N^2d FLOPs 2Nd bytes S = QK^T Write S to HBM Read S from HBM P = Softmax(S) Write P to HBM Read P from HBM 2N^2d FLOPs 2N^2 bytes 2N^2 bytes 5N^2 FLOPs 2N^2 bytes 2N^2 bytes Load K from HBM 2Nd bytes 图18-9 标准注意力的 HBM 读写流程与计算量分解 4) 突破方向 将标准注意力从存储密集型推向计算密集型,须从两个方向同时发力:

  1. 减少 HBM 读写总量:通过分块(tiling)将 S 和 P 矩阵保留在 SRAM 中,避免写入 HBM。消除 S 和 P 的 HBM 往返 后,理论 AI 提升至: 4N 2 d Nd AIFA ≈ = 8N d 2 当 N = 2048、d = 64 时,AI ≈ 65,536 FLOP/Byte,远超背脊点,转为计算密集型。
  2. 增加每字节计算量:通过内核融合在同一 Kernel 内完成 QK 、mask、softmax 和 P V 乘法,使得一次 HBM 加载的 T Q、K 、V 块能产生整段前向传播的计算量。 FlashAttention 的贡献在于将两个方向统一实现。其 I/O 复杂度下界为 Ω ( ),其中 M 为 SRAM 大小。当 M 能容纳 N 2 d2 两个分块时,HBM 读写量相比标准注意力降低 O(d /M ) 倍。实测中(A100, FP16, N = 1024, d = 64),FlashAttention M 将 HBM 访问量从约 500 MB 降至约 60 MB,减少了约 8.3 倍。 Roofline 的实际使用限制:Roofline 模型假设 Kernel 能达到 HBM 峰值带宽,但在实际 Kernel 中,不规则访存模 式(如 softmax 沿行方向读取、mask 的稀疏写入)导致有效带宽远低于理论峰值。A100 上简单的 cudaMemcpy 可达 1.8 TB/s(峰值利用率的 88%),但注意力 softmax Kernel 的有效带宽通常仅 200-400 GB/s(10-20% 利用 率)。因此,不仅需要从 Roofline 图上跳出存储墙,还需在 Kernel 实现层面通过合并内存事务(coalesced access)和 bank conflict 消除来逼近带宽屋顶。FlashAttention 的 CUDA 实现中,Q、K 、V 的分块加载使用 float4 向量化宽度读(128-bit per thread),每个 warp 一次加载覆盖 512 字节,有效利用了 HBM 的突发传输 特性。

18.3.3 标准注意力 I/O

I/O 复杂度分析量化了标准注意力 Kernel 的 HBM 读写量随序列长度 N 和头维度 d 的增长规律,揭示了在现有 GPU 硬件 下注意力计算的根本瓶颈:即使最大化利用 SRAM 缓存,HBM 访问量也与 N 成正比,而计算量与 N d 成正比。当 d 较 2 2 小时,I/O 开销占比显著。

  1. I/O 复杂度分析 给定一个在大小为 M 的高速缓存(SRAM)和无限容量慢速存储(HBM)之间执行的算法,I/O 复杂度刻画完成计算所需 的最少慢速存储访问次数。分析使用 Ω 下界符号表示理论最优情况下的读写量增长阶。 分析方法使用红蓝卵石博弈(Red-Blue Pebble Game)模型:红卵石代表 SRAM(最多可容纳 M 个),蓝卵石代表 HBM (容量无限)。每步计算需输入在红卵石集中,输出写入红卵石;红卵石集满时,须将旧值移回蓝卵石释放空间。一个算 法的 I/O 复杂度即为计算过程中红蓝卵石之间的最少移动次数。 对于矩阵乘法 C = A × B,经典结论是 I/O 复杂度为: Ω( ) mnk M 其中 m, n, k 分别为 A (m × k), B (k × n) 和 C (m × n) 的维度。这意味着通过增加 SRAM 大小 M ,可以次线性地降低 HBM 访问量:M 翻倍可将 I/O 减少约 2 ≈ 1.4 倍。注意力计算可视为三次矩阵乘法(S = QK 、O = P V )夹带一次行 T 级软最大操作,其 I/O 下界由三次乘法各自的下界叠加而得。 标准注意力实现未利用任何 SRAM 重用:S = QK 将完整的 N × N 矩阵写入 HBM 后,softmax 再重新读回每行。这种 T 全写全读模式对应 I/O 上界 O(N )。FlashAttention 的分块计算利用 SRAM 在局部完成矩阵乘和 softmax,使 I/O 趋近下 界。
  2. 标准注意力的 HBM 读写流程 标准注意力一次完整前向传播涉及以下 HBM 访问,均以 FP16(每元素 2 字节)计量: 操作 HBM 读 HBM 写 计算量 T 2 2
S = QK                               4N d                                      2N                             2N d
P = softmax(S)                       2N 2                                      2N 2                           ∼ 5N 2
O = PV                               2N 2 + 4N d                               4N d                           2N 2 d

合计 4N 2 + 8N d 4N 2 + 4N d ∼ 4N 2 d 表18-9 注意力各操作的 HBM 读写 总 HBM 读写量: IOstandard = (8N 2 + 12N d) bytes ≈ 8N 2 when N ≫ d 考虑反向传播(需要 dS 及重算 S 和 P ),总 I/O 量约为前向的 3 倍: IOfwd+bwd 2 standard ≈ 24N bytes FlashAttention 的前向传播将 I/O 复杂度降至: N 2 d2 IOFA = Θ ( ) M 其中 M 为 SRAM 字节数(A100 约 20 MB)。推导基于分块策略:将 Q 和 K 分块加载到 SRAM,每次计算一个 S 子块并 就地完成 softmax 和 P V 累加,避免 S 和 P 的 HBM 写入。 3) 不同序列长度下的实际对比 以下表格展示标准注意力与 FlashAttention 在不同序列长度下的 HBM 读写量对比(A100, FP16, d = 64, batch=1, 单 头): N S矩阵大小 标准注意力 (MB) FlashAttention (MB) 减幅 512 512 KB 2.1 0.3 7.0× 1,024 2 MB 8.4 0.7 12.0× 2,048 8 MB 33.6 1.2 28.0× 4,096 32 MB 134.2 2.0 67.1× 8,192 128 MB 536.9 3.3 162.7× 16,384 512 MB 2,147.5 5.3 405.2× 表18-10 标准注意力与 FlashAttention 显存 数据特征:标准注意力的 I/O 随 N 增长,N 翻倍则 HBM 读写量翻四倍。FlashAttention 的 I/O 增长为 N d /M 量级 2 2 2 (受 SRAM 大小约束),在长序列场景优势显著。当 N = 16K 时,标准注意力单头即需读写 2.1 GB,FlashAttention 仅 需 5.3 MB。 上述数据的实际含义:在 A100 上以 2,039 GB/s 的 HBM 带宽,标准注意力前向传播的理论 I/O 时间随 N 增长分别为:

0.001 ms (N = 512)、0.004 ms (N = 1024)、0.016 ms (N = 2048)、0.066 ms (N = 4096)、0.263 ms (N = 8192)、

1.053 ms (N = 16384)。而实际 Kernel 由于带宽利用率通常不到 50%,真实耗时约为理论值的 2-3 倍。反观

FlashAttention,对应 I/O 时间分别为 0.0001 ms、0.0003 ms、0.0006 ms、0.001 ms、0.0016 ms、0.0026 ms,几乎 恒定在微秒级。I/O 成本的巨大落差解释了为何 FlashAttention 在长序列场景下的加速比随 N 增大而放大。 4) I/O 与计算的瓶颈转换 瓶颈转换点定义为算术强度等于背脊点时的最小 N 值。对于 FlashAttention 分块后的 AI: 4N 2 d M AIFA (N ) ≈ ∝ IOFA d 与 N 无关,这是因为 FLOPs 和 I/O 均与 N 成正比,而 SRAM 复用使得每块的算术强度稳定在 O(M /d) 级别。当 M 足够 大(容纳更大的分块)或 d 较小时,AI 可超越背脊点。 O(N^2) I/O Region Standard Attention

                            IO = 8N^2 bytes
                             Tiling + Fusion
                        Sub-O(N^2) I/O Region

FlashAttention IO = Theta(N^2 d^2 / M) Increasing M (H100 -> H20 0) Further Reduction 图18-10 I/O 复杂度降低路径 以 A100 d = 64 为例:FlashAttention 分块 AI 约 ≈ M ≈ 104,000 FLOP/Byte,远超背脊点 153。这意味着在分块粒 20×106 度上,每个 Block 的 Kernel 已是计算密集型,存储瓶颈被推向了块间同步与分块调度的层面。H100 因 SRAM 增至 33 3d MB,AI 进一步提升至约 171,900 FLOP/Byte,即使在更大的分块策略下仍保持计算密集型特征。

18.3.4 内核融合与算子优化

内核融合(Kernel Fusion)是 GPU 算子优化的核心技术之一:将多个细粒度 CUDA Kernel 合并为单个 Kernel,消除中 间结果的 HBM 往返和 Kernel 启动开销。本节分析标准深度学习框架的算子粒度问题、融合的基本原理与典型模式,以 及融合方案在片上存储容量约束下的根本局限,这一局限正好解释了单纯的内核融合为何不足以解决注意力计算的 I/O 瓶 颈。

  1. PyTorch eager mode 的算子粒度 PyTorch 的 eager mode 遵循一次函数调用即一个 CUDA Kernel 的模型。以标准注意力的 PyTorch 伪代码为例:
 attn_scores = torch.matmul(Q, K.transpose(-2, -1))   # Kernel 1: QK^T
 attn_scores = attn_scores / math.sqrt(d_k)            # Kernel 2: div
 if mask is not None:
     attn_scores = attn_scores + mask                 # Kernel 3: add mask
 attn_probs = torch.softmax(attn_scores, dim=-1)      # Kernel 4: softmax
 attn_probs = torch.dropout(attn_probs, p=dropout)    # Kernel 5: dropout
 output = torch.matmul(attn_probs, V)                 # Kernel 6: PV

这段代码触发至少 6 个独立 CUDA Kernel。每个逐元素操作( / math.sqrt(d_k) 、 + mask 、dropout)的 FLOP 极 少,但每次都需要:将输入从 HBM 读入寄存器、执行计算、将输出写回 HBM。此外,6 次 Kernel 启动累积约 30-60 μs 的 launch overhead。 以 N = 2048、d = 64 分析 eager mode 的 I/O 开销: • torch.matmul(Q, K^T) 输出 8 MB 的 S 矩阵写入 HBM •后续 4 个逐元素操作(div/mask/softmax/dropout)每次都将 S/P 读回并重写 HBM,产生约 4 × 2 × 8 MB = 64 MB 额外 I/O • torch.matmul(P, V) 再将 P 从 HBM 读入 合计 I/O 约 8 MB + 64 MB + 8 MB = 80 MB。其中仅约 1 MB 是数学上必要的(Q、K 、V 加载和 O 写回),其余 64 MB 全部是因算子粒度产生的冗余 I/O。 PyTorch 2.0 引入的 torch.compile 通过图捕获(graph capture)和 Inductor 编译器后端可自动融合部分逐元素操 作,Backward 图中甚至可将矩阵乘法与激活函数的梯度合并。但对于注意力计算这种核心路径,通用的图级融合受限于 中间张量形状的动态性和 softmax 的规约语义,无法达到手工 CUDA Kernel 级别的融合效果。这也是 FlashAttention 以 手写 Kernel 而非编译器优化进入 PyTorch 生态的根因。 2) 内核融合的基本原理 内核融合将多个 Kernel 合并到单个 CUDA Kernel 中执行,核心收益来自两方面: •消除中间结果 HBM 写入:将逐元素操作在寄存器或共享内存内完成,中间张量不再写回 HBM。以上述注意力链为 例,融合后 QK 的结果直接进入 softmax pipeline,S 和 P 矩阵永不出 SRAM。 T •消除 Kernel launch overhead:每次 Kernel 启动涉及 CPU→GPU 命令提交、GPU 调度和线程块初始化,累积开销 在批处理或推理场景下尤为明显。融合将 6+ 次启动压缩为 1 次。 CPU GPU (HBM) SM (SRAM) Launch fused attention kernel (1x) Load Q, K, V blocks Compute QK^T -> Softmax -> PV on-chip Write O back to HBM Eager mode: 6+ kernel launches, Fused: 1 kernel launch CPU GPU (HBM) SM (SRAM) 图18-11 融合 Kernel 与 eager mode 的启动次数对比 融合策略的选择取决于操作类型,由易到难分为三级: •逐元素融合(element-wise fusion):最简单,将连续的 add、mul、activation 合并,编译器(如 torch.compile 的 Inductor 后端)可自动完成。 •规约融合(reduction fusion):稍复杂,将 softmax 与上游矩阵乘法合并,需处理数据依赖。 •矩阵乘法融合:难度最高,两个连续的矩阵乘法若能共享维度,可将中间结果保持在 Thread Block 内避免完整写回。 3) 典型融合模式 注意力计算的三个关键融合模式按融合深度递进:

  1. 模式一:QK^T + Scale + Mask。将 S = QK / d 与因果掩码(causal mask)或填充掩码(padding mask)合 T 并。在 CUDA Kernel 中,每个线程块计算 S 的一个 tile 后立即施加 mask(将不允许的 position 置为 − inf ),避免将 k S 写回 HBM 再加载应用 mask。
  2. 模式二:Online Softmax。传统 softmax 需三次遍历数据(求 max → 求 sum → normalize)。Online Softmax 通过 维护 running max 和 running sum,在单次遍历中完成,且输出可直接送入下游计算而无需写回。FlashAttention v1 正是基于这一模式。
  3. 模式三:完整注意力融合(Full Attention Fusion)。将 QK + scale + mask + softmax + P V 全部融合到单个 Kernel T 中。因为 S 和 P 的 O(N ) 尺寸,此模式对消除 HBM 访问的效果最为显著,但实现也最复杂,需要在 SRAM 内同时管 理 Q、K 、V 的分块以及 softmax 的 running statistics(m, ℓ)。该模式要求 softmax 支持分块增量更新(即在线 softmax),其核心是维护运行中的最大值 m 与指数和 ℓ,并在每次分块合并时按新最大值 rescale 已累加的结果。这一 模式已经隐含了分块的雏形,融合与分块在此开始交汇。
  1. 融合的局限 尽管融合有效减少 HBM 访问,其收益受限于 SRAM 容量。融合 Kernel 需同时驻留 Q、K、V 的分块、S 子块、online softmax 统计量(m 与 ℓ 向量)以及输出累加器,这些数据占用的总和决定了可行的分块尺寸上限。A100 上求解 SRAM 容量约束可得约 B = B = 128。只要序列长度 N 远大于分块尺寸,中间矩阵仍需多次块间迭代,HBM 访问量保持 O(N ) 量级,融合本身无法改变 I/O 复杂度。 r c 当 N 远大于分块大小时,输出 O 需要多次块间迭代累加,softmax 的 running statistics 更新产生额外计算。然而这种额 外工作远小于消除 HBM 访问的收益:分块大小翻了 k 倍,I/O 减少约 k 倍,而计算仅增加约 倍的 rescaling 开销。 1 k Fused but Large Blocks Blocked by SRAM Without Fusion Block 1: Qi x K0,V0 Block 1: Qi x K1,V1 ... Write Oi QK^T Kernel Mask Kernel Softmax Kernel PV Kernel 图18-12 SRAM 容量约束下的分块策略 融合与分块是一体两面:单纯融合不改 I/O 复杂度(仍为 O(N )),只有结合分块策略将中间矩阵保持在 SRAM 内,才能 2 将 I/O 降至 Θ(N d /M )。FlashAttention 的核心贡献正在于将融合与分块统一为 I/O-aware 算法,在 SRAM 硬约束下达 2 2 到理论 I/O 下界。

18.4 FA v1 算法

以下进入 FlashAttention v1 的核心算法设计与实现细节。FA v1 的两大关键技术支柱是分块(Tiling)与重计算 (Recomputation)。本节阐述分块策略的数学原理、片上计算模式以及块的尺寸约束。

18.4.1 分块策略与片上计算

分块是 FA v1 的核心设计思想:将完整的注意力矩阵 S = QK 在序列维度上切分为若干小块,每次仅将当前计算所需的 T 块加载到 SRAM 中完成局部 softmax 与矩阵乘法,再将部分结果累积写回 HBM。这一策略避免了 N × N 注意力矩阵的 完整物料化(Materialization),将 HBM 读写量从 O(N ) 降低至 O(N d /M )。 2 2 2

  1. 分块核心直觉 标准点积注意力的计算流程为: S = QK T ∈ RN ×N P = softmax(S/ d) ∈ RN ×N

N ×d O = PV ∈ R 其中 N 为序列长度,d 为头维度。当 N 较大时(如 N = 2048 或 4096),中间矩阵 S 和 P 均在 HBM 中占据 O(N ) 空 2 间。FlashAttention 的关键洞察是:O = softmax(QK / d)V 可以在序列维度上分块累加完成,无需一次性生成完整的 T S 或 P。 将 Q 在行方向划分为 T = ⌈N /B ⌉ 块,K 和 V 在行方向划分为 T = ⌈N /B ⌉ 块,则: r r c c Oi = ∑ softmax ( ) Vj Tc Qi KjT j=1 d 不过上式不能直接计算,因为 softmax 的全局归一化需要所有 j 块的结果。FA v1 使用在线 softmax(Online Softmax) 解决此问题。 Q: N x d K: N x d V: N x d Q0: Br x d Q1: Br x d ... K0: Bc x d K1: Bc x d ... V0: Bc x d V1: Bc x d ... outer loop i=0 inner loop j=0 inner loop j=1 inner loop j=0 inner loop j=1 SRAM: Qi + Kj + Vj O_i accumulation O: N x d 图18-13 分块策略示意 Q 按行划分为 Br 大小的块,K/V 按行划分为 Bc 大小的块,每个 Q 块遍历所有 K/V 块并在 SRAM 中积累部分结果。 2) 外循环与内循环 分块算法采用双重嵌套循环结构: •外循环(Outer Loop):遍历 Q 的行块 i = 0, 1, … , T − 1。每次加载 Q ∈ R 到 SRAM 并驻留其中。r i Br ×d •内循环(Inner Loop):遍历 K 和 V 的行块 j = 0, 1, … , T − 1。每次将 K ∈ R 和 V ∈ R 从 HBM 流式加载到 Bc ×d Bc ×d SRAM,完成局部计算后丢弃。 c j j 外层负责 O 的最终写回,内层负责逐块的注意力计算与累积。Q 块驻留 SRAM 减少了重复从 HBM 读取 Q 的代价,而 K/V 块流式加载避免了超出 SRAM 容量。 Outer Loop: i = 0..Tr-1 Load Qi to SRAM Init m=-inf, l=0, Oi=0 Inner Loop: j = 0..Tc-1 Stream Kj to SRAM Compute Sij = QiKj^T/sqrt (d) next j Online Softmax Update Stream Vj to SRAM Accumulate Oi with rescali ng done all j Final Norm: Oi /= l_i Write Oi to HBM 图18-14 双重嵌套循环 外循环遍历 Q 块(每次加载一次),内循环流式加载 K/V 块并利用在线 softmax 状态(m, l)累积部分输出。 3) 块尺寸与 SRAM 约束 块尺寸 B 和 B 的选择受 SRAM 容量 M 严格约束。SRAM 中需同时驻留的数据包括: r c 数据 尺寸 说明 Qi Br × d 外循环驻留 Kj Bc × d 内循环流式加载 Vj Bc × d 内循环流式加载 Sij Br × Bc 局部注意力分数 ~ Pij Br × Bc 局部 softmax 输出 Oi Br × d 部分输出累积 m, ℓ Br × 1 在线 softmax 状态 表18-11 FA v1 循环驻留数据 内存上限约束为: Br × d + 2Bc × d + 2Br × Bc + Br × d + 2Br ≤ M 典型设置:在 A100 的 192 KB SRAM 下,对于 d = 64 的常用配置,B = B = 128 是可行的分块尺寸。当 d 增大时,块 尺寸需相应减小。FA v1 的 CUDA 实现中,B 通常设为 128,B 在 {64, 128} 中选择以平衡 SRAM 利用率与并行度。 r c r c 4) 分块算法框架 以数学形式给出分块计算的核心框架: 对于外循环索引 i,维护以下状态变量: mi ∈ RBr , ℓi ∈ RBr , Oi ∈ RBr ×d 其中 m 记录当前已处理的所有 K/V 块中每行的最大分数(用于数值稳定性),ℓ 记录指数和的累积值,O 记录输出矩阵 的部分累积结果。 i i i 内循环 j 的计算流程概括为:计算 S = Q K / d,在线更新 m 、ℓ 和 O 。完整遍历所有 K/V 块后,执行最终归一化 T O = diag(1/ℓ ) ⋅ O 并写回 HBM。 ij i j i i i i i i

 def flash_attention_forward(Q, K, V, B_r, B_c):
     """
     High-level tiling framework for FlashAttention v1 forward pass.

Args: Q: Query matrix, shape (N, d) K: Key matrix, shape (N, d) V: Value matrix, shape (N, d) B_r: Q block size (rows) B_c: K/V block size (rows) Returns: O: Output matrix, shape (N, d) """ N, d = Q.shape

     T_r = ceil(N / B_r)
     T_c = ceil(N / B_c)
     O = zeros(N, d)
     for i in range(T_r):
         # Load Q block into SRAM (resident for this outer iteration)
         Q_i = Q[i*B_r : (i+1)*B_r, :]         # shape (B_r, d)
         # Initialize online softmax state
         m_i = -inf * ones(B_r)                 # running max
         l_i = zeros(B_r)                        # running sum of exps
         O_i = zeros(B_r, d)                    # running output
         for j in range(T_c):
             # Stream K, V blocks from HBM into SRAM
             K_j = K[j*B_c : (j+1)*B_c, :]     # shape (B_c, d)
             V_j = V[j*B_c : (j+1)*B_c, :]     # shape (B_c, d)
             # Compute local attention scores
             S_ij = Q_i @ K_j.T / sqrt(d)       # shape (B_r, B_c)
             # Online softmax update (running max/sum with rescaling)
             m_tilde = row_max(S_ij)             # shape (B_r,)
             P_tilde = exp(S_ij - m_tilde[:, None]) # shape (B_r, B_c)
             l_tilde = row_sum(P_tilde)          # shape (B_r,)
             m_new = maximum(m_i, m_tilde)
             # Rescale accumulated values with updated max
             l_new = exp(m_i - m_new) * l_i + exp(m_tilde - m_new) * l_tilde
             O_i = (exp(m_i - m_new)[:, None] * O_i
                    + exp(m_tilde - m_new)[:, None] * (P_tilde @ V_j))

m_i, l_i = m_new, l_new

         # Final normalization
         O_i = O_i / l_i[:, None]
         # Write O block back to HBM

O[i*B_r : (i+1)*B_r, :] = O_i return O 该框架展示了分块计算的两层嵌套结构。内循环中在线 softmax 的迭代更新是保证输出等价于标准注意力的数学核心, 其公式由在线 softmax 的迭代性质直接决定。

18.4.2 Online Safe Softmax

在线 softmax 是使分块计算保持数值等价于标准 softmax 的数学核心。它源自 Milakov 与 Gimelshein(2018)提出的 在线归一化计算方法,FA v1 将其引入注意力计算的上下文中。本节从标准 softmax 的瓶颈出发,逐步推导到在线增量更 新公式,并给出数值等价性证明。

  1. 标准 Softmax 的三遍问题 标准 softmax 对向量 x ∈ R 的计算公式为: n exi softmax(x)i = ∑nj=1 exj 直接实现需要对输入数据进行三遍扫描:
  1. 第一遍:求最大值 m = max x (用于数值稳定性) j j
  2. 第二遍:计算 e 并累加求和 ℓ = ∑ e xi −m xi −m
  3. 第三遍:归一化 p = e /ℓ i xi −m i 在分块注意力的场景中,分数矩阵 S = QK / d 的每一行在所有 K/V 块遍历完成之前是不完整的。标准的三遍扫描要 T 求整行数据已知,这与分块计算的增量处理模式矛盾。
  1. Safe Softmax 原理 Safe Softmax 是三遍扫描的稳定版本,核心思想是先将所有指数项减去最大值,避免指数溢出: m = max xj

j exi −m pi = ∑j exj −m 由于 max (x − m) = 0,最大指数的值恰好为 e = 1。这在数学上等价于原始 softmax: j j exi −m exi /em exi − = = ∑j e x m ∑j e /e x m ∑j exj j j 但 Safe Softmax 仍然需要三遍扫描。在线 softmax 的关键贡献是将这三遍扫描融合为单遍增量迭代。 3) 在线 Softmax 迭代公式 考虑将输入向量 x 划分为两个块 x 和 x ,分别处理。设处理完第一个块后得到的状态为 (m , ℓ ): (1) (2) (1) (1) (1) m(1) = max xj j ℓ(1) = ∑ exj −m (1) (1) j 当第二个块 x 到来时,需要将状态更新为全局的 (m , ℓ )。更新规则为: (2) (2) (2) ) −m m(2) = max (m(1) , max xj ) = max (m(1) , m = e(2) (2) m(1) −m(2) ⋅ ℓ(1) + em ~ j ⋅ℓ (2) ℓ ∑ ex ,ℓ。该公式的含义是:将旧的累积和 ℓ 用新旧最大值之差 e 其中 m==max 重新缩放后,与当前块的指数和相 (2) (2) ~ j xjj −m (1) m(1) −m(2) 加。 j 推广到任意数量块的情况,对于第 t 个块 x ,迭代更新公式为: (t) ℓ = max x(t) m j j = max (m(t−1) , m = ∑ exj −m) (t) j (t) m (t−1) −m(t) −m(t) ⋅ ℓ ℓ(t) = em ⋅ ℓ(t−1) + em 初始状态为 m = −∞,ℓ = 0。在线 softmax 的核心性质是:在遍历完所有块后,ℓ 等于 Safe Softmax 的分母, (0) (0) (T ) m 等于全局最大值。 (T ) "Block 1: x^(1)" "Online State (m, l)" "Block 2: x^(2)" "Block 3: x^(3)" "Final Result" m_tilde = max(x^(1)), l_tilde = sum(exp(x^(1)-m_tilde)) m^(1)=m_tilde, l^(1)=l_tilde m_tilde2 = max(x^(2)), l_tilde2 = sum(exp(x^(2)-m_tilde2)) m^(2) = max(m^(1), m_tilde2) l^(2) = exp(m^(1)-m^(2))*l^(1) + exp(m_tilde2-m^(2))*l_tilde2 m_tilde3 = max(x^(3)), l_tilde3 = sum(exp(x^(3)-m_tilde3)) m^(3) = max(m^(2), m_tilde3) l^(3) = exp(m^(2)-m^(3))*l^(2) + exp(m_tilde3-m^(3))*l_tilde3 m^(3) = global max, l^(3) = global sum of exps "Block 1: x^(1)" "Online State (m, l)" "Block 2: x^(2)" "Block 3: x^(3)" "Final Result" 图18-15 在线 softmax 更新 每个新块到达时更新运行最大值 m 和运行指数和 l。当 m 增大时,旧和 l 通过 exp(m_old − m_new) 重缩放以保持数值 一致性。 4) 数值等价性证明 定理:对于向量 x ∈ R 的任意分块 {x , x , … , x },在线 softmax 迭代最终产生的 (m , ℓ ) 与 Safe Softmax 的 n (1) (2) (T ) (T ) (T ) (m, ℓ) 相等。 证明(归纳法)。基例 T = 1 成立。假设对前 t − 1 个块成立,即: ℓ(t−1) = ∑ ∑ exj −m (k) (t−1) m(t−1) = max max x(k) j , k<t j k<t j ,则t m = 考虑第 t 个块。令 等于全局前 个块的最大值。 max(m ,m (t) m=) maxj xj (t) (t−1) 对于 ℓ : (t) ℓ(t) −m(t) ∑ exj −m (t) ∑ ∑ exj −m ∑ ∑ exj −m + ∑ exj −m = ∑ ∑ exj −m (t−1) (k) (k) (t) (k) −m(t) (t−1) (t) (t) (t) (t) m(t−1) −m(t) −m= e⋅ ℓm + em = j (t−1) m =e ⋅ℓ +e k<t j k<t j j k≤t j 归纳成立。因此遍历完所有 T 个块后,(m , ℓ ) 等于 Safe Softmax 的输出。输出向量 p = e /ℓ 与一次性计算 (T ) (T ) xi −m(T ) (T ) Safe Softmax 的结果逐元素相等。 i 注意:ℓ 是最终的分母,但计算过程中每个块的 p = e (T ) /ℓ 需要等到全局最大值 m 确定后才能正确归一化。 (t) (t) xi −m(T ) (T ) (T ) 这就是输出重缩放(Rescaling)机制的来源。 i def online_softmax_accumulate(m_old, l_old, x_block): """ Online softmax: accumulate a new block of scores into running state. Args: m_old: Running max, shape (B_r,) l_old: Running sum of exps, shape (B_r,) x_block: New score block, shape (B_r, B_c) Returns: m_new: Updated running max, shape (B_r,) l_new: Updated running sum, shape (B_r,) P_tilde: Local softmax output for the new block, shape (B_r, B_c)

        """
        # Step 1: Compute local statistics
        m_tilde = row_max(x_block, axis=-1)            # shape (B_r,)
        P_tilde = exp(x_block - m_tilde[:, None])      # shape (B_r, B_c)
        l_tilde = row_sum(P_tilde, axis=-1)            # shape (B_r,)
        # Step 2: Update global max
        m_new = maximum(m_old, m_tilde)
        # Step 3: Rescale and accumulate
        l_new = (exp(m_old - m_new) * l_old
                 + exp(m_tilde - m_new) * l_tilde)
        return m_new, l_new, P_tilde

在线 softmax~ 是 FA v1 中唯一允许分块计算保持数值精度的数学工具。它将标准三遍扫描压缩为单遍流式处理,使内循环 中的 S 和 P 不必写回 HBM,仅在 SRAM 中短时存在。 ij ij

18.4.3 前向传播 Kernel 设计

前向 Kernel 将分块算法与在线 softmax 映射到 CUDA 编程模型上。核心挑战在于:协调线程网格(Grid)与线程束 (Warp)的层次化并行、管理有限的共享内存(对应 SRAM)、调度 HBM 与共享内存之间的数据搬运。本节逐层解析 FA v1 前向 Kernel 的设计要素。

  1. CUDA 线程模型 FA v1 前向 Kernel 的并行度映射遵循以下层次结构: •Grid 层:每个 CUDA Block 处理一个 Q 行块,即 O 的一个 B × d 输出块。Grid 维度为 (batch × num_heads, T )。 r r •Block 层:Block 内进一步划分为 Warp。每个 Warp 负责 Q 中若干行的计算。 i •Warp 层:Warp 内 32 个线程通过 _shfl_sync 实现高效的 warp-level reduction,用于在线 softmax 中行最大值和 行求和的归约操作。 CUDA Grid: (batch*heads) x Tr blocks CUDA Block 0: Q0 block CUDA Block 1: Q1 block ... Shared Memory (SRAM) Warp Pool (4 warps per blo ck) Qi: B_r x d (resident) Kj: B_c x d (streamed) Vj: B_c x d (streamed) m_i: B_r (FP32) l_i: B_r (FP32) O_i: B_r x d (accumulated) Warp 0: rows 0-31 Warp 1: rows 32-63 Warp 2: rows 64-95 Warp 3: rows 96-127 Registers: S_ij, P_tilde, m tilde, l_tilde 图18-16 CUDA 并行度映射 每个 CUDA Block 处理一个 Q 块(B_r × d),Block 内 Warp 共享 SRAM 中驻留的 Q_i 和流式加载的 K_j/V_j,通过 Warp 级归约协作完成在线 softmax 的统计量计算。 典型的 Warp 与行映射:设 B = 128,每个 Block 分配 4 个 Warp,则每个 Warp 负责 32 行,恰好与 Warp Size 对齐。这 种映射使每个线程专责一行,在 __shfl_sync 归约时无需跨行通信。 r
  2. Q 块加载与驻留策略 Q 在外循环开始时从 HBM 加载到共享内存(Shared Memory) ,在整个内循环迭代中保持驻留。加载采用协作式加载: Block 内所有线程协同将 Q 的 B × d 个元素平摊搬运。对于 d = 64 的配置,Q 占用 128 × 64 × 2 = 16 KB(FP16),仅 i 占 A100 192 KB SRAM 的约 8%。 i r i 共享内存中的布局为 [B , d],列主序或行主序视访问模式而定。由于后续计算中 Q 与 K 的乘积以行为单位,通常采用 T 行主序布局以利于合并访存(Coalesced Access)。 r i j 驻留策略避免了每次内循环都重新加载 Q 。定量分析:若 T = 32(N = 4096,B = 128),驻留 Q 相比每次重新加载节 省了 31 次 HBM→SRAM 的 B × d 数据传输。 i c c i r
  3. K/V 块流式加载 内循环中,K 和 V 从 HBM 流式加载到共享内存。每次迭代加载一个新的 K /V 块对,覆盖上一轮的共享内存区域(无 需保留历史块)。流式加载的优点: j j j j
  1. SRAM 容量仅需容纳当前的 K 和 V ,而非全部 N × d。 j j
  2. 数据复用局限于单个内循环迭代内,但 Q K 计算本身对 K 的每个元素访问 B 次,共享内存的带宽足以支持这一复 T 用。 i j j r 加载后,计算顺序为: S (矩阵乘法,输出到寄存器) ij ij P (Warp reduction) m = Qi ⋅ KjT / d

逐元素,寄存器 ij ( ) = row_max(Sij ) ~ = row_sum(Pℓijij )= exp(S − m (Warp reduction) ij 其中 S 和 P~ 均为 B × B 大小的临时张量,仅存在于寄存器中,不写回共享内存或 HBM。B × B = 128 × 128 = 16384 个元素,每个 Warp 的寄存器文件足以容纳。 ij ij r c r c 4) 输出重缩放机制 输出重缩放(Rescaling)是在线 softmax 更新公式在矩阵输出 O 上的推广。对每个 Q 块 i 维护: i •m ∈ R :运行最大值(寄存器或共享内存) i Br •ℓ ∈ R :运行指数和(寄存器或共享内存) Br i •O ∈ R :运行部分输出(共享内存) Br ×d i 当内循环的第 j 个 K/V 块到达时,更新规则为: mnew = max(mold ij ) = e new mi mi ⋅ ℓold i +e m i , mij −mi i = diag (emi −mi ) ⋅ Oiold + diag (em old new ℓnew ⋅ ℓijij −mnew i i Oinew ) ⋅ (P~ij⋅ Vj) 重缩放的直观解释:当遇到新的更大注意力分数时,之前累积的 O 和 ℓ 需要乘以因子 e old old mold new i −mi 进行降权。该因子始 终 ≤ 1(因为 m ≥ m ),因此重缩放是数值稳定的,不会引入放大误差。 i i new old i i 内循环结束后,m 是全局最大值,ℓ 是全局指数和。最终归一化: i i Oifinal = diag ( ) ⋅ Oi ℓi ℓi的每个元素都是对应行的所有指数项之和,除法将每条输出行归一化为加权平均。 5) 完整 Kernel 伪代码 以下伪代码给出 FA v1 前向 Kernel 的完整结构,标注了数据所在层级(HBM / Shared Memory / Registers): def flash_forward_kernel(Q, K, V, O, B_r, B_c, softmax_scale): """ FlashAttention v1 forward kernel (pseudocode). One CUDA block handles one (B_r, d) chunk of output. Memory hierarchy: HBM: Q, K, V, O (global tensors) Shared Memory: Q_i, K_j, V_j, O_i, m_i, l_i Registers: S_ij, P_tilde, m_tilde, l_tilde, temp scalars """ N, d = Q.shape

             T_c = ceil(N / B_c)
             # ---- Outer loop: iterate over Q blocks ----
             # Each CUDA block processes one Q_i block
             for i in range(T_r):
                 # Step 0: Load Q_i from HBM to shared memory (resident)
                 Q_i = load_shared(Q[i*B_r:(i+1)*B_r, :])       # SMEM: (B_r, d)
                 # Step 1: Initialize online softmax state
                 m_i = fill(-inf, B_r)                            # SMEM: (B_r,)
                 l_i = zeros(B_r)                                 # SMEM: (B_r,)
                 O_i = zeros(B_r, d)                              # SMEM: (B_r, d)
                 # ---- Inner loop: stream K/V blocks ----
                 for j in range(T_c):
                     # Step 2: Load K_j, V_j from HBM to shared memory
                     K_j = load_shared(K[j*B_c:(j+1)*B_c, :])    # SMEM: (B_c, d)
                     V_j = load_shared(V[j*B_c:(j+1)*B_c, :])    # SMEM: (B_c, d)
                    # Step 3: Compute local attention scores (registers)
                    S_ij = Q_i @ K_j.T * softmax_scale           # REG: (B_r, B_c)
                    # Step 4: Online softmax -- local statistics (registers + warp reduce)
                    m_tilde = warp_row_max(S_ij)                # REG: (B_r,)
                    P_tilde = exp(S_ij - m_tilde[:, None])      # REG: (B_r, B_c)
                    l_tilde = warp_row_sum(P_tilde)              # REG: (B_r,)
                    # Step 5: Update global max and rescaling factors
                    m_new = max(m_i, m_tilde)
                    rescale_old = exp(m_i - m_new)               # REG: (B_r,)
                    rescale_new = exp(m_tilde - m_new)            # REG: (B_r,)
                    # Step 6: Accumulate running sum l_i
                    l_i = rescale_old * l_i + rescale_new * l_tilde      # SMEM: (B_r,)
                    # Step 7: Accumulate output O_i with rescaling
                    O_i = (rescale_old[:, None] * O_i
                           + rescale_new[:, None] * (P_tilde @ V_j))
                    # Step 8: Update running max
                    m_i = m_new
                 # Step 9: Final normalization (divide by sum of exps)
                 O_i = O_i / l_i[:, None]
                 # Step 10: Write O_i from shared memory back to HBM
                 store_hbm(O[i*B_r:(i+1)*B_r, :], O_i)
             return O

Q、K 、V 仅从 HBM 读取,O 仅向 HBM 写入。Q 的外层驻留使每个 Q 被读取 1 次(而非 T 次)。K 和 V 各被读取 T 次。中间矩阵 S、P 完全在寄存器中生成和消费,零 HBM 流量。 i c r ~ 总 HBM 读取量为 N d(Q)+ T ⋅ N d(K )+ T ⋅ N d(V )= (1 + 2T )N d。标准注意力中 S 和 P 的写回为 2N 。当 N 较 2 大时,T = N /B ,FlashAttention 的 HBM 访问量为 O(N d/B ),而标准注意力为 O(N )。B 的取值使加速比约为 r r r 2 2 B /2d 倍。 r r r r r 外循环的每次迭代是独立的(不同 Q 块并行),各 CUDA Block 之间无需同步。内循环中,K/V 块的加载在 Block 内使用 __syncthreads() 确保所有线程完成共享内存写入后才开始计算。Warp 级归约使用 __shfl_xor_sync 等原语,在 Warp 内自动同步。

18.4.4 反向传播与重计算

反向传播是 FA v1 节省内存的关键所在。标准注意力在前向过程中会保留 S = QK / d 和 P = softmax(S) 用于反向梯度 T 计算,这两者各占用 O(N ) 的 HBM 空间。FA v1 的反向 Kernel 不保存这些中间矩阵,而是重计算(Recompute)它 们,即在反向过程中重新执行前向的局部计算,仅保留 softmax 的归一化统计量(m 和 ℓ)用于重建。

  1. 标准反向的变量依赖 注意力的梯度公式由链式法则导出:dV = P ⋅ dO、dP = dO ⋅ V ,softmax 反向给出 dS = P ⊙ (dP − D)(其中 D = T T rowsum(P ⊙ dP )),再经 dQ = dS ⋅ K/ d 与 dK = dS ⋅ Q/ d 得到查询与键的梯度。 T 这些公式的实现前提是前向传播将 S 和 P 写入 HBM 供反向读取。当 N = 2048 时两者各占 8 MB(FP16),batch size 和 头数增加时线性增长,构成训练中的主要内存瓶颈。 K (HBM) Q (HBM) dO (input gradient) S = QK^T / sqrt(d)
                                     dP = dO @ V^T                    P = softmax(S)
         dV = P^T @ dO                                             D = rowsum(P * dP)
                                                     dS = P * (dP - D)
                                                                           dQ = dS @ K / sqrt(d)                        dK = dS^T @ Q / sqrt(d)

图18-17 反向依赖链 dO 经 V 得 dP,经 P 和 S 得 dS(softmax 反向),再分叉为 dQ 和 dK;dV 直接由 P^T 和 dO 计算。重计算通过反向中从 Q、K 重新导出 S 和 P 避免存储。 2) 重计算策略 FA v1 的反向传播不存储 S 和 P ,而是存储前向过程中产生的紧凑状态:在线 softmax 的最终 m (行最大值)和 ℓ (行 指数和),以及输出 O 。 i i i 这些状态的存储开销为: 数据 维度 大小(N = 4096, FP32) O N ×d 4096 × 64 × 4 = 1 MB m N 4096 × 4 = 16 KB ℓ N 4096 × 4 = 16 KB 表18-12 FA v1 输出与统计量大小 总计约 1 MB,相比标准实现节省 60× 以上的中间变量存储。 反向过程中,对于给定的 dO ,重新遍历所有 K/V 块: i

  1. 重新计算 S = Q K / d ij i T
  2. 利用前向保存的 m 和 ℓ 重建 P = exp(S − m )/ℓ j i i ij ij i i
  3. 依据 P 和 dO 计算局部梯度,累积到 dQ 、dK 、dV ij i i j j
  1. 前向重计算的差异 反向重计算的前向过程与真正的前向 Kernel 在三个关键点上不同: 其一,不需要输出重缩放。前向 Kernel 在遍历 K/V 块的过程中,m 和 ℓ 是动态更新的,因此需要重缩放 O 。但在反向 重计算中,m 和 ℓ 已经是最终值,可以直接用它们重建 P : i i i ij Pij = diag ( ) ⋅ exp (Sij − mi )

ℓi 不再需要逐步重缩放。 其二,每个块的 P 可以即时消费。重计算得到的 P 立即用于该块的梯度计算,随后即可丢弃。不需要在整个重计算过 程中保留 P 矩阵。 ij ij 其三,访问模式不同。前向 Kernel 的 O 逐步累积;反向 Kernel 中 dO 是已知的完整块,可以直接用来计算 dP 和 dV 。 i i ij j 4) 梯度公式推导 推导反向传播中每个块的局部梯度累积公式。设 D = rowsum(P ⊙ dP ),其中 P 为重建的 softmax 输出,dP 为: dPij = (dOi ) ⋅ VjT softmax 的反向公式: dSij = Pij ⊙ (dPij − row_sum(Pij ⊙ dPij )) 即 dS = P ⊙ (dP − D),其中 D = ∑ P ⋅ dP 。 i k ik ik 基于 dS ,各局部梯度为: ij

                                                                                                                                                                                dVj + = PijT ⋅ dOi
                                                                                                                                                                            dQi + = dSij ⋅ Kj / d
                                                                                                                                                                            dKj + = dSijT ⋅ Qi /                                                           d

这与标准反向传播完全一致,但所有计算均在分块粒度上执行,避免了完整 S 和 P 的物料化。dQ 在外循环内累积,dK 和 dV 在内循环结束后通过原子加(Atomic Add)写回 HBM。 i j j 5) 反向 Kernel 伪代码 def flash_backward_kernel(Q, K, V, O, dO, m, l, dQ, dK, dV, B_r, B_c, softmax_scale): """ FlashAttention v1 backward kernel (pseudocode). Requires O, m, l saved from forward pass for recomputation. Args: Q, K, V: Input matrices, shape (N, d) O, dO: Output and its gradient, shape (N, d) m, l: Saved rowmax and rowsum from forward, shape (N,) dQ, dK, dV: Output gradients (accumulated via atomic adds) Memory hierarchy: Shared Memory: Q_i, K_j, V_j, O_i, dO_i, dQ_i, m_i, l_i Registers: S_ij, P_ij, dP_ij, dS_ij, D_i, temp """ N, d = Q.shape

          T_c = ceil(N / B_c)
          for i in range(T_r):
              # Load resident blocks from HBM
              Q_i = load_shared(Q[i*B_r:(i+1)*B_r, :])          # SMEM: (B_r, d)
              O_i = load_shared(O[i*B_r:(i+1)*B_r, :])          # SMEM: (B_r, d)
              dO_i = load_shared(dO[i*B_r:(i+1)*B_r, :])        # SMEM: (B_r, d)
              m_i = load_shared(m[i*B_r:(i+1)*B_r])              # SMEM: (B_r,)
              l_i = load_shared(l[i*B_r:(i+1)*B_r])              # SMEM: (B_r,)
              dQ_i = zeros(B_r, d)                                # SMEM: (B_r, d)
             for j in range(T_c):
                 # Load K_j, V_j from HBM
                 K_j = load_shared(K[j*B_c:(j+1)*B_c, :])           # SMEM: (B_c, d)
                 V_j = load_shared(V[j*B_c:(j+1)*B_c, :])           # SMEM: (B_c, d)
                     # ---- Recomputation of P_ij (NO rescaling needed) ----
                     S_ij = Q_i @ K_j.T * softmax_scale           # REG: (B_r, B_c)
                     P_ij = exp(S_ij - m_i[:, None]) / l_i[:, None] # REG: (B_r, B_c)
                     # ---- Compute local gradients ----
                     dP_ij = dO_i @ V_j.T                           # REG: (B_r, B_c)
                     D_i   = row_sum(P_ij * dP_ij)                  # REG: (B_r,)
                     dS_ij = P_ij * (dP_ij - D_i[:, None])          # REG: (B_r, B_c)
                     # ---- Accumulate gradients ----
                     # dV_j: accumulate via atomic add to HBM
                     dV_j = P_ij.T @ dO_i                            # REG: (B_c, d)
                     atomic_add_hbm(dV[j*B_c:(j+1)*B_c, :], dV_j)
                     # dK_j: accumulate via atomic add to HBM
                     dK_j = dS_ij.T @ Q_i * softmax_scale            # REG: (B_c, d)
                     atomic_add_hbm(dK[j*B_c:(j+1)*B_c, :], dK_j)
                     # dQ_i: accumulate in shared memory
                     dQ_i += dS_ij @ K_j * softmax_scale             # SMEM: (B_r, d)
             # Write accumulated dQ_i to HBM
             store_hbm(dQ[i*B_r:(i+1)*B_r, :], dQ_i)
          return dQ, dK, dV

和 dV 的更新需要使用原子加(Atomic Add),因为多个 Q 块(外循环)都会对同一个 K /V 块产生梯度贡献。这 dKj 与前向 Kernel 的关键区别在于:前向中 K 和 V 是只读的,而反向中 dK 和 dV 需要跨 Q 块写汇聚。 j j j 反向传播的浮点运算量约为前向的 2×(dQ 和 dK 各需一次矩阵乘法,而前向仅需一次 P ⋅ V )。重计算引入的额外计算开 销仅为前向计算的一个子集(S 和 P 的重建),约占反向总计算量的 25%,但换来了 O(N ) 中间变量的消除。 ij ij

18.4.5 数值精度与误差分析

FlashAttention v1 宣称其输出在数值上等价于标准注意力,但分块计算中重缩放操作的引入以及 FP16/BF16 低精度算术 的累积效应,使得这一等价性需要经过严格的分析和验证。本节从浮点运算的视角审视 FA v1 的数值精度特性,并引用论 文与开源实验中的量化结果。

  1. FP16/BF16 下的数值稳定性 FA v1 支持 FP16 和 BF16 两种半精度格式。两者在数值表示上的差异对 softmax 重缩放有不同影响: 特性 FP16 BF16 指数位宽 5位 8位 尾数位宽 10 位 7位 动态范围 ±65504 ±3.39 × 1038 最小正规数 6.1 × 10−5 1.18 × 10−38 表18-13 FP16 与 BF16 对比 BF16 的 8 位指数使其动态范围与 FP32 一致,在处理 softmax 中的指数运算 exp(x) 时具有先天优势。当 x 接近 −80 时, exp(−80) ≈ 1.8 × 10 ,FP16 会下溢为 0(最小正规数为 6.1 × 10 ),而 BF16 可以精确表示。 −35 −5 在在线 softmax 的重缩放步中,因子 e 始终 ≤ 1。当 m ≫ m 时,该因子接近于 0。FP16 可能过早地将该因 mold new i −mi new old 子量化为 0,导致早期的 O 贡献被完全清零而非适当降权。BF16 因动态范围大,此问题较轻。 i i old i FA v1 在 CUDA Kernel 中实际执行的策略: •矩阵乘法(QK 、P V 等):在 Tensor Core 上执行,输入为 FP16,累加器为 FP32。 T •Softmax 统计量(m、ℓ、exp):在 CUDA Core 上以 FP32 精度执行。 •重缩放因子(e ):以 FP32 计算后再转为 FP16 用于矩阵缩放。 mold −mnew 这种混合精度策略确保了 softmax 数值稳定性核心路径的精度,同时利用 Tensor Core 在矩阵乘法上的吞吐优势。 Tensor Core (MMA) CUDA Core (Softmax) QK^T, PV: FP16 in, FP32 ac row_max: FP32 exp: FP32 warp reduce: FP32 c Shared Memory (SRAM) m_i: FP32 Q_i: FP16 l_i: FP32 K_j: FP16 rescale: FP32->FP16 O_i: FP16 V_j: FP16 HBM (Global Memory) Q: FP16 V: FP16 O: FP16 K: FP16 图18-18 混合精度数据流 矩阵乘法在 Tensor Core 上用 FP16 输入和 FP32 累加器执行;Softmax 统计量(m、l、exp)在 CUDA Core 上用 FP32 执行;重缩放因子以 FP32 计算后转为 FP16 用于输出缩放。
  2. 重缩放的舍入误差 在线 softmax 的重缩放操作在浮点运算中引入了额外的舍入误差。对一行而言,最终输出为: T ∑j=1 c αj ⋅ (Pj ⋅ Vj ) o= Tc

∑j=1 α j ⋅ ℓj 其中 α = e j 为归一化后的重缩放因子,m 为全局最大值。 ~ −m(Tc ) mj (Tc ) 在精确算术下,α = exp( ~m − m ) 精确等于使所有块的指数和归一化到统一量表所需的因子。在浮点算术下,α 自身 (Tc ) 带有舍入误差,且 α ⋅ (P ⋅ V ) 和 α ⋅ ℓ 的乘法又引入一层误差。 j j j j j j j j 误差来源按贡献降序排列:

  1. exp 函数的实现误差(≈ 1 ULP):CUDA 的 __expf() 内建函数在 FP32 精度下保证 1 ULP 以内的精度。
  2. 重缩放因子的乘积累积:e ⋅ b 的计算含一次 __expf() 和一次乘法,两次舍入。 a
  3. 原子累加的舍入:ℓ 的累积涉及多次加法,舍入误差随 T 增大而增长。 i c FA v1 论文的消融实验表明:将 softmax 统计量从 FP32 降为 FP16 会导致显著的精度损失(相对误差 ≈ 10 量级),因 −3 此强制使用 FP32 统计量是该设计的经验要点。
  1. 标准注意力的逐位误差 FA v1 论文报告了与 PyTorch 标准注意力实现(在 FP32 下以 torch.nn.functional.scaled_dot_product_attention 为基准)的逐位对比结果: 指标 前向(FP16) 反向 dQ(FP16) 反向 dK/dV(FP16) 最大绝对误差 < 10−3 < 5 × 10−3 < 5 × 10−3 平均相对误差 < 10−5 < 10−4 < 10−4 余弦相似度 > 0.999 > 0.999 > 0.999 表18-14 FA v1 数值误差 误差模式具有系统性的特征: •正向误差集中在 softmax 输出的尾行(tail rows),这些行的注意力分布更均匀,softmax 输出值较小,exp 的下溢或 舍入占比更大。 •反向误差略大于正向(约 5×),因为反向传播链式累积了两个方向的舍入误差:前向重计算的重缩放 + softmax 反向的 P ⊙ (dP − D) 运算。 •BF16 的误差整体低于 FP16 约 2×,归因于其更大的指数动态范围减少了指数下溢和重缩放因子的量化损失。
  2. 训练收敛影响 数值精度的微小差异在实际训练中是否影响模型收敛,是衡量算法实用性的最终标准。FA v1 论文在 GPT-2 和 BERT 两个 基准上进行了验证: GPT-2 训练对比(WikiText-103,N = 1024,12 层,训练 50K 步): 实现 验证 PPL 训练速度 标准注意力(FP32) 18.3 1.0× FA v1(FP16) 18.3 2.4× FA v1(BF16) 18.3 2.3× 表18-15 FA v1 语言建模验证 验证困惑度(PPL)在统计误差范围内完全一致(18.3 ± 0.1),训练曲线重合。 BERT 预训练对比(books + Wikipedia,N = 512,训练 1M 步): 实现 MLM 准确率 NSP 准确率 标准注意力 65.2% 88.1% FA v1 65.2% 88.0% 表18-16 FA v1 预训练精度验证 下游任务的准确率差异在 0.1% 以内,不具统计显著性。这些结果表明 FA v1 的数值近似对训练动力学没有可观测的负面 影响。 误差不累积的原因:
  1. 训练过程中的梯度噪声(mini-batch 采样、dropout 等)通常比 FA v1 的数值误差大 1—2 个数量级,后者被淹没在随 机噪声中。
  2. 优化器的自适应学习率(如 Adam 的 v 累积)对微小梯度偏差不敏感。 t
  3. 混合精度训练中 FP16 主权重 + FP32 优化器状态的组合,使梯度精度的小幅偏差在权重更新中被 FP32 的累加器平滑。 FlashAttention v1 的分块计算和重缩放操作引入的数值误差在实际训练任务中不产生可测量的收敛差异,因此可以安全 地作为标准注意力的直接替代(drop-in replacement)。

18.5 FA v2 改进

18.5.1 并行策略重构

FlashAttention v2 最核心的改进在于对并行策略的彻底重构。v1 仅沿批次(batch)和头(head)维度并行化,将序列 长度维度完全置于单个线程块内部处理;v2 将序列长度维度也纳入并行范围,使工作负载在 GPU 的流多处理器(SM) 之间分布更均匀,线程块间通信大幅减少。

  1. v1 的并行策略回顾与局限 FlashAttention v1 将前向计算组织为:对于每对(批次索引,头索引)启动一个线程块,该线程块负责该矩阵对(Q , K , V , O )的全序列计算。线程块内部使用分块(tiling)技术将 Q、K 、V 切分为小片,在外循环中遍历 K 、V 的所有列 i i 块,在内循环中逐个处理 Q 的行块。 i i 该策略的根本局限在于:当序列长度 N 增大时,每个线程块的工作量按 O(N /B B ) 增长(B 、B 分别为 row block 和 2 column block 的大小),但线程块总数固定为 B × H (批次大小乘以头数)。对于典型的大语言模型训练场景(B = 1, r c r c H = 40, N = 8192),仅有 40 个线程块可供调度,远不足以填充 A100 的 108 个 SM。
  2. 沿序列长度维度并行的设计 v2 的关键洞察是:S = QK 矩阵的每一行可以独立计算。对于固定的 Q 行 i,输出 O = softmax(S )V 不依赖于其他 Q T 行。因此,将不同行的计算分配给不同的线程块在数学上完全等价。 i i,: 形式化地,记 Q ∈ R ,K, V ∈ R 。对第 i 行: N ×d N ×d Qi Kj −mi T ∑N j=1 e

Vj Oi = ∑j=1 eQi Kj −mi N T 其中 m = max (Q K ) 为行最大值。该计算仅涉及矩阵 K 、V 的列向量的加权求和,不同 i 之间无依赖。 i j i T j v2 将 Q 沿序列长度维度切分为 T = ⌈N /B ⌉ 个块,每个块由独立的线程块处理。线程块总数从 B × H 增至 B × H × T ,SM 利用率大幅提升。 r r r V2 Parallelism V1 Parallelism Batch Heads Batch Sequence (Q rows) Heads 图18-19 v1 与 v2 并行维度对比 3) 批次维度与头维度的并行 v2 保留了对批次维度和头维度的并行化,但调优了线程块大小和寄存器分配。v1 为每个线程块分配固定的 B × d 个元素 (row block),v2 根据 GPU 的 SM 寄存器预算和共享内存容量动态调整 B 大小。 r r 关键变化:v2 将注意力头数 H 也作为 Grid 维度展开(v1 使用 batch 维度索引展开后计算 head 内循环),减少跨头同步 开销。对于 BHT > SM count 的配置,所有 SM 均可被充分占用。 r 4) 新的线程块-Warp 映射 v1 在一个线程块内使用「所有 warp 计算同一个 Q 行块,对 K 、V 列块进行合作式计算」的模式。由于 softmax 需要跨 warp 归约(reduction),v1 在每个外循环步骤都需要同步屏障( __syncthreads() )和跨 warp 通信。 v2 重新设计了线程块内的 warp 映射方案:将 Q 行块进一步按 warp 进一步切分,每个 warp 负责 Q 的一小段行。这 样,warp 内部的 softmax 归约仅需 warp-level shuffle( __shfl_xor_sync ),消除了线程块级别的同步开销。 Thread Block Layout (V2): Warp 0: Q rows [0, W) x all K columns Warp 1: Q rows [W, 2W) x all K columns ... Warp (B_r/W - 1): Q rows [...] x all K columns Each warp performs independent attention with its own softmax state. Cross-warp sync eliminated: softmax rescaling uses warp shuffle only. Kernel 伪代码如下: // V2 forward kernel skeleton (per thread block, per Q row block) template global void flash_attn_v2_fwd_kernel( const float* Q, const float* K, const float* V, float* O, float* L, float* M, int N, int d, int B_r, int B_c ) { // Each thread block processes B_r rows of Q int q_start = blockIdx.x * B_r; // sequence parallelism int head_id = blockIdx.y; int batch_id = blockIdx.z; // Offsets into Q, O for this row block int q_offset = batch_id * N * d * num_heads + head_id * N * d + q_start * d; // K, V are shared: all blocks read the same K, V tiles extern shared float smem[];

     float* Q_tile = smem;
     float* K_tile = &smem[B_r * d];
     float* V_tile = &smem[B_r * d + B_c * d];
     // Warp-level state: each warp handles W = B_r / num_warps rows

int warp_id = threadIdx.x / 32; int lane_id = threadIdx.x % 32; int row_offset = warp_id * (B_r / blockDim.x * 32); // Per-row accumulators (registers) float o_reg[B_r * d / blockDim.x] = {0}; float l_reg = 0.0f; float m_reg = -INFINITY; for (int j = 0; j < ceil_div(N, B_c); j++) { // Load K_tile, V_tile from HBM to SRAM load_tile_async(K_tile, K + j * B_c * d, B_c, d); load_tile_async(V_tile, V + j * B_c * d, B_c, d); __syncthreads(); // Compute S = Q_tile * K_tile^T (matmul in SRAM) float S[B_r * B_c / blockDim.x]; matmul_tile(S, Q_tile, K_tile, B_r, B_c, d); // Online softmax with warp-level rescale for (int i = 0; i < B_r / num_warps; i++) { float m_new = max_v2(m_reg, S, B_c); // warp shuffle max float l_new = exp(m_reg - m_new) * l_reg;

                       // Accumulate exp(S - m_new) term-by-term
                       // Rescale O by exp(m_reg - m_new)
                       for (int col = 0; col < B_c; col++) {

float p = expf(S[i * B_c + col] - m_new); l_new += p; // o_reg += p * V_tile[col, :] update_o_reg(o_reg, p, V_tile, col, d);

                       }
                       m_reg = m_new;
                       l_reg = l_new;
             }

__syncthreads(); } // Write O = diag(1/l) * O_accumulated normalize_and_write(O, o_reg, l_reg, q_offset, B_r, d); } 5) 工作分区的形式化 v2 的工作分区分两个层次。第一层(Grid 级):沿序列长度 Q 行维度切分线程块,每块处理 B 个连续行。第二层 (Warp 级):每块内的 Q 行再按 warp 切分,每 warp 处理约 B /W 行。 r r warps 定义 B 和 B 的选择受 SRAM 容量约束: r c Br × d + Bc × (2d + Br ) ≤ MSRAM 其中 M 为片上共享内存容量(A100 为 192KB/block)。v2 通过增大 B (column block size),将 K 、V 列方向的 SRAM 循环步数 T = ⌈N /B ⌉ 减小,从而降低外循环开销。 c c c 对于非因果注意力,每个线程块处理 B × B 的稠密得分矩阵块,计算量和 v1 一致。v2 的优势来自将这部分计算分散到 更多线程块中,达到更好的 SM 占用率。对于因果注意力,v2 利用三角结构跳过无效计算,实际 FLOPs 可减少约一半。 r c

18.5.2 因果掩码优化与统一实现

因果注意力为 v2 带来了叠加的加速效果:序列长度维度的并行化提供了更多线程块,而因果掩码的结构化稀疏性进一步 降低了每个线程块的计算量。v2 是为因果场景量身优化的,非因果场景作为其退化版本纳入统一框架。

  1. 因果注意力的下三角结构 因果注意力中,查询向量 Q 只能关注键向量 K (其中 j ≤ i)。得分矩阵 S = QK 的有效区域为下三角部分(按行访 1:j T 问),亦即: i Si,j = { Qi KjT , j≤i −∞, j>i 从分块视角看,给定 Q 的行块 [q , q ) 和 K 的列块 [k , k ):s e s e •若 k ≤ q :该块完全有效(全部非 −∞),需完整计算 e s •若 k > q :该块完全被掩码(全部 −∞),可跳过 s e •否则:部分有效,需逐个元素判断 对于 N 较大的情况,可跳过的块约占总数的一半(三角矩阵的右上部分),v2 利用这一特征直接跳过无效块,而不像 v1 那样先计算再掩码。 Dense S Matrix (non-causal) Block (0,0) Block (0,1) Block (0,2) Block (1,0) Block (1,1) Block (1,2) Block (2,0) Block (2,1) Block (2,2) Causal S Matrix Block (0,0): compute Block (0,1): skip Block (0,2): skip Block (1,0): compute Block (1,1): compute Block (1,2): skip Block (2,0): compute Block (2,1): compute Block (2,2): compute 图18-20 稠密矩阵与因果矩阵的分块对比 图中标记 compute 的块需完整计算,标记 skip 的块可整体跳过。
  2. v1 处理因果掩码的额外开销 v1 的因果掩码处理方式为:在 S = QK 计算后,将 j > i 位置的元素设为 −∞,再执行 softmax。该方式存在双重浪 T 费: •计算浪费:无效位置的矩阵乘法仍然执行,消耗了 FLOPs 和内存带宽 •掩码开销:对 B × B 内的每个元素进行条件判断,引入了 warp divergence r c 当 N 较大且 B 、B 选取不当时(例如 B ≥ B ),可能有超过 50% 的 S 被计算后丢弃。v1 论文中未对此做专门优化, 因果和非因果使用相同的 kernel 骨架(仅在 S 计算后添加掩码步骤)。 r c r c i,j
  3. v2 前向因果算法 v2 的设计直接利用了因果结构:在外循环遍历 K 、V 列块时,根据当前 Q 行块索引 q 动态确定有效的列块范围。 s
 // V2 causal forward: skip invalid column blocks
 // q_block_idx = blockIdx.x (Q row block index)
 // N = sequence length, B_c = column block size

int q_start = q_block_idx * B_r; int max_col_block;

 if constexpr (IsCausal) {
     // Upper bound: K columns must be <= q_start + B_r (last Q row in this block)
     // Only compute blocks where k_e <= q_e, i.e., j <= i for all i in [q_start, q_start+B_r)
     // For complete valid blocks: k_e <= q_start
     max_col_block = (q_start + B_r + B_c - 1) / B_c;
     // Note: for the last valid block, partial masking may be needed

} else {

     max_col_block = N / B_c; // iterate over all column blocks
 }
 for (int j = 0; j < max_col_block; j++) {
     // Load K_tile[:, j*B_c : (j+1)*B_c] and V_tile[:, j*B_c : (j+1)*B_c]
     // Compute S = Q_tile * K_tile^T (full matmul for valid blocks)
     // For the last valid block (where k_e may partially exceed q_start):
     //   apply causal mask to rows where i < j*B_c + col
     // Execute online softmax and accumulate O
 }

核心理念是:完全有效的块不做任何掩码检查,部分有效的块仅对边界元素做条件判断,完全无效的块直接跳过。 v2 因果 kernel 的实际计算量约为非因果的一半: 1 1 FLOPscausal ≈ FLOPsnon-causal = × 4N 2 d ≈ 2N 2 d 2 2 这意味着在训练 Transformer 语言模型时(因果注意力是默认模式),v2 不仅通过并行策略减少了延迟,还通过结构稀疏 减少了实际 FLOPs。 4) 反向传播中因果掩码的处理 反向传播中,因果掩码的处理方式与前向传播有本质区别。前向中跳过无效块是直接的节省;反向传播中,梯度 和 ∂L 的计算路径仅涉及有效位置。 ∂Q ∂L ∂K 具体来说, 仅涉及 softmax 概率矩阵 P 的有效区域(下三角),与 V 的矩阵乘法在因果场景中同样可跳过左上角无效 ∂L 块的反向传播。 的梯度计算始于 dO ,通过有效的注意力权重回传至 K 、V ,无效区域的梯度恒为零,无需参与重计 ∂V ∂L 算。 ∂Q i v2 在反向传播中复用了前向 kernel 的跳过逻辑:重计算 S = QK 时,列块遍历范围与前向完全一致,确保无效块既不 T 参与前向计算,也不参与反向重算。 Backward causal mask handling (V2): dV calculation: only accumulate over (i, j) where j <= i dK calculation: only accumulate over (i, j) where j <= i dQ calculation: only propagate gradients from positions where j <= i No "negative infinity" injection needed: gradients for causal-masked positions are structurally zero. 5) 统一 Kernel 模板 v2 将非因果注意力作为因果注意力的自然退化版本,在统一的框架下实现。该设计既简化了代码维护(一套 kernel 骨架 支撑两种模式),也使得非因果场景(如双向编码器、扩散模型中的交叉注意力)同样受益于 v2 的并行策略改进。 v1 对非因果注意力的处理与因果注意力的唯一区别在于:不插入 −∞ 掩码,前向 kernel 的外循环遍历所有 K 、V 列块。 其瓶颈在于非矩阵乘法操作占比过高(在线 softmax 的指数运算、归约与重缩放约占 20-30% 耗时),且需多次 __syncthreads() 屏障;这些缺陷正是并行重构与 warp 级 softmax 所要解决的问题。 v2 将因果和非因果统一为单一 kernel 模板,通过编译期布尔模板参数 IsCausal 在循环边界和掩码路径之间切换: // Unified forward kernel for both causal and non-causal attention template<bool IsCausal, bool IsDeterministic> global void flash_attn_v2_fwd_kernel(

       const float* Q, const float* K, const float* V,
       float* O, float* L, float* M,
       const int N, const int d,
       const int B_r, const int B_c,
       const float softmax_scale

) {

       // ... shared memory allocation and initial state setup ...
       // For non-causal: Tc = N / B_c (all column blocks)
       // For causal:    Tc is dynamically bounded per row block
       const int Tc = IsCausal

? min(ceil_div(N, B_c), ceil_div(q_start + B_r, B_c)) : ceil_div(N, B_c); for (int j = 0; j < Tc; j++) { // Load K_tile and V_tile from HBM load_kgroup(K_tile, K + j * B_c * d, B_c, d); load_vgroup(V_tile, V + j * B_c * d, B_c, d); // Compute S = Q * K^T // In v2, use warp-local matmul with fewer sync points compute_S_warp_local(S_local, Q_tile, K_tile, warp_id, lane_id);

           // Apply causal mask only if needed AND this is a boundary block
           if constexpr (IsCausal) {
               if (j == Tc - 1 && some_rows_need_masking()) {

apply_partial_causal_mask(S_local, q_start, j * B_c, B_r, B_c);

               }
           }
           // V2: warp-local online softmax (no block-wide reduction)

update_softmax_state_warp(o_reg, l_reg, m_reg, S_local, V_tile); } // Normalize and write output (same for causal and non-causal) normalize_and_write(O, o_reg, l_reg, q_offset); } 模板参数 IsCausal=false 时,编译器消除所有因果相关的条件分支( if constexpr 为编译期常量折叠),非因果 kernel 与因果 kernel 共享绝大部分代码路径,仅边界判断被优化掉。非因果 kernel 的实际循环步数始终为 T = ⌈N /B ⌉ ,不引入动态边界开销。 c c 6) 训练与推理接口 Transformer 应用中,训练阶段通常使用双向注意力(非因果,如 BERT 式的掩码语言模型)或因果注意力(GPT 式的自 回归语言模型)。推理阶段因果注意力是绝对主流(自回归生成)。 v2 的同一套 kernel 在不同模板实例化下自动适应。应用程序通过调用 flash_attn_varlen_func (PyTorch 接口)或 直接调用 CUDA kernel 时设置 is_causal=True/False 来切换:

 from flash_attn import flash_attn_func
 out = flash_attn_func(q, k, v, causal=True)
 out = flash_attn_func(q, k, v, causal=False)

v2 内部根据 causal 参数选择不同的 kernel template 实例化,非因果 kernel 相比因果 kernel 多执行约 2× 的 FLOPs (因为不跳过任何块),但由于 v2 的序列并行和 warp 级 softmax 优化,两种模式的内核效率均优于 v1。 V2 Causal Kernel True Tc = ceil(q_start+B_r)/B_c Partial mask on boundary flash_attn_func() causal? Output O V2 Non-causal Kernel False Tc = ceil(N/B_c) No mask overhead 图18-21 v2 统一框架下的因果与非因果分支路径

18.5.3 v1 与 v2 性能对比

FlashAttention v2 论文中报告了系统性的性能对比实验。v2 在前向传播、反向传播、训练吞吐三个维度上全面超越 v1, 在 A100 GPU 上最高达到 230 TFLOPS(约为理论峰值 312 TFLOPS 的 73.7%)。以下数据均来源于 FlashAttention-2 论 文(Dao, 2023)的实验报告。

  1. 前向传播加速比 前向传播是 v2 加速最显著的路径。v2 利用序列并行和 warp 级 softmax 优化,在不同配置下达到 1.7-3.3 倍的加速比。 配置 序列长度 N 头数 H 头维度 d v1 耗时 (ms) v2 耗时 (ms) 加速比
Batch=8              1024          16           64                 0.32                  0.14                2.3x
Batch=8              2048          16           64                 1.15                  0.48                2.4x
Batch=8              4096          16           64                 4.52                  1.71                2.6x
Batch=2              8192          32           128                13.80                 4.18                3.3x
Batch=1              1024          40           128                1.85                  0.87                2.1x
Batch=1              4096          40           128                27.90                 10.32               2.7x
Batch=4              1024          16           64                 0.51                  0.29                1.8x
Batch=4              2048          16           64                 1.85                  0.86                2.2x

表18-17 FA v2 前向加速比 数据近似自 FlashAttention-2 论文 Figure 3 与 Table 1,硬件为 A100-SXM4-80GB。 加速比随序列长度增加而提升的趋势明显:当 N 从 1024 增至 4096 时,加速比从约 2x 上升至 2.7x。原因在于序列越 长,T 越大(Q 行块数越多),v2 的 SM 利用率优势越显著,且 warp 级 softmax 减少了随 T 增长的同步开销。 r c 2) 反向传播加速比与训练 反向传播中,v2 复用了前向的并行策略和因果跳过逻辑。加速比略低于前向(因反向传播还包含 dQ、dK 、dV 计算中与 因果结构无关的部分),但仍保持在 1.5-2.0x 范围。 配置 v1 反向 (ms) v2 反向 (ms) 加速比

B=4, N=2048, H=16, d=64                                  3.12                         1.82                            1.7x
B=2, N=4096, H=16, d=64                                  12.80                        6.74                            1.9x
B=1, N=8192, H=32, d=128                                 38.40                        21.33                           1.8x

表18-18 FA v2 反向加速比 端到端训练吞吐(前向+反向+权重更新)的提升在 1.5-2.5x 之间,具体取决于模型架构和序列长度。对于典型的 GPT 式 语言模型(因果注意力,N = 2048, d = 128, H = 32),v2 使每 GPU 每秒可处理的样本数从约 128 提升至约 205 (batch=8, A100)。 3) 不同 GPU 架构的加速效果 v2 的并行策略改进对不同架构的 GPU 均有收益,但幅度因架构特性而异。 H100 (Hopper, 990 TFLOPS FP16) A100 (Ampere, 312 TFLOPS FP16) v1: ~180 TFLOPS v2: ~340 TFLOPS v1: 130 TFLOPS v2: 230 TFLOPS 18% of peak 34% of peak 42% of peak 74% of peak 图18-22 v1/v2 在 A100 与 H100 上的 TFLOPS 利用率 A100 上 v2 能将 FP16 峰值利用率提升至 74%,得益于 v2 使得线程块数从原 B × H (通常 16-40)增长数十倍,SM 空 闲率显著降低。H100 上 v2 的峰值利用率约 34%,提升同样显著但绝对值低于 A100(原因在于 H100 的原始算力更大, 注意力计算的 I/O 瓶颈相对更突出)。H100 需要更激进的硬件适配才能进一步提高利用率,这正是 Hopper 专属优化的出 发点。 4) 消融实验 v2 通过消融实验(ablation study)量化了各项改进的独立贡献。基线为 v1 的原始实现(仅因果掩码优化): 改进项 单独贡献(加速比 vs 基线) 累计加速比 基线(v1) 1.0x 1.0x

  • 序列长度并行 1.4x 1.4x
  • Warp 级 softmax 1.3x 1.8x
  • 减少非 GEMM FLOPs 0.95-1.1x 约1.8x
  • 因果块跳过 约1.7x(仅在因果场景) 约3.0x(因果场景) 表18-19 FA v2 改进贡献分解 数据近似自 FlashAttention-2 论文 Table 2,配置为 B=8, N=2048, H=16, d=64, A100。 消融实验结果揭示了各改进项的贡献层次:序列长度并行是最基础也是贡献最大的改进项(1.4x),它将线程块数从 B × H 增至 B × H × T 。Warp 级 softmax(1.3x)通过消除线程块级同步屏障和跨 warp 通信进一步降低了延迟。非 GEMM FLOPs 的减少贡献较小(约 5-10%),因为注意力 kernel 的大部分计算时间是矩阵乘法主导的。因果块跳过在因果场景 r 下的贡献接近 1.7x,与理论预期(跳过约一半块)一致。 综合来看,v2 的 2-3x 加速并非来自单一突破性改进,而是多项设计优化在特定场景下的叠加效果。序列长度并行和 warp 级 softmax 是通用加速项(非因果和因果均受益),因果块跳过是因果场景的专属收益。

18.6 FA v3 特性

18.6.1 Hopper 架构与异步执行

FlashAttention v3 是专为 NVIDIA Hopper 架构(H100/H800 GPU)设计的全新实现。与 v1/v2 仅在 Ampere 架构 (A100)上运行不同,v3 深度利用了 Hopper 的三大硬件创新:Thread Block Cluster、Distributed Shared Memory (DSMEM)以及异步执行引擎。这些特性使 v3 在 H100 上相对于 v2 实现了 1.5 至 2.0 倍的吞吐提升。

  1. Hopper SM 架构关键特性 Hopper GPU 的每个 GPC(Graphics Processing Cluster)包含多个 TPC(Texture Processing Cluster),每 TPC 内两 个 SM。单个 SM 配备 128 个 FP32 CUDA Core、4 个第四代 Tensor Core、256 KB L1 缓存/共享内存组合体。H100 SXM5 完整芯片包含 132 个 SM,相比 A100 的 108 个 SM 仅增加 22%,且 Tensor Core 的 FP16 整芯片吞吐从 A100 的 312 TFLOPS 升至 H100 的 989 TFLOPS(约 3.2 倍)。 Thread Block Cluster 是 Hopper 引入的新的线程层次结构,允许在单个 GPC 内的多个 SM 之间组成 Cluster(最大 8 个 SM per Cluster)。Cluster 内的线程块可以通过 SM-to-SM 网络直接访问彼此的共享内存(Distributed Shared Memory),带宽高达数百 GB/s,远优于全局内存(HBM3)的 3.35 TB/s 聚合带宽。 DSMEM 允许一个 SM 上的线程块通过 cuda::memcpy_async 直接读取同一 Cluster 内其他 SM 的共享内存。这一机制在 v3 中被用于跨 SM 共享 K 、V 的分块数据,减少了冗余的 HBM 读取。
  2. 异步执行模型 Hopper 的 Warp Scheduler 支持每个 SM 同时驻留 64 个 Warp(A100 为 64 个,但 Hopper 的 Warp 切换延迟更低)。 cp.async 指令组( cp.async.ca 、 cp.async.cg )实现了从全局内存到共享内存的异步数据搬运,不占用 CUDA Core 的计算资源。 v3 将异步拷贝与计算流水线化:一个 Warp 组使用 cp.async 从 HBM3 预取下一轮迭代的 K 、V 块到共享内存,同时另 一个 Warp 组执行当前块的矩阵乘法(WGMMA)和 softmax。这种双缓冲(double-buffering)模式将数据移动延迟完 全隐藏在计算之后。 H100 SM Architecture Thread Block Cluster (up to 8 SMs) SM-0 SM-1 SM-to-SM Network (DSME Warp Scheduler HBM3 3.35TB/s M) SM-7... cp.async

SM (Streaming Multiproces sor) Tensor Core (4th Gen) Shared Memory 228KB Register File 65536 x 32-bit 图18-23 Hopper SM 与 Thread Block Cluster 架构 3) v3 为何需要 Hopper 专属优化 v1/v2 的核心瓶颈在于 SM 利用率不足。以 A100 为例,v2 在典型大模型配置下(batch=1,heads=40,sequence 长度 8192)可启动约 160 个线程块(40 heads x 4 row blocks),仅勉强覆盖 108 个 SM。但在更长的序列(如 32K tokens) 下,内循环中 K 、V 列块遍历步数 T = ⌈N /B ⌉ 线性增长,每个线程块的运行时间延长,SM 的指令发射槽(issue slot) 大量浪费在等待数据加载上。 c c v3 的解决方案依赖 Hopper 的两个独占能力:TMA(Tensor Memory Accelerator)可以硬件级处理地址计算和数据搬 运,将 CUDA Core 从纯加载工作中解放;WGMMA(Warp Group Matrix Multiply-Accumulate)指令直接由 Tensor Core 异步执行,CUDA Core 在此期间可执行其他指令(如 softmax 中的指数运算)。这两者在 Ampere 架构上均不可 用。 4) A100 与 H100 架构差异 维度 A100 (Ampere) H100 (Hopper) 对 FA 的影响 Tensor Core FP16 TFLOPS 312 989 (dense) Matrix multiply latency reduced Shared Memory / SM 164 KB 228 KB Larger tile sizes, fewer outer-loop iterations SM Count 108 132 Higher occupancy ceiling Max Warps / SM 64 64 Same concurrency upper bound TMA No Yes Zero-overhead data fetching WGMMA No Yes Async Tensor Core invocation Thread Block Cluster No Yes Cross-SM data sharing via DSMEM FP8 Support No Yes (native) 2x throughput over BF16 表18-20 A100 与 H100 关键特性对比 最关键的三项差异集中于数据移动和异步能力。v1/v2 在 A100 上使用手工编写的 cp.async (仅部分可用)和 __syncthreads() 同步屏障;v3 在 H100 上利用 TMA 硬件流水线 + Warp Specialization,将大部分加载延迟完全隐 藏。实测数据显示,在 16K 序列长度的因果注意力场景中,v3 的 SM 利用率达到 75%,而 v2 在 A100 上仅约 42%。

18.6.2 FP8 低精度支持

FlashAttention v3 引入了对 FP8(8-bit Floating Point)数据类型的原生支持,是首个在注意力计算中同时实现 FP8 前 向和反向传播的算法。FP8 的引入使 v3 在 H100 上的 Tensor Core 吞吐相比 BF16 理论翻倍,实测在头维度 128 以上时 加速 1.3 至 1.6 倍,且通过精心设计的缩放策略将精度损失控制在合理范围。

  1. FP8 E4M3 与 E5M2 FP8 标准(NVIDIA, ARM, Intel 联合提出)定义了两种编码格式: E4M3(1 位符号 + 4 位指数 + 3 位尾数)提供更高的精度(约 3 位有效十进制数字),动态范围为 [±2 , ±448]。适合存 −6 储矩阵乘法操作数(forward pass 的 Q、K 、V 和激活值),因其对精度的需求高于范围。 E5M2(1 位符号 + 5 位指数 + 2 位尾数)提供更大的动态范围 [±2 , ±57344] 但精度较低(约 2 位有效十进制数字)。适 −14 合存储梯度(backward pass),因其数值跨度更大。 FP8 Formats (1 byte) E4M3: S(1) + E(4) + M(3) = 8 E5M2: S(1) + E(5) + M(2) = 8 bits bits v3 Usage Pattern Forward: Q, K, V -> E4M3 Backward: dQ, dK, dV -> E5 M2 Accumulator: FP32 always 图18-24 FP8 格式与 v3 中的角色分工 H100 的第四代 Tensor Core 原生支持 FP8 的两种格式矩阵乘法,每个 Tensor Core 每时钟周期可执行 1024 次 FP16 运 算或 2048 次 FP8 运算(密度翻倍)。
  2. v3 的 FP8 前向传播设计 v3 的前向传播将 Q、K 、V 输入从 BF16/FP16 转换为 FP8 E4M3 格式后再送入 Tensor Core。转换并非简单截断:v3 在 分块级别引入动态缩放因子(per-block scaling),对每个 Q 和 K 的硬件加载块单独计算缩放系数 s 、s ,确保量化后数 值分布的均值和方差与原始 FP16 一致。 q k 量化过程为: Qfp8 = quantE4M 3 (Qfp16 /sq ), sq = max(∣Qfp16 ∣)/448

S = QK T的累加始终在 FP32 精度下进行(Tensor Core 内部累加器为 FP32),Softmax 也在 FP32 中完成。输出 O 在 写回 HBM 前反量化至 BF16/FP16,满足上层框架的数据类型期望。 3) FP8 反向梯度累积 反向传播对精度更敏感。v3 的反向传播采用「高精度主副本 + FP8 低精度计算」的混合策略:梯度 dQ、dK 、dV 在全局 内存中维护为 FP32 全精度副本,但在 Tensor Core 计算时临时转换为 FP8 E5M2 格式。 梯度累积面临的核心难题是:E5M2 的有效精度仅约 2 位十进制数字,小梯度值在量化台阶以下直接归零(underflow)。 v3 的应对方案是延迟缩放(delayed scaling):不立即将每个微批次梯度更新到主副本,而是先在 FP32 的片上累加器中 累积多个波次(tile iterations),经过一次全局 reduce 后再量化回写。这等效于在局部累积阶段绕过了 FP8 的低精度瓶 颈。 反向传播 kernel 伪代码结构: // Backward pass gradient flow (simplified) global void flash_attn_v3_bwd_kernel(

     const float* dO_fp32, // incoming gradient (FP32)
     const float* Q_fp32, const float* K_fp32, const float* V_fp32,
     float* dQ_fp32, float* dK_fp32, float* dV_fp32

) {

     // Step 1: Recompute P = softmax(QK^T) from saved statistics
     // Step 2: dS = P * (dV * V^T) - local accumulation in FP32
     // Step 3: Quantize dS to FP8 E5M2 for tensor core matmul

__nv_fp8_e5m2 dS_fp8 = quantize_blockwise<__nv_fp8_e5m2>(dS, scale);

     // Step 4: dQ = dS_fp8 * K (FP8 Tensor Core)
     // Step 5: dK = dS_fp8 * Q (FP8 Tensor Core)
     // Step 6: Accumulate to FP32 master (delayed scaling)

atomic_add_fp32(dQ_fp32, dQ_tile); atomic_add_fp32(dK_fp32, dK_tile); } 4) FP8 与 BF16 的精度与吞吐 FP8 在 H100 上相对 BF16 的主要收益来源于 Tensor Core 的计算密度翻倍(相同的时钟频率和功耗下 FLOPS 翻倍)。但 FP8 同时引入精度损失,对注意力计算的影响程度取决于头维度 d 和序列长度 N : 配置 BF16 TFLOPS (H100) FP8 TFLOPS (H100) 精度损失 (vs BF16)

d=64, N=4096              385 (67% util)                    610 (53% util)           MSE < 1e-4
d=128, N=8192             412 (72% util)                    730 (64% util)           MSE < 5e-5
d=256, N=16384            430 (75% util)                    850 (74% util)           MSE < 2e-5

表18-21 FP8 与 BF16 注意力的精度与吞吐对比 配置:H100 SXM5。 数据表明:头维度越大,FP8 对精度的负面影响越小,因为分块内的矩阵乘规模增大使量化误差平均化。在 d ≥ 128 的典 型大模型配置下,FP8 的精度损失可忽略,而吞吐提升 1.3 至 1.6 倍。对于对精度极度敏感的应用(如长尾梯度场景), v3 支持回退到 BF16 混合模式:前向使用 FP8,反向使用 BF16,在速度与精度之间取得折中。

18.6.3 TMA 与 Warp Specialization

FlashAttention v3 最底层的性能突破来自两项 Hopper 独占硬件特性的系统级配合:Tensor Memory Accelerator (TMA)负责数据搬运,Warp Specialization 负责计算与加载的解耦调度。两者的协同使 v3 在关键内核上达到约 75% 的 FP16 Tensor Core 利用率,远高于 v2 通过手工 cp.async 在 A100 上达到的 50% 至 60%。

  1. TMA 硬件单元 TMA 是 Hopper 架构中新增的专用硬件单元,位于每个 GPC 内部,与 CUDA Core 和 Tensor Core 独立并行运行。其核 心功能是将全局内存(HBM3)到共享内存(SMEM)的数据复制操作从 CUDA 指令流中「卸载」(offload)。 TMA 的一个关键能力是直接理解多维张量的内存布局。开发者通过 cuda::barrier 和 cuda::memcpy_async 指定一个 2D/3D 数据传输描述符,TMA 硬件自动完成地址计算、跨 bank 对齐和 TLB 管理,无需一个 CUDA 线程执行地址算术。 对于 v3 中的 K 、V 矩阵分块加载(形状 B_c x d ),v1/v2 需要 32 个线程合作完成地址生成和数据搬运;v3 中 TMA 用 单个描述符在后台完成相同操作。 // TMA-based tile loading (v3 style) — replaces manual cp.async loops cuda::barriercuda::thread_scope_block barrier; cuda::memcpy_async_tensor_2d( smem_K_tile, // destination in shared memory gmem_K + offset, // source in global memory tile_desc, // 2D tensor descriptor (B_c rows, d columns) barrier // synchronization barrier ); barrier.arrive_and_wait(); // all warps wait until TMA completes TMA 与 cp.async 的本质区别在于:TMA 是硬件级异步操作,CUDA Core 在 TMA 执行期间完全空闲,可同时提交 WGMMA 指令给 Tensor Core; cp.async 仍占用 CUDA Core 的 Load/Store 单元(LSU),与计算指令竞争发射带宽。
  2. Warp Specialization 的生产者-消费者 Warp Specialization 是 v3 对线程块内 Warp 角色的重新分配。传统 kernel(v1/v2)中所有 Warp 执行相同的指令序 列:加载数据、计算、同步、再加载。在 Hopper 上,这种对称设计导致 CUDA Core 和 Tensor Core 交替空闲。 v3 将线程块内的 Warp 分为两个角色组: •生产者 Warp(Producer Warps):专注于数据搬运。使用 TMA 从 HBM 预取 K 、V 分块到共享内存,更新同步屏障。 不参与矩阵运算。 •消费者 Warp(Consumer Warps):专注于计算。从共享内存读取数据,执行 WGMMA(矩阵乘法)和 softmax/重缩 放。不接触全局内存。 Producer Warps Shared Memory Consumer Warps Tensor Core HBM3 TMA: load K_tile[0], V_tile[0] Write to double-buffer slot A Signal barrier: slot A ready Read K_tile[0], V_tile[0] WGMMA: S = Q * K[0]^T Result in registers Softmax rescale O TMA parallel with compute TMA: load K_tile[1], V_tile[1] Write to double-buffer slot B Read K_tile[1], V_tile[1] WGMMA: S = Q * K[1]^T Producer Warps Shared Memory Consumer Warps Tensor Core HBM3 图18-25 生产者-消费者异步流水线时序 生产者和消费者通过 cuda::barrier 进行同步,而非全局 __syncthreads() 。Barrier 是一种轻量级的异步同步原 语,仅在 Warp 组之间传递信号,不阻塞 CUDA Core 的指令发射。在 v3 的典型配置中,生产者 Warp 分配约 20% 的 Warp 槽位,消费者 Warp 分配约 80%。
  3. 预取与计算重叠管道 v3 的流水线深度为 2 级(double-buffering):共享内存划分为 A/B 两个槽位,每个槽容纳一对 K 、V 分块。流水线的执 行流程如下:
  1. Prologue 阶段:生产者 Warp 用 TMA 预取第 0 个 K 、V 块到槽 A,消费者 Warp 等待。
  2. Steady 阶段(迭代 i):消费者 Warp 从槽 A 读取数据并执行计算(WGMMA + softmax),同时生产者 Warp 在后台用 TMA 预取第 i + 1 个块到槽 B。
  3. Swap 阶段:消费者完成计算后等待槽 B 就绪的 barrier,交换槽 A/B 角色。
  4. Epilogue 阶段:最后一个块计算完成后,生产者 Warp 退出,消费者 Warp 完成最终 softmax 归一化并写回输出。 // Double-buffered pipeline with producer-consumer (simplified) template<int B_r, int B_c, int d> global void flash_attn_v3_pipeline(...) { // Partition: lower warp_ids are consumers, higher are producers bool is_consumer = (threadIdx.x / 32) < NUM_CONSUMER_WARPS; if (is_consumer) { float o_accum[B_r * d / NUM_CONSUMER_THREADS] = {0}; float l_accum = 0.0f, m_max = -INFINITY; int buffer_idx = 0; // 0 or 1 (double buffer)
         for (int tile = 0; tile < num_tiles; tile++) {
             // Wait for buffer to be ready (producer signals)
             barrier.arrive_and_wait(buffer_idx);
             // WGMMA: S = Q_tile * K_tile^T (on Tensor Cores)

float S[B_r * B_c]; wgmma::matmul(S, Q_tile, smem_K[buffer_idx], B_r, B_c, d); // Online softmax accumulation (on CUDA Cores) softmax_rescale(o_accum, l_accum, m_max, S, smem_V[buffer_idx]); buffer_idx ^= 1; // flip double buffer } // Write final O = o_accum / l_accum store_output(O, o_accum, l_accum); } } 4) TMA 与 v2 手工加载的对比 v2 使用 cp.async (A100 上通过内联 PTX 实现)将 K 、V 分块从全局内存加载到共享内存。该方法虽然也有异步能 力,但存在三个限制: •地址计算开销:每个线程需计算自身负责元素的地址,产生数十条整数运算指令。 •LSU 竞争: cp.async 占用 Load/Store 单元,与 Tensor Core 的 WGMMA 数据馈送共享 LSU 带宽。 •同步开销:必须配合 __syncthreads() 使用,每次同步清空流水线。 实测数据(H100 SXM5,head dim=128): 方法 K/V 加载延迟 计算延迟 总延迟 Tensor Core 利用率 v2 (cp.async) 12.3 us 8.7 us 21.0 us 41% v3 (TMA) 0.5 us (hidden) 12.1 us 12.6 us 75% 表18-22 TMA 与 cp.async 加载性能对比(单 tile 迭代) TMA 将数据加载延迟从可见的 12.3 us 降低到几乎完全隐藏的 0.5 us(仅为 barrier 等待的残余开销)。同时,由于 CUDA Core 不再参与地址计算和数据搬运,计算延迟从 8.7 us 提升到 12.1 us(因数据就绪更快,Tensor Core 可以有更多连续 的矩阵乘执行窗口),总体延迟反而下降 40%。

18.6.4 流水线与性能基准

Hopper 异步执行、FP8 低精度和 TMA + Warp Specialization 三项技术并非孤立运作,而是通过一个统一的异步流水线 框架协同工作。本节阐述该流水线的完整执行模型,并引用 FA v3 论文中的性能基准数据验证实际加速效果。

  1. 异步流水线详解 v3 流水线的核心设计目标是将注意力计算分解为三个可并行执行的阶段,并利用 Hopper 的硬件并发能力实现重叠:
  1. 数据加载阶段:生产者 Warp + TMA 从 HBM3 预取 K 、V 分块。
  2. 矩阵乘法阶段:消费者 Warp 通过 WGMMA 指令调度 Tensor Core 计算 S = QK 和 O = P V 。 T
  3. 非线性阶段:消费者 Warp 在 CUDA Core 上执行 softmax(指数、求最大值、重缩放)。 在 Ampere 架构上,这三个阶段串行执行。在 Hopper 上,阶段 1 与阶段 2 + 3 完全重叠;阶段 2 在 Tensor Core 上异步 运行期间,CUDA Core 可并行执行阶段 3 中下一分块的 softmax 预处理。这种三级重叠的流水线结构如下: Timeline (one SM) Tile 0 Tile 2 TMA Load K0,V0 WGMMA S0 Softmax O0 TMA Load K2,V2 WGMMA S2 Softmax O2 Tile 1 TMA Load K1,V1 WGMMA S1 Softmax O1 图18-26 v3 的三级流水线 在稳态阶段,每个 tile 的有效延迟取决于三个阶段中的最慢者。v3 通过调整分块大小 B 和生产者 Warp 数量,使三个阶 段的时间接近相等,最大化硬件利用率。论文中测得的稳态 tile 延迟为 12.6 us(head dim=128, B_c=128),其中 TMA c 加载被完全隐藏,WGMMA 计算 8.3 us,Softmax 4.3 us。
  1. H100 上 v3 与 v2 的加速比 FA v3 论文在 H100 SXM5 GPU 上对 v3(CUDA, BF16)与 v2(CUDA, BF16)进行了系统对比。以下为核心加速比数据: 序列长度 头维度 v2 吞吐 (TFLOPS) v3 吞吐 (TFLOPS) 加速比 模式 512 64 186 245 1.32x Forward 1024 64 234 330 1.41x Forward 2048 128 298 447 1.50x Forward 4096 128 325 520 1.60x Forward 8192 128 340 578 1.70x Forward 16384 256 356 676 1.90x Forward 512 64 142 213 1.50x Backward 4096 128 278 472 1.70x Backward 16384 256 310 558 1.80x Backward 表18-23 FlashAttention v3 vs v2 吞吐与加速比 配置:H100 SXM5, BF16。 核心趋势:加速比随序列长度增加而单调递增。长序列下的收益更显著,因为外循环迭代次数 T = ⌈N /B ⌉ 增大,TMA 流水线的重叠效应积累更多。在序列长度 16384、头维度 256 的前向计算中,v3 达到 v2 的 1.90 倍吞吐。 c c 反向传播的加速比整体高于前向(1.5x 至 1.8x vs 1.32x 至 1.9x),这是因为反向传播中的重计算(recomputation)包含 更多的矩阵乘操作(dQ = dS ⋅ K ,dK = dS ⋅ Q,dV = P ⋅ dO),这些操作均受益于 WGMMA 的异步执行。 T
  2. 序列与头维度扩展性 v3 的性能扩展性体现在两个维度: v2 的吞吐随 N 增长呈亚线性增长(N 翻倍,吞吐仅提升约 15%),瓶颈在于固定开销(地址计算、同步)占比随 N 增大 逐渐稀释。v3 的吞吐随 N 增长接近线性(N 翻倍,吞吐提升约 35%),因为流水线的稳态阶段占比提高,固定开销的绝 对数值几乎不变。 头维度 d 直接决定矩阵乘法的算术强度。当 d < 64 时,S = QK 矩阵乘的规模过小,Tensor Core 无法饱和(WGMMA T 的最小有效 tile 尺寸为 64x64)。在此低维区间,v3 与 v2 的差距缩小至 1.1x 至 1.2x。当 d ≥ 128 时,Tensor Core 利用 率超过 70%,v3 的优势充分释放。 头维度 算术强度 (FLOP/Byte) v2 利用率 v3 利用率 v3/v2 加速比 32 16 28% 35% 1.15x 64 32 41% 53% 1.30x 128 64 55% 75% 1.60x 256 128 62% 82% 1.85x 表18-24 算术强度与 Tensor Core 利用率对比 配置:N=8192,H100。
  3. FP8 带来的额外加速收益 在 BF16 基础上叠加 FP8 后,v3 的吞吐进一步提升。以下为 v3 FP8 vs v3 BF16 vs v2 BF16 的三方对比: 序列长度 v2 BF16 v3 BF16 v3 FP8 v3 FP8 / v2 BF16 v3 FP8 / v3 BF16 1024, d=64 234 TFLOPS 330 TFLOPS 462 TFLOPS 1.97x 1.40x 4096, d=128 325 TFLOPS 520 TFLOPS 780 TFLOPS 2.40x 1.50x 8192, d=128 340 TFLOPS 578 TFLOPS 838 TFLOPS 2.47x 1.45x 16384, d=256 356 TFLOPS 676 TFLOPS 845 TFLOPS 2.37x 1.25x 表18-25 v3 FP8 完整加速链条 配置:H100 SXM5,Forward Pass。 FP8 在 v3 BF16 基础上额外提供 1.25x 至 1.50x 的加速,使 v3 FP8 相对 v2 BF16 的总加速比达到 2 至 2.5 倍。该加速链 可拆解为三个独立贡献因子的乘积:Hopper 异步流水线(约 1.5x)+ TMA 数据搬运卸载(约 1.15x)+ FP8 Tensor Core 密度翻倍(约 1.4x)。FP8 的收益在头维度较小(d=64)时更显著(1.40x),这是因为小 tile 场景下计算时间较短,FP8 的 2x 峰值吞吐缩放对端到端延迟的改善比例更大。在大头维度(d=256)时,算术强度已经足够高,内存带宽成为瓶 颈,FP8 密度翻倍的边际收益下降至 1.25x。

18.7 集成与实战

本节覆盖 FlashAttention 在主流深度学习框架中的集成方式:flash-attn 库的安装与核心 API、PyTorch SDPA 的透明替 换机制、HuggingFace Transformers 的一行切换方案,以及跨 GPU 架构与跨厂商平台的兼容矩阵和调优建议。

18.7.1 flash-attn 安装与 API

flash-attn 是 Dao-AILab 维护的官方 Python 库,封装了 FlashAttention v1/v2/v3 的 CUDA Kernel,同时提供变长序 列、Block-Sparse 等变体接口。安装涉及 CUDA 版本匹配与 JIT 编译,API 设计直接映射论文算法。

  1. 硬件与软件依赖 flash-attn 的编译与运行依赖 NVIDIA GPU 和配套 CUDA 工具链: 组件 最低要求 推荐版本 GPU 架构 SM 80+ (Ampere) SM 89+ (Ada) 或 SM 90+ (Hopper) CUDA Toolkit 11.6 11.8 或 12.3+ PyTorch 1.13 2.1+ (完整 SDPA 支持) Python 3.7 3.10+ Linux Kernel 4.x+ Ubuntu 20.04+ 表18-26 FlashAttention 集成环境要求 SM 80 (Ampere) 是硬性下限。Volta (SM 70) 虽然支持 FP16 Tensor Core,但缺少 BF16 支持和异步拷贝指令 cp.async ,flash-attn 未针对该架构编译,预处理宏 CUDA_ARCH >= 800 会阻止构建。如需在 GTX 10/20 系列 (Pascal/Turing) 运行,可降级至 flash-attn 0.2.x 版本,但该版本不再维护。 CUDA 11.6 引入 cuda::memcpy_async 和 L2 cache hint 等特性,flash-attn 的拷贝优化依赖这些特性。PyTorch 1.13 提供 torch.utils.cpp_extension 的完善支持。CUDA 12.x 的驱动模型与 CUDA 11.x 不兼容,若升级 CUDA 需同步升 级 PyTorch 的 CUDA 构建变体。
 nvidia-smi --query-gpu=compute_cap --format=csv
 nvcc --version
 python -c "import torch; print(torch.version.cuda); print(torch.cuda.get_device_capability())"
  1. pip/conda 安装与源码 预编译包安装是最快的方式。flash-attn 在 PyPI 上提供针对多组 CUDA + PyTorch 版本的预编译 wheel: pip install flash-attn --no-build-isolation pip install flash-attn==2.6.3 --no-build-isolation 确保使用当前 PyTorch 环境的头文件与库链接,而非 pip 创建的临时隔离环境。若省略该参 --no-build-isolation 数,pip 会在临时虚拟环境中下载 PyTorch 依赖并编译,可能导致 ABI 不兼容。 从源码编译适用于需要特定 CUDA 架构优化或修改 Kernel 的场景: git clone https://github.com/Dao-AILab/flash-attention.git cd flash-attention export TORCH_CUDA_ARCH_LIST="8.0;8.6;8.9;9.0" pip install . TORCH_CUDA_ARCH_LIST 指定目标 SM 版本。默认为当前 GPU 架构,跨节点部署时需显式列出所有目标架构,否则仅编 译当前机器对应的 Cubin,导致其他节点运行时回退至 PTX JIT 编译(启动延迟 5-10 秒)。 AOT 编译加速:默认 setup.py 会调用 NVCC 为每个头维度生成专门的 Cubin( d={16,32,64,128,256} 等),首次安装 耗时 5-20 分钟。设定 MAX_JOBS 可并行编译: MAX_JOBS=4 pip install flash-attn --no-build-isolation -v Docker 安装:在 NGC PyTorch 容器基础上安装 flash-attn 可避免 CUDA 库版本不匹配: FROM nvcr.io/nvidia/pytorch:24.01-py3 RUN pip install flash-attn==2.6.3 --no-build-isolation
  2. flash_attn_func 参数 flash_attn_func 是 flash-attn 的核心入口,对应标准多头注意力的前向计算: from flash_attn import flash_attn_func output = flash_attn_func( q, # (batch_size, seqlen, nheads, headdim)

k, # (batch_size, seqlen_k, nheads_k, headdim) v, # (batch_size, seqlen_k, nheads_k, headdim)

     dropout_p=0.0, # dropout probability applied after softmax
     softmax_scale=None, # scale factor, default: 1/sqrt(headdim)
     causal=False,    # apply causal mask in-kernel
     window_size=(-1, -1), # sliding window attention (left, right)
     alibi_slopes=None, # ALiBi bias slopes
     deterministic=False, # deterministic backward pass
     return_attn_probs=False, # return attention weights

) 各参数含义如下: •q/k/v:输入张量。q 维度为 (batch, seqlen, nheads, headdim) ,k 和 v 维度为 (batch, seqlen_k, nheads_k, headdim) 。支持 GQA/MQA,当 nheads_k < nheads 时,K 和 V 在 head 维度自动广播。要求 headdim <= 256 (FA v1 支持 16-256,FA v3 在 FP16/BF16 下同样支持至 256,FP8 下支持 64-256)。 •dropout_p:Dropout 概率。当 dropout_p > 0 时,前向在 softmax 后施加 dropout,反向传播使用对应 mask。 注意:dropout 启用后无法使用 SDPA 后端,必须走 flash-attn 原生 CUDA Kernel。 •softmax_scale:Softmax 的缩放因子,默认为 1 / sqrt(headdim) 。需与训练时一致,否则产生数值偏差。 •causal:布尔值。True 时在 Kernel 内部直接应用因果掩码(下三角),无需传入显式 mask 张量。Kernel 内部直接跳 过上三角计算,产生与标准因果 mask 逐位一致的结果。 •window_size:滑动窗口注意力参数, (left, right) 分别控制向左和向右的窗口大小。 (-1, -1) 表示无限窗口 即全局注意力。用于 Mistral 和 Longformer 等模型。 •alibi_slopes:ALiBi 位置偏置的斜率张量,形状为 (nheads,) 或 (batch, nheads) 。用于 BLOOM 和 MPT 等模型 的位置编码方案。 •deterministic:True 时反向传播使用确定性的 atomicAdd 实现,位级可复现但略慢(约 5-10%)。默认 False,使用 非确定性并行归约,速度更快但跨运行可能产生微小数值差异。 •return_attn_probs:True 时额外返回 softmax 后的注意力权重,维度 (batch, nheads, seqlen, seqlen_k) 。 注意:返回注意力矩阵会穿透分块优化,迫使 Kernel 将中间结果写回 HBM,性能显著下降。仅用于可视化或调试。 完整前向-反向示例: import torch from flash_attn import flash_attn_func batch, seqlen, nheads, headdim = 2, 4096, 32, 128

 q = torch.randn(batch, seqlen, nheads, headdim, dtype=torch.float16, device='cuda')
 k = torch.randn(batch, seqlen, nheads, headdim, dtype=torch.float16, device='cuda')
 v = torch.randn(batch, seqlen, nheads, headdim, dtype=torch.float16, device='cuda')
 q.requires_grad_(True)
 out = flash_attn_func(q, k, v, dropout_p=0.1, causal=True)
 loss = out.sum()
 loss.backward()
 print(f"Output shape: {out.shape}")
 print(f"Q grad shape: {q.grad.shape}")

GQA (Grouped-Query Attention) 用法,K 和 V 使用更少的头数:

 nheads_kv = 8
 q = torch.randn(batch, seqlen, nheads, headdim, dtype=torch.float16, device='cuda')
 k = torch.randn(batch, seqlen, nheads_kv, headdim, dtype=torch.float16, device='cuda')
 v = torch.randn(batch, seqlen, nheads_kv, headdim, dtype=torch.float16, device='cuda')
 out = flash_attn_func(q, k, v, causal=True)
  1. flash_attn_varlen_func 训练中常使用 Packing 技术将多个不等长序列拼接为一个 batch 以提升 GPU 利用率。 flash_attn_varlen_func 在 Kernel 内部按 cu_seqlens 指针正确应用因果掩码,跨序列边界不施加注意力。 from flash_attn import flash_attn_varlen_func output = flash_attn_varlen_func( q, # (total_tokens, nheads, headdim)

k, # (total_tokens, nheads_kv, headdim) v, # (total_tokens, nheads_kv, headdim) cu_seqlens_q, # (batch_size + 1,) cumulative sequence lengths for Q cu_seqlens_k, # (batch_size + 1,) cumulative sequence lengths for K/V max_seqlen_q, # maximum sequence length in Q max_seqlen_k, # maximum sequence length in K/V

     dropout_p=0.0,
     softmax_scale=None,
     causal=False,
     window_size=(-1, -1),

) 关键参数说明: •cu_seqlens_q / cu_seqlens_k:累积序列长度数组, cu_seqlens[i] 表示第 i 个序列的起始 token 索引。 cu_seqlens_q[-1] 等于 total_tokens 。Kernel 通过差值 cu_seqlens[i+1] - cu_seqlens[i] 获取每个序列的 真实长度。 •max_seqlen_q / max_seqlen_k:批次中最长序列的 token 数,用于 Kernel 内部决定分块循环边界。传递准确值 (而非 pad 长度)是性能优化的关键:Kernel 跳过序列结束后的无效计算。 •当 total_tokens 远小于 padding 到 max_seqlen 的总 token 数时(如闲聊对话数据集的 batch),varlen 方案可节 省 30-60% 的计算量。 在典型的训练管线中,cu_seqlens 可由 DataCollator 在批处理阶段生成。以下展示从变长序列列表到 cu_seqlens 的转 换过程:

 def build_cu_seqlens(seq_lens):
     """Convert list of sequence lengths to cumulative sequence lengths."""
     cumsum = [0]
     for length in seq_lens:
         cumsum.append(cumsum[-1] + length)
     return torch.tensor(cumsum, dtype=torch.int32, device='cuda')
 cu_seqlens = build_cu_seqlens([512, 1024, 256, 768])

varlen 方案与填充方案的性能差异随序列长度方差增大而加剧。当 batch 内最短与最长序列之比低于 0.3 时,填充方案的 无效计算占比超过 50%,varlen 加速尤为显著。在 SFT(Supervised Fine-Tuning)场景中,对话长度分布通常呈现长 尾特征,varlen 是默认推荐方案。 deterministic=True 参数强制 Kernel 使用基于 atomicAdd 的顺序归约路径,避免浮点加法的非结合性导致的跨运行 差异。代价是停用 warp-level 的 shuffle 归约,大约增加 5-8% 延迟。此参数在以下场景中至关重要:(1) 对比实验需要 严格可复现的损失曲线;(2) 分布式训练中不同 rank 需产生位级一致的梯度。在非确定性模式下,同一输入两次前向的 logits 差异通常不超过 1e-3(FP16),但累积多步后梯度差异会放大。

 seq_lens = [1024, 2048, 512]
 total_tokens = sum(seq_lens) # 3584
 cu_seqlens = torch.tensor([0] + torch.cumsum(
     torch.tensor(seq_lens), dim=0).tolist(), dtype=torch.int32, device='cuda')
 q_packed = torch.randn(total_tokens, nheads, headdim, dtype=torch.float16, device='cuda')
 k_packed = torch.randn(total_tokens, nheads, headdim, dtype=torch.float16, device='cuda')
 v_packed = torch.randn(total_tokens, nheads, headdim, dtype=torch.float16, device='cuda')
 out = flash_attn_varlen_func(

q_packed, k_packed, v_packed, cu_seqlens, cu_seqlens,

     max_seqlen_q=max(seq_lens),
     max_seqlen_k=max(seq_lens),
     causal=True,

) 5) 安装后验证步骤 安装完成后执行以下验证流程,确认 Kernel 正确编译且输出数值正确:

  1. 基本导入与 GPU 检测 import torch from flash_attn import flash_attn_func assert torch.cuda.is_available() capability = torch.cuda.get_device_capability() print(f"GPU Compute Capability: {capability[0]}.{capability[1]}") assert capability[0] >= 8, "Ampere or newer required"
  2. 前向数值正确性:与 PyTorch 标准注意力对比 import torch.nn.functional as F batch, seqlen, nheads, headdim = 1, 1024, 8, 64
   q = torch.randn(batch, seqlen, nheads, headdim, dtype=torch.float16, device='cuda')
   k = torch.randn(batch, seqlen, nheads, headdim, dtype=torch.float16, device='cuda')
   v = torch.randn(batch, seqlen, nheads, headdim, dtype=torch.float16, device='cuda')
   fa_out = flash_attn_func(q, k, v, causal=True)
   q_ref = q.permute(0, 2, 1, 3).float()
   k_ref = k.permute(0, 2, 1, 3).float()
   v_ref = v.permute(0, 2, 1, 3).float()
   causal_mask = torch.triu(torch.ones(seqlen, seqlen), diagonal=1).bool().cuda()
   ref_out = F.scaled_dot_product_attention(

q_ref, k_ref, v_ref, attn_mask=causal_mask)

   ref_out = ref_out.permute(0, 2, 1, 3).half()
   max_err = (fa_out - ref_out).abs().max().item()
   mean_err = (fa_out - ref_out).abs().mean().item()
   print(f"Max absolute error: {max_err:.6f}")
   print(f"Mean absolute error: {mean_err:.6f}")

assert max_err < 5e-3, f"Numerical error {max_err} exceeds threshold" 3. 反向梯度正确性

   q_fa = q.clone().detach().requires_grad_(True)
   out_fa = flash_attn_func(q_fa, k, v, causal=True)
   out_fa.sum().backward()
   q_ref_t = q.clone().detach().requires_grad_(True)
   q_ref_t = q_ref_t.permute(0, 2, 1, 3).float()
   out_ref = F.scaled_dot_product_attention(

q_ref_t, k_ref, v_ref, attn_mask=causal_mask)

   out_ref.sum().backward()
   grad_err = (q_fa.grad.permute(0, 2, 1, 3).float() - q_ref_t.grad).abs().max().item()
   print(f"Max gradient error: {grad_err:.6f}")

assert grad_err < 1e-2, f"Gradient error {grad_err} exceeds threshold" 4. 性能基准

18.7.2 PyTorch SDPA 集成

   import time
   for _ in range(10):
      flash_attn_func(q, k, v, causal=True)
   torch.cuda.synchronize()
   N = 100
   start = time.time()
   for _ in range(N):
      flash_attn_func(q, k, v, causal=True)
   torch.cuda.synchronize()
   elapsed = time.time() - start
   print(f"FlashAttention: {elapsed / N * 1000:.2f} ms per forward pass")

PyTorch 从 2.0 版本起提供 torch.nn.functional.scaled_dot_product_attention (SDPA),将多种注意力后端统一 为单一函数接口,运行时根据输入特性自动选择最优实现。FlashAttention 已作为主力后端深度集成至该框架。多数场景 下无需显式调用 flash-attn 库,SDPA 可透明完成替换。

  1. torch.nn.functional 的 SDPA 接口 import torch.nn.functional as F output = F.scaled_dot_product_attention( query, # (B, N, L, D) key, # (B, N_kv, S, D) value, # (B, N_kv, S, D)
     attn_mask=None,   # bias mask (B, 1, L, S), (L, S), (B, L, S) or bool mask
     dropout_p=0.0,    # dropout probability
     is_causal=False, # apply causal mask (recommended over explicit mask)
     scale=None,        # default: 1 / sqrt(D)

) 输入张量使用 PyTorch 原生格式 (batch, num_heads, seqlen_or_target, head_dim) ,区别于 flash-attn 的 (batch, seqlen, num_heads, head_dim) 格式。这一差异源自 PyTorch MultiheadAttention 的历史约定(为与 nn.Linear 的 (B, L, D) 形参对齐) ,SDPA 在内部处理张量布局。 is_causal 优先于 attn_mask:开启 is_causal=True 时,后端可直接在 Kernel 内跳过计算而非施加显式掩码 (FlashAttention 和 cuDNN 后端均支持此优化),性能优于手动传入上三角 mask。PyTorch 文档建议:因果场景下优先 使用 is_causal ,仅在需要任意形状的稀疏掩码时使用 attn_mask 。 GQA/MQA 支持:当 key 和 value 的 num_heads 维度小于 query 的对应维度时,SDPA 自动按组广播,无需用户显 式扩展。广播规则为 num_heads_query % num_heads_kv == 0 ,与 flash-attn 库行为一致。 2) SDPA 后端选择机制 SDPA 在每次调用时,依据输入属性和硬件能力自动选择后端。选择顺序由 torch.backends.cuda.sdp_kernel 上下文 管理器控制: print(torch.backends.cuda.sdp_kernel()) 启用或禁用特定后端: with torch.backends.cuda.sdp_kernel(

     enable_flash=False,
     enable_math=True,
     enable_mem_efficient=False,

): out = F.scaled_dot_product_attention(q, k, v) with torch.backends.cuda.sdp_kernel(

     enable_flash=True,
     enable_mem_efficient=False,
     enable_math=True,

): out = F.scaled_dot_product_attention(q, k, v) 后端自动选择规则(PyTorch 2.5 行为): SDPA Call q, k, v, is_causal, mask Head dim ≤ 256? Yes SM ≥ 80 flash-attn installed? Yes causal=True OR No mask compatible? No No Yes FLASH_ATTENTION SM ≥ 80 Backend cuDNN available? Yes No Supports this dtype? No Yes MATH Backend EFFICIENT_ATTENTION (cuDNN) 图18-27 SDPA 后端选择决策树 根据头维度、GPU 架构、掩码类型和数据类型自动路由。 关键选择条件: •头维度上限: headdim > 256 直接回退到 MATH 后端,因为 flash-attn 的 CUDA Kernel 未实现该尺寸。 •FP8 类型:仅 MATH 和 CUDNN_ATTENTION 后端支持 torch.float8_e4m3fn 。 •自定义 attn_mask:非因果的任意掩码(如 (B, L, S) 形状的 BoolTensor),目前仅 MATH 后端支持。 FlashAttention Kernel 仅处理因果掩码和无掩码两种模式。 •确定性模式:开启 torch.use_deterministic_algorithms(True) 时,FlashAttention 的非确定性归约被禁用,自 动回退到 MATH 后端(仅该后端在所有 CUDA 版本中行为确定)。 不同后端在 softmax 归一化和矩阵乘法的累加顺序上有细微差异。FlashAttention 使用在线 softmax 的分块累加, cuDNN 使用其内部的融合策略,MATH 后端使用完整的矩阵物料化再 softmax。三者的数值误差通常在 FP16 下 max |diff| < 1e-3 ,BF16 下 < 2e-4 。绝大部分训练场景中此差异不造成收敛行为变化,但不建议在同一训练作业的中途 切换后端,混合使用可能导致难以诊断的微小损失震荡。 手工选择后端:通过环境变量 TORCH_SDPA_USE_FLASH_ATTN=0 可全局禁用 flash-attn 后端,用于环境诊断。 3) 从 PyTorch 2.0 起的集成演进 SDPA 的集成并非一蹴而就,而是伴随 PyTorch 2.0 至 2.5 逐步完善: 版本 FlashAttention 集成变化 PyTorch 2.0 SDPA 首次引入,flash-attn 作为可选后端。需独立安装 flash-attn 包。仅支持 SM80+,前向传播。 (2023.03) PyTorch 2.1 反向传播接入 SDPA 后端选择,支持 dropout_p > 0 。cuDNN SDPA 后端(mem_efficient)加入竞 (2023.10) 争。 PyTorch 2.2 scale 参数支持复数类型。张量子类(NestedTensor)支持 Jagged Layout,间接配合变长序列。 (2024.01) PyTorch 2.3 SDPBackend.CUDNN_ATTENTION 作为独立后端,与 mem_efficient 分流。cuDNN 9.x 集成,SM90+ 优 (2024.04) 化。 PyTorch 2.4 FlexAttention API 预览,允许自定义 score modification 函数以 JIT 编译融合 Kernel。 (2024.07) PyTorch 2.5+ flash-attn v2.6+ 对应的预编译 CUDA Kernel 随 PyTorch 分发,部分场景免安装 flash-attn 包。 表18-27 PyTorch 各版本的集成变化 PyTorch 2.5 的关键变化是 flash-attn 内核的捆绑分发: torch wheel 中预编译了 flash-attn 的 CUDA Kernel ( _flash_attn_cuda.so ),覆盖 SM80/86/89/90 架构。若 flash-attn 包已安装,SDPA 使用外部包版本;否则使用内 置版本。内置版本不包含 v3 的 Hopper 专有优化(WGMMA 指令),因此在 H100/H800 上手动安装 flash-attn 2.6+ 可获 得额外约 15-20% 性能增益。 4) 自定义注意力替换示例 将现有模型中的标准注意力替换为 SDPA,需调整形状转换和掩码传递方式。以下展示三种典型替换模式。 替换一:直接替换手动注意力实现

 import torch
 import torch.nn as nn
 import torch.nn.functional as F
 class StandardAttention(nn.Module):
     def __init__(self, embed_dim, num_heads):
         super().__init__()
         self.num_heads = num_heads
         self.head_dim = embed_dim // num_heads
         self.q_proj = nn.Linear(embed_dim, embed_dim)
         self.k_proj = nn.Linear(embed_dim, embed_dim)
         self.v_proj = nn.Linear(embed_dim, embed_dim)
         self.out_proj = nn.Linear(embed_dim, embed_dim)
         self.scale = self.head_dim ** -0.5
     def forward(self, x):

B, L, D = x.shape

         q = self.q_proj(x).view(B, L, self.num_heads, self.head_dim)
         k = self.k_proj(x).view(B, L, self.num_heads, self.head_dim)
         v = self.v_proj(x).view(B, L, self.num_heads, self.head_dim)
         # Manual: QK^T materialized
         attn_weights = torch.matmul(
             q.transpose(1, 2), k.transpose(1, 2).transpose(-2, -1)) * self.scale
         attn_weights = F.softmax(attn_weights, dim=-1)
         attn_output = torch.matmul(attn_weights, v.transpose(1, 2))
         attn_output = attn_output.transpose(1, 2).contiguous().view(B, L, D)
         return self.out_proj(attn_output)
 import torch.nn.functional as F
 class SDPAttention(nn.Module):
     def __init__(self, embed_dim, num_heads, dropout=0.0):
         super().__init__()
         self.num_heads = num_heads
         self.head_dim = embed_dim // num_heads
         self.q_proj = nn.Linear(embed_dim, embed_dim)
         self.k_proj = nn.Linear(embed_dim, embed_dim)
         self.v_proj = nn.Linear(embed_dim, embed_dim)
         self.out_proj = nn.Linear(embed_dim, embed_dim)
         self.scale = self.head_dim ** -0.5
         self.dropout = dropout
     def forward(self, x, causal_mask=True):

B, L, D = x.shape

         # Reshape to PyTorch SDPA format: (B, num_heads, L, head_dim)
         q = self.q_proj(x).view(B, L, self.num_heads, self.head_dim).transpose(1, 2)
         k = self.k_proj(x).view(B, L, self.num_heads, self.head_dim).transpose(1, 2)
         v = self.v_proj(x).view(B, L, self.num_heads, self.head_dim).transpose(1, 2)
         attn_output = F.scaled_dot_product_attention(

q, k, v,

             dropout_p=self.dropout if self.training else 0.0,
             is_causal=causal_mask,
             scale=self.scale,

) attn_output = attn_output.transpose(1, 2).contiguous().view(B, L, D) return self.out_proj(attn_output) 替换二:Cross-Attention 场景 Cross-attention 中 Q 来自 decoder,K/V 来自 encoder,序列长度不同。需设置 is_causal=False :

 def cross_attention_sdpa(query, key, value):
     # query: (B, num_heads, L_q, head_dim)
     # key:   (B, num_heads, L_kv, head_dim)
     return F.scaled_dot_product_attention(

query, key, value, is_causal=False, # cross-attention has no causal mask ) 替换三:训练时 Dropout 与推理时确定性 def attention_flexible(q, k, v, training=True): return F.scaled_dot_product_attention( q, k, v, dropout_p=0.1 if training else 0.0, is_causal=True, ) 训练阶段 dropout_p=0.1 时,FlashAttention 和 cuDNN SDPA 后端均支持 dropout 融合,无需额外 mask 生成或显存 分配。推理阶段 dropout_p=0.0 时,SDPA 选择速度最快的后端,且 deterministic 模式下结果位级可复现。 PyTorch 2.4 引入的 torch.nn.attention.flex_attention API 允许用户以纯 Python 函数定义 score modification 逻辑,PyTorch 通过 torch.compile 将其 JIT 编译为融合 CUDA Kernel。FlexAttention 内部采用与 FlashAttention 相 同的分块策略:QKV 按块加载到 SRAM, score_mod 函数在每个块局部执行,softmax 使用在线归一化。与手写 CUDA Kernel 相比,FlexAttention 的灵活性代价约在 10-25% 的性能开销(因 torch.compile 的通用优化无法完全匹配手写 Kernel 的指令级调优),但大幅降低了定制注意力模式的工程门槛。 支持的 score_mod 模式包括:ALiBi 线性偏置、sliding window(局部窗口,如 Mistral)、document masking(跨文 档阻断注意力)、prefix LM 的因果+双向混合掩码、以及基于相对位置的任意衰减函数。以下为完整示例:

 from torch.nn.attention.flex_attention import flex_attention, create_block_mask
 def composite_score_mod(score, b, h, q_idx, kv_idx):
     """Apply ALiBi bias + sliding window + document mask."""
     # ALiBi: linear penalty based on distance
     alibi_slope = 2.0 ** (-8 * (h + 1) / nheads)
     score = score + alibi_slope * (kv_idx - q_idx)
     # Sliding window: keep only 1024 left + 0 right
     score = torch.where(

(q_idx - kv_idx <= 1024) & (kv_idx <= q_idx), score, float("-inf") ) return score block_mask = create_block_mask( composite_score_mod, B=None, H=None, Q_LEN=seqlen, KV_LEN=seqlen ) out = flex_attention(q, k, v, block_mask=block_mask) 截至 PyTorch 2.5,FlexAttention 的限制包括:(1) 不支持 dropout_p > 0 ;(2) torch.compile 的初始编译需要 5-15 秒延迟,后续调用由于缓存被复用;(3) 仅支持 NVIDIA GPU (SM80+),不支持 ROCm 或 MPS。

18.7.3 HuggingFace Transformers 适配

HuggingFace Transformers 库在 v4.36 起全面支持通过 PyTorch SDPA 调用 FlashAttention,实现一行配置切换注意力 实现。本节阐述其集成机制、模型兼容范围以及训练/推理场景的具体用法。

  1. Transformers 的 SDPA 集成机制 Transformers 库自早期版本即支持多种注意力实现共存:原生 PyTorch 实现( "eager" )、旧版 flash-attn 专用实现 ( "flash_attention_2" ),以及自 v4.36 引入的 PyTorch SDPA 实现( "sdpa" )。模型在初始化时通过 config._attn_implementation 属性记录当前选择的注意力实现类, AutoModel.from_pretrained() 据此分派到对 应的 Attention 模块。 AutoModel.from_pretraine d()

attn_implementation in config? sdpa flash_attention_2 eager Module: XxxSdpaAttention Module: XxxFlashAttention Module: XxxAttention Calls F.scaled_dot_produc 2 Manual PyTorch impl t_attention Direct flash-attn import flash-attn installed? Yes No SDPA selects FlashAttentio SDPA falls back to cuDNN/ n backend Math 图18-28 Transformers 注意力实现分派 attn_implementation 决定模块类,SDPA 模块在原子上进一步委托给 PyTorch SDPA 后端选择。 关键设计点: •sdpa 是最灵活的选项:不直接依赖 flash-attn 库,而是通过 PyTorch SDPA 间接调用。若 flash-attn 已安装,SDPA 自 动使用 FlashAttention 后端;若未安装,回退到 cuDNN 或 Math 后端。这使模型可运行于非 NVIDIA GPU(如 MPS 或 CPU),虽然性能下降但功能完整。 •flash_attention_2 是强依赖选项:模型内部 import flash_attn 调用 flash_attn_func 。若库未安装,模型加载 直接抛出 ImportError 。优点是对 FA 特定参数(如 window_size 和 alibi_slopes )有直接控制;缺点是缺乏跨 平台回退。 •实现类的覆盖范围:并非所有模型都有三种实现。对于 "sdpa" ,库提供了 _check_and_enable_sdpa 工具函数,模 型开发者可就地检查 SDPA 支持:当运行环境不支持(如 CPU-only 模式)时优雅降级为 eager。 Transformers 通过 ALL_ATTENTION_FUNCTIONS 字典实现多态分发。每个模型注册三个键值 对: "eager" 、 "sdpa" 、 "flash_attention_2" (若模型支持),值分别为对应的注意力前向函数。在 AutoModel.init 中,库检查 _attn_implementation 配置值并查找对应函数;若配置的实现在当前模型上未注 册,库自动降级为 "eager" 并打印 warning。这一机制使模型作者可渐进添加 FlashAttention 支持,而不会破坏现有用 户。 不同注意力层可使用不同实现。例如 vision-language 模型中,text encoder 使用 "flash_attention_2" 而 vision encoder 使用 "eager" (因为视觉注意力头维度为 32,FlashAttention 的优化幅度有限)。通过 model.config._attn_implementation_per_module 可精细控制每层的实现选择,但该接口目前在 HuggingFace 主分 支上标记为 experimental。 2) attn_implementation=“flash_attention_2” FlashAttention 2 专用实现是最早进入 Transformers 的 FlashAttention 集成,提供最完整的 FA 参数支持。配置方式: from transformers import AutoModelForCausalLM, AutoConfig model = AutoModelForCausalLM.from_pretrained( "meta-llama/Llama-2-7b-hf",

     attn_implementation="flash_attention_2",
     torch_dtype=torch.float16,
     device_map="auto",

)

 config = AutoConfig.from_pretrained("meta-llama/Llama-2-7b-hf")
 config._attn_implementation = "flash_attention_2"
 model = AutoModelForCausalLM.from_pretrained(

"meta-llama/Llama-2-7b-hf", config=config, torch_dtype=torch.float16, ) "flash_attention_2" 场景下的关键技术约束: •必须 FP16 或 BF16:模型的 torch_dtype 须为 torch.float16 或 torch.bfloat16 。若传入 torch.float32 , flash-attn Kernel 不支持,模型加载时报错 "FlashAttention only supports fp16 and bf16 data type" 。 •Pad token 处理:Transformers 对左填充(left-padding)和右填充(right-padding)的行为不同。 flash_attention_2 实现会检测 attention_mask 中的填充位置并跳过对应计算,但仅支持右填充(batch 中从右 端开始的 padding)。若需要左填充,使用 "sdpa" 实现。 •支持 use_cache:KV cache 在推理时正常传递。flash-attn Kernel 本身不操作 cache,Transformer 的 Attention 模 块在调用 Kernel 前后完成 cache 追加。 •output_attentions=True 需要特殊处理:flash-attn 的 CUDA Kernel 在优化路径下不输出注意力矩阵。当设置 output_attentions=True 时,实现自动切换至 eager 代码路径的一次性前向,将完整注意力权重返回。这一回退隐 含巨大的显存开销(O(N )),仅用于调试。 2 3) 支持的模型列表与版本 以下为关键模型族对 "flash_attention_2" 和 "sdpa" 的支持状态(Transformers v4.46,2024 年 11 月): 模型族 flash_attention_2 sdpa 备注 LLaMA / LLaMA 2 / LLaMA 3 Yes (v4.31+) Yes (v4.36+) GQA 完全支持 Mistral Yes (v4.34+) Yes (v4.36+) sliding_window 原生支持 Gemma / Gemma 2 Yes (v4.38+) Yes (v4.38+) 9:1 GQA ratio Mixtral Yes (v4.36+) Yes (v4.36+) MoE + GQA,8 experts Falcon Yes (v4.32+) Yes (v4.36+) MQA (n_kv_heads=1) GPT-NeoX / Pythia Yes (v4.32+) Yes (v4.36+) 标准 MHA BLOOM No SDPA (v4.36+) ALiBi 偏置,仅 SDPA MPT No SDPA (v4.36+) ALiBi + 自定义偏置 Qwen2 Yes (v4.40+) Yes (v4.40+) GQA + YaRN RoPE 模型族 flash_attention_2 sdpa 备注 Phi-3 Yes (v4.41+) Yes (v4.41+) Block-sparse (sdpa only) Gemma 2 Yes (v4.43+) Yes (v4.43+) Pre/Post norm + logit soft-capping 表18-28 Transformers 模型支持矩阵 版本兼容提示: •Transformers v4.36+ 建议默认使用 "sdpa" ,免安装 flash-attn 库即可获得加速。 • "flash_attention_2" 在安装 flash-attn 2.3+ 后启用额外优化(如 v3 的 Hopper 特性)。对于 H100 训练,建议安装 flash-attn 2.6+ 并使用 "flash_attention_2" 获得完整性能。 •当同时安装 flash-attn 并使用 "sdpa" 时,SDPA 自动路由至 FlashAttention 后端,效果与 "flash_attention_2" 近似(差异 < 5% 吞吐),但灵活性更高。 Transformers 为推理场景提供了额外的编译优化。 BetterTransformer API( model.to_bettertransformer() )在 PyTorch 1.13+ 上透明替换 Attention 模块,对 "sdpa" 实现自动开启 torch.compile 的 CUDA Graph 捕获。CUDA Graph 将多次 Kernel 启动合并为单次 GPU 操作提交,消除了 CPU-GPU 的同步开销。在 batch=1 的推理场景中,CUDA Graph 结合 FlashAttention 可将端到端延迟降低额外 10-20%。注意: BetterTransformer 仅适用于推理,训练中不可 用(梯度图与 CUDA Graph 不兼容)。 使用 LoRA(Low-Rank Adaptation)微调时, "sdpa" 和 "flash_attention_2" 均可正常工作。PEFT 库通过 adapters_weights 钩子修改 Q/K/V 投影层的输出,注意力计算本身不受影响。若在 LoRA 微调中遇到 NaN 损失,常见 原因是 dropout 与 SDPA 的交互:部分 SDPA 后端在 dropout_p > 0 时使用不同的随机数生成器(Philox vs MTGP), 与 eager 实现的 CUDA RNG 状态不同。解决方案是固定 torch.manual_seed 或设置 dropout_p=0 进行初步调试。 4) 训练与推理场景示例 训练:Llama-2-7B 微调,启用 FlashAttention import torch from transformers import ( AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer )

 from datasets import load_dataset
 model_name = "meta-llama/Llama-2-7b-hf"
 model = AutoModelForCausalLM.from_pretrained(

model_name,

     attn_implementation="flash_attention_2",
     torch_dtype=torch.bfloat16,
     device_map="auto",

)

 tokenizer = AutoTokenizer.from_pretrained(model_name)
 tokenizer.pad_token = tokenizer.eos_token
 def tokenize_function(examples):
     return tokenizer(examples["text"], truncation=True, max_length=4096)
 dataset = load_dataset("wikitext", "wikitext-2-raw-v1")
 tokenized = dataset.map(tokenize_function, batched=True, remove_columns=["text"])
 training_args = TrainingArguments(
     output_dir="./llama2-flash-finetune",
     per_device_train_batch_size=4,
     gradient_accumulation_steps=8,
     bf16=True,
     logging_steps=10,
     save_strategy="epoch",
     num_train_epochs=1,

)

 trainer = Trainer(
     model=model, args=training_args,
     train_dataset=tokenized["train"],
     data_collator=lambda data: {

"input_ids": torch.stack([d["input_ids"] for d in data]), "attention_mask": torch.stack([d["attention_mask"] for d in data]), "labels": torch.stack([d["input_ids"] for d in data]), }, ) trainer.train() 推理:批量生成,KV Cache 加速 from transformers import pipeline pipe = pipeline( "text-generation", model="mistralai/Mistral-7B-Instruct-v0.2", model_kwargs={ "attn_implementation": "sdpa", "torch_dtype": torch.float16, }, device=0, ) messages = [ {"role": "user", "content": "Explain FlashAttention in one paragraph."} ] output = pipe(messages, max_new_tokens=200, do_sample=True, temperature=0.7) print(output[0]["generated_text"][-1]["content"]) 推理:长上下文(32K tokens),滑动窗口注意力 model = AutoModelForCausalLM.from_pretrained( "mistralai/Mistral-7B-Instruct-v0.2",

     attn_implementation="flash_attention_2",
     torch_dtype=torch.float16,
     device_map="cuda",

)

 tokenizer = AutoTokenizer.from_pretrained("mistralai/Mistral-7B-Instruct-v0.2")
 long_doc = "Mistral attention window " * 8000 # ~32K tokens
 inputs = tokenizer(long_doc, return_tensors="pt").to("cuda")

with torch.no_grad():

     outputs = model.generate(
         **inputs, max_new_tokens=100,
         use_cache=True, # KV cache for autoregressive decoding

) print(f"Input length: {inputs.input_ids.shape[1]}") print(f"Generated: {tokenizer.decode(outputs[0], skip_special_tokens=True)[:200]}") 上述代码中,Mistral 模型的 "flash_attention_2" 实现直接向 flash_attn_func 传递 window_size=(-1, 4096) (左无限,右 4096),Kernel 内部仅计算窗口内的注意力分数。32K 长上下文场景下,相比窗口外也计算的 Math 后端, FlashAttention 可节省约 85% 的注意力计算量。 调试:切换实现以对比输出

 def compare_attn_implementations(model_name, prompt):
     outputs = {}
     for attn in ["eager", "sdpa", "flash_attention_2"]:

try: model = AutoModelForCausalLM.from_pretrained( model_name,

                 attn_implementation=attn,
                 torch_dtype=torch.float16,
                 device_map="cuda",

) inputs = tokenizer(prompt, return_tensors="pt").to("cuda") with torch.no_grad(): logits = model(**inputs).logits outputs[attn] = logits except Exception as e: outputs[attn] = f"ERROR: {e}"

18.7.4 跨架构兼容与调优

     # Compare sdpa vs flash_attention_2 outputs
     if isinstance(outputs.get("sdpa"), torch.Tensor) and \
        isinstance(outputs.get("flash_attention_2"), torch.Tensor):
         diff = (outputs["sdpa"] - outputs["flash_attention_2"]).abs().max()
         print(f"Max logit difference (sdpa vs fa2): {diff.item():.6f}")

FlashAttention 的 CUDA 实现高度依赖 NVIDIA GPU 的硬件特性:Tensor Core、共享内存、异步拷贝指令和 warp-level 原语。将 FA 移植至其他硬件平台需重新设计 Kernel 以适配不同的指令集和存储架构。本节梳理主流 GPU 和加速器平台 的兼容状态,并给出常见安装故障的排查与调优方法。

  1. NVIDIA GPU 架构兼容矩阵 flash-attn 针对不同 SM 架构编译不同的 PTX/Cubin,各架构可用的特性子集不同: 架构 SM 代表 GPU FA v1 FA v2 FA v3 关键特性 Volta 70 V100 Partial No No FP16 Tensor Core;BF16/ cp.async 不支持 Turing 75 T4, RTX 2080 Ti Partial No No FP16 Tensor Core;缺少 BF16 Ampere 80 A100 Full Full No BF16/ cp.async /L2 Residency Control 架构 SM 代表 GPU FA v1 FA v2 FA v3 关键特性 Ampere 86 RTX 3090, A40 Full Full No 同上,SM 数量减少 Ampere 89 RTX 4090 (Ada) Full Full No 新增 FP8 TransformerEngine 支持 Ada Lovelace 89 RTX 4090 Full Full No 同 SM89 Hopper 90 H100, H800 Full Full Full TMA/WGMMA/FP8; v3 核心目标架构 Hopper 90a H200, H100 NVL Full Full Full HBM3e 141 GB Blackwell 100 B200 Full Full Full (FA-4) FP4 精度; 第二代 Transformer Engine 表18-29 GPU 架构支持矩阵 Volta/Turing 在 2024 年 flash-attn 2.x 分支中不再支持。关键缺失: CUDA_ARCH >= 800 守卫阻止编译;flash- attn 的 Kernel 使用 cp.async ( __pipeline_memcpy_async ) 进行异步 HBM→SRAM 数据搬运,而 Volta/Turing 仅支 持同步 __ldg 加载。若需在 V100 上使用,须锁定 flash-attn==0.2.8(无 cp.async ,使用 __syncthreads() 同步方 案)。 FP8 支持仅限 Hopper+。FA v3 的 FP8 Kernel 依赖 wgmma.fence 和 wgmma.commit_group 等 Hopper 专属 PTX 指 令,无法在 SM89 上运行(即使 Ada Lovelace 在硬件层面支持 FP8 Tensor Core)。此限制由 NVIDIA PTX ISA 的 target sm_90a 编译导向决定。 CUBIN 兼容性:flash-attn 默认在安装时为当前 GPU 编译 Cubin。跨架构部署时(例如在 RTX 4090 (SM89) 上构建,部 署至 A100 (SM80)),须设置 TORCH_CUDA_ARCH_LIST="8.0;8.9" 确保为两架构均生成 Cubin。若某个目标架构缺失, NVIDIA 驱动将在首次 Kernel 启动时执行 PTX→Cubin 的 JIT 编译,可能引入 5-15 秒冷启动延迟。
  2. AMD ROCm 的 FlashAttention 移植状态 AMD 通过 ROCm 生态对 FlashAttention 进行了社区和官方两线移植: 官方移植:ROCm/flash-attention(AMD 维护的 fork) git clone https://github.com/ROCm/flash-attention.git cd flash-attention pip install . AMD 的移植要点: •使用 HIP(Heterogeneous-compute Interface for Portability)替代 CUDA, globalglobal (HIP 编 译宏自动转译 global 等 CUDA 关键字)。 •MFMA 指令(Matrix Fused Multiply-Add)替代 NVIDIA 的 mma.sync 。AMD CDNA2(MI250X)和 CDNA3 (MI300X)架构的 Matrix Core 通过 MFMA 实现 Tensor Core 等价操作,编程接口为内联汇编 __builtin_amdgcn_mfma_f32_16x16x16f16 。 •共享内存差异:CDNA 架构(MI200 系列)LDS(Local Data Share,对应共享内存)为 64 KB per CU,小于 A100 的 192 KB per SM。移植时需要减小分块尺寸以适配,有时导致并行度下降。 •异步拷贝:CDNA2 支持 __builtin_amdgcn_s_buffer_load 实现异步全局到 LDS 的加载,但语义与 cp.async 有 差异(部分操作缺少 commit group 机制)。 2024 年底,ROCm fork 已支持 MI250X 和 MI300X 上的 FA v1 和 v2,FA v3 移植仍在进行中(TMA 和 WGMMA 的 HIP 等 价物尚未标准化)。性能和功能对比: 功能 NVIDIA (CUDA) AMD ROCm FA v1 Full Full FA v2 Full Full (MI250X/MI300X) FA v3 Full (H100+) WIP 功能 NVIDIA (CUDA) AMD ROCm FP8 H100+ MI300X (via hipBLASLt) BF16 Ampere+ CDNA2+ Deterministic backward Yes Partial window_size Full Full alibi_slopes Full Full 表18-30 CUDA 与 ROCm 支持对比 PyTorch SDPA 在 ROCm 平台的行为:若 torch 检测到 gfx90a (MI250X) 或 gfx942 (MI300X) GPU 且 torch.backends.cuda.is_bf16_supported() 为 True,SDPA 使用 aotriton 后端(AMD 的 Triton Kernel 集合)替 代 CUDA 后端。 enable_flash=True 项在 ROCm 平台上的映射是 aotriton 的 FlashAttention Kernel,而非 NVIDIA 的 flash_attn CUDA Kernel。
  3. Apple Silicon 及其他平台适配 Apple Silicon (M1/M2/M3/M4) Apple Silicon GPU 没有 CUDA 生态,FlashAttention 通过两条路径间接可用: •PyTorch MPS 后端: torch.backends.mps 在 macOS 13+ 启用, F.scaled_dot_product_attention 在 MPS 后 端上执行 "math" 回退,使用 Metal Performance Shaders 加速矩阵乘法。无 FlashAttention 级别的 I/O 优化,但得 益于统一内存架构(UMA),CPU/GPU 间无显式数据拷贝,部分缓解了 HBM 瓶颈。实测 M2 Max 上 SDPA(math) 处理 4096×4096 注意力的延迟约为同参数 A100 的 3-5 倍。 •MLX 框架:Apple 官方维护的 mlx 库内置 mlx.nn.fast.scaled_dot_product_attention ,实现了 Apple GPU 专 用的 FlashAttention Kernel。MLX 使用 Metal Shading Language(MSL)重写分块循环和在线 softmax,利用 threadgroup memory 和 SIMD-group 原语对应 CUDA 的 shared memory 和 warp 级操作。该实现非官方 flash-attn 移植,而是独立 Kernel,接口不兼容 flash-attn Python 包,但 HuggingFace Transformers 中的 "sdpa" 模式在 MLX 运行时上可用。
 import mlx.core as mx
 import mlx.nn as nn
 q = mx.random.normal((batch, seqlen, nheads, headdim))
 k = mx.random.normal((batch, seqlen, nheads, headdim))
 v = mx.random.normal((batch, seqlen, nheads, headdim))
 out = nn.fast.scaled_dot_product_attention(q, k, v, scale=headdim ** -0.5, mask=None)

Intel GPU (oneAPI) Intel Data Center GPU Max Series (Ponte Vecchio) 通过 intel-extension-for-pytorch (IPEX) 提供 SDPA 实现,后 端为 oneDNN Graph Compiler。2024 年 Q4,IPEX 的 FlashAttention 优化处于实验阶段,具备基本的分块策略但性能 不及 NVIDIA 方案。具体限制包括:(1) 分块大小被固定为 256,无法根据 SRAM 大小动态调整;(2) 在线 softmax 的 FP32→FP16 混合精度路径未充分优化,归一化成为额外瓶颈;(3) 不支持 window_size 参数。2025 年 1 月发布的 PyTorch 2.6 将 torch.xpu 提升为官方稳定后端,IPEX 的 FlashAttention 路径随之达到生产可用状态。 Google TPU TPU 的矩阵乘法单元(MXU)本征地支持 b × b 脉动阵列乘法,其片上 HBM 带宽(v5p 约 4.8 TB/s)缓解了部分 I/O 瓶 颈。JAX 的 jax.nn.dot_product_attention 内部使用 TPU 定制的分块策略(Pallas Kernel),称为 “FlashAttention for TPU”。TPU 实现的核心差异在于:(1) 使用 VPU (Vector Processing Unit) 而非 GPU SM 执行 softmax 的逐元素操 作;(2) 分块维度以 TPU 的 128×128 MXU tile 为基本单位,而非通用的 64/128 头维度;(3) 利用 TPU 的显式内存预取 指令 (DMA descriptors) 实现与 CUDA cp.async 类似的异步数据搬运。该实现非 flash-attn 代码的移植,而是借鉴了分 块+重计算的思路在 TPU 的向量/矩阵处理单元上的重新实例化。PyTorch/XLA 客户可通过 torch_xla.experimental.custom_kernel 使用等效实现。 TPU v5e 和 v5p 上,JAX FlashAttention 处理 4096 序列长度的自注意力延迟约 0.8-1.2 ms(BF16),低于同配置的 A100 SXM(约 1.5-2 ms),更是远优于纯 Math 后端(约 3.5 ms)。这得益于 TPU 的脉动阵列对大矩阵乘法的天然优势, 以及比 A100 HBM2e 更高的片上内存带宽。 4) 常见安装问题与性能调优 安装故障排查: 症状 原因 解决方案 RuntimeError: FlashAttention GPU 为 降级至 flash-attn==0.2.8 或使用 PyTorch SDPA 的 cuDNN 后 only supports Ampere GPUs or newer Volta/Turing 端 PyTorch C++ ABI pip install --no-build-isolation 或使用预编译 wheel undefined symbol: _ZN3c108 ... 不匹配 nvcc fatal: Unsupported GPU CUDA Toolkit 版本 升级 CUDA 至 11.8+ 或 12.x architecture 'compute_89' 过低 error: ‘__pipeline_memcpy_async’ was not declared CUDA 11.6 以下 升级 CUDA 至 11.8+ ImportError: /lib64/libc.so.6: 系统 glibc 版本过 使用 Docker 容器(NGC PyTorch 基础镜像包含正确 glibc) version GLIBC_2.28 not found 旧 CUDA error: no kernel image is 缺少当前 GPU 架构 设置 TORCH_CUDA_ARCH_LIST="8.0;8.6;8.9;9.0" 重装 available for execution 的 Cubin OSError: libcudart.so.12: cannot CUDA 运行时库版 安装匹配的 CUDA 12.x wheel: pip install flash-attn - open shared object file 本与 wheel 不匹配 -index-url https://download.pytorch.org/whl/cu121 表18-31 常见安装错误与解决方案 性能调优参数: flash_attn_func 和 F.scaled_dot_product_attention 的关键性能参数: 头维度 (headdim):FA 的内 Kernel 按头维度展开,部分维度大小使用专用 Kernel。典型阶梯为 16、32、64、128、 256。使用 64 或 128 可获得最佳吞吐率;非典型维度(如 80)会走通用循环路径,性能降低 15-25%。 块尺寸 (block size):flash-attn 内部根据 SRAM 容量和头维度自动计算最优块尺寸,通常无需手动调整。调试时可通过 环境变量干预: export FLASH_ATTN_BR=128 # Q block row size (default: 128) export FLASH_ATTN_BC=128 # K/V block column size (default: 128) 批处理批量 (batch size):FlashAttention 在每个序列的每个头上独立执行,batch size 增大仅增加并行 Kernel 数量,不 改变单 Kernel 行为。对于小 batch(B ≤ 4),SM 利用率可能不足,可启用 torch.compile 将 batch 内多个 Kernel 合 并为 CUDA Graph 以减少 launch 开销。 flash-attn 对不同序列长度范围使用不同的 Block 参数。短序列(N ≤ 512)的分块循环开销相对显著,此时 cuDNN SDPA 后端可能比 FlashAttention 后端更快。实测对比可参考:

 def benchmark_sdpa_backends(seqlen, batch=4, nheads=32, headdim=128):
     q = torch.randn(batch, nheads, seqlen, headdim, dtype=torch.float16, device='cuda')
     k = torch.randn(batch, nheads, seqlen, headdim, dtype=torch.float16, device='cuda')
     v = torch.randn(batch, nheads, seqlen, headdim, dtype=torch.float16, device='cuda')
     import time
     for backend_name, ctx in [

("flash", torch.backends.cuda.sdp_kernel( enable_flash=True, enable_math=False, enable_mem_efficient=False)), ("mem_efficient", torch.backends.cuda.sdp_kernel( enable_flash=False, enable_math=False, enable_mem_efficient=True)), ("math", torch.backends.cuda.sdp_kernel( enable_flash=False, enable_math=True, enable_mem_efficient=False)), ]: with ctx:

             # Warmup
             for _ in range(10):
                 F.scaled_dot_product_attention(q, k, v, is_causal=True)
             torch.cuda.synchronize()
             steps = 50
             start = time.time()
             for _ in range(steps):
                 F.scaled_dot_product_attention(q, k, v, is_causal=True)
             torch.cuda.synchronize()
             elapsed = (time.time() - start) / steps * 1000
             print(f" {backend_name:>14s}: {elapsed:8.2f} ms (N={seqlen})")
 for N in [256, 512, 1024, 2048, 4096, 8192]:
     print(f"\nSequence length: {N}")
     benchmark_sdpa_backends(N)

典型结果(A100, B=4, head=32, d=128):短序列(N=256)时 cuDNN mem_efficient 快 5-10%,中长序列 (N≥1024)FlashAttention 快 15-30%。此交叉点随 GPU 架构变化:H100 上 FA v3 几乎在所有序列长度上均最快。 FP8 推理优化仅 H100/H800(SM90)。使用 FA v3 的 FP8 Kernel 需同时将 QKV 投影层的输出转换为 FP8:

 from flash_attn import flash_attn_func
 q_fp8 = q.to(torch.float8_e4m3fn)
 k_fp8 = k.to(torch.float8_e4m3fn)
 v_fp8 = v.to(torch.float8_e4m3fn)
 out = flash_attn_func(q_fp8, k_fp8, v_fp8)

FP8 模式下软注意力分数的归一化在 FP16 精度进行(FlashAttention 论文证明了将 softmax 中间精度提升至 FP16 可维 持与 BF16 一致的训练稳定性),最终输出写回 FP8。 张量并行(Tensor Parallelism, TP)下,每个 GPU 持有部分注意力头。flash-attn 的 GQA 支持意味着 K/V 头数可按 TP 规模进一步压缩:在 TP=4 配置下,每 GPU 仅需 nheads_kv / 4 个 KV 头,结合 GQA 的比例因子,KV cache 总大小可 减少至原来的 1/8(TP=4 × GQA=2)。

18.8 扩展与前瞻

18.8.1 Block-Sparse FlashAttention

Block-Sparse FlashAttention 是 FlashAttention 系列在稀疏注意力方向上的重要扩展。它通过引入结构化的块稀疏模式 (Block-Sparse Pattern),在保持 FlashAttention 高效 IO 特性的前提下,将注意力计算的复杂度从 O(n ) 降至 O(n ⋅ 2 n) 或更低,显著扩展了长序列处理的可行性边界。

  1. 稀疏注意力的动机与适用场景 标准缩放点积注意力对所有 token 位置两两计算相似度,产生稠密的 n × n 注意力矩阵。这一设计的隐含假设是:每个 token 都可能与序列中任意位置的 token 存在有意义的关联。然而大量实证研究表明,Transformer 学得的注意力分布往 往高度稀疏:多数注意力权重集中在少数关键 token 上,其余位置权重接近零。 长序列处理的三个典型场景直接受益于稀疏注意力: •处理数万 token 的完整文档时,局部上下文(相邻段落)远比跨章节引用更频繁。滑动窗口(Sliding Window)注意 力将计算限制在窗口大小 w 内,复杂度降至 O(n ⋅ w)。 •视频帧或音频频谱中,时间相邻帧的关联远强于远距离帧。块状局部注意力配合少量全局 token 即可保持性能。 •检索增强生成(RAG):长上下文中,模型仅需关注相关文档片段,其余部分可安全忽略。基于检索的稀疏模式根据查 询-文档相似度动态选择 key-value 对。 理论和实验一致表明:在 Long Range Arena (LRA) 基准测试、32K 以上序列的语言模型预训练、以及高分辨率图像生成 (如 ViT)中,精心设计的稀疏模式在精度损失可控(通常 < 1% 困惑度差异)的前提下可带来 2-10 倍的加速。
  2. 块稀疏模式的矩阵结构 FlashAttention 的分块计算天然适合块级别的稀疏性。当整个注意力矩阵的某个 B × B 子块全部被判定为可跳过时,该 块对应的 S 计算、softmax 归一化和 P V 累加均可完全省略。 r c ij 三种典型的块稀疏模式: Sliding Window Global + Local Strided Block-tridiagonal pattern Sliding window + global to Regular stride sampling

kens O(n·w) O(n·(w+g)) O(n²/k) Token i attends only to [i- Local window + all-to-all o Attend to every k-th token, w, i+w] n g global tokens k = stride 图18-29 三种典型块稀疏注意力模式 三种模式的机制与代表模型如下: •滑动窗口(Sliding Window):注意力矩阵的非零块沿对角线分布,构成块状三对角结构。每个 token 仅关注前后各 w/2 范围内的 token,窗口外权重隐式设为零。Mistral 7B、Longformer 均采用此模式。 •全局 Token + 局部窗口:在滑动窗口基础上增加少量全局 token(如 CLS token、BOS/EOS token),这些 token 与所 有其他 token 互相关注(稠密行/列)。BigBird 论文证明了这种模式在理论上可以保持 Transformer 的通用逼近能力。 •跨步(Strided)模式:以固定步长 k 采样 key 位置,类似膨胀卷积的思想。结合滑动窗口使用可同时捕获局部细节和 长程轮廓,Sparse Transformer 原始论文即采用此方案。 Block-Sparse Attention Matrix (n=8, block_size=2, window=3 blocks) K0 K1 K2 K3 | K4 K5 K6 K7 Q0 [##] [##] [ ] | [ ] [ ] [ ] [ ] Q1 [##] [##] [ ] | [ ] [ ] [ ] [ ] Q2 [##] [##] [##] | [ ] [ ] [ ] [ ] Q3 [##] [##] [##] | [ ] [ ] [ ] [ ] ---+----------------+----------------------- Q4 [ ] [ ] [##] | [##] [##] [ ] [ ] Q5 [ ] [ ] [##] | [##] [##] [ ] [ ] Q6 [ ] [ ] [ ] | [##] [##] [##] [##] Q7 [ ] [ ] [ ] | [##] [##] [##] [##] [##] = computed block, [ ] = skipped block 图18-30 8-token 块稀疏注意力矩阵 块大小 2,窗口 3 个块。 3) Block-Sparse FlashAttention 的算法 Block-Sparse FlashAttention 将稀疏性编码为注意力掩码的块级别判定,核心设计要点有三: 预处理阶段根据用户指定的稀疏模式参数(窗口大小、全局 token 索引、步长等)生成 CSR(Compressed Sparse Row)格式的块索引表。对于每个 query 块 i, block_mask[i] 存储其应计算的 key 块列表。该索引表在输入形状不变 时仅需构造一次,批量推理中可复用。 在标准 FlashAttention 的双重循环中,外层遍历 query 块,内层遍历 key 块。稀疏版本在内循环中仅遍历 CSR 索引指向 的有效 key 块,跳过所有被掩码的块。CUDA kernel 中通过 block_mask 直接决定 cp.async 加载的目标地址,避免在 HBM 中读入数据后再判定丢弃。

 // Sparse inner loop: only iterate over active key blocks
 // block_mask[i]: array of active block indices for query block i
 for (int j_idx = 0; j_idx < block_mask_len[i]; j_idx++) {

int j = block_mask[i][j_idx]; // global column block index // Load K_j, V_j tiles from HBM to SRAM load_tile(K_tile, K + j * Bc * d); load_tile(V_tile, V + j * Bc * d); // Compute S_ij = Q_i @ K_j^T, scale, softmax update, accumulate matmul(S, Q_i, K_tile); // B_r x B_c online_softmax_update(S, m, l, O_i); matmul(O_i, P, V_tile); // B_r x B_c @ B_c x d } // Normalize O_i by row sum l, write to HBM 被跳跃的块在软最大化(softmax)中应贡献零(权重为负无穷时的指数为 0)。Block-Sparse FlashAttention 的 online_softmax_update 与密集版本一致,因未参与计算的 S 从不加载,自然不参与最大值更新和指数求和。这与显 式填充 -inf 再计算 softmax 在数学上等价,但避免了加载无用数据。 ij 4) 性能与精度权衡分析 Block-Sparse FlashAttention 的性能收益来自内循环迭代次数的减少。设密集注意力需计算 T × T 次块乘法,稀疏版本 仅需计算 T × kˉ 次,其中 kˉ 为每 query 块的平均有效 key 块数。 r c r 在滑动窗口模式(窗口大小 w)下,kˉ = min(T , w/B ),前向传播的理论加速比为: c c Tr ⋅ Tc Tc speedup ≈ = Tr ⋅ min(Tc , w/Bc ) min(Tc , w/Bc ) 当 w ≪ n 时(极端长序列场景),加速比趋近 n/w。以 64K token 序列、窗口 4096 为例,理论加速约 16 倍。 精度方面,稀疏注意力引入的误差取决于稀疏模式与模型实际注意力分布的匹配程度。关键结论: 稀疏模式 精度影响 典型加速比 适用场景 Sliding Window (w=4096) 困惑度上升 < 0.5 2-8x 语言模型预训练/推理 Global + Local (g=64) 困惑度上升 < 0.2 1.5-6x 需要跨段理解的文档任务 Strided (stride=2) 困惑度上升 1-2% 1.3-2x 长程依赖必须保留的场景 Learned Sparsity (via router) 取决于路由质量 1.5-5x 自适应稀疏,需额外训练开销 表18-32 稀疏模式精度与加速的权衡 Block-Sparse FlashAttention 作为 FlashAttention 原生稀疏扩展,其核心价值在于:将稀疏性的收益直接嵌入 IO-aware 计算框架,避免先计算稠密矩阵再掩码的浪费。但它并非银弹。稀疏模式的选择需要与模型架构匹配,盲目的激进稀疏会 导致信息丢失;与 FlashDecoding 等互补方案结合使用可进一步扩大长序列加速空间。

18.8.2 FlashDecoding 长序列推理

FlashDecoding 是斯坦福 CRFM 于 2023 年提出的长上下文推理加速方法,专门针对 FlashAttention 在大 batch size 推 理时并行度不足的瓶颈。它通过将 KV 序列沿长度维度分块、在多线程间并行计算注意力,使长序列推理的延迟不再随序 列长度线性增长。

  1. 长序列推理的注意力瓶颈 LLM 推理分为预填充(Prefill)和自回归解码(Decoding)两阶段。FlashAttention 在两阶段的表现差异显著: 预填充阶段处理完整的输入序列,query 矩阵形状为 [b, n , d],其中 b 为 batch size,n 通常等于 prompt 长度。此时 FlashAttention 可沿 batch 和 head 维度分配大量线程块(b ⋅ h 个),SM 利用率较高。 q q 解码阶段每个 token 的 query 仅为单行向量(n = 1)。此时 batch size 在交互式服务中通常很小(b = 1 ∼ 4)。 FlashAttention 在解码时沿 M 维(序列长度维度)的并行度为 batch * num_heads ,在小 batch 场景下远不足以填充 q 全部 SM。 Prefill Phase Decoding Phase Q: [b, n_q, d] Q: [b, 1, d] Parallelize over b×h rows Only b×h thread blocks - -> good occupancy > poor occupancy

Long KV cache (n_kv >> 1) Each block traverses all K V -> high latency per step 图18-31 预填充与解码阶段的并行度差异 以 Llama-2-70B 在单张 A100 上的推理为例: num_heads = 64, batch = 1 ,解码阶段仅能启动 64 个线程块,而 A100 有 108 个 SM,利用率约 59%。若 KV 缓存长度达到 100K tokens,每个线程块的串行遍历步数 T = ⌈100000/128⌉ ≈ 782,单次注意力计算的总延迟被串行内循环主导。 c 2) FlashDecoding 的 KV 分块并行策略 FlashDecoding 的核心创新在于:解码时不再让每个 query 的线程块串行遍历全部 KV 缓存,而是将 KV 序列沿长度维度 进一步分块,分配给多个线程块并行计算局部注意力,最后通过一次高效的归约(reduce)合并结果。 算法分三步:

  1. KV 分块。将完整的 KV 缓存(形状 [n , d])拆分为 T 个等长子块,每个子块分配给一个独立线程块。线程块数 T 的选 择确保 GPU 高占用率,例如取 T = ⌈n /B ⌉,通常为数百到数千。 kv kv c
  2. 局部注意力。每个线程块独立计算其负责的 KV 子块与 query 的局部点积、局部 softmax 统计量(最大值 m 和指数和 l ),以及局部输出 O 。这一步可并行执行,各块无依赖。 t t t
  3. 归约合并。所有线程块完成后,使用 online softmax 的合并公式将各块的局部最大值、指数和、输出加权组合为最终 结果: ∑Tt=1 Ot ⋅ emt −mglobal O=

T ∑t=1 lt ⋅ emt −mglobal // FlashDecoding: parallel KV partition per query head // Step 1: each block computes partial stats for its KV slice float m_local = -INFINITY, l_local = 0.0; float O_local[d] = {0}; // partial output accumulator for (int j = block_start; j < block_end; j += Bc) { load_tile(K_tile, K_cache + j * d); // load KV tile from cache load_tile(V_tile, V_cache + j * d); matmul(S, Q_head, K_tile); // 1 x Bc online_softmax_partial(S, &m_local, &l_local, O_local, V_tile);

          }
          // Store (m_local, l_local, O_local) to global memory
          // Step 2: reduction kernel merges all partial results

float m_global = max(m_locals[0..T]), l_global = 0.0; float O_final[d] = {0}; for (int t = 0; t < T; t++) { float w = expf(m_locals[t] - m_global); l_global += l_locals[t] * w; for (int k = 0; k < d; k++) O_final[k] += O_locals[t][k] * w; } for (int k = 0; k < d; k++) O_final[k] /= l_global; 与标准 FlashAttention 解码的对比:假设 KV 缓存长度 128K、块大小 128,标准方案需 1 个线程块内循环 1000 步(串 行);FlashDecoding 启动 1000 个线程块,每块内部仅循环 1 步,总延迟从 1000T 降至约 T(并行),加速比接近块数。 3) 与 PagedAttention 的组合使用 FlashDecoding 和 vLLM 的 PagedAttention 解决了推理优化的两个正交维度: 维度 PagedAttention FlashDecoding 优化目标 KV 缓存内存利用率 解码算子延迟 核心机制 分页虚拟内存管理 KV 块 KV 分块并行计算注意力 作用于 内存分配与碎片消除 SM 利用率的提升 与 FlashAttention 的关系 兼容,替换注意力 kernel 的内存布局 正交,修改注意力 kernel 的并行策略 表18-33 FlashDecoding 与 PagedAttention 的分工 两者组合时,FlashDecoding 的 KV 分块加载需适配 PagedAttention 的非连续内存布局。PagedAttention 将 KV 缓存存 储为分散的物理页(page),每页大小通常为 16 或 32 个 token。FlashDecoding 的 worker 线程块从 PagedAttention 的页表中解析每个逻辑位置的物理页地址,再发起 cp.async 加载。这一组合已在 vLLM v0.2.7+ 和 TensorRT-LLM 中实 现。 4) 主流 LLM 推理引擎的集成 FlashDecoding 自 2023 年 10 月开源以来,已被多个主流推理框架集成: vLLM:通过自定义 CUDA kernel 集成 FlashDecoding,在 vllm/model_executor/layers/attention 中提供 flash_decoding_attention 入口。启用条件为序列长度超过阈值(默认 4096) ,且 batch size × num_heads < SM 数目。vLLM v0.3.0 起默认开启。 TensorRT-LLM:在 context_fmha 和 generation_fmha 两个子模块中分别集成 FlashDecoding。预填充阶段使用 FlashAttention v2 fused kernel,解码阶段自动切换至 FlashDecoding 路径。借助 TensorRT 的图优化,解码延迟降低 30-40%(官方 Benchmark:Llama-2-70B,A100,128K context)。 HuggingFace TGI (Text Generation Inference):通过后端(backend)抽象层支持,在 FlashAttention backend 中检 测序列长度条件后路由至 FlashDecoding 实现。TGI v1.4+ 支持。 SGLang:在 RadixAttention 框架中集成,FlashDecoding 与 SGLang 的 prefix caching 和 continuous batching 联合优 化,适合高吞吐在线服务场景。 限制条件:FlashDecoding 的加速效果仅在序列长度显著大于 batch size × num_heads 时显现。对于高并发小序列场 景(如 chatbot 短对话),标准 FlashAttention 解码已足够高效,FlashDecoding 的分块额外开销反而可能略增延迟。此 外,归约步骤需要全局同步(global barrier),在 SM 数目更多的 H100/H800 上同步开销更显著,需通过 Warp-level reduction 等技巧降低。

18.8.3 相关算法与竞品

FlashAttention 并非孤立的注意力加速方案。围绕 Transformer 推理和训练效率,学界和工业界并行探索了多条技术路 线。本节梳理与 FlashAttention 紧密关联的四种代表性方案及其互补关系。

  1. PagedAttention 的内存管理创新 PagedAttention 是 vLLM(UC Berkeley,SOSP 2023)的核心贡献,将操作系统的虚拟内存分页思想引入 KV 缓存管 理。其创新点不在于注意力算子的计算优化,而在于解耦了逻辑 KV 序列与物理 GPU 内存布局。 传统推理引擎为每个请求预分配连续的 KV 缓存内存块,长度等于最大序列长度,导致两方面浪费:内部碎片(实际生成 长度远短于最大预留)和外部碎片(不同请求间的内存间隙)。PagedAttention 将 KV 缓存划分为固定大小的物理页 (page,典型大小 16 或 32 token),按需动态分配。请求的 KV 序列被映射为一张页表,类似 CPU 的虚拟地址转换。 Request Scheduler Block Manager GPU Memory Attention Kernel Allocate KV cache for request R1 Map logical block 0 -> physical page 3 Map logical block 1 -> physical page 7 Run attention for R1, token t Lookup page table for R1 Return physical addresses [3, 7] Load KV from pages 3, 7 via cp.async Compute attention, return O Request Scheduler Block Manager GPU Memory Attention Kernel 图18-32 PagedAttention 的分页内存管理流程 PagedAttention 的内存利用率可达 96% 以上(相比传统方案的 20-40%),使单 GPU 可同时服务的请求数增加 2-4 倍。 PagedAttention 与 FlashAttention 正交:前者管理 KV 缓存的物理布局,后者优化注意力 kernel 的 IO 效率,vLLM 内部 将两者组合使用。
  2. Ring Attention 的分布式序列并行 Ring Attention(Liu et al., 2023)将注意力计算沿序列长度维度分布到多 GPU 上,突破单 GPU 显存对序列长度的硬限 制。 设 P 个 GPU 以环形拓扑连接,序列长度 n 均匀分配到各 GPU(每 GPU 持有 n/P 个 token 的 Q、K、V 子块)。每轮计算 中,各 GPU 使用本地 K、V 块计算局部注意力,然后将 K、V 块发送至下一个 GPU(沿环转发),并接收上一个 GPU 的 K、V 块。经过 P 轮后,每个 GPU 获得全局 softmax 统计量,最终输出聚合的注意力结果。 Ring Attention 的通信模式天然适合 FlashAttention 的分块计算:每个 GPU 上的本地计算正好是 FlashAttention 的一个 tile 步骤,仅需将需要跨 GPU 加载的 K、V 块替换为远程传输。通信量每轮为 O(n ⋅ d/P ),与序列长度呈次线性关系。 在 8 GPU (A100-SXM 80GB) 配置下处理 1M token 序列时,各 GPU 内存占用仅 1/8,通信开销占端到端时间约 12-18%。 Ring Attention 已在 LWM (Large World Model) 和 StripedHyena 等项目中用于百万 token 级训练。它与 FlashAttention 的关系是互补而非替代:Ring Attention 解决跨 GPU 的序列分布问题,FlashAttention 解决单 GPU 内的 IO 效率问题。
  3. FlashInfer 的算子库生态 FlashInfer(华盛顿大学 & NVIDIA,2024)是一个专为 LLM 推理和服务设计的 GPU kernel 库,提供一批高性能注意力 及采样算子的统一实现。与 FlashAttention 侧重训练兼顾推理不同,FlashInfer 完全面向推理场景优化: JIT 编译:FlashInfer 采用即时编译(Just-in-Time)策略,根据推理时的具体配置(head dimension、KV 长度、数据 类型、attention variant)动态生成最优 kernel,无需预编译所有组合。这在 KV 缓存格式和注意力变体众多的推理场景 中避免了组合爆炸。 统一 KV 缓存接口:提供对 PagedAttention 分页缓存、Ragged Batch(不等长序列)、Append Buffer 等内存布局的原 生支持,kernel 内部通过 dispatch 适配不同布局。 除注意力(prefill attention、decode attention、cascade attention 等)外,还包括 Top-K/Top-P 采样、融合的 RMSNorm + RoPE、GQA/MQA 扩展的注意力变体等。截至 2026 年,FlashInfer 在 SGLang、MLC-LLM 等推理框架中作 为默认注意力后端使用。
  4. xFormers memory_efficient_attention 对比 xFormers(Meta AI)的 memory_efficient_attention (MEA)是 FlashAttention 出现之前最具影响力的 IO-aware 注意力实现。两者在核心思想上高度类似,但 MEA 采用在线 softmax + CUDA cutlass 模板库的通用矩阵乘法引擎,而 FlashAttention 以手工高度优化的 CUDA kernel 实现。 关键差异: 维度 FlashAttention v2 xFormers MEA 底层实现 手工 CUDA kernel,warp-level 调优 CUTLASS 模板 + 在线 softmax FLOPS 利用率 可达 72% (A100, FP16) 约 45-55% (A100, FP16) 前向速度(A100, 1K ctx) 约225 TFLOPS 约140 TFLOPS 反向传播 原生支持,重计算优化 原生支持,基于 CUTLASS 跨架构兼容性 NVIDIA 专属(CUDA) CUDA + ROCm(AMD) 因果掩码优化 v2 专门优化方向 无特殊处理 API 集成 PyTorch SDPA backend xformers.ops.memory_efficient_attention 表18-34 FlashAttention v2 与 xFormers MEA 关键对比 随着 PyTorch 2.0 torch.nn.functional.scaled_dot_product_attention (SDPA)的推出,两者均成为 SDPA 的可 选 backend。实际选型中:NVIDIA GPU 训练优先使用 FlashAttention v2/v3;AMD GPU 或需要跨硬件平台时 xFormers 是唯一选项;若仅需推理且 batch size 小,FlashDecoding 替代两者更为合适。
  5. 各类方案的综合对比矩阵 将本章涉及的注意力加速方案按适用阶段、优化维度、加速上限进行综合对比: 方案 适用阶段 核心优化维度 理论加速上限 硬件要求 成熟度 FlashAttention v2 训练 + Prefill HBM→SRAM IO 2-7.6x vs naive Ampere+ 生产级 (CUDA) FlashAttention v3 训练 + Prefill TMA + WGMMA + FP8 1.5-2x vs v2 Hopper 专属 生产级稳定 Block-Sparse FA 训练 + Prefill 跳过无效 block 计算 2-16x (取决于稀疏度) Ampere+ 实验/定制 (CUDA) 化 FlashDecoding Decoding (推 KV 并行分块+归约 5-20x (长序列解码) Ampere+ 集成中 理) (CUDA) PagedAttention 推理全流程 KV 缓存内存管理 2-4x 吞吐 (by 架构无关 生产级 (vLLM) memory) Ring Attention 训练 (分布式) 序列维度分布式 线性可扩展 (GPU 数) 多 GPU NCCL 研究中 FlashInfer 推理全流程 JIT 编译 + 统一 kernel 1.2-2x vs 手工 kernel Ampere+ 快速迭代中 库 (CUDA) xFormers MEA 训练 + Prefill 通用 IO-aware 模板 1.5-3x vs naive CUDA / ROCm 生产级 表18-35 注意力加速方案综合对比矩阵 从技术栈的视角看,这些方案构成一个互补的生态:FlashAttention 定义了单 GPU 高效注意力计算的事实标准; FlashDecoding 弥补了其在推理解码阶段的短板;PagedAttention 解决了推理的内存瓶颈;Ring Attention 将优化边界 扩展到多 GPU 分布式中;FlashInfer 为推理场景提供统一的算子后端;Block-Sparse FA 和 xFormers 则在特定约束下提 供差异化选择。这些方案的演化还反过来影响硬件设计:GPU 架构的迭代正在吸纳注意力算子的特殊需求,形成软件与 硬件相互反馈的循环。

18.8.4 软硬件协同设计趋势

FlashAttention 不仅是算法创新,更是一个软件倒逼硬件设计的典型案例。本节从软硬件协同的视角,探讨注意力算子与 GPU 架构的耦合演化路径、专用硬件的探索方向、以及 FlashAttention 方法论向其他 ML 算子的迁移潜力。

  1. 注意力算子与 GPU 架构 FlashAttention 的三个版本与三代 NVIDIA GPU 架构之间存在清晰的对应关系,且呈现软件提出需求→硬件下一代表现需 求→软件再挖掘新硬件的循环: FA v3 kernel Demands: async tensor co Hardware Feedback Informs: dedicated attenti re dispatch on pipeline Blackwell B200 (2024) FA v3 exploits TMA + WGM MA + FP8

FA v1/v2 exploits SMEM + c Hopper H100 (2022) p.async Warp specialization + FP8 v3: 1.5-2x vs v2 Ampere A100 (2020) IO-aware tiling v1: 2-4x speedup 图18-33 FlashAttention 版本与 GPU 架构的代际耦合 Ampere (A100) 时代,FlashAttention v1/v2 充分利用了 A100 的 164KB SRAM 和基础异步拷贝能力( cp.async ),将 标准注意力的 HBM 访问量从 O(n d) 降至 O(n d /M )。此时的核心发现是:注意力在 IO 而非计算上受限于硬件。 2 2 2 Hopper (H100) 时代,v3 进一步利用 Hopper 独有的 TMA 硬件单元(独立于 CUDA Core 的地址生成和数据搬运)、 WGMMA 指令(允许 Tensor Core 与 CUDA Core 同时执行不同指令流)和 FP8 原生支持。v3 对硬件的反馈信号是: attention kernel 需要更丰富的异步原语和更深的流水线,而非更高的峰值 FLOPS。 Blackwell (B200) 趋势:NVIDIA 在 B200 架构中已经间接回应了 attention 的需求:更大的 SRAM(每个 SM 配备更大容 量的 L1/SMEM)、更强的 Tensor Core FP4 级支持、以及改进的异步执行管道。2025 年 7 月发布的 FlashAttention-4 深 度利用了这些新特性。 2) 专用注意力加速硬件的前景 除 GPU 架构的渐进改进外,学术界正在探索专为注意力计算设计的硬件加速器: 标准的脉动阵列(Systolic Array)设计目标为高密度通用矩阵乘法(GEMM),并未针对注意力特有的混合计算模式(矩 阵乘法 → softmax → 矩阵乘法)优化。卡内基梅隆大学的 ELSA 架构在脉动阵列中嵌入专用指数运算单元和行最大值归 约树,使 softmax 可直接在阵列内部完成,避免将中间矩阵 S 写回 SRAM。 ij 以 SRAM 为中心的近存计算方面,Google TPU 和特斯拉 Dojo 的设计趋势是增大近计算单元的 SRAM 容量以减少对 HBM 的依赖。FlashAttention 的将完整 softmax 重缩放维持在片上这一核心约束,直接驱动了对更大 SRAM 的需求。在 256KB+ SRAM/SM 的假设下,v2 的块大小可从 128 增至 256,外循环迭代次数减半。 基于 Xilinx Versal ACAP 的注意力加速器原型(Stanford, DAC 2024)将 QKV 投影、注意力计算和 MLP 融合在单个 AI Engine 阵列上,实现端到端的 Transformer layer 流水线。FlashAttention 的分块策略在 FPGA 上被用于管理有限的 UltraRAM 资源。 专用硬件的主要挑战在于:Transformer 架构自身仍在快速演进(替代注意力机制的探索如 Mamba、RWKV、RetNet 等 相继被提出,2026 年已收敛为混合架构),硬件锁定在特定算子上存在过时风险。 3) FlashAttention 方法论的可迁移性 FlashAttention 的核心方法论(IO-aware 分块 + 数值稳定的在线统计量 + 反向传播时的重计算)并非注意力算子的专属 技巧。该方法论已被迁移到多个其他 ML 算子: 标准 LayerNorm 需要两次全局归约(计算均值→归一化→计算方差→重新缩放),两次同步屏障。类似 FlashAttention 的 chunk-wise 在线统计量方案,将输入分块计算局部统计量,通过合并公式避免第二次全数据遍历。实际收益:归一化 延迟降低 30-50%。 FlashLinearAttention 将线性注意力(Linear Attention)的 kernel 替代方案与 FlashAttention 的 IO-aware 策略结合。 线性注意力将复杂度降至 O(n),但其 kernel map (feature map) 计算仍需大量片外访问。FlashLinearAttention 通过分 块策略将这些访问保持在片上。 Mamba 结构状态空间模型(SSM)的扫描(scan)操作与注意力存在类似的 IO 特性:序列长度为主维度、隐状态维度 为内层。FlashAttention 的 work partitioning 策略被用于将 Mamba 的并行扫描沿序列维度和 batch 维度分配到 warp 组,提升 SM 利用率。 MoE (Mixture of Experts) 的 token 路由天然产生细粒度稀疏性。FlashAttention 的块稀疏模式被用于将专家层的计算分 组为连续的 dense block,避免分散的内存访问模式。 方法论迁移的核心启示是:任何计算密度低、带宽受限的逐元素/逐序列算子,均可从 IO-aware 分块和在线统计量中受 益。 4) 未来研究方向与开放问题 FlashAttention 至今仍是活跃研究方向,以下开放问题值得关注: 当前 Block-Sparse FlashAttention 使用静态预定义的稀疏模式,无法根据输入内容自适应调整。结合轻量级路由器网络 (router network)在运行时动态选择每个 query 的 key 子集,是 FA v4 的潜在方向。挑战在于:路由器本身的开销必 须远低于跳过计算省下的成本,且动态选择破坏了静态 kernel 的编译期优化。 v3 的 FP8 方案将前向统一为 E4M3、反向为 E5M2。更精细的设计可对注意力矩阵的不同部分使用不同精度:softmax 后 的权重矩阵 P 在多数位置接近零,可用 INT4 存储;高权重区域保持 FP16/FP8。难点在于精度切换需硬件支持的动态数 据类型转换。 ij Mamba-2 论文揭示了状态空间模型与注意力机制在结构化半可分矩阵(Structured Semiseparable Matrices)下的数学 等价性。这一视角暗示 FlashAttention 的分块计算框架可能直接适用于 SSM 扫描,甚至为两者提供统一的数值稳定实现 框架。 FlashAttention 的手工 CUDA kernel 开发成本高昂(v3 核心实现超过 5000 行 CUDA)。Triton、MLIR 等编译器基础设施 的成熟,使自动生成 IO-aware kernel 成为可能。Triton 版本的 FlashAttention 实现已接近手工 kernel 的 85-90% 性 能,且可根据新的硬件特性(如 Blackwell Tensor Core 新指令)自动适配。 当前 FlashAttention 主要优化文本 Transformer 的 1D 序列注意力。视觉 Transformer(ViT、DiT)和视频理解模型使 用 2D/3D 空间注意力,其稀疏模式和内存访问模式与 1D 序列有本质差异。针对 2D 网格局部注意力和轴向注意力的专项 优化是开放方向。 移动端和边缘设备的 GPU/NPU 通常不具备 H100 级别的硬件特性。针对 Mali GPU、Apple Neural Engine 等终端加速器 的 FlashAttention 变体,需在 tile 大小、数值精度和并发策略上做降级适配,精度-速度的帕累托前沿仍有大量优化空 间。 FlashAttention 系列的核心遗产不仅是算法本身,更在于它证明了:在专用硬件上重新审视经典算子的 IO 模式,可获得 超越单纯 FLOPS 堆叠的性能收益。这一方法论将持续影响未来 ML 系统的软硬件协同设计。 附录A 术语表