详细解释
FlashAttention(FA,如今主流是 v2,v3 也已发布) 是 2022 年之后整个大模型训练和推理的地基性工程突破,没有它跑 70B 模型要用的显卡数会翻几倍。
它解决的核心痛点是:Transformer 里 Attention 的计算本身对 GPU 算力来说其实不重,重的是数据搬运(显存带宽瓶颈)。标准 Attention 流程里,N 乘 N 大小的 大小的注意力矩阵(N 是上下文长度)要在显存 HBM 和计算单元 SRAM 之间来回读写很多次——长上下文时 N=64K,N 平方 量级的 是 40 亿,显存直接爆了。
FlashAttention 的三个核心技巧(论文原文):
- 分块计算(Tiling):把 Q、K、V 按序列维度切成小块(一般 128 或 256 个 Token 一块),每次只算一小块的注意力,这样 N 平方 量级的 的临时矩阵不用全部落显存。
- 在线 Softmax(Online Normalization):Softmax 需要先知道整行的 max 和 sum——FA 用”分块的缩放系数 + 归并校正”技巧,在不需要看到全部 K 的情况下逐块累加最终结果。
- 重计算(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 / BF16 | 2× 左右 |
| FlashAttention-2(当前主流) | Ampere + Ada + Hopper(H100) | FP16 / BF16 / FP8 | 2× ~ 4×(H100 上 FP8 最高 6×) |
| FlashAttention-3 | Hopper 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」——这是供应方和框架的事。但你会从它身上间接受益:
- 长上下文更便宜:N 从 8K 拉到 128K 时,标准 Attention 的显存是 O(N 平方) 复杂度,会贵几十倍;有了 FA 变成接近线性 O(N) 线性复杂度,所以 Context Window 128K 的模型 能以合理价格对外提供。
- TTFT(首字延迟) 更快:Prefill 阶段 Q 乘 K 转置 乘积是大头,FA 把 Prefill 时间砍一半以上。
- 相同硬件能接更大并发:显存省了,意味着同一张 A100 上驻留的 KV Cache / 请求数更多——间接降低厂商成本 → 传导到你端的 每百万 Token 价格下降。
- GQA(分组查询注意力) + FA 组合拳:两者都是 2023 年之后”同等效果下成本砍一半”的大杀器,新模型基本都同时启用(Llama 3.x、Qwen3、DeepSeek-V3 等)。