跳转至

二十四:Flash Attention

来源:http://mp.weixin.qq.com/s?__biz=MzYyNTk3Njg1NA==&mid=2247484249&idx=1&sn=809a652ddfae59386eb2d973a5effed7&chksm=f01eb020c76939366daf782c72ecc736cde7346bd1311d2d58e9c3f08341963534fea41df9e6#rd

1. 学习范围

本日主题是 Flash Attention,重点是 Flash Attention 的加速原理。学习目标不是只会说“Flash Attention 更快”,而是理解它为什么更快、它没有改变什么、它怎样保持 exact attention、它如何利用 GPU 存储层次结构,以及在实际大模型训练/推理中如何使用和排查问题。 包含笔者对Flash Attention的数学原理的手推公式。 本日覆盖以下知识: - 标准 scaled dot-product attention 的公式、张量形状和显存瓶颈。

  • GPU HBM、SRAM/shared memory/register 的存储层次和 IO 瓶颈。

  • FlashAttention 的核心思想:tiling、kernel fusion、online softmax、避免物化完整注意力矩阵。

  • FlashAttention forward 与 backward 的关键流程。

  • causal mask、padding mask、dropout、mixed precision 的处理要点。

  • FlashAttention 1、FlashAttention 2、FlashAttention 3 的演进方向。

  • PyTorch SDPA、flash-attn 包、Hugging Face 模型中的使用方式。

  • 常见性能收益、限制条件、数值误差、OOM 和调试路径。

先把手推稿放在上面,后续各处可以对应手推的数学证明来进行理解: 首先是Forward部分:

img

img

然后是Backward部分:

img

img

2. 标准 Scaled Dot-Product Attention

Transformer 中常见的 scaled dot-product attention (标准缩放点积注意力)为:

S = Q K^T / sqrt(d)
P = softmax(S + mask)
O = P V

img

其中: - Q 是 query,形状常写为 [B, H, Nq, D]

  • K 是 key,形状常写为 [B, H, Nk, D]

  • V 是 value,形状常写为 [B, H, Nk, Dv]

  • S 是 attention score,形状为 [B, H, Nq, Nk]

  • P 是 attention probability,形状同 S

  • O 是输出,形状为 [B, H, Nq, Dv]

img

在 self-attention 中通常 Nq = Nk = ND = Dv = head_dim。在 decoder causal attention 中,mask 会禁止当前位置看到未来 token。 标准 attention 的计算复杂度仍然是:

O(B * H * N^2 * D)
这来自 QK^TPV 两次矩阵乘。Flash Attention 并没有把数学复杂度从二次变成线性,它优化的是显存访问和中间张量存储。

3. 标准 Attention 的显存瓶颈

标准实现通常会显式生成并保存 SP

Q, K, V -> S = QK^T -> P = softmax(S) -> O = PV
当序列长度很长时,SP 的大小是 O(N^2)。例如:

B = 1
H = 32
N = 8192
dtype = fp16, each element = 2 bytes

attention matrix size = B * H * N * N * 2 bytes
                      = 1 * 32 * 8192 * 8192 * 2
                      ≈ 4 GB
这还只是一个 attention matrix。训练时还要保存激活用于 backward,多个层叠加后显存压力非常大。 标准 attention 的另一个瓶颈是 HBM 访问(访存瓶颈)。GPU 的矩阵乘本身很快,但把巨大的 SP 写入 HBM、再从 HBM 读回,会消耗大量内存带宽。长序列训练中,attention 往往受 memory bandwidth 限制,而不只是受 FLOPs 限制。 这也是 FlashAttention 针对性优化的两大方向:切块计算、不持久存储完整注意力矩阵,减少 HBM 读写与显存占用。

4. GPU 存储层次与 IO-Aware 思想

GPU 存储层次可以粗略理解为:

register / SRAM / shared memory: 容量小,速度快
HBM / global memory: 容量大,速度慢
CPU memory / disk: 更大,更慢
FlashAttention 论文强调 IO-aware,即不仅统计 FLOPs,还要统计不同存储层之间的数据搬运量。对长序列 attention 来说,瓶颈经常不是算不了 QK^T,而是中间矩阵太大,反复读写 HBM。 优化目标是: - 尽量把小块 Q/K/V 放入高速 SRAM/shared memory。

  • 在片上完成 score、softmax、加权求和。

  • 不把完整 SP 写到 HBM。

  • 用更少 HBM IO 换取更高实际吞吐。

这就是 FlashAttention 和普通 fused attention 的关键差异:它是围绕 attention 计算图和 GPU 存储层次一起设计的。

5. FlashAttention 的定位

FlashAttention 是 exact attention。它计算的仍然是:

softmax(QK^T / sqrt(d)) V
它不是稀疏 attention,不是线性 attention,也不是低秩近似。它的结果与标准 attention 在数学上等价,实际实现中只会因为浮点计算顺序、精度和 kernel 细节产生可接受的数值差异。 FlashAttention 的核心收益: - 显著减少 attention 中间矩阵的 HBM 读写。

  • 不显式物化完整 [N, N] attention probability。

  • 前向和反向中使用 block-wise 计算。

  • 训练时降低 attention 激活显存。

  • 在长序列、大 batch、多头场景下提升速度和可训练长度。

FlashAttention 的核心限制: - 计算复杂度仍是 O(N^2D)

  • 需要 CUDA GPU 和特定 dtype/shape 支持。

  • 对任意复杂 attention mask 的支持可能不如普通 attention 灵活。

  • 对短序列或小模型不一定显著更快。

  • 具体收益依赖 GPU 架构、head_dim、batch、num_heads、causal、dropout、dtype 和框架后端。

6. Tiling 与 Block-wise Attention

FlashAttention 把 QKV 按 block 切分。典型计算方式是:

for each Q block:
    load Q_i into SRAM
    initialize running softmax stats
    for each K/V block:
        load K_j, V_j into SRAM
        compute S_ij = Q_i K_j^T
        update online softmax stats
        update partial output O_i
    write final O_i to HBM
关键点是:对一个 Q block,逐块扫描所有 K/V block,并在扫描过程中维护 softmax 所需的统计量和输出累积值。完整的 SP 不需要写回 HBM。 block size 的选择需要平衡: - shared memory 容量。

  • register 压力。

  • head_dim。

  • GPU occupancy。

  • matmul tile 的效率。

  • causal mask 下跳过无效块的能力。

这也是 FlashAttention 需要专门 CUDA kernel 或 Triton kernel 的原因:普通 PyTorch 算子组合很难表达这种片上融合计算。

7. Online Safe Softmax

softmax 的数值稳定形式为:

softmax(x_i) = exp(x_i - m) / sum_j exp(x_j - m)
m = max_j x_j

img

正是因为这个最大值偏移,导致必须同时拿到整行才可以计算最大值。 如果一次性拿到整行 score,可以直接计算最大值 m 和分母 l。但 FlashAttention 是分块扫描 score,因此需要 online softmax:每读入一个新 block,就更新当前行的最大值、归一化分母和输出累积。 对同一行 attention,假设旧统计量为:

m_old: 已扫描 score 的最大值
l_old: 已扫描 score 的 exp 归一化和
acc_old: 已扫描 value 的加权累积,未必已经除以 l
新 block 的 score 为 s_new,则:

m_new = max(m_old, max(s_new))
l_new = exp(m_old - m_new) * l_old
        + sum(exp(s_new - m_new))

acc_new = exp(m_old - m_new) * acc_old
          + exp(s_new - m_new) @ V_new

O = acc_new / l_new

img

这个公式保证了分块计算和一次性 softmax 等价。m 的更新用于数值稳定,l 的更新用于正确归一化,acc 的重缩放用于保证旧 block 和新 block 在同一个 softmax 基准下相加。

8. Forward 计算流程

FlashAttention forward 的逻辑可以简化为:

Input: Q, K, V
Output: O

