
摘要:提前预测下一层 KV 访问范围,优化稀疏选择与数据预取。
在长上下文推理中,稀疏注意力可以让每个 Query 只访问一部分重要 KV,从而减少注意力计算和数据读取。但随着 Context 持续增长,完整的 KV Cache 会占用越来越多的显存;如果将其卸载到 CPU,解码时还得把选中的 KV 再传回 GPU。此外,在执行稀疏注意力之前,模型要先确定当前 Query 应该访问哪些 KV,稀疏选择本身也会带来额外开销。
在论文「SparDA: Sparse Decoupled Attention for Efficient Long-Context LLM Inference」中,NVIDIA、Thinking Machines Lab、ByteDance Seed 和 MIT 的研究者提出了 SparDA。它的核心设计是将“当前层需要哪些 KV**”的选择过程从 Query 的执行路径中拆出,并提前一层完成预测**。
在具体的设计上,SparDA 会在每层的 Q、K、V 之外增加了一组 Forecast,用当前层的信息预测下一层需要访问的 KV Block。这样,SparDA 可以利用提前一层得到的选择结果,简化稀疏选择,并提前发起下一层 KV 的预取,让数据传输尽可能与当前层计算重叠。
稀疏注意力的两个问题
论文建立在 InfLLM-V2 这类块稀疏注意力之上。对于每个 Query,参与注意力计算的 KV 主要由三部分组成:固定保留的初始块、附近的本地块,以及根据相关性动态选出的 Top-k 块。通过只计算这部分 KV,可以减少注意力计算量和数据读取。
不过,这部分 KV 并非全部预先确定,其中的 Top-k 块需要在每层根据当前 Query 动态选择。原来的流程中,第 l 层先生成 Query(Qₗ),再用它对压缩后的 Key 进行打分,选出当前层需要访问的 Top-k 块。完成选块之后,系统才能确定需要读取哪些 KV,并继续执行后续的稀疏注意力和 FFN。
如果 KV Cache 放在 CPU,选块结果还决定了这一层需要传回 GPU 的数据。因此,稀疏选择既会产生计算开销,也会影响 KV 传输的启动时机。
整个解码路径大致可以表示为:生成 Query → 选择 KV 块 → 获取对应 KV → 稀疏注意力 → FFN。
这里有两层开销叠在一起。一方面,稀疏选择本身需要计算。论文指出,在它采用的块稀疏结构中,注意力主体的复杂度可以从 O(T²) 降到 O(T),而负责选块的部分仍保持 O(T²)。随着 Context 的增长,稀疏选择在整层计算中的占比会越来越高。
另一方面,稀疏选择还决定了 KV 传输什么时候开始。在选块结果出来之前,系统不知道应该从 CPU 获取哪些 KV,CPU→GPU 的传输也就无法提前发起。稀疏选择因此既占用计算时间,也限制了后续 KV 预取能够提前多少。
论文在 MiniCPM4.1-8B、Batch Size 4 上进一步拆分了逐层耗时。预填充阶段,选块的耗时会随序列长度增加,到 128K 时接近块稀疏注意力本身;解码阶段每步只有一个 Query Token,注意力计算更轻,选块的成本更加突出。
SparDA 在 128K 预填充下将选块的耗时最高降低 2.50 倍,解码阶段也明显压低了这部分开销。

