跳转到内容

输入关键词开始搜索

    FlashAttention

    概念更新 2026-07-26置信度 medium#概念#基础设施#注意力#进阶#长青

    一种保持标准 attention 精确结果、却通过分块和在线 softmax 大幅减少 GPU 显存 I/O 的实现算法。

    FlashAttention 不改变 Transformer 的注意力定义:

    O=softmax(QKd)VO=\operatorname{softmax}\left(\frac{QK^\top}{\sqrt{d}}\right)V

    它改变的是 GPU 上的执行顺序。朴素实现会把大小为 T×TT\times T 的 score 矩阵和 softmax 概率矩阵反复写入、读出 HBM 显存;FlashAttention 将 Q,K,VQ,K,V 分为适合片上 SRAM 的小块,边计算边累积最终输出,避免物化完整注意力矩阵。

    因此它是 IO-aware(关注内存读写) 算法,而非稀疏注意力、近似注意力或新模型架构。

    1. 训练显存更省:不保存完整 T×TT\times T 注意力中间结果,长序列时收益尤其显著。
    2. 注意力更快:减少 HBM 往返;在许多 GPU 工作负载中,数据搬运而非 FLOPS 才是瓶颈。
    3. 成为默认后端:现代 PyTorch 的 scaled_dot_product_attention 在兼容的 CUDA 环境会选择融合 attention 内核;许多 LLM 训练/推理框架也直接集成 FlashAttention 系列。
    4. 与 KV 优化互补:GQA/MQA 减少 KV Cache 的容量;FlashAttention 降低注意力计算本身的读写开销,两者通常叠加使用。

    对单个头,先算 S=QKS=QK^\top,再算 P=softmax(S)P=\operatorname{softmax}(S),最后算 PVPVSSPP 都是 T×TT\times T,长度翻倍时其体积变为四倍。

    即使乘法本身很快,写入 HBM、读回作 softmax、再读回与 VV 相乘仍会消耗大量带宽。

    FlashAttention 固定一小块 QQ 在片上 SRAM,顺序流过多个 K,VK,V 块。对每行维护:

    • mm:截至当前块的最大 score,保证指数计算稳定;
    • \ell:重标定后的指数和;
    • OO:已经归一化的输出累积值。

    读取一个新块 Sb=QbKb/dS_b=Q_bK_b^\top/\sqrt d 时:

    m=max(m,maxSb)=emm+eSbmO=emmO+eSbmVb\begin{aligned} m' &= \max(m,\max S_b)\ \ell' &= e^{m-m'}\ell+\sum e^{S_b-m'}\ O' &= \frac{e^{m-m'}\ell O+e^{S_b-m'}V_b}{\ell'} \end{aligned}

    所有块处理完毕,OO' 就是标准 softmax attention 的输出(除浮点舍入误差外相同)。因果掩码、padding mask 和多头计算可在对应 tile 内处理。

    通常不需要手写 tile 或 online softmax;把 Q/K/V 交给框架即可:

    out = torch.nn.functional.scaled_dot_product_attention(
    q, k, v,
    is_causal=True,
    enable_gqa=True, # 仅在使用 GQA 时需要
    )

    是否实际选中 FlashAttention 取决于 CUDA GPU、dtype、head dimension、张量布局和 PyTorch 版本。应以 profiler 或框架日志确认,而不是仅凭这一行 API 推断。

    • FlashAttention(2022):Tri Dao 等提出 IO-aware 精确 attention;通过 tiling 和 online softmax 避免完整注意力矩阵物化。
    • FlashAttention-2(2023):改进 GPU 并行切分与工作分配,提升训练和长序列吞吐。
    • 后续内核:针对 Hopper 等硬件以及 decode 场景继续优化;其具体支持范围随软件版本、GPU 架构变化。
    • “FlashAttention 改变模型输出。” 不改变数学公式;它是精确重排,只有常规浮点舍入差异。
    • “它把 O(T2)O(T^2) attention 变成线性。” 不会。计算的理论二次复杂度仍在;降低的是中间存储和实际 I/O。
    • “它解决 KV Cache 显存。” 不直接解决。GQA(分组查询注意力)、MQA、MLA 等决定 KV Cache 的体积。
    • “生成一个 token 时收益和训练一样大。” 不一定。单 token decode 往往受读取历史 KV Cache 的带宽限制,需使用面向 decode 的内核/调度优化。
    概念 主要优化对象 与 FlashAttention 的关系
    GQA(分组查询注意力) / MQA KV 头数、KV Cache 容量 可叠加;GQA 使需读取/缓存的 KV 更少
    KV Cache 避免重复计算历史 K/V FlashAttention 不替代缓存,而是更高效地使用它
    PagedAttention KV Cache 的分页管理 解决缓存碎片/调度;也可配合高效 attention 内核
    稀疏注意力 限制可见 token 对 改变计算图;FlashAttention 仍计算所有允许的 token 对