PRR: Predict, Reuse, and Repair
Accelerating Dynamic Sparse Attention for Long-Context LLM Decoding
PRR: Predict, Reuse, and Repair — Accelerating Dynamic Sparse Attention for Long-Context LLM Decoding
一、论文概述
| 项目 | 内容 |
|---|---|
| 标题 | Predict, Reuse, and Repair: Accelerating Dynamic Sparse Attention for Long-Context LLM Decoding |
| 作者 | Tianyu Wang¹, Gourav Rattihalli², Aditya Dhakal², Junbo Li¹, Zhiwei Ren¹, Dejan Milojicic², Longfei Shangguan¹ |
| 机构 | ¹University of Pittsburgh, Pittsburgh, PA, USA; ²HPE Labs, Milpitas, CA, USA |
| 论文 | arXiv:2606.30389 |
| 代码 | github.com/Tianyu9748/Incremental_FlashAttention |
| 发布 | 29 Jun 2026 |
| 篇幅 | 9 pages body + 3 pages appendix = 13 pages total |
二、核心思想
问题定义
动态稀疏注意力(Dynamic Sparse Attention, DSA)通过在每个解码步仅选择与当前查询最相关的 top-K KV blocks 来加速长上下文 LLM 推理。然而,DSA 引入了一个新的延迟瓶颈:selection-to-attention 依赖。在标准 DSA 管线中,每个 transformer layer 的解码步骤由 selection → attention → FFN 三阶段组成。DSA 必须先运行压缩注意力(compressed attention)来识别 top-K KV blocks,然后才能获取这些 blocks 并执行稀疏注意力。由于 block 身份在 selection 完成前未知,attention 严格串行在 selection 之后。随着上下文长度增长,这种选择到注意力的依赖成为越来越严重的瓶颈——selection 占到了每 token 生成延迟的 41%(在 512K 上下文时 selection 从 4ms 增长到 8ms)。
解决方案概述
PRR(Predict, Reuse, and Repair)是一个 speculate-reuse-repair 运行时,利用 DSA 选择在连续解码步之间的 时间局部性(temporal locality) 来预测可能的 blocks,在 selection 还在进行时对预测的 blocks 执行推测注意力(speculative attention),然后在已知真实选择集合后增量修复遗漏的 blocks。
PRR 包含三个核心组件:
- 轻量级 EMA 预测器:跟踪 per-block importance scores 的时间轨迹,预测 upcoming top-K 集合
- profile-guided speculation budget:控制推测集合大小,确保推测计算不超出 selection 窗口
- FlashAttention-based incremental repair kernel:使用 online-softmax 统计将遗漏 blocks 合并到部分注意力状态中
实验结果:在多个长上下文 benchmark 和 DSA 方法上,PRR 实现了最高 40% 的 per-token 解码延迟降低(1.42×–1.64× speedup),同时保持下游任务准确率不变。
三、技术架构
Motivation 关键发现
Observation One: DSA 显著减少了注意力计算,但暴露了 selection-to-attention 串行依赖。在 16K tokens 时,selection+attention 路径消耗约 60% 的解码延迟,在 512K 时上升到 71%。
Observation Two: DSA 的 top-K blocks 在相邻解码步之间表现出强局部性。在五个代表性长上下文 benchmark 上,Quest 和 InfLLM-V2 一致地重新选择超过 65% 的 blocks(平均约 68%)。
Observation Three: 长上下文 DSA 解码留下了充足的空闲 GPU 资源。即使在 512K 上下文长度下,SM、L2-bandwidth 和 DRAM-bandwidth 利用率均低于 40%,表明有充足的空闲资源用于预计算推测注意力。
设计空间分析
PRR 面对两个不对称的设计空间:
Missed Blocks (M = A \ P):位于关键路径上——每个 missed block 必须在 true selection 完成后被迁移和参与注意力,partial attention output 必须修正。
Wasted Blocks (N = P \ A):消耗额外带宽和计算资源,但不延长关键路径。
优化目标是最小化 |M|,约束为 speculation budget:
其中 是 budget ratio,限制 speculator 可以 fetch 的 blocks 数量。 越大能覆盖更多 true top-K blocks,但也增加带宽压力和推测注意力计算成本。如果 过大,speculative attention 本身可能超过 block selection 窗口,延迟 FFN 执行。
Motivation 分析

