CRISP: Cliff-awaRe Input-adaptive Sparse Prefilling with Structural-Mass-Motivated Routing
长上下文 LLM 推理的 attention prefilling 阶段计算量随序列长度二次增长,自注意力成为严重瓶颈。传统稀疏注意力方法通过固定模式或离线分析缓解该问题,但无法适应输入相关的注意力结构。近期动态方法通过实时将头路由至稀疏模式,但依赖间接路由代理,带来额外开销,且其预算分配机制忽视了 softmax 后的质量层级。 我们提出 CRISP(Cliff-awaRe Input-adaptive Sparse Prefilling),识别并解决该动态路由范式中的两个结构性问题。第一,我们证明路由决策可直接从代理注意力图的结构中读取。我们用 Cstruct(一种结构代理,度量 Vertical-Slash 兼容位置的质量)替代 Jensen-Shannon Divergence 路由,在复现 JSD 路由决策的同时,消除了池化 matmul 及随后的 KL 散度开销。第二,我们形式化 softmax 后质量悬崖,并从理论上证明:在长上下文下,严格累积覆盖率阈值会累积 O(n) 背景噪声。CRISP 通过基于噪声底限的 sink-aware 阈值 规避该问题。 实验方面,在 InfiniteBench、RULER 和 LongBench 上的两个模型家族中,CRISP 整体是最强的稀疏方法,在检索密集型基准上达到或超过精确稠密注意力:相对基线在检索任务上提升高达 +28.0 pp,在 512k token 下实现最高 5.30x 注意力加速。其主要来源是选择阶段 O(n) 噪声消除,同时保持结构完整性。
论文精读
TL;DR CRISP 用结构代理 C_struct 替代 JSD 路由,并通过 sink-aware 阈值消除 softmax 质量悬崖下的 O(n) 噪声,在长上下文 LLM 推理中实现高达 5.30 倍注意力加速和检索精度大幅提升。
问题
问题背景
长上下文 LLM 推理中,自注意力预填充阶段的计算复杂度随序列长度二次增长,成为严重瓶颈。
现有方法局限
- 固定模式或离线剖析的稀疏注意力无法适应输入相关的注意力结构。
- 动态路由方法(如 FlexPrefill)虽然能实时分配稀疏模式,但依赖 JSD 散度 等间接代理,引入额外的池化矩阵乘和 KL 散度计算开销。
- 其预算分配机制忽视 post-softmax 质量层级,使用严格累积覆盖阈值会在长上下文中积累
O(n)背景噪声,导致检索任务精度下降。
为什么这个问题难/重要
- 技术挑战:直接从代理注意力图结构读取路由决策可消除间接开销,但如何定义结构度量并应对 softmax 后的质量悬崖(mass cliff)是核心难点。
- 现有动态方法在检索密集型基准上相比稠密注意力精度损失明显,稀疏加速与精度保持的平衡是超长上下文推理落地的关键。
- 支持 512k token 等超长上下文的模型需求日益增长,需要计算高效且无损的稀疏预填充方案。
行业类比
类似于实时视频分析中,从海量帧中动态筛选关键帧进行目标检测,若选取策略不当,背景噪声帧会显著拖慢处理速度且降低检测精度。
核心洞察
- 结构性代理替代间接路由信号。CRISP 提出 `C_struct` 直接度量注意力图中垂直斜杠兼容位置的 mass,而非通过 JSD 等间接散度。这揭示了动态稀疏路由的信号可以直接从注意力结构读取,无需额外池化矩阵乘法和 KL 散度开销。与 FlexPrefill 的 JSD 路由相比,`C_struct` 在复现路由决策的同时消除了计算瓶颈,证明路由决策本身并不需要复杂的度量,结构特征已足够,为高效输入自适应稀疏提供了新思路。
- Sink-aware threshold 处理质量悬崖与 O(n) 噪声。论文指出严格累积覆盖阈值在长上下文下会累积 O(n) 背景噪声,因为 post-softmax 质量分布存在悬崖式层级,传统阈值无法导航。CRISP 通过基于噪声底部的阈值,识别并保留结构相关 token,同时消除背景噪声,从而在检索任务上匹配甚至超过 dense attention。这一角度独特在于:之前的稀疏注意力预算分配忽略了 softmax 后的质量层次,单纯按覆盖比例分配预算导致噪声污染,而 CRISP 的噪声地板阈值提供了新的分配准则,在 512k tokens 下实现 5.30x 加速并恢复检索精度。
方法
CRISP 的输入是长上下文 LLM 在 prefilling 阶段 的注意力计算需求:给定输入序列,每个注意力头需要实时决定该头应采用何种稀疏模式,以跳过大量低质量注意力计算。
关键模块 1:结构代理路由 C_struct
传统动态稀疏方法(如 FlexPrefill)先计算代理注意力图,再通过 Jensen-Shannon 散度 (JSD) 与固定模式对比进行路由。CRISP 证明该路由决策可直接从代理注意力图的结构中读取。具体做法:
- 在代理注意力图中,测量 Vertical-Slash 兼容位置 的质量总和,得到标量 C_struct。
- 该结构代理无需额外的池化 matmul 和后续 KL 散度计算,但能复现 JSD 路由决策,显著降低路由开销。
- 路由结果:每个头被分配到与其注意力结构最匹配的稀疏模式。
关键模块 2:质量悬崖与 sink-aware threshold
后 softmax 注意力质量呈 悬崖状分布:少数 token 贡献绝大部分注意力质量,其余 token 为近似均匀的背景噪声。传统累积覆盖阈值(设定累积质量比例 γ)会随序列长度线性累积 O(n) 背景噪声,因为长上下文中噪声总量与 n 成正比。CRISP 的解决方案:
- 基于注意力图的 噪声底限(如背景噪声水平的中位数)确定动态阈值,仅保留质量显著高于噪声底限的 token。
- 该阈值是 sink-aware 的,能自动适应不同头、不同层的噪声分布,避免将背景噪声误判为重要信号。
最终输出为每个头的稀疏 attention mask,用于 prefilling 阶段跳过非关键计算。
与同类的动态稀疏注意力方法(如 FlexPrefill 的 JSD 路由 + 预算分配)相比,CRISP 使用直接的结构代理消除间接计算开销,并用 sink-aware 阈值替代盲目累积阈值,从根本上抑制了长上下文下的 O(n) 噪声累积。
实验
实验设计
- 在 InfiniteBench、RULER、LongBench 三个长上下文基准上,评估两个不同模型家族。
- 对比基线包括 FlexPrefill 等动态稀疏方法,以及精确稠密注意力作为上限参考。
- 重点考察检索密集型任务与一般长上下文任务,并在 512k tokens 序列长度下测量加速。
关键发现
- CRISP 在检索密集型基准上精度恢复最高达 +28.0 pp,匹配甚至超过稠密注意力。
- 在 512k tokens 时实现 5.30x 注意力加速,主要得益于
O(n)噪声消除。 - 整体而言,CRISP 在稀疏注意力方法中取得最佳综合表现。
与基线对比解读
- 相比 FlexPrefill 使用的 JSD 路由,CRISP 的
C_struct直接读取代理注意力图的结构,省去池化矩阵乘与 KL 散度计算,降低路由开销且更具可解释性。 - sink-aware threshold 取代累计覆盖阈值,避免了长序列下背景噪声的
O(n)累积,这是检索精度大幅回升的关键。 - 该工作表明:动态稀疏注意力的路由决策应基于注意力图本身的几何结构,而非间接统计代理;同时预算分配需考虑 softmax 后的质量层级(mass hierarchy)。
行业影响
落地场景
CRISP 适用于需要长上下文预填充的实时应用,如 RAG 系统、长文档问答、代码库分析、多轮对话记忆等。例如,在金融研报分析中,对百页 PDF 进行摘要与问答,可显著降低首 token 延迟;在医疗记录综合中,处理跨多次就诊的 EHR 数据,提高检索准确性。
商业价值
- 降本:注意力加速最高 5.30x,减少 GPU 占用,降低按 token 计费成本,提升单卡并发吞吐。
- 体验提升:支持 512k tokens 长输入,避免截断导致的信息丢失,恢复检索任务准确率(最高 +28.0 pp)。
- 增收:为长上下文场景提供高性能推理,使高端模型能力可下沉到实时 API。
与现有产品/工作流的接口
- 作为自定义 attention 算子集成到 vLLM / TensorRT-LLM 等推理引擎,通过模型配置启用稀疏预填充,无需重新训练。
- 路由信号基于预填充阶段的 proxy attention map,可直接替换 FlexPrefill 等动态稀疏方法,适配各类 Transformer 架构。
- 阈值参数(sink-aware)可从噪声 floor 估计,提供默认值,降低调参负担。
局限
- **任务范围局限**。CRISP 主要处理 long-context prefilling 阶段的 self-attention 计算瓶颈,其 routing 决策基于 proxy attention map 上的 **Vertical-Slash 结构**。该结构假设注意力质量集中在特定的行/列模式,对于某些 head(如全局均匀注意力、对角线外长程依赖)可能失效,从而降低路由准确性。论文在单模型家族和 newer architecture 上做了初步迁移实验,但缺乏对更广泛模型(如混合架构、非 Transformer)的评估,泛化边界尚不清晰。
- **阈值机制依赖噪声底估计**。sink-aware threshold 依赖对噪声 floor 的估计(如 median 或 mean),该估计对 attention score 分布的形状敏感。当输入长度、batch size 或温度系数变化时,噪声底可能漂移,导致阈值设置不理想,影响稀疏度与精度平衡。论文附录虽有超参数鲁棒性实验,但未探索自适应估计噪声底的方法,实际部署时需要针对硬件和模型重新校准,增加了工程成本。
- **与 subquadratic 替代方案的定位差异**。CRISP 依然保留了 exact attention 的二次计算路径,仅在 routing 后对选中的位置进行稀疏注意力计算,因此理论复杂度未达到线性。相比 Linear Attention、RWKV 等 subquadratic 方案,CRISP 在极端长序列(如 1M+ tokens)下可能仍受限于 proxy 计算和稀疏化后的有效 token 数量。论文报告的 5.30x attention speedup 主要来自注意力模块本身,未给出端到端推理吞吐和内存占用的完整对比,实际系统级收益需要进一步验证。