H20上128K预填充,V2比FA2快47.26倍

FlashPrefill V2: Block-Sparse Prefill Attention for Long-Context LLM Serving

Qihang Fan, Huaibo Huang, Zhiying Wu, Bingning Wang, Ran He

cs.CL

2026-08-20

FlashPrefill V2用均值修正和Hopper内核把块稀疏预填充接到SGLang。H20上128K相对FA2达47.26×(FP8),RULER平均只掉约1分。

这篇在解决什么

长上下文大模型的预填充阶段,注意力仍是平方级计算。训练无关的稀疏注意力先粗估分数、再只算关键块,能把算量砍下来,但多数还停在算法原型:稀疏一狠精度就掉;内核还停在 FlashAttention-2,吃不满 Hopper 上 TMA 和异步流水;连续 KV 布局也接不进 vLLM、SGLang 的分页缓存和连续批处理。

同一组作者前作 FlashPrefill 已经用即时模式发现和基于最大值的动态阈值,避开了 Top-k / Top-p 的全局排序。它仍过不了生产这三关。FlashPrefill V2 的目标很明确:把块稀疏预填充做成能塞进现代推理框架的注意力后端。

方法

预填充仍分两段。第一段用探测 query 和块均值 Key 给每个 query 块打分,阈值取该块最高分乘以 \(\alpha\)(默认 0.1),超过阈值的 Key 块留下,再强制保留 sink、局部窗和最近块。分数矩阵不必物化成 \(L \times (L/B)\),内存从 \(O(L^2/B)\) 收到 \(O((L/B)^2)\)。

极端稀疏时被丢掉的 softmax 质量不可忽略。均值修正给每个被剪块补一项零阶代理:用块内均值 \(\bar kJ,\bar vJ\) 代表整块,往分子分母各加 \(|\mathcal{B}J| e^{\bar sJ} \bar vJ\) 和 \(|\mathcal{B}J| e^{\bar sJ}\)。质量代理是二阶精度,分子还留一项块内协方差;被剪块的质量份额被最大值阈值压住,所以这项误差可控。内核里修正块当作一次额外迭代,logit 平移 \(\log|\mathcal{B}J|\),复用同一套 MMA 和在线 softmax,不再另开 kernel。

Hopper 对齐的稀疏算子是这篇真正能跑进生产的部分。PackGQA 把同一 KV 组的 \(g\) 个 query 头折进同一块,KV 只装一次,稀疏索引也按 KV 头存,元数据缩小 \(g\) 倍。持久化 kernel 拆成生产者 warpgroup(TMA / cp.async 搬 K/V)和两个消费者 warpgroup(wgmma),组内 QK GEMM、PV GEMM、在线 softmax 乒乓重叠。选中块用 CSR 倒序遍历,生产者和消费者各自算同一序列,跳块不必同步。分页地址走 page table,变长请求由持久调度器按 \((b,h,I)\) 枚举,这就是连续批处理的执行模型。FP8 路径在线反量化,概率映射到 e4m3 的 \([0,256]\) 再在比值里消掉。解码阶段单 token 没有块稀疏空间,回退到稠密注意力。

接到 SGLang 时只替换 extend(预填充)后端,模型定义、KV 布局、调度逻辑都不用改。索引阶段和工作空间按流缓存,稳态不再打主机同步。

结果

评测在 NVIDIA H20 上,模型是 Llama-3.1-8B-Instruct、Qwen3-4B-Instruct-2507、Qwen3-30B-A3B-Instruct-2507。精度配置统一:块大小 128、sink 256、局部窗 512、\(\alpha=0.1\),不按模型调参。算子测单卡;端到端 TTFT 用 SGLang、张量并行 4、四张 H20,解码后端固定为对齐 FA3/4 的稠密核。

设置指标FlashPrefill V2对照
Qwen3-30B-A3B, 128K, BF16算子相对 FA227.19×V1 为 18.67×
同上, FP8算子相对 FA247.26×相对 FA3/4 稠密核 30.49×
三模型 RULER 平均相对 Full差 0.29–1.03 分128K 密度约 5%
Qwen3-30B, BS=16, 128KTTFT36.21 s / 25.51 s (FP8)FA3/4 为 123.23 s

RULER 上 V2 平均 87.79 / 86.23 / 91.76,Full 是 88.82 / 87.06 / 92.05。128K 密度掉到 4.6%–4.9%,V2 仍距 Full 不超过 1.8 分;Llama 的 FP8 是在线量化、没校正权重,128K 从 73.82 掉到 67.78,这是精度代价最大的一格。LongBench 21 项平均,V2 是稀疏方法里最高的一档(49.31 / 46.96 / 50.73),距 Full 0.45–0.90 分。

消融把均值修正拿掉:Qwen3-4B 的 RULER 上,BF16 平均只少 0.46 分,FP8 平均少 2.33 分,128K 少 6.2 分。修正本身在 64K、90% 稀疏时额外延迟约 3.6–4.5 ms。\(\alpha\) 从 0.1 推到 0.2,64K FP8 密度从 9.2% 降到 5.2%,带修正仍有 80.12,Full 是 82.81。

开环服务(泊松到达、prompt 4K–128K 混合)里,FA3/4 的请求吞吐卡在 0.31–0.37 req/s,V2 大约翻倍,FP8 到 0.88–1.34 req/s。Chunked prefill 会侵蚀加速,因为每块都要重做索引、短块强制尾部块推高有效密度;论文建议 chunk 至少 8K。

为什么重要

这是稀疏注意力从论文数字走向推理框架的一步,而且步子主要在系统,不在再发明一种选块启发式。H20 是国内线上很常见的推理卡,47× 是相对 FA2 的算子墙钟;更公允的对照是他们自己对齐 FA3/4 的稠密核,128K FP8 仍有 30.49×,端到端 TTFT 最高约 4.8×。4K 附近密度还在 70% 左右,BF16 几乎和 FA3/4 打平,加速真正出现在长上下文、注意力占预填充大头的时候。

能直接当 SGLang 后端、认分页 KV 和连续批,这比再发一个只认连续布局的 CUDA kernel 有用。解码仍走稠密,所以它解决的是预填充墙,不是全阶段稀疏。

局限与存疑

论文没有单独写 Limitations。数字全部来自 H20,没有 A100 / H100 / Blackwell 对照,换架构内核优势可能缩水。速度主表用针在草堆里找针的输入,注意力图比真实杂乱文档更「有结构」,密度 5% 未必能搬到代码仓库或多跳检索。Llama 的 FP8 精度带星号,是在线量化而非校正权重,和 BF16 比 128K 掉分更陡。均值修正是零阶,块内 logit 与 Value 相关时分子仍有一阶残差;作者用低块内方差给这个近似撑腰,没有在高方差注意力图上单独压测。Chunk 开小会把加速吃掉,和线上为了保 TPOT 常用的小 chunk 是冲突的。仓库公开了,生产环境能不能无痛替换,还取决于具体模型和量化方案。

术语

原文与代码

相关论文

全部论文解读