DashAttention:可微分自适应稀疏分层注意力

DashAttention: Differentiable and Adaptive Sparse Hierarchical Attention

arXiv: 2605.18753v1

论文信息

标题: DashAttention: Differentiable and Adaptive Sparse Hierarchical Attention

作者: Yuxiang Huang, Nuno M. T. Gonçalves, Federico Alvetreti, et al.

发布日期: 2026-05-18

arXiv ID: 2605.18753v1

PDF 链接: 下载 PDF

3 分钟速览

  • 研究问题:现有分层稀疏注意力方法(如 NSA、InfLLMv2)采用固定的 top-k 块选择策略,无法根据查询动态调整所需块的数量,且因硬截断操作阻断了梯度流动。论文旨在解决查询依赖的灵活性与严格的稀疏性并存的问题。
  • 核心方法:提出 DashAttention,用可微分的自适应稀疏变换 α-entmax 替代 top-k 选择,在粗粒度阶段生成变长的稀疏支持集,并将其作为先验信息注入第二阶段细粒度的 softmax 注意力计算中,保持全流程端到端可微。
  • 关键结果:在约 75% 的稀疏度下,DashAttention 在长文本基准 RULER 和 HELMET 上的准确率与全注意力相当,并在推理速度上最高可实现 FlashAttention-3 的 3.36 倍加速(见表 5)。
  • 主要局限:论文作者指出,DashAttention 的定制化 Triton 内核尚未集成到 vLLM、SGLang 等主流 LLM 推理框架中,目前无法直接开箱即用地在生产环境中实现端到端加速。
  • 适合读者:对大型语言模型高效推理、长上下文处理、以及注意力机制稀疏化与差异化训练感兴趣的研究者和工程师。

论文背景和研究动机

处理长上下文的核心挑战在于信息量、信息模糊性及其在上下文中的分布复杂性。标准的 softmax 全注意力机制会给每个可见 token 分配非零的概率权重,这保证了模型不会遗漏关键信息。然而,Veliˇckovi´c 等人的研究表明,在极长上下文中,softmax 的注意力分布存在弥散现象:随着序列长度 nn 增加,其分布的香农熵 H(p)H(p) 以 log⁡n\log n 的速度增长(lim⁡n→∞H(p)/log⁡n=1\lim_{n\to\infty} H(p) / \log n = 1),导致模型将注意力分散到过多无关 token 上,难以精准聚焦。

为应对此问题,以 NSA 和 InfLLMv2 为代表的现代方法采用了层次化稀疏注意力设计. 它们先将上下文分成块并压缩,用粗粒度分值选出 top-k 个最相关的块,再在这些块内部进行精确的 softmax 注意力计算。但这引入了两个新的矛盾:

  1. 预算固定,缺乏灵活性:top-k 操作为所有查询预设了相同数量的获选块,忽略了不同查询对信息量的需求差异。
  2. 梯度断裂,影响优化:硬性的 top-k 截断使得粗粒度路由决策(哪些块被选中)与最终的损失函数之间失去了可微的连接,梯度信息无法反向传播来优化块选择和压缩表征。

DashAttention 的动机正是融合两种互补的归纳偏置:利用 α-entmax 在粗粒度层面抑制不相关的块以实现稀疏,而在被选中的精细区域内用 softmax 保留语义相关性,从而同时满足 “足够的筛选力” 和 “充分的灵活性”。

核心方法和技术细节

DashAttention 按三个阶段进行,设计上确保整个流程是完全可微的。

Stage 0:局部块摘要 首先将键值(KV)缓存分割成大小固定为 BB 的连续块。与 NSA 使用可训练的 MLP 或 InfLLMv2 使用平均池化不同,DashAttention 引入了一个可学习的局部注意力机制。它初始化一个与主模型无关的零向量查询 qˉ\bar{q},使其在每个块内部执行缩放点积注意力,从而聚合出一个块摘要向量 kˉc\bar{k}_c。这一设计的巧妙之处在于,训练开始时,qˉ\bar{q} 为零向量,局部 softmax 退化为均匀的平均池化。这确保了 DashAttention 能够平稳地从预训练的全注意力模型开始微调,然后逐渐学会更富表现力的压缩。

Stage 1:Entmax 块路由 对于给定的查询 qiq_i,它首先与所有块的摘要向量 kˉc\bar{k}_c 计算粗粒度相似度分数。关键创新在于,这些分数不再是经由 top-k 筛选,而是通过 α\alpha-entmax 变换(α>1\alpha > 1)转化为一个概率分布 w^i\hat{w}_i。其数学形式为 α-entmax(s)=[(α−1)s−τ1]+1α−1\alpha\text{-entmax}(s) = [(\alpha-1)s - \tau \mathbf{1}]_+^{\frac{1}{\alpha-1}},其中 τ\tau 是确保概率总和为 1 的阈值,[⋅]+[\cdot]_+ 代表 ReLU 函数。这一操作的输出天然具备动态稀疏性:得分低于阈值 τ\tau 的块的概率会被精确置零,从而形成一个对当前查询敏感的动态支持集。这种稀疏模式及其大小完全由输入数据本身的几何性质所决定,赋予了不同查询、注意力头乃至不同层自适应分配稀疏度的能力。

