Back to blog

DashAttention: Differentiable and Adaptive Sparse Hierarchical Attention

用可微、自适应稀疏的 α-entmax 变换替代分层稀疏注意力中的 top-k 路由,让第一阶段按 query 自适应选择可变数量的 KV 块并为第二阶段 softmax 提供 prior,使整个层级端到端可微且非弥散(non-dispersive)。75% 稀疏度下精度媲美 full attention,Pareto 前沿优于 NSA/InfLLMv2,推理较 FlashAttention-3 最高 3.36x 加速

DashAttention: Differentiable and Adaptive Sparse Hierarchical Attention

一、论文概述

项目内容
标题DashAttention: Differentiable and Adaptive Sparse Hierarchical Attention
作者Yuxiang Huang, Nuno M. T. Gonçalves, Federico Alvetreti, Lei Li, Xu Han, Edoardo M. Ponti, André F. T. Martins, Marcos V. Treviso
机构Tsinghua University、University of Edinburgh、Instituto de Telecomunicações / SARDINE Lab(葡萄牙)等
论文arXiv:2605.18753(NeurIPS 2026 preprint)
发布2026 年 5 月 18 日
实现Triton(三个融合 kernel),Stage 1 用 AdaSplash-2,Stage 2 用 FlashAttention
基座模型MiniCPM-4 的 1B / 3B / 8B 变体(长上下文续训 + SFT)

二、核心思想

问题定义

长上下文任务的难度取决于待检索信息的数量、混淆度(能否与噪声区分)、以及在上下文中的分布(分散或集中)。要在这类任务上表现好,模型必须:

  1. 足够选择性(selective):忽略无关 token;
  2. 足够灵活(flexible):无论相关 token 的数量、位置、与其他内容的相似度如何,都能恢复对当前 query 重要的位置。

现有方法无法同时满足两点:

  • Dense softmax attention:满足灵活性(每个可见 token 都有非零权重),但违反选择性——所有 token 都分到概率质量,在长上下文中导致弥散(dispersion)(注意力分布的香农熵随序列长度 nn 增长,lim⁡n→∞H(p)/log⁡n=1\lim_{n\to\infty}\mathcal{H}(\bm p)/\log n = 1)。
  • 硬稀疏路由(top-k 块选择,如 NSA、InfLLMv2):满足选择性,但通过固定预算 kk 实现,牺牲了灵活性;且 top-k 操作切断了粗粒度路由决策与细粒度 token 注意力之间的可微路径(梯度无法直接指导路由分数如何改变块成员)。

核心矛盾:同时实现 query 相关的灵活性 + 严格的 token 级选择性,仍是开放挑战。

解决方案概述

DashAttention(DA) 的关键创新:粗粒度路由器本身就是一个稀疏注意力机制——用 α-entmax(一个自适应稀疏分布,其 support 从输入本身学习、非零质量保持可微)替代对 dense 分数施加的 top-k 截断。

三阶段层级设计(图 1):

  • Stage 0(局部块摘要):用可学习的 summary head 对每个 chunk 做局部 SDPA 生成紧凑摘要;
  • Stage 1(Entmax 块路由):query 用 α-entmax 注意 chunk 摘要,产生稀疏路由分布,support 大小由分数几何自适应决定(信息丰富的 query 路由到多个 chunk,尖锐的 query 路由到极少),实现跨 token / head / layer 的动态稀疏分配;
  • Stage 2(prior 诱导的稀疏 softmax):仅对路由到的 chunk 展开回 token 分辨率,用一个 logits 被 Stage 1 路由权重偏置的 softmax 精炼。

由此模型粗粒度地学习「看哪里」「看多少」,精细地读「读什么」,且整个层级端到端可微。这一设计同时受理论(sparse alternatives 保持集中度、改善长上下文能力)与系统(分层分解,避免直接对全部注意力分数选 token 的昂贵 QK 乘法)双重驱动。