图注:不同上下文长度下的单层注意力耗时,其中绿色表示块选择,蓝色表示块稀疏注意力
解耦选块与注意力计算
在第 l 层中,Query(Qₗ)继续负责当前层的注意力计算,而 Forecast(Fₗ)会与第 l+1 层的压缩 Key 进行匹配,提前选出下一层需要访问的 Top-k 块。等模型进入第 l+1 层后,注意力计算仍由这一层自己的 Query(Qₗ₊₁)完成。
整个执行关系可以概括为:
当前层的 Query → 负责当前层注意力
当前层的 Forecast → 负责下一层选块
原本发生在同一层的选块和注意力计算由此被错开一层。模型在执行当前层计算时,就能先得到下一层的选块结果,为后续的数据准备和 KV 预取留出时间。
这里还有一个关键的数据依赖问题:第 l 层为什么能够提前选出第 l+1 层需要访问的 KV**?**Top-k 选择只负责动态选块。固定保留的初始块和局部块由位置规则确定,不需要 Forecast 预测;参与 Top-k 打分的压缩 Key 则保存在 GPU,并在解码过程中持续更新。因此,第 l 层可以利用 Fₗ 和下一层已有的压缩 Key 完成动态选块,无需等到第 l+1 层开始执行。
对于边界层,SparDA 也做了单独处理。第 0 层没有来自上一层的 Forecast,因此会额外训练一个用于当前层选块的 Forecast;最后一层没有下一层需要预测,对应的 Forecast 不再使用。解码阶段,第 0 层的 KV Cache 保留在 GPU,后续各层再进入预取流水线。

图注:SparDA 的整体执行流程
轻量化设计与训练
Forecast 从 Query 中解耦之后,负责稀疏选择的 Indexer 也可以进一步简化。论文采用 GQA,一个 KV Head 对应一组 Query Head。原来的块稀疏选择器需要让同一 GQA Group 中的多个 Query Head 分别参与打分,再汇总这些结果;SparDA 则为每个 GQA Group 只保留一个 Forecast Head,也就是每个 KV Head 对应一个 Forecast Head,同时省去了跨多个 Query Head 汇总时使用的 Softmax。
这一改动在 Prefill 阶段尤其重要。Prefill 时 Key 保存在 GPU 上,不涉及 CPU→GPU 的 KV 预取,因此这一阶段的性能提升主要来自 Forecast Indexer 本身的轻量化。Forecast 除了把选块提前一层,还让负责选块的计算路径变得更短。
训练方面,SparDA 无需重新训练整个基础模型。在 MiniCPM4.1-8B 和 NOSA-8B 上,作者冻结模型主体,只训练新增的 Forecast 投影,共增加 33.5M 参数,约占 8B 模型的 0.41%。训练使用 32 张 H100、2000 个优化步和 ProLong-64K 数据,MiniCPM4.1-8B 在 64K 序列长度下的训练可以在 48 小时内完成。
Forecast 的训练目标,是学习原有选择器给出的块重要性分布。作者使用 Query 计算得到的块重要性作为监督,并通过 KL 散度让 Forecast 的预测结果与其对齐。这里有两个训练细节:第一,监督信号取自最终最大池化之前的块重要性分数,以保留更细粒度的排序信息;第二,在计算 KL 散度时,目标选中的 Top-k 块分别保留各自的分数,其余概率质量合并为一个其余项,组成 k+1 维分布。这样既能学习入选块之间的相对排序,也能让未入选部分参与训练。
作者还提高了监督信号的分辨率。预测端继续使用推理阶段的 (kernel=32, stride=16) 压缩窗口,目标端则使用更细的 (kernel=2, stride=1),计算完成后再通过最大池化映射回相同的块网格。更细的窗口能够提供区分度更高的块重要性信号。