Stage 2:先验引导的稀疏 softmax 注意力 最后进入细粒度阶段。前一步产生的块路由概率 wiw_i 作为先验分布 gσ(wi)g_\sigma(w_i) 被注入 softmax 计算中。论文从 softmax 的变分形式出发,将 KL 散度中的基准分布由均匀分布替换为该先验。其中,一个关键的超参数 σ\sigma 用于控制先验的强度。当 σ→∞\sigma \to \infty 时,先验在被选中的块上趋于均匀分布,该方法就退化为仅在被选中的块上执行标准的 softmax 注意力,保证了与现有高效注意力内核(如 FlashAttention)的兼容性。最终注意力权重的计算被等价地转换为在标准的 softmax logits zi,jz_{i,j} 上添加一个路由偏置项 di,jd_{i,j},即 softmax(zi,j+di,j)\text{softmax}(z_{i,j} + d_{i,j}),这极大地便利了工程实现。

创新点和贡献

  1. 端到端可微的分层注意力框架:论文首次将 α-entmax 变换引入层次化注意力设计的粗粒度路由阶段,解决了此前方法因 top-k 硬截断导致的梯度断裂问题。整个三级流水线(局部摘要、块路由、精细注意力)均在训练时保持完整的梯度链。
  2. 理论上的非弥散性证明:论文形式化地证明了在分组查询注意力框架下,基于 softmax 的头聚合操作具有弥散性,而 DashAttention 采用的 entmax 头聚合因为能将支持集控制在 O(nβh),βh∈(0,1)O(n^{\beta_h}), \beta_h \in (0,1) 量级,所以是非弥散的(见论文定理 1)。这一理论特性解释了为什么 DashAttention 在需要多跳检索的复杂长文本任务上表现更优。
  3. 动态稀疏度分配:与固定预算的方法不同,DashAttention 无需繁琐地调参指定 top-k 值,其路由机制能根据几何结构自动学习分配稀疏度。论文实验观察到,早期层倾向于更密集的注意力模式,而中间层变得更稀疏,这与许多手动设计的金字塔式预算分配策略不谋而合,但它是自动涌现的(见图 3)。
  4. 高效的 GPU 感知实现:论文提供了定制化的 Triton 核心实现。其位压缩掩码和单次融合遍历的设计,避免了像 InfLLMv2 那样在评分和注意力阶段之间进行显式的索引物化,使得在预填充和解码阶段的速度均超越了所有基准。解码阶段最高达 FlashAttention-3 的 3.36 倍(上下文长度 96K,稀疏度 93.75%,见表 5)。

实验结果分析

论文在长上下文持续预训练设定下,基于 MiniCPM-4 的 1B、3B 和 8B 参数量模型,与 NSA 和 InfLLMv2 进行了全面对比。

  • 长上下文性能:在综合基准 RULER 和 HELMET 上,DashAttention 在所有模型规模下均显著超越两种基线方法,平均稀疏度也略高(约 75.4% 对 75%)。尤其在 RULER 的 MK1-MK3 这类需要从上下文中提取多跳关联信息的任务上,优势尤其明显。例如,在 8B 模型的 MK2 任务上,FullAttn 为 100%,DashAttention 为 96%,而 NSA 仅为 34%,InfLLMv2 为 82%(见表 1)。
  • 通用性能保持:在 MMLU、GSM8K 等短文本通用任务上,DashAttention 的得分与进行全注意力推理的模型相当甚至略优(8B 模型平均分 59.4 对 FullAttn 的 59.5,见表 3),证明引入稀疏注意力并未损害模型的基础能力。
  • 推理效率:在 NVIDIA GH200 GPU 上的内核级基准测试表明,DashAttention 是所有被测方法中最快的(见表 5)。
  • 成本效益的帕累托前沿:论文通过调节稀疏度绘制了准确性-稀疏度曲线。DashAttention 的帕累托前沿完全主导了 NSA 和 InfLLMv2,在高稀疏度(~90%)时仍能保持 39.4% 的 HELMET 整体准确率,远超 InfLLMv2(~30%)和 NSA(~20%)(见图 2)。

实践建议

对于希望在生产环境中应用 DashAttention 或类似技术的读者,以下几点具有实践指导意义:

  1. 平滑迁移预训练模型:DashAttention 的 Stage 0 设计确保了从标准 softmax 全注意力模型开始的零障碍初始化。如果你计划对现有的密集模型进行长上下文微调,采用可学习摘要查询 qˉ\bar{q} 的方法比直接引入 MLP 压缩器或使用平均池化更稳妥,它能自然地保留模型的初始表征能力,然后渐进学习。
  2. 关注超参数 α\alpha 和 σ\sigma 的调优:α\alpha 控制 α-entmax 的稀疏度,越大越稀疏。论文采用 1.251.25 到 1.51.5 的渐增训练策略,并使用温度 γ\gamma 调节推理稀疏度。σ\sigma 控制先验强度,极大的值将使路由先验失效。建议根据任务对信息检索的粒度需求来调节这些参数:对于高度聚焦的检索任务,可保持适中稀疏度;对于需要上下文概述的任务,可适当降低稀疏度或增大 σ\sigma。
  3. 进行上下文长度微调而非从头预训练:论文的方法是在小型长文本数据集(如 InfLLM-5B)上做轻量级持续预训练,而非从头训练。这大大降低了计算成本。对于自研模型,推荐沿用此思路,并使用 WSD 学习率调度器进行高效的知识注入。
  4. 预留工程集成时间:尽管 DashAttention 的内核性能优异,但目前其 Triton 实现尚未像 FlashAttention 那样被深度集成到 vLLM 或 SGLang 等推理框架。从算法验证到高吞吐服务部署,还需在内存管理、连续批处理和分页注意力机制对接上进行额外的工程投入。至少在这个集成完成前,部署时可采用论文提到的折中方案:训练使用 DashAttention,推理时直接切换到全注意力,这在实验中甚至能获得更高性能。