for each block of Q:
    m = -inf
    l = 0
    acc = 0

    for each block of K, V:
        S = Q_block @ K_block.T * scale
        S = S + mask

        m_new = max(m, rowmax(S))
        P = exp(S - m_new)
        alpha = exp(m - m_new)

        acc = alpha * acc + P @ V_block
        l = alpha * l + rowsum(P)
        m = m_new

    O_block = acc / l
真实 kernel 会使用更复杂的 tile、warp 分工、向量化、shared memory 管理和数值优化,但主线就是 block-wise matmul + online softmax + block-wise output accumulation。 该流程的重点是:SP 只在片上以 block 形式短暂存在,不会以完整 [Nq, Nk] 矩阵写入 HBM。

9. Mask、Causal Attention 与 Dropout

FlashAttention 支持常见的 causal attention。causal mask 的含义是第 i 个 query 只能关注 j <= i 的 key:

S[i, j] = -inf, if j &gt; i
在 block-wise 计算中,causal mask 可以按 tile 处理: - 完全位于未来的 K/V block 可以跳过。

  • 与对角线相交的 block 需要在 tile 内应用 mask。

  • 完全合法的 block 可以正常计算。

padding mask 和变长序列通常通过 unpadding / cu_seqlens / varlen kernel 支持,把有效 token 压缩后计算,避免大量 padding 浪费。 - 借助 cu_seqlens(序列累积长度)标记真实有效 token 区间;

  • 先做去填充(unpadding),压缩有效序列后再执行注意力计算;

  • 计算完成后还原形状,全程避开 padding 区域,提升硬件利用率。

训练时的 dropout 通常作用在 attention probability 上。FlashAttention 可以在 fused kernel 中处理 dropout,但需要保存或可重建 dropout mask 的随机状态,保证 backward 与 forward 一致。

10. Backward 与重计算

标准 attention 训练中,backward 往往需要保存 PS,显存开销大。FlashAttention 的策略是保存更小的中间量,例如输出 O、每行 softmax 的 logsumexp 或相关统计量,然后在 backward 中按 block 重新计算局部 score 和 probability。 这是一种 compute-memory tradeoff: - 节省 HBM 中 O(N^2) attention matrix 存储。

  • backward 需要重算部分 QK^T 和 softmax。

  • 由于减少 HBM IO,整体仍可能更快。

反向传播需要计算:

dV = P^T dO
dP = dO V^T
dS = P * (dP - rowsum(dP * P))
dQ = dS K
dK = dS^T Q
FlashAttention 不会完整保存 P,而是在每个 block 内重构 P,并累积 dQ/dK/dV

11. 复杂度与显存收益

FlashAttention 的计算复杂度仍然是二次:

QK^T: O(N^2D)
PV: &nbsp; O(N^2D)
它降低的是 HBM IO 和 activation memory。标准 attention 需要显式保存 S/P,attention matrix 是 O(N^2)。FlashAttention 不保存完整 attention matrix,attention 部分的额外激活可以接近 O(N) 级别,主要保存输出和每行统计量。 面试中应避免两个误区: - FlashAttention 不是把 attention 算法复杂度从 O(N^2) 变成 O(N)

  • FlashAttention 不是近似算法,它是 exact attention 的 IO-aware 实现。

它的实际收益通常在长序列上更明显,因为 N^2 attention matrix 的 HBM 读写随着序列长度迅速增加。

12. 张量形状与实现约定

不同库对 Q/K/V 的布局约定不同。 PyTorch scaled_dot_product_attention 常见输入布局:

q: [B, Hq, L, D]
k: [B, H, &nbsp;S, D]
v: [B, H, &nbsp;S, Dv]
flash-attn 包中的部分 API 常见布局:

qkv: [B, N, 3, H, D]
q: &nbsp; [B, Nq, Hq, D]
k: &nbsp; [B, Nk, H, &nbsp;D]
v: &nbsp; [B, Nk, H, &nbsp;Dv]
实际使用时必须看具体 API 文档。布局不匹配会导致 silent wrong result、shape error 或性能退化。 常见注意点: - head_dim 通常需要满足 kernel 支持范围。

  • dtype 通常使用 fp16 或 bf16。

  • Q/K/V 需要在 CUDA device 上。

  • tensor stride 和 contiguous 状态可能影响性能。

  • MQA/GQA 中 HqHkv 可以不同,但需要 API 支持。

13. FlashAttention 1、2、3 的演进

FlashAttention 1 的核心贡献是 IO-aware exact attention:通过 tiling 和 online softmax 避免物化完整 attention matrix,降低 HBM IO 和 attention 激活显存。 FlashAttention 2 主要改进并行性和 work partitioning: - 减少非矩阵乘部分的额外 FLOPs。

  • 改进 block/warp 之间的工作划分。

  • 在 batch 和 head 数较小时,也能更好利用 GPU。

  • 进一步提高训练和推理吞吐。

FlashAttention 3 面向 NVIDIA Hopper 架构进一步优化: - 利用 Hopper 的异步数据搬运和新矩阵乘能力。

  • 更好地重叠数据加载、matmul 和 softmax。

  • 支持更高性能的 FP16/BF16,并探索 FP8 attention。

  • 重点解决 H100 等新硬件上的利用率问题。

版本演进的主线不是改变 attention 数学定义,而是不断贴近 GPU 硬件特性,减少 IO、提高并行度、提升 kernel 利用率。

14. PyTorch SDPA 与 flash-attn 使用方式

PyTorch 提供 torch.nn.functional.scaled_dot_product_attention,会根据设备、dtype、shape、mask、dropout、is_causal 等条件选择可用后端,包括 math、memory-efficient attention、FlashAttention 等。 典型写法:

import torch
import torch.nn.functional as F

q = torch.randn(2, 16, 1024, 64, device="cuda", dtype=torch.float16)
k = torch.randn(2, 16, 1024, 64, device="cuda", dtype=torch.float16)
v = torch.randn(2, 16, 1024, 64, device="cuda", dtype=torch.float16)

out = F.scaled_dot_product_attention(
&nbsp; &nbsp; q, k, v,
&nbsp; &nbsp; attn_mask=None,
&nbsp; &nbsp; dropout_p=0.0,
&nbsp; &nbsp; is_causal=True,
)
flash-attn 包提供更直接的 CUDA kernel 接口,具体函数名和参数随版本变化,需要以官方仓库文档为准。常见模型框架如 Hugging Face Transformers 也支持通过配置启用 SDPA 或 FlashAttention 后端。 实际项目中应优先使用框架稳定接口,只有在需要极致性能、变长序列或特殊布局时再直接调用底层 flash-attn API。

15. MHA、MQA、GQA 与 KV Cache

Multi-Head Attention 中,每个 query head 通常对应自己的 key/value head。MQA 和 GQA 减少 key/value heads 数量,从而降低 KV cache 和 memory bandwidth。 在自回归推理中,decode 阶段每次生成一个或少量 token,query length 很短,key/value length 随上下文增长。此时瓶颈常是读 KV cache 的带宽,而不是构造完整 N x N 矩阵。FlashAttention 对 prefill 长上下文阶段收益通常更明显;decode 阶段还需要专门的 paged attention、KV cache 管理或 decode kernel 优化。 因此面试中要区分: - prefill:一次处理完整 prompt,attention 矩阵大,FlashAttention 收益明显。

  • decode:逐 token 生成,Q 很短,KV cache 访问和调度更关键。

16. 数值精度与稳定性

FlashAttention 使用 online safe softmax 保证数值稳定。它通常在 fp16/bf16 输入下,对部分累积量使用更高精度或稳定公式,避免 exp 溢出。 由于计算顺序与标准 attention 不同,输出不一定 bitwise identical,但应在合理误差范围内接近。比较结果时应使用 torch.allclose 之类的容差比较,而不是要求完全相等。 常见数值问题来源: - 学习率过大导致 logits 极端。

  • attention mask 错误导致整行全是 -inf

  • dtype 不支持或混合精度配置错误。

  • dropout 随机状态不一致。

  • 自定义 mask 与 kernel 支持范围不匹配。