图注:消融实验结果显示,这个设置让 MiniCPM4.1-8B 的 RULER 提升 3.0 分、Reasoning 提升 2.2 分,NOSA-8B 的四类基准也都有提升。
KV 预取的运行时实现
前面解决的是“如何提前知道下一层要访问哪些 KV”。要把这个选块结果转化成解码阶段的性能收益,还需要运行时及时把对应数据搬到 GPU,并尽可能让数据传输与当前层计算重叠。
SparDA 在 Offload 配置下,将完整的 KV Cache 放在锁页 CPU 内存中,GPU 上只保留压缩后的 Key 和第 0 层 KV Cache。压缩后的 Key 体积约为原始 K 的 1/16,主要用于选块。底层实现基于 NOSI 推理引擎,NOSI 会缓存每层上一个解码步骤使用过的 Top-k 块,因此每一步只需要获取本轮新增、且 GPU 中尚未缓存的部分。
以第 l 层为例,模型先生成 Q、K、V 和 F,并将当前 Token 的 K、V 写入 Cache;随后更新压缩 Key,用 Fₗ 计算第 l+1 层需要访问的块,并发起下一层 KV 的预取。与此同时,当前层等待上一层提前发起的 KV 传输完成,再继续执行稀疏注意力和 FFN。
这样一来,第 l 层执行注意力和 FFN 的时间,可以用来传输第 l+1 层需要的 KV。等模型进入下一层时,对应的数据可能已经完成大部分传输,从而减少 PCIe 等待时间。论文的 Algorithm 2 展示的就是这套跨层流水线。

不过,Top-k 选块得到的是一组离散的小块。如果为每个块频繁发起小尺寸数据拷贝,Kernel 启动和同步本身也会产生额外成本。SparDA 因此使用基于 UVA 的 Persistent Triton Kernel:Kernel 启动后,让一小组 CTA 常驻 GPU,持续处理后续的 KV 搬运任务,减少反复启动和同步带来的开销。
CTA 数量还需要在数据传输和主计算之间做平衡。分配更多 CTA 可以提高 KV 搬运吞吐,但也会占用注意力和 FFN 所需的 SM 资源。小 Batch 下传输量较少,只需要少量 CTA;随着 Batch 增大,KV 搬运量同步增加,就需要提高预取侧的并行度。
论文在 H100 上采用的规则是:Batch Size 小于 32 时使用 16 个 CTA,其余使用 32 个;A100 上的分界点为 64。

Table 7 的扫描结果显示,这套策略在不同 Batch Size 下都能达到最优配置,或将差距控制在 4% 以内。
从运行时实现来看,Forecast 提供的是下一层的访问范围,真正把这份信息转化成吞吐收益,还需要 KV Cache Offload、跨层异步预取、Persistent Kernel 和 CTA 资源调度共同配合。
Batch 规模与性能收益
SparDA 在 H100 和 A100 上测试了 MiniCPM4.1-8B 与 NOSA-8B,两个硬件平台上的整体趋势一致。以 H100 的结果为例,在 128K Context 下,MiniCPM4.1-8B 的 Prefill 吞吐达到 17,087.6 tok/s,相比带 Offload 的 Sparse 基线提升 1.25 倍,相比 Dense 提升 2.11 倍;NOSA-8B 相比 Sparse 的最高提升为 1.16 倍。
Decode 阶段的提升更明显。在同样使用 CPU Offload 的条件下,MiniCPM4.1-8B 最高提升 1.69 倍,NOSA-8B 最高提升 1.40 倍。
如果进一步比较单卡在各自最大 Batch 下的峰值吞吐,MiniCPM4.1-8B 相比不使用 Offload 的 Sparse 基线最高达到 5.28 倍,相比 Dense 最高达到 9.21 倍。这里要区分下:1.69 倍反映的是相同 Offload 条件下的执行效率提升,5.28 倍则同时包含 Offload 释放显存后带来的 Batch 扩展收益。

图注:MiniCPM4.1-8B 在不同 Batch Size 下的 Decode 吞吐