核心结果:75% 稀疏度下精度媲美 full attention,Pareto 前沿全面优于 NSA/InfLLMv2(尤其高稀疏区间);推理较 FlashAttention-3 最高 3.36× 加速,较 InfLLMv2 1.35×。

三、技术架构

整体框架图

DashAttention 总览

图 1:DashAttention 高层总览。Stage 0 用局部 SDPA 构建 chunk 摘要;Stage 1 用 α-entmax 路由得到自适应稀疏 support;Stage 2 在 token 分辨率精炼,其 logits 被从路由权重导出的 di,jd_{i,j} 偏置,保持全可微与 FlashAttention 兼容。

背景公式

标准缩放点积注意力(SDPA):给定 Q,K,V∈Rn×d\bm Q, \bm K, \bm V \in \mathbb{R}^{n\times d},

Z=QK⊤d,P=π(Z+M),O=PV\bm Z = \frac{\bm Q\bm K^\top}{\sqrt d}, \quad \bm P = \pi(\bm Z + \bm M), \quad \bm O = \bm P\bm V

其中 π:Rn→△n\pi:\mathbb{R}^n\to\triangle_n 逐行映射 logits 到概率单纯形,softmax 最常用。

α-entmax 变换(可微稀疏替代):

α-entmax(s)=[(α−1)s−τ1]+1α−1\alpha\text{-entmax}(\bm s) = \left[(\alpha-1)\bm s - \tau\bm 1\right]_+^{\frac{1}{\alpha-1}}

  • [⋅]+[\cdot]_+ 为 ReLU,τ\tau 为使输出为有效概率分布的唯一归一化常数;
  • α→1\alpha\to 1 恢复 softmax,α=2\alpha=2 得 sparsemax,稀疏度随 α\alpha 单调递增;
  • (α−1)si≤τ(\alpha-1)s_i \le \tau 的坐标精确置零,故 entmax 产生动态稀疏——零的模式与数量都取决于输入 s\bm s。

核心公式(三阶段)

Stage 0 — 局部块摘要(Eq. 4):引入初始化为零的可学习 query 向量 qˉ∈Rhkv×dh\bar{\bm q}\in\mathbb{R}^{h_{kv}\times d_h},对每个 chunk Cc\mathcal{C}_c 做局部 SDPA:

kˉc(r)=∑t∈Ccexp⁡(⟨qˉ(r),kt(r)⟩/dh)∑u∈Ccexp⁡(⟨qˉ(r),ku(r)⟩/dh) kt(r)\bar{\bm k}_c^{(r)} = \sum_{t\in\mathcal{C}_c} \frac{\exp(\langle\bar{\bm q}^{(r)}, \bm k_t^{(r)}\rangle/\sqrt{d_h})}{\sum_{u\in\mathcal{C}_c}\exp(\langle\bar{\bm q}^{(r)}, \bm k_u^{(r)}\rangle/\sqrt{d_h})}\, \bm k_t^{(r)}

初始化为零使内积为零 → 局部 softmax 退化为均匀 mean pooling,随训练平滑过渡到加权平均。比 MoBA/InfLLMv2 的 mean pooling 更具表达力,比 NSA 的 MLP 更易从预训练 softmax 模型适配。chunk 摘要一旦 chunk 生成完毕即固定,推理无需重算。

Stage 1 — Entmax 块路由(Eq. 6):query head hh 关联 KV head r=⌊h/gq⌋r=\lfloor h/g_q\rfloor,计算 chunk 级 logits zˉi,c(h)=⟨qi(h),kˉc(r)⟩/dh\bar z_{i,c}^{(h)} = \langle\bm q_i^{(h)}, \bar{\bm k}_c^{(r)}\rangle/\sqrt{d_h},用缩放因子 γ\gamma 做 entmax:

w^i(h)=α-entmax(γzˉi(h))∈△⌊n/B⌋\hat{\bm w}_i^{(h)} = \alpha\text{-entmax}(\gamma\bar{\bm z}_i^{(h)}) \in \triangle_{\lfloor n/B\rfloor}

support S^i(h)={c∣wi,c(h)>0}\hat{\mathcal{S}}_i^{(h)} = \{c\mid w_{i,c}^{(h)}>0\} 决定保留的 chunk。GQA 处理:对每组 rr 内成员 head Gr\mathcal{G}_r,平均 head 级概率 wi(r)=∑h∈Grw^i(h)/gq\bm w_i^{(r)} = \sum_{h\in\mathcal{G}_r}\hat{\bm w}_i^{(h)}/g_q,support 为并集 Si(r)=∪h∈GrS^i(h)\mathcal{S}_i^{(r)}=\cup_{h\in\mathcal{G}_r}\hat{\mathcal{S}}_i^{(h)}。

Stage 2 — Prior 诱导的稀疏 softmax(Eq. 9):从 softmax 的变分形式出发,将 KL 项的参考分布从均匀 u\bm u 换为由 Stage 1 分数导出的 g=gσ(w)\bm g = g_\sigma(\bm w):

piroute=arg max⁡p∈△nzi⊤p−KL⁡(p ∥ gσ(wi)),pi,jroute=gi,jexp⁡(zi,j)∑t≤igi,texp⁡(zi,t)\bm p_i^{\text{route}} = \argmax_{\bm p\in\triangle_n}\bm z_i^\top\bm p - \operatorname{KL}(\bm p\,\|\,g_\sigma(\bm w_i)), \quad p_{i,j}^{\text{route}} = \frac{g_{i,j}\exp(z_{i,j})}{\sum_{t\le i}g_{i,t}\exp(z_{i,t})}

当 gi,j=0g_{i,j}=0(即 wi,j=0w_{i,j}=0)时,−KL⁡→+∞-\operatorname{KL}\to+\infty 自动 mask,softmax 天然稀疏;同时动态稀疏与全可微性通过 entmax 分数 w\bm w 保持。最终输出:

oi=∑j∈Siroutegi,jexp⁡(zi,j)∑t∈Siroutegi,texp⁡(zi,t)⋅vj\bm o_i = \sum_{j\in\mathcal{S}_i^{\text{route}}}\frac{g_{i,j}\exp(z_{i,j})}{\sum_{t\in\mathcal{S}_i^{\text{route}}}g_{i,t}\exp(z_{i,t})}\cdot\bm v_j

该形式对 q,k,v\bm q,\bm k,\bm v 和 prior g\bm g 全可微,梯度可回传至 Stage 1 的 entmax 分数、进而使 Stage 0 摘要可训练。

Prior 强度控制与对角块处理(Eq. 11-12):对角/近对角区域可能不足 BB token(无块摘要)。将 w\bm w 的质量分为 routed 分支 Ri\mathcal{R}_i 与 diagonal 分支 Di\mathcal{D}_i,引入超参 σ\sigma 构造强度削弱的 prior wi,j′=1B⋅wi,c(j)1/σ1⊤wi1/σw'_{i,j}=\frac{1}{B}\cdot\frac{w_{i,c(j)}^{1/\sigma}}{\bm 1^\top\bm w_i^{1/\sigma}}。σ→∞\sigma\to\infty 时 prior 在 routed support 上趋于均匀(仅对选中 token 做 softmax,无额外 prior)。分配因子:

λi=sigmoid⁡(KL⁡(uRi ∥ wRi′)+log⁡∣Ri∣∣Di∣)\lambda_i = \operatorname{sigmoid}\left(\operatorname{KL}(\bm u_{\mathcal{R}_i}\,\|\,\bm w'_{\mathcal{R}_i}) + \log\frac{|\mathcal{R}_i|}{|\mathcal{D}_i|}\right)

直觉:路由器接近均匀(KL≈0)时按 ∣Ri∣/(∣Ri∣+∣Di∣)|\mathcal{R}_i|/(|\mathcal{R}_i|+|\mathcal{D}_i|) 分配;路由器高度信息化(大 KL)时给 routed 分支更多质量。

Proposition 4.1(等价注意力偏置形式):上述计算等价于先算 μi=mean⁡j∈Ri{log⁡wi,c(j)}\mu_i=\operatorname{mean}_{j\in\mathcal{R}_i}\{\log w_{i,c(j)}\},再加偏置:

di,j={log⁡wi,c(j)−μiσ,j∈Ri0,j∈Di,oi=∑j∈Ri∪Diexp⁡(zi,j+di,j)vj∑texp⁡(zi,t+di,t)d_{i,j}=\begin{cases}\frac{\log w_{i,c(j)}-\mu_i}{\sigma}, & j\in\mathcal{R}_i\\ 0, & j\in\mathcal{D}_i\end{cases}, \quad \bm o_i=\sum_{j\in\mathcal{R}_i\cup\mathcal{D}_i}\frac{\exp(z_{i,j}+d_{i,j})\bm v_j}{\sum_t\exp(z_{i,t}+d_{i,t})}

即 prior 折叠为对注意力 logits 的简单加性偏置 di,jd_{i,j},与 FlashAttention kernel 完全兼容。

GPU-Aware 实现(三个融合 Triton kernel)

Kernel作用关键优化
Stage 0可学习 summary token 对其 keys 做 online softmaxkeys 同时作 values,同一 K\bm K tile 一次读取双用,无额外 HBM 往返;chunk 完成即写回 chunk-representation cache 复用
Stage 1注意缓存的 chunk 表示形成 chunk logitschunk 数 Tc=⌊n/B⌋T_c=\lfloor n/B\rfloor 小(16K/B=64→256),整行常驻寄存器;AdaSplash-2 就地解 entmax 阈值 τ\tau;GQA 组内平均后剪枝;support 存为 bitpacked block mask Mi∈{0,1}hkv×Tc\bm M_i\in\{0,1\}^{h_{kv}\times T_c}(32 列/int32)
Stage 2masked FlashAttention 遍历 M\bm M 中 active chunk每选中 chunk 加 per-chunk 路由偏置 di,jd_{i,j} 后 online softmax,单次融合;decoding 用 split-KV 变体沿 KV 维切分暴露并行;训练反向复用 M\bm M(entmax 稀疏 Jacobian)

四、理论分析:非弥散性(Non-Dispersion)

弥散问题:softmax 长上下文注意力的香农熵满足 lim⁡n→∞H(p)/log⁡n=1\lim_{n\to\infty}\mathcal{H}(\bm p)/\log n = 1,使长程依赖建模愈发困难。top-k 稀疏将熵界定在 log⁡k\log k,缓解弥散。

但现有分层方法(NSA/InfLLMv2)在 top-k 选择前用 post-softmax head 聚合,使弥散在聚合阶段重新出现。

定义 4.1(Head aggregation):给定 f:Rn→△nf:\mathbb{R}^n\to\triangle_n、HH 个有界序列 {z(h)}\{\bm z^{(h)}\} 与聚合权重 θ∈△H\bm\theta\in\triangle_H:

aggr⁡f(z(1),…,z(H);θ)=∑h=1Hθh⋅f(z(h))\operatorname{aggr}_f(\bm z^{(1)},\dots,\bm z^{(H)};\bm\theta) = \sum_{h=1}^H\theta_h\cdot f(\bm z^{(h)})

定理 4.2(非正式):任意有限 HH 与 θ∈△H\bm\theta\in\triangle_H 下:

  1. softmax head 聚合是弥散的(构造上如此);
  2. 若 p(h)=α-entmax(z(h))\bm p^{(h)}=\alpha\text{-entmax}(\bm z^{(h)}) 且 ∥p(h)∥0=O(nβh)\|\bm p^{(h)}\|_0=\mathcal{O}(n^{\beta_h})(βh∈(0,1)\beta_h\in(0,1)),则 entmax head 聚合非弥散。

结论:softmax head 聚合会破坏 top-k 稀疏的非弥散性、导致噪声选择;DashAttention 直接在稀疏 entmax 分数上做 head 聚合,规避此问题,在 MK1–MK3 等困难检索任务上表现更好。

α-entmax 映射可视化

图 4:不同 α 值的映射与 top-k softmax (k=1,2) 可视化——展示 entmax 如何随输入自适应产生稀疏 support,区别于固定预算的 top-k。

五、核心创新

创新点说明理论/实验依据
路由器即稀疏注意力用 α-entmax 替代 top-k 截断,support 从数据学习、可变大小Eq. 6;端到端可微
全可微层级Stage 2 的 prior 诱导 softmax 让梯度回传至 Stage 1/0Eq. 9-10;Proposition 4.1
动态稀疏分配跨 token/head/layer 自适应稀疏,摆脱固定预算 kk图 3 层稀疏度;RULER MK 任务
非弥散性entmax head 聚合非弥散,优于 softmax 聚合的 top-k定理 4.2;MK1-MK3 提升
可学习块摘要零初始化 summary head,mean-pooling→加权平均平滑过渡Eq. 4;比 MLP 易适配预训练模型
偏置等价形式prior 折叠为加性 logits 偏置 di,jd_{i,j},兼容 FlashAttentionProposition 4.1
GPU-aware 实现三融合 kernel + bitpacked mask + AdaSplash-2最高 3.36× vs FA-3

六、实验结果

6.1 长上下文性能(RULER 16K,75% 稀疏度)

模型方法Avg (%)SparsityMK2MK3
1BFullAttn66.6078.048.0
NSA48.375.020.06.0
InfLLMv262.475.052.016.0
DashAttention64.975.770.026.0
3BFullAttn69.2094.060.0
NSA49.075.024.010.0
InfLLMv262.875.056.028.0
DashAttention67.775.488.040.0
8BFullAttn85.30100.096.0
NSA55.075.034.012.0
InfLLMv278.975.082.052.0
DashAttention83.675.796.086.0

DA 在所有模型尺寸上全面超越 NSA/InfLLMv2,尤其在困难多键检索(MK2/MK3)上大幅领先,逼近 full attention。

HELMET 16K(Overall %):1B 31.2(vs Full 32.5 / Iv2 29.5)、3B 34.3(vs 37.4/34.2)、8B 46.9(vs 47.7/45.9);Recall 子任务领先明显(8B DA 88.3 vs Iv2 81.4)。

6.2 效率基准(较 FullAttn+FlashAttention 的 wall-clock 加速)

chunk size 64,稀疏度 s∈{75%,87.5%,93.75%}s\in\{75\%, 87.5\%, 93.75\%\}。Prefill batch=1,Decoding batch=24。

场景上下文稀疏NSAInfLLMv2DashAttn
Prefill16K75%0.710.961.34
96K93.7%2.323.063.09
Decoding96K75%0.801.731.96
96K93.7%1.343.103.36
  • Prefill:DA 在每个操作点均最快,较 dense FA 加速 1.34×–3.09×,最密设置下相对 InfLLMv2/NSA 优势最大(其 top-k 开销未被摊薄)。
  • Decoding(内存受限):DA 在 96K、93.75% 稀疏下达 3.36×(vs InfLLMv2 3.10×);优势随上下文与稀疏度单调增长——Stage 2 单次遍历 bitpacked mask,避免 InfLLMv2 在评分与注意力阶段间的显式 top-k 与 per-query 索引物化。

6.3 Cost-Effectiveness / Pareto 前沿

HELMET Pareto 前沿

图 2:HELMET 上精度-稀疏度 Pareto 前沿(8B)。DashAttention 全程支配 NSA/InfLLMv2,低-中稀疏度下略超 full attention。

8B 模型扫描稀疏度(DA 调温度 γ\gamma,NSA/InfLLMv2 调 kk):DA 全程支配基线;~90% 稀疏度下 DA 保持 39.4% overall accuracy,超 InfLLMv2 约 9%、超 NSA 约 19%。凸显自适应性——固定 top-k 会过度分配简单 query 或欠分配困难 query,而 entmax 自适应重塑 support。

6.4 动态稀疏分析

逐层稀疏度

图 3:逐层注意力稀疏度(16K RULER-SG1 输入)。早期层更密集,中间层更稀疏,自动产生类似预算分配策略(PyramidKV/PyramidInfer/MoBA)的效果——但无需人工设计。

6.5 其他结果

  • General tasks(8B 短上下文):DA Avg 59.4,与 FullAttn 59.5 持平,略超 NSA(59.2)/InfLLMv2(59.1),验证不损短上下文能力。
  • DA + softmax 推理(DA+FA):DA 训练的模型用 full softmax 推理,甚至优于 FA 训练模型(1B RULER 66.6→70.4,8B 85.3→86.7),说明可无缝回退到 vLLM/SGLang 的高度优化 softmax 内核。
  • Chunk size:减小 chunk size 提升精度但降低效率(chunk=1 时退化为 entmax+softmax,承担两者成本)。

七、相关工作

  • 注意力稀疏化:静态模式(attention sinks、sliding window)→ 随机(BigBird)/动态(H2O、MInference、block sparsity)→ head 异构稀疏。DA 通过训练消除 train-inference mismatch。
  • KV cache 优化(正交方向):eviction(H2O/SnapKV)、offloading(InfLLM/ShadowKV)、量化(KVQuant/KIVI)——压缩 KV 而非优化注意力稀疏。
  • 可训练稀疏注意力:SeerAttn / NSA / MoBA 用 top-k 选压缩块;InfLLMv2 统一末阶段 kernel;FSA 扩展到更小 GQA 组;HSA 用 local encoder 但加参数多。这些方法受固定 top-k 限制,无动态性。DA 将 entmax 引入分层稀疏注意力,桥接「可训练稀疏」与「entmax 加速」两方向,且易从预训练 softmax 模型适配。

八、总结

核心贡献

  1. 分析 top-k 稀疏注意力的局限,提出 DashAttention——端到端可微、跨 head 自适应分配稀疏的方法;
  2. 集成到长上下文续训,在匹配稀疏度下超越现有分层稀疏方法(NSA/InfLLMv2),并媲美 full attention 精度;
  3. 高效 GPU-aware Triton 实现,较 FlashAttention-3 加速 3.36×、较 InfLLMv2 1.35×。

技术影响

  • 首次将**可微自适应稀疏(α-entmax)**引入分层稀疏注意力的路由阶段,打破 top-k 固定预算范式;
  • 非弥散性理论为长上下文稀疏注意力设计提供新视角——head 聚合方式(softmax vs entmax)直接影响长程建模能力;
  • 提供从预训练 softmax 模型平滑适配的路径(零初始化 summary head、prior 强度可调),且训练后可回退 full softmax 推理。

局限性

  • DashAttention 的 kernel 尚未集成到 vLLM/SGLang 等现代 LLM serving 框架(future work);
  • 未探索应用于其他架构(如混合模型 Nemotron 等);
  • chunk size 与效率/精度权衡需针对场景调优。

九、参考资源

  • arXiv 论文:https://arxiv.org/abs/2605.18753
  • 依赖实现:AdaSplash / AdaSplash-2(entmax GPU kernel)、FlashAttention(Stage 2)、Triton
  • 基线方法:NSA、InfLLMv2、MoBA、SeerAttention、FSA、HSA
  • 评测基准:RULER、HELMET、MMLU、GSM8K、MATH、HumanEval 等