17. 性能评测方法

评测 FlashAttention 不能只看单次 wall time。推荐做法: - 固定 GPU 型号、CUDA、PyTorch、flash-attn 版本。

  • 使用相同的 B/H/N/D/dtype/causal/dropout

  • 做 warmup,避免首次 kernel 编译或缓存影响。

  • 使用 torch.cuda.synchronize() 包住计时。

  • 分别记录 forward、backward、端到端训练 step。

  • 记录 torch.cuda.max_memory_allocated()

  • 使用 profiler 检查是否真的走了 FlashAttention 后端。

示例计时结构:

torch.cuda.synchronize()
start = time.time()
for _ in range(iters):
&nbsp; &nbsp; out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
torch.cuda.synchronize()
elapsed = time.time() - start
短序列、CPU tensor、fp32、复杂 mask、head_dim 不支持、dropout 条件不匹配等情况,都可能导致没有使用 FlashAttention 后端。

18. 常见限制与故障排查

FlashAttention 常见使用限制包括: - CUDA/GPU 架构要求。

  • dtype 通常要求 fp16/bf16,部分后端支持 fp8。

  • head_dim 支持范围有限。

  • mask 类型支持有限,任意 dense mask 可能回退到 math kernel。

  • dropout、causal、GQA、varlen 支持随版本变化。

  • Windows 环境安装 flash-attn 可能比 Linux 更麻烦。

常见排查路径: - 输出 shape 错误:检查 Q/K/V layout 和 head 维度位置。

  • 速度没提升:确认是否走 FlashAttention 后端,检查序列长度是否足够长。

  • 显存没下降:确认没有在外层保存 attention weights,检查模型是否仍返回 attentions。

  • 结果异常:检查 mask、scale、causal、dropout、chat padding 和 dtype。

  • OOM:减小 batch、序列长度、head_dim,启用 gradient checkpointing,检查是否回退到普通 attention。

19. 面试表达要点

FlashAttention 的高分表达可以组织成四句话: - 标准 attention 会物化 N x N score/probability,长序列下 HBM IO 和激活显存成为瓶颈。

  • FlashAttention 是 exact attention,不是近似;它通过 tiling 把 Q/K/V 分块搬到片上 SRAM,并用 online softmax 分块计算输出。

  • 它不保存完整 attention matrix,backward 通过保存少量统计量并重算局部 score 来节省显存。

  • 它降低的是 HBM IO 和 activation memory,数学计算复杂度仍是 O(N^2D),收益依赖序列长度、dtype、GPU 架构和后端支持。

常见扣分点: - 说 FlashAttention 把 attention 复杂度变成线性。

  • 说 FlashAttention 是稀疏 attention 或近似 attention。

  • 只知道“更快”,说不出 HBM/SRAM 和 online softmax。

  • 忽略 backward 重计算。

  • 不会解释 causal mask 和 padding/varlen 的处理。

  • 不知道 PyTorch SDPA 可能自动选择后端,也可能因为条件不满足而回退。

20. 参考资料

  • FlashAttention paper: https://arxiv.org/abs/2205.14135

  • FlashAttention-2 paper: https://arxiv.org/abs/2307.08691

  • FlashAttention-3 paper: https://arxiv.org/abs/2407.08608

  • FlashAttention GitHub: https://github.com/Dao-AILab/flash-attention

  • PyTorch scaled_dot_product_attention: https://docs.pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html

  • PyTorch SDPA tutorial: https://docs.pytorch.org/tutorials/intermediate/scaled_dot_product_attention_tutorial.html

  • 推荐资料:Flashattention 1/2/3 讲解: https://blog.csdn.net/v_JULY_v/article/details/133619540

  • 推荐资料:FlashAttention 的加速原理: https://blog.csdn.net/asd8705/article/details/140136587

            预览时标签不可点
    

    <div class="