主结果之外,附录 Table 10 进一步拆分了 Forecast Indexer 和异步预取各自带来的收益。
在 Batch 4 时,只使用轻量 Forecast Indexer、采用同步搬运的版本表现更高,说明此时性能提升主要来自选块开销的下降;从 Batch 16 开始,完整 SparDA 超过无预取版本;到 Batch 64,完整版本相比无预取版本快约 40%。
这组结果也对应了两种不同的性能瓶颈:小 Batch 下,选块计算占比更高,Forecast Indexer 的轻量化贡献更明显;随着 Batch 增大,KV 搬运量上升,异步预取与计算重叠带来的收益逐渐扩大。
精度影响与适用边界
提前一层预测下一层需要访问的 KV 块,还需要确认预测误差是否会影响模型效果。论文在 HELMET、LongBench、RULER,以及 MATH-500、AIME 2024、AIME 2025 等任务上测试了 MiniCPM4.1-8B 和 NOSA-8B。
从综合结果来看,SparDA 与原有稀疏基线基本持平或略有提升。MiniCPM4.1-8B 的平均分从 61.4 提升到 61.7,NOSA-8B 则从 49.4 提升到 51.7。在 RULER 的长度扩展测试中,SparDA 在两个模型的 32K~128K 设置下均高于稀疏基线,其中 NOSA-8B 在 128K 上从 40.7 提升到 45.0。
不过,不同任务上的变化并不完全一致。MiniCPM4.1-8B 的 HELMET 平均分从 38.9 降到 38.3,其中 Recall 从 67.8 降到 61.5。因此,从论文当前结果来看,提前一层预测没有带来明显的整体精度损失,但部分具体任务仍会出现下降。
长上下文结果还需要结合位置编码外推的设置来看。MiniCPM4.1-8B 原生支持 64K Context,论文在更长长度上使用官方 LongRoPE 配置;NOSA-8B 原生支持 32K,扩展测试则将 rope_theta 从 10,000 调整到 40,000。Dense、Sparse、InfiniGen 和 SparDA 使用相同的外推设置,因此方法之间的对比口径一致,但 96K、128K 下的绝对分数仍会受到位置编码扩展的影响。
SparDA 本身建立在现有块稀疏注意力之上,修改的是 Top-k 选块路径,底层的稀疏模式保持不变。论文目前完整验证的模型只有 MiniCPM4.1-8B 和 NOSA-8B,都是 8B 规模的块稀疏模型。作者认为同样的解耦思路还可以扩展到 Token 级稀疏选择,以及 DeepSeek-V4 的 CSA 路径,但这些方向目前还没有给出完整实验结果。
不同底座的 KV 访问方式也会影响 SparDA 的收益。NOSA-8B 本身带有与 Query 无关的驱逐头,可以先减少一部分 KV 读取,因此留给预取重叠的优化空间更小,这也是它的加速幅度低于 MiniCPM4.1-8B 的一个原因。
此外,SparDA 在 Decode 阶段需要让第 0 层的 KV Cache 常驻 GPU,会额外占用一部分显存。在部分长上下文配置下,这会让它更早遇到 OOM。整体来看,SparDA 的效果仍与底座的稀疏结构、KV 访问策略以及硬件内存条件密切相关。
从稀疏计算到访存调度
SparDA 的工程价值,在于它把原本依赖当前层 Query 才能完成的稀疏选择提前了一层。模型侧只增加一组规模很小的 Forecast 投影,就能更早得到下一层的 KV 访问范围;选择器因此可以进一步简化,运行时也能提前安排对应的数据搬运。
在此基础上,SparDA 再通过独立 CUDA Stream、Persistent Kernel 和 CTA 资源调度,让下一层的 KV 预取与当前层计算尽可能重叠。小 Batch 下,主要收益来自选块计算的降低;随着 Batch 增大,KV 搬运压力上升,提前获取访问范围带来的预取收益也更加明显。
这篇论文提供了一个很具体的思路:如果模型能够更早给出后续需要访问哪些数据,运行时就能围绕这些信息提前做预取、缓存和资源调度。 对长上下文推理来说,稀疏注意力除了减少计算量,也可以进一步为后续的数据访问提供更明确的调度依据。
参考资料:
论文:SparDA: Sparse Decoupled Attention for Efficient Long-Context LLM Inference
arXiv:2606.04511
代码:github.com/NVlabs/SparDA