FIA / FlashAttention
FIA = FlashInfer‑Attention,是 FlashInfer 库的注意力算子;FlashAttention(FA) 是Tri Dao提出的IO‑Aware精确注意力算法库。
- FlashAttention:偏向训练,兼顾Prefill推理
- FIA(FlashInfer‑Attention):专门面向LLM推理(Prefill+Decode),内部复用FA2/FA3分块+Online‑Softmax核心,但针对KV‑Cache、分页内存、变长batch做整套推理侧改造。
一、FlashAttention(FA‑1/2/3)核心原理
标准Attention:
Attention(Q,K,V)=softmax(QKTd)V\\text{Attention}(Q,K,V)=\\text{softmax}\\left(\\frac{QK^\\mathrm{T}}{\\sqrt d}\\right)VAttention(Q,K,V)=softmax(dQKT)V
直接算出完整 N×NN\\times NN×N score矩阵存入HBM,显存 O(N2)O(N^2)O(N2),大量HBM读写,内存带宽瓶颈。
FlashAttention三件套:
| FA‑1 | 基础IO‑aware分块+online softmax | Ampere及以上 | 训练、prefill |
| FA‑2 | 更好线程块划分、降低非GEMM开销 | A100最优,支持Ada/Hopper | 训练主力,广泛集成 |
| FA‑3 | Hopper专属,WGMMA、TMA异步流水线、FP8支持 | 仅H100/H200(Hopper) | 长序列训练、prefill,FP8低精度 |
✅ FlashAttention是精确注意力,不做近似,数学等价标准SDPA;原生对Decode(单token生成)支持弱,KV‑cache需要上层自己处理。
二、FIA(FlashInfer‑Attention)
FlashInfer是伯克利+字节推出的推理专用kernel库,FIA就是它的注意力算子集合,底层复用FlashAttention分块算法,但解决线上推理真实痛点:分页KV‑Cache、变长batch、大量并发Decode请求、负载不均衡、MLA/GQA、稀疏KV、CUDAGraph适配。
FIA关键能力(FA原生缺少)
- Prefill:直接调用FA‑2/FA‑3模板做Prompt编码;
- Flash‑Decoding:专门Decode算子,QQQ仅1行(新生成token),分块遍历全部KV‑Cache,online‑softmax做归约,不需要重建完整注意力矩阵。
- Plan阶段:预处理变长序列,做负载均衡、block分配;
- Run阶段:批量执行;解决不同请求长度不一导致GPU线程空闲。
三、FlashAttention vs FIA(FlashInfer‑Attention)对比
| 定位 | 训练优先,兼顾prefill推理 | 推理服务专用:prefill + decode完整链路 |
| Decode支持 | 很差,需要上层自己封装KV‑Cache | 原生Flash‑Decoding高性能decode kernel |
| KV‑Cache内存布局 | 只支持连续tensor | 支持分页(非连续)KV Cache(vLLM/SGLang) |
| Batch形态 | 偏向固定shape训练batch | 高度优化变长、参差不齐推理batch,负载均衡调度 |
| MLA/GQA | 基础GQA支持 | 深度优化MLA、Head‑Query融合、前缀共享 |
| FP8 | FA3仅Hopper | Hopper/Ampere都有FP8推理kernel |
| 反向传播 | 完整、高性能(训练必需) | 几乎不用反向,只做前向推理 |
| 典型使用场景 | LLM预训练、微调、prefill | 线上LLM服务、高并发推理、长上下文生成 |
简单理解:
- 做训练:优先FlashAttention‑2/3
- 做线上推理部署:优先FlashInfer(FIA),内部会复用FA的prefill计算逻辑,但decode、分页KV、调度层全部增强。
四、代码极简示例
FlashAttention
import flash_attn
# 训练/prefill,支持causal掩码
out = flash_attn.flash_attn_qkvpacked_func(qkv, dropout_p=0.0, causal=True)
FlashInfer(FIA)推理(Decode阶段,分页KV Cache)
import flashinfer
# plan阶段做调度
plan = flashinfer.BatchDecodePlan()
plan.plan(...)
# run:批量解码,直接读分页KV‑Cache
out = flashinfer.batch_decode_with_paged_kv_cache(q, kv_cache, plan)
五、面试核心考点总结
如果你需要,我可以进一步推导Online‑Softmax数学公式,或者Flash‑Decoding算法详解。
网硕互联帮助中心




评论前必须登录!
注册