DSA 每个 decoding step 包含三个阶段:selection → attention → FFN。Selection 阶段需要运行 compressed attention 来识别 top-K KV blocks。由于 block identities 在 selection 完成前不可知,attention 严格串行等待。

关键发现:即使在高 context 长度下,GPU SM、L2-BW 和 DRAM-BW 利用率均低于 40%,说明有充足的空闲资源执行 speculative attention。
Temporal Locality 分析

在五个 long-context benchmarks 上测量连续 decoding steps 间的 block overlap rate:Quest 平均 ~67%,InfLLM-v2 平均 ~69%,跨模型和任务一致。
Speculation Budget 分析

整体架构图

标准 DSA:selection → attention(A) → FFN(串行) PRR:prediction(P) + selection(A) 并行 → speculative attention(P) + selection(A) 并行 → incremental repair(A\P) → FFN
EMA 校准流程

核心公式
EMA 预测器
Prediction(预测):在第 步开始时,在 DSA selection 产生真实分数之前,PRR 用历史分数预测每个 block 的分数:
其中 控制预测器外推近期趋势的激进程度。
Update(更新):在 DSA selection 产生真实分数 后,更新状态:
这里 控制 level 跟随新观测分数的速度, 控制 trend estimate 适应分数变化的速度。
初始化:block 首次进入 KV cache 时(第 步),,。
Online Softmax Incremental Repair
Speculative attention 输出 后,repair kernel 对每个 missed block 应用一次 recurrence:
其中 、 和 是 missed block 的 running statistics。
校准目标函数
在 prefill 阶段校准超参数 :
其中 是 prefill tokens 数量, 是 compressed-attention score matrix, 是 prefill token 的 ground-truth top-K block selection set。
搜索网格: (step 0.2), (step 0.1), (step 0.25),共 80 个候选。
四、核心创新
| 创新点 | 说明 | 理论/实验依据 |
|---|---|---|
| 发现 selection-to-attention 依赖瓶颈 | DSA 虽然减少算术量,但引入串行依赖,在长上下文中占 41% 延迟 | Section 2 profiling:512K 时 selection 占 41%,selection+attention 占 71% |
| EMA-based predictor | 替代 naive “reuse previous top-K” 启发式,利用 block importance scores 的时间轨迹预测 upcoming 集合,平均命中率 98% | Table 4:Quest 下 98.05%,InfLLM-v2 下 97.52% |
| Online-softmax incremental repair kernel | 基于 FlashAttention 的自定义 CUDA kernel,只 attend 到 missed blocks 并通过 log-sum-exp rescaling 合并到已有 accumulator,post-selection 工作从 $ | A |
| Profile-guided dynamic budget | 离线 profile 不同 context length 下的 selection latency 和 speculative attention latency,runtime lookup table 选择最大 safe | Figure 5 展示 过小/过大的影响 |
五、代码实现分析
GitHub 仓库结构
仓库: https://github.com/Tianyu9748/Incremental_FlashAttention
基于 FlashAttention-2 fork,关键目录:
csrc/— CUDA kernel 源码(sparse FlashAttention fork)hopper/— Hopper (H100) 特定优化,variable-kBlockN sparse FA-3 pathsparse_flash_attn_2/— 主 sparse attention 实现benchmarks/— 性能基准测试tests/— 正确性测试roofline/— roofline 分析和利用率指标training/— 训练相关代码
关键实现细节
Instrumented Speculative Kernel:修改 FlashAttention forward,额外写入 denominator 和 running maximum 到 HBM。改动是 surgical 的——inner-loop arithmetic 不变,额外存储仅为每 head 两个 scalars。
Repair Kernel:给定 和 missed blocks 的 KV entries,对每个 missed block 应用一次 recurrence。missed blocks 以 tiled fashion stream through on-chip memory,不 materialize 中间 score matrix。FLOPs 和 memory traffic 仅随 缩放。
GQA Extension for Quest(Appendix A.1):原始 Quest 针对 Multi-Head Attention,每个 query head 独立选 top-K pages。扩展至 GQA:以 KV-head 粒度操作,形成 representative query ,criticality estimation 为 ,每个 KV head 一个 score,选出的 top-K pages 在 group 内共享。
六、实验结果
实验设置
- 硬件: NVIDIA H100 GPUs, CUDA 12.8, tensor parallelism = 2
- 基线后端: FlashInfer’s BlockSparseAttention
- 模型: GLM-4-9B-1M, GLM-Z1-9B, DeepSeek-R1-8B, Llama3-8B-1M, Qwen3-14B, Qwen3-32B
- DSA 方法: Quest, InfLLM-v2(NSA 排除因其需要从头训练)
- Benchmarks: LongBench, InfiniteBench, RULER, AIME, MATH500
端到端加速效果
表1:PRR 相对串行 DSA 执行的解码加速比
| Model | LongBench | InfiniteBench | RULER | AIME | MATH500 | Avg |
|---|---|---|---|---|---|---|
| Quest | ||||||
| GLM-4 | 1.50× | 1.40× | 1.45× | 1.30× | 1.42× | 1.42× |
| GLM-Z1 | 1.48× | 1.39× | 1.43× | 1.33× | 1.45× | 1.42× |
| DeepSeek-R1 | 1.39× | 1.42× | 1.47× | 1.35× | 1.39× | 1.41× |
| Llama3 | 1.36× | 1.40× | 1.41× | 1.39× | 1.34× | 1.38× |
| Qwen3-14B | 1.44× | 1.27× | 1.41× | 1.44× | 1.45× | 1.40× |
| Qwen3-32B | 1.26× | 1.34× | 1.27× | 1.27× | 1.26× | 1.28× |
| InfLLM-v2 | ||||||
| GLM-4 | 1.64× | 1.38× | 1.58× | 1.62× | 1.59× | 1.56× |
| GLM-Z1 | 1.61× | 1.35× | 1.61× | 1.64× | 1.55× | 1.55× |
| DeepSeek-R1 | 1.30× | 1.48× | 1.61× | 1.52× | 1.28× | 1.44× |
| Llama3 | 1.32× | 1.45× | 1.55× | 1.55× | 1.30× | 1.43× |
| Qwen3-14B | 1.56× | 1.43× | 1.51× | 1.53× | 1.53× | 1.51× |
| Qwen3-32B | 1.35× | 1.27× | 1.29× | 1.32× | 1.35× | 1.31× |
PRR 相对于 Quest 平均 1.35× 加速,相对于 InfLLM-v2 平均 1.47× 加速。
消融实验
表2:GLM-4-9B + InfLLM-v2 各组件贡献(相对 Serial 执行)
| Benchmark | S1 (Reuse prev K) | S2 (+Repair kernel) | S3 (PRR full) |
|---|---|---|---|
| LongBench | 1.02× | 1.34× | 1.64× |
| InfiniteBench | 1.00× | 1.19× | 1.38× |
| RULER | 1.01× | 1.34× | 1.58× |
| AIME | 1.02× | 1.34× | 1.62× |
| MATH500 | 1.01× | 1.31× | 1.59× |
| Avg | 1.01× | 1.30× | 1.56× |
关键发现:
- S1(naive reuse):几乎无改善(1.01×)。虽然 blocks 重叠约 68%,但单个 missed block 触发完整 attention recomputation,抵消了推测的收益
- S2(+repair kernel):提升到 1.30×。customized kernel 允许 repair 缺失的 32% blocks
- S3(+EMA predictor):达到 1.56×。EMA 预测器使 overlap rate 从 68% 提升到 98%
Kernel 级加速
表3:PRR sparse kernel vs FlashInfer BlockSparseAttention
| Block Size | 1024 tokens | 2048 | 4096 | 6144 | 8192 | Avg |
|---|---|---|---|---|---|---|
| 16 | 1.17× | 1.28× | 2.00× | 2.83× | 3.68× | 2.19× |
| 32 | 1.17× | 1.38× | 2.09× | 2.78× | 3.69× | 2.22× |
| 64 | 1.17× | 1.38× | 2.26× | 2.91× | 3.64× | 2.27× |
| 128 | 1.17× | 1.41× | 2.23× | 3.00× | 3.71× | 2.31× |
| 256 | 1.03× | 1.11× | 2.00× | 2.75× | 3.64× | 2.11× |
PRR kernel 在所有配置下均优于 BlockSparseAttention,平均加速 2.22×,峰值 3.71×(8K budget)。在 1K tokens 时加速有限(1.03–1.17×)因为绝对延迟小且固定 overhead 主导;在 8K 时 bandwidth-saturation 优势累积。
EMA 预测命中率
表4:GLM-4-9B 上 EMA 预测与真实 top-K 集合的 overlap rate
| DSA Method | LongBench | InfiniteBench | RULER | AIME | MATH500 | Avg |
|---|---|---|---|---|---|---|
| Quest | 98.65% | 96.86% | 98.36% | 98.15% | 98.25% | 98.05% |
| InfLLM-v2 | 97.86% | 97.14% | 96.46% | 98.15% | 98.01% | 97.52% |
准确性验证
PRR preserves 标准 DSA 的语义:虽然推测 likely top-K blocks,但最终注意力输出始终覆盖真实 DSA-selected blocks(通过 repair 确保 full coverage)。因此 PRR 保持与原始 DSA 实现相同的准确率——所有 benchmark 上的下游任务准确率与 baseline 完全一致。
七、相关工作
KV Cache Retrieval for DSA
- H2O (Zhang et al. 2023): KV-dropping 方法
- FlexiCache (Takbir et al. 2026): 利用跨 heads 的 temporal stability
- LouisKV / AsyncSpade: 预测 query state prefetch relevant tokens,但 approximate estimation 可能丢失重要 tokens
PRR 基于历史 importance scores 的 temporal locality prefetch,并通过 attention repair 保证零精度损失。PRR 与 FlexiCache 可组合:PRR 的 EMA predictor 可指导 FlexiCache stable heads 中保留哪些 pages。
Temporal Locality for KV Cache Management
- Infinigen (Lee et al. 2024): 同一 tokens 倾向于在连续解码步被选中或获得高注意力分数
Inference Engines & Kernels
- vLLM, SGLang: 未暴露 incremental incorporation interface
- FlashAttention: 未支持 incremental repair
- FlashInfer BlockSparseAttention: PRR 的 customized kernel 在其基础上快 2.22×
八、总结
核心贡献
- 发现新瓶颈:识别 selection-to-attention dependency 作为 DSA 机制中的新关键路径瓶颈
- PRR 运行时设计:设计 speculate-reuse-repair 运行时,重叠 predicted blocks 的 speculative attention 与 selection,复用正确预测的 block states,仅 repair missed blocks
- Online-softmax incremental repair kernel:基于 FlashAttention 的自定义 CUDA kernel,确保 full top-K coverage(与标准 DSA 语义一致),post-selection 工作 bound 在 missed blocks 数量上
- 广泛验证:在 6 个 LLM 和 5 个 long-context benchmarks 上验证,平均 1.35×–1.56× 加速,零精度损失
局限性
- 仅限 training-free DSA:仅评估 Quest 和 InfLLM-v2,不包含需要从头训练的 NSA
- 仅 Hopper 架构 + half precision:kernel 针对 NVIDIA Hopper 优化,未来可扩展到 Blackwell 和其他低精度(如 fp8)
- 仅 batch size = 1:未探索 batch 内多 request 间的 stage 协调(block selection、speculation 等)
未来方向
- 扩展到可训练 DSA 方法(如 NSA)
- 适配其他 GPU 架构和低精度格式
- 多 request batch 内的 stage coordination