FlashAttention 加速算子

FlashAttention

FlashAttention 是 Tri Dao 等人提出的 I/O 感知型 Attention 算子实现:通过分块计算(Tiling)+ SRAM 复用 + 反向重计算,在不改变输出数值的前提下把 Attention 速度提升 2–4 倍,显存占用减少数十倍,已被 vLLM 和 PyTorch 2.x 内置。

详细解释

FlashAttention(FA,如今主流是 v2,v3 也已发布) 是 2022 年之后整个大模型训练和推理的地基性工程突破,没有它跑 70B 模型要用的显卡数会翻几倍。

它解决的核心痛点是:Transformer 里 Attention 的计算本身对 GPU 算力来说其实不重,重的是数据搬运(显存带宽瓶颈)。标准 Attention 流程里,N 乘 N 大小的 大小的注意力矩阵(N 是上下文长度)要在显存 HBM 和计算单元 SRAM 之间来回读写很多次——长上下文时 N=64K,N 平方 量级的 是 40 亿,显存直接爆了。

FlashAttention 的三个核心技巧(论文原文):

  1. 分块计算(Tiling):把 Q、K、V 按序列维度切成小块(一般 128 或 256 个 Token 一块),每次只算一小块的注意力,这样 N 平方 量级的 的临时矩阵不用全部落显存。
  2. 在线 Softmax(Online Normalization):Softmax 需要先知道整行的 max 和 sum——FA 用”分块的缩放系数 + 归并校正”技巧,在不需要看到全部 K 的情况下逐块累加最终结果。
  3. 重计算(Recomputation / Stitch Backward):反向传播的时候不保存前向过程中那些 N 平方 量级的 的中间 Softmax 分数,而是重新算一遍(用 20% 的算力开销换 10× 显存下降)。

参考:FlashAttention-2 技术博客PyTorch scaled_dot_product_attention(PyTorch 2.0 开始内置 FA)。

各版本与硬件对应表(2025 年生产部署)

算子适用硬件支持数据类型相对标准 Attention 的加速
FlashAttention-1(已基本淘汰)Ampere (A10/A100) + Ada (L40/L40S/RTX 4090)FP16 / BF162× 左右
FlashAttention-2(当前主流)Ampere + Ada + Hopper(H100)FP16 / BF16 / FP82× ~ 4×(H100 上 FP8 最高 6×)
FlashAttention-3Hopper H100/H200 专用(优化 TMA 异步拷贝)FP8 / FP16在 v2 基础上再 +30–50%
FlashDecoding(FA 的推理变种)长上下文推理场景同上N=64K 时比 FA v2 再 +1.5×(batch 维度并行)

一个关键结论:如果你在 H100 上还没开 FA-3 + FP8,你的显卡利用率可能只有理论性能的 1/3 至 1/2。主流推理框架(vLLM、SGLang、TensorRT-LLM)默认都会选当前硬件最优的算子。唯元智创 Weimeta 的大模型推理网关会自动根据目标硬件路由到最优算子配置,无需业务方手动调优。

对终端 API 消费者意味着什么

作为调用 API 的业务方,你不需要手动「打开 FlashAttention」——这是供应方和框架的事。但你会从它身上间接受益:

  1. 长上下文更便宜:N 从 8K 拉到 128K 时,标准 Attention 的显存是 O(N 平方) 复杂度,会贵几十倍;有了 FA 变成接近线性 O(N) 线性复杂度,所以 Context Window 128K 的模型 能以合理价格对外提供。
  2. TTFT(首字延迟) 更快:Prefill 阶段 Q 乘 K 转置 乘积是大头,FA 把 Prefill 时间砍一半以上。
  3. 相同硬件能接更大并发:显存省了,意味着同一张 A100 上驻留的 KV Cache / 请求数更多——间接降低厂商成本 → 传导到你端的 每百万 Token 价格下降。
  4. GQA(分组查询注意力) + FA 组合拳:两者都是 2023 年之后”同等效果下成本砍一半”的大杀器,新模型基本都同时启用(Llama 3.x、Qwen3、DeepSeek-V3 等)。

常见问题

FlashAttention 会改变输出数值吗?
在数值等价意义上不改变(数学上一模一样)。但在 fp16/bf16 浮点下,分块的求和顺序不同会有 1e-6 至 1e-5 量级的浮点误差——这是浮点数加法交换律”近似成立”的本质,任何 GPU 算子都会有。对实际业务输出完全无感知,除非你做字节级可复现(此时要固定算子版本+硬件)。
消费级 RTX 4090 跑自部署模型,FA 能跑吗?
可以。FlashAttention-2 已支持所有 Ada Lovelace 架构(RTX 40xx、L4、L40、L40S)。vLLM 和 HuggingFace Transformers 默认就是开的,不用你手改代码。前提是:你的 CUDA 版本 >= 11.8,PyTorch >= 2.0。
还有类似 FA 的推理加速算子吗?
有,常见同方向的三个:(a) FlashDecoding:FA 针对推理场景(prefill 长 + decode batch 小)做了重拆分,长上下文 N=64K+ 时比 FA v2 还快 1.5×;(b) PagedAttention(vLLM 作者提出):借鉴 OS 虚拟内存的分页思想,把 KV Cache 分块存在非连续显存页里,解决显存碎片问题,和 FA 是互补关系;(c) xFormers Memory-Efficient Attention:Meta 实现的同类算子,效果接近 FA,适用更广泛硬件(AMD/TPU 等)。