MiniMax Sparse Attention 全文翻译
本文译自 MiniMax 的 MiniMax Sparse Attention1,2026 年 6 月 11 日提交于 arXiv,30 页 14 图。原文提出 MSA:一种建立在 Grouped Query Attention(GQA)之上的块状 sparse attention。超长上下文能力正在成为前沿 LLM 不可或缺的一环——agentic 工作流、仓库级代码推理、持久记忆,都要求模型在几十万到上百万 token 上联合做 attention,而 softmax attention 的平方级开销让这件事在部署规模下无法承受。MSA 用一个轻量的 Index Branch 给 key-value 块打分,为每个 GQA group 独立选出 Top-$k$ 子集,从而实现 group 特有的稀疏检索,同时保持块级执行的高效;Main Branch 随后只在被选中的块上做精确的块稀疏 attention。MSA 围绕「简单、可扩展」这条原则设计,被刻意做得很精简,因此可以直接在很宽的 GPU 范围上高效部署。为了把稀疏度真正变成加速,作者把 MSA 和一条 GPU 执行路径协同设计:用 exp-free 的 Top-$k$ 选择和 KV-outer 的 sparse attention,改善块粒度访存下的 tensor core 利用率。在一个 109B 参数、原生多模态训练的模型上,MSA 表现与 GQA 相当,而 1M 上下文下的单 token attention 计算量降低到原来的 1/28.4;配合协同设计的 kernel,MSA 在 H800 上取得 14.2 倍 prefill、7.6 倍 decoding 的墙钟加速。推理 kernel 已开源2,由 MSA 驱动的生产级原生多模态模型 也已公开发布3。
1 引言
大语言模型(LLM)正在从短的单轮交互,迅速转向跨越数百个交错的推理与行动步骤的长程 agentic 工作流——编写并部署生产代码、浏览开放网络、编排各种工具、产出结构化文档[OpenAI,2025;Anthropic,2025;Google DeepMind,2025;DeepSeek-AI,2026;Moonshot AI,2026;Zhipu AI,2026]。然而这些任务所需的超长上下文,对训练和推理都造成了严重的计算与显存瓶颈,平方级开销的 softmax attention 是主要元凶,生产规模部署的延迟与吞吐约束进一步放大了这个问题。
上下文长度是 LLM 的一个关键 scaling 维度,在这个维度上权衡模型质量与效率仍是一个艰巨的挑战。社区正在积极推进这条战线上的 Pareto 前沿。混合架构[MiniMax,2025b;Qwen,2026]把一部分 softmax attention 层替换成 linear attention[Team 等,2025a;Yang 等,2025;Gu 和 Dao,2023]或 sliding window attention[OpenAI 等,2025;MiMo 等,2026]这类更高效的替代品。另一条路线则尝试把 softmax attention 本身稀疏化[DeepSeek-AI 等,2025;DeepSeek-AI,2026;Team 等,2025b;Lu 等,2025],以突破计算瓶颈。
我们提出 MiniMax Sparse Attention(MSA),设计遵循奥卡姆剃刀:在做了大量消融之后,只保留真正必要的组件。MSA 沿用 sparse softmax attention 这一范式,以最大程度复用已有的软硬件基础设施。我们采用块状 token 选择配合较小的 top-$k$,使其能在更广的 GPU 架构上高效执行,同时放宽了此前设计所施加的 head 维度约束。具体地,一个超轻量的 Index Branch 通过 max-pooling 打分,为每个 attention group 选出 top-$k$ 个块,同时始终保留最近的那个块以保证训练稳定。
要把 MSA 理论上的稀疏度转化为实际的端到端加速,需要把算法和它的 GPU 执行路径协同设计。为此我们设计了一个专门面向小 k 场景的 exp-free TopK kernel,借助块状 indexer 绕开选择之前那些不必要的 softmax 计算。对于主 attention 分支,我们把 sparse attention 组织成 KV-outer 的顺序:被选中的 KV 块把关联到它的 query 收集起来、拼接以填满 tensor core 的 MMA,用预调度的分块加两阶段合并来应对块热度高度倾斜的情况,且不需要原子更新。对训练,我们进一步把 sparse KL loss 所需的辅助 LSE 计算融合进前向,并在反向里采用 persistent 的负载均衡。
为了验证 MSA 同时保住了文本与多模态能力,我们在一个 109B 参数、以 3T token 预算从零训练的 Mixture of Experts(MoE)模型上,把它与 Grouped Query Attention 做了对比。MSA 在下游 benchmark 上与 GQA 相当,同时在 1M 上下文长度下带来 14.2 倍 prefill 和 7.6 倍 decoding 加速。
主要贡献。
- 我们提出 MSA,一种极简、可扩展且经过加速的块状 sparse attention 机制,既支持从零训练,也支持从预训练好的 GQA checkpoint 近乎无损地转换过来。
- 我们协同设计了高效的训练与推理 kernel,把 MSA 理论上的计算节省在大规模下变成真实的墙钟加速。
- 我们做了大量消融,一直扩展到 109B 参数、原生多模态训练的 MoE 模型,剖析了 MSA 在不同规模与模态下的行为。
2 预备知识
2.1 Causal Attention 与 GQA
我们用 $N$ 表示序列长度,$d_{\rm model}$ 表示 hidden 维度,$d_h$ 表示 head 维度。对每个 query 位置 $t$ 和 head $h$,causal Softmax Attention 计算
$$ {\bm{o}}_t^{(h)} \;=\; \sum_{i \le t} \alpha_{t,i}^{(h)} \, {\bm{v}}_i^{(h)}, \qquad \alpha_{t,i}^{(h)} \;=\; \frac{\exp\!\big(\langle {\bm{q}}_t^{(h)}, {\bm{k}}_i^{(h)} \rangle / \sqrt{d_h}\big)}{\sum_{j \le t} \exp\!\big(\langle {\bm{q}}_t^{(h)}, {\bm{k}}_j^{(h)} \rangle / \sqrt{d_h}\big)}. \tag{1} $$式 (1) 的开销是 $\Theta(2H_q N^2 d_h)$ FLOPs,随序列长度 $N$ 平方增长。Grouped-Query Attention[Ainslie 等,2023]使用 $H_q$ 个 query head,并把 key-value head 的数量减少到 $H_{kv}$,把相邻的 $G=H_q/H_{kv}$ 个 query head 绑到一个共享的 key-value head 上。因此,每个 key-value head 定义了一个 GQA group。
译注:attention 与 GQA 的基础可参见 当计算撞上内存墙:Attention!注意力机制及其优化算法浅析,KV cache 侧的优化手段见 KV-Cache 优化。
2.2 稀疏 attention 的两阶段视角
一个 sparse attention 层把 causal attention 分解成两部分:一个 indexer 负责选出要 attend 哪些 key,以及一次在被选中的 key 上做的 sparse attention 计算。对每个 query 位置 $i$,
$$ {\mathcal{I}}_i \;=\; \mathrm{Index}_\phi\!\big({\bm{q}}_i, {\bm{K}}_{\le i}\big), \qquad {\bm{o}}_i \;=\; \mathrm{Attn}\!\big({\bm{q}}_i, {\bm{K}}[{\mathcal{I}}_i], {\bm{V}}[{\mathcal{I}}_i]\big), \tag{2} $$其中 $\mathrm{Index}_\phi$ 由 $\phi$ 参数化(固定规则的 indexer 为空,可训练的则是学出来的),${\mathcal{I}}_i \subseteq \{1, \dots, i\}$ 表示被选中的索引集合,$\mathrm{Attn}$ 表示限制在这个索引集合上的标准 scaled dot-product softmax attention。我们把第一阶段称为 Index Branch,第二阶段称为 Main Branch。在 multi-head attention 里,由位置 $i$ 和 query head $h$ 共同确定的每个 query,都可以选择不同的 key/value 索引集合,记作 ${\mathcal{I}}_i^{(h)}$;式 (2) 省略 head 下标只是为了记号简洁。
2.3 基于 GQA 的块稀疏 attention
per-head 的 token 级选择粒度最细,但这种细粒度计算很难高效映射到 GPU 的矩阵运算上。为了效率,建立在 GQA 之上的 sparse attention 可以在每个 GQA group 内部共享索引结果。令 $\mathcal{H}_r$ 表示由第 $r$ 个 key-value head 服务的那 $G$ 个 query head,group 共享的索引集合可以写成
$$ {\mathcal{I}}_i^{(r)} = {\mathcal{I}}_i^{(h)} = {\mathcal{I}}_i^{(h')}, \qquad h,h' \in \mathcal{H}_r . \tag{3} $$选择 key/value 块而不是单个 token,能减少路由开销,也让 sparse attention 更规整。对块大小 $B_k$,定义
$$ {\mathcal{B}}_b
{(b{-}1)B_k+1,\dots,\min(bB_k,N)}, \qquad b=1,\dots,B,\quad B=\lceil N/B_k\rceil . \tag{4} $$
对 query 位置 $i$ 和 GQA group $r$,集合 ${\mathcal{I}}_i^{(r)} \subseteq \{1,\dots,B\}$ 表示被选中的块索引集合。group $r$ 中任一 query head 的 sparse attention 输出,就是在被选中块内那些因果可见的 token 上、用同一 group 的 key-value head 算出来的。MSA 遵循这套基于 GQA 的块稀疏形式,具体的 indexer 架构与训练目标见下一节。
3 MSA
我们提出 MiniMax Sparse Attention(MSA),一种带两个分支的、基于 GQA 的 sparse attention 机制,如图 1 所示。对每个 query token,一个轻量的 Index Branch 从因果上下文里选出一小组 key 块,Main Branch 则在这些块内的 token 上计算 softmax attention。Index Branch 在标准 GQA 之上只增加两个投影矩阵,工作在块粒度上,并且为每个 GQA group 独立做选择。我们在 3.1 节描述架构,在 3.2 节描述训练过程。
3.1 架构
MSA 把 2.2 节的两阶段 sparse attention 形式,实例化到 GQA group 与块这两个粒度上(图 1)。对每个 query token,Index Branch 为每个 GQA group 选出 $k$ 个大小为 $B_k$ 的 key 块,Main Branch 只 attend 被选中块内的 token,其预算最多是 $kB_k$。令 ${\bm{X}} \in \mathbb{R}^{N \times d_{\rm model}}$ 为输入 hidden states。沿用 2.1 节的记号,我们用 $H_q$ 和 $H_{kv}$ 分别表示 query head 数和 key-value head 数,于是每个 key-value head 服务 $G = H_q/H_{kv}$ 个 query head。
Index Branch。Index Branch 为每个 GQA group 引入一个 index query head,以及一个跨 group 共享的 index key head:
$$ {\bm{Q}}^{\rm idx} = {\bm{X}} {\bm{W}}_q^{\rm idx} \in \mathbb{R}^{N \times H_{kv} \times d_{\rm idx}}, \qquad {\bm{K}}^{\rm idx} = {\bm{X}} {\bm{W}}_k^{\rm idx} \in \mathbb{R}^{N \times 1 \times d_{\rm idx}}. \tag{5} $$对 query token $i$ 和 group $r$,Index Branch 先给可见的 key token 打分,再把这些分数聚合到块级。用 2.3 节定义的块划分 ${\mathcal{B}}_1,\dots,{\mathcal{B}}_B$,
$$ {\bm{S}}^{\rm idx,(r)}_{i,j} \;=\; \frac{\bigl({\bm{Q}}^{\rm idx}\bigr)^{(r)}_i \,\bigl({\bm{K}}^{\rm idx}\bigr)_j^{\top}} {\sqrt{d_{\rm idx}}}, \qquad M^{\rm idx,(r)}_{i,b} \;=\; \max_{\substack{j \in {\mathcal{B}}_b \\ j \le i}} {\bm{S}}^{\rm idx,(r)}_{i,j}. \tag{6} $$这里 $r$ 索引 GQA group,$j \le i$ 保证因果性,没有可见 token 的块被赋予分数 $-\infty$。Index Branch 随后选出 top-$k$ 个块索引:
$$ {\mathcal{I}}_i^{(r)} \;=\; \mathrm{TopK}_{b \in \{1,\dots,B\}}\!\bigl(M^{\rm idx,(r)}_{i,\cdot},\, k\bigr). \tag{7} $$这里 $\mathrm{TopK}(\cdot, k)$ 返回在 $M^{\rm idx,(r)}_{i,\cdot}$ 下最大的 $k$ 个块的索引。我们总是把包含位置 $i$ 的 local 块包含进来,且 ${\mathcal{I}}_i^{(r)}$ 由 group $r$ 中全部 $G$ 个 query head 共享。
Main Branch。给定 Index Branch 选出的块索引集合 ${\mathcal{I}}_i^{(r)}$,Main Branch 只 attend 被选中块里那些因果可见的 token。对任意 query head $h \in \mathcal{H}_r$,它在这些 token 上应用标准的 scaled dot-product attention,使用与 GQA group $r$ 关联的 key-value head:
$$ {\bm{O}}_{i}^{(h)} \;=\; \mathrm{softmax}\Biggl( \frac{{\bm{Q}}_{i}^{(h)}\, \bigl({\bm{K}}^{(r)}\!\bigl[{\mathcal{I}}_i^{(r)}\bigr]\bigr)^{\top}} {\sqrt{d_h}} \Biggr) {\bm{V}}^{(r)}\!\bigl[{\mathcal{I}}_i^{(r)}\bigr], \tag{8} $$其中 ${\bm{Q}}_{i}^{(h)}$ 表示位置 $i$、query head $h$ 上的 query 向量,${\bm{K}}^{(r)}$ 和 ${\bm{V}}^{(r)}$ 表示第 $r$ 个 GQA group 的 key 与 value 矩阵。记号 ${\bm{K}}^{(r)}[{\mathcal{I}}_i^{(r)}]$ 和 ${\bm{V}}^{(r)}[{\mathcal{I}}_i^{(r)}]$ 表示从被选中的块里收集因果可见的 token。块索引集合 ${\mathcal{I}}_i^{(r)}$ 由 $\mathcal{H}_r$ 中所有 query head 共享,而每个 head 保留自己的 query 投影。由于被选中的块最多包含 $kB_k$ 个因果可见 token,单 query 的 attention 开销从 $O(N)$ 降到 $O(kB_k)$,且不随序列长度增长。
3.2 训练
式 (7) 里的 top-$k$ 选择不可微,因此语言建模 loss 无法直接训练 index 的 Q/K 投影 ${\bm{W}}^{\rm idx}_q, {\bm{W}}^{\rm idx}_k$。我们因此用一个 KL 对齐 loss 来训练 Index Branch,并用三个机制来稳定稀疏训练:Gradient Detach、Indexer Warmup,以及一个强制的 Local Block。下面逐个说明。
KL Loss。KL loss 通过在被选中的 token 上把 Index Branch 的分数匹配到 Main Branch,给 Index Branch 一个直接的学习信号。记 ${\mathcal{I}}_{i,\mathrm{tok}}^{(r)}=(\bigcup_{b\in{\mathcal{I}}_i^{(r)}}{\mathcal{B}}_b)\cap\{1,\dots,i\}$ 为被选块索引所诱导出的因果可见 token,对每个 query 位置 $i$ 和 GQA group $r$,我们在这个 token 索引集合上定义 Index Branch 分布 $P^{\rm idx}$ 与 Main Branch 的 teacher 分布 $P$:
$$ P^{{\rm idx},(r)}_{i,j} = \frac{\exp(S^{{\rm idx},(r)}_{i,j})} {\sum_{u\in{\mathcal{I}}_{i,\mathrm{tok}}^{(r)}}\exp(S^{{\rm idx},(r)}_{i,u})}, \qquad P^{(r)}_{i,j} = \frac{1}{G}\sum_{\ell\in\mathcal{H}_r} \frac{\exp(S^{(\ell)}_{i,j})} {\sum_{u\in{\mathcal{I}}_{i,\mathrm{tok}}^{(r)}}\exp(S^{(\ell)}_{i,u})}, \qquad j\in{\mathcal{I}}_{i,\mathrm{tok}}^{(r)}, \tag{9} $$其中 $S^{{\rm idx},(r)}_{i,j} = ({\bm{Q}}^{\rm idx})^{(r)}_i({\bm{K}}^{\rm idx})_j^{\top}/\sqrt{d_{\rm idx}}$ 是 token 级的 index 分数,$S^{(\ell)}_{i,j} = {\bm{Q}}^{(\ell)}_i({\bm{K}}^{(r)}_j)^{\top}/\sqrt{d_h}$ 是 query head $\ell\in\mathcal{H}_r$ 的 Main Branch 分数。teacher 分布 $P$ 是在概率层面对各 head 的 Main Branch 分布做平均。indexer 随后被训练去匹配 $P$,在所有 query 位置和 GQA group 上取平均:
$$ \mathcal{L}_{\rm KL} = \frac{1}{NH_{kv}} \sum_{i=1}^{N}\sum_{r=1}^{H_{kv}} D_{\mathrm{KL}}\bigl(\mathrm{stopgrad}(P^{(r)}_{i,\cdot}) \,\|\, P^{{\rm idx},(r)}_{i,\cdot}\bigr), \tag{10} $$其中 $N$ 是序列长度,teacher 分布 $P^{(r)}_{i,\cdot}$ 从梯度计算中断开。这个辅助 loss 把 index 分布对齐到 Main Branch 的 attention pattern,使后续的块选择在语义上是有意义的。
Gradient Detach。为了把辅助目标与骨干网络隔离开,我们对 Index Branch 的输入施加 stop-gradient:
$$ {\bm{Q}}^{\rm idx} \;=\; \mathrm{stopgrad}({\bm{X}}){\bm{W}}^{\rm idx}_q, \qquad {\bm{K}}^{\rm idx} \;=\; \mathrm{stopgrad}({\bm{X}}){\bm{W}}^{\rm idx}_k. \tag{11} $$式 (9) 中的 teacher $P$ 是断开的,所以 $\mathcal{L}_{\rm KL}$ 不会碰 Main Branch 的投影;式 (11) 进一步阻止它经由 ${\bm{X}}$ 传到骨干。在这条规则下,$\mathcal{L}_{\rm KL}$ 只更新 ${\bm{W}}^{\rm idx}_q$ 和 ${\bm{W}}^{\rm idx}_k$,使 KL 成为一个干净的、只作用于 indexer 的对齐信号。
Indexer Warmup。我们用一个两阶段的训练调度来初始化 Index Branch,避免早期的随机选择。在最初若干轮迭代里,模型在两个分支上都跑 full attention,并用 $\mathcal{L}_{\rm KL}$ 训练新加入的 index 投影。warmup 结束后,模型切换到 sparse attention,$\mathcal{L}_{\rm KL}$ 改为在 top-$k$ 选中的位置上计算。把预训练好的 full-attention checkpoint 稀疏化时也用同一套调度,这有助于在新加的 index 投影开始控制 Main Branch 路由之前先把它对齐好。
Local Block。对每个 query 位置 $i$ 和 GQA group $r$,包含 $i$ 的那个 local 块在训练和推理时都始终作为 ${\mathcal{I}}_i^{(r)}$ 的一部分被选中。这个固定分配占掉一个块槽位,把剩下的槽位留给 Index Branch 去挑,避免出现漏掉 query 紧邻邻域的退化选择。
完整的层级训练流程总结在算法 1 中。
算法 1:一个 MSA 层:训练前向与辅助 KL loss。该层返回它的输出和这一层的 $\mathcal{L}_{\rm KL}$;模型总 loss $\mathcal{L} = \mathcal{L}_{\rm LM} + \lambda\sum_{\rm layers}\mathcal{L}_{\rm KL}$ 由训练循环组装。
输入:hidden states ${\bm{X}} \in \mathbb{R}^{N \times d_{\rm model}}$;块大小 $B_k$,被选块数 $k$。
- ${\bm{Q}}, {\bm{K}}, {\bm{V}} \leftarrow {\bm{X}}{\bm{W}}_q,\, {\bm{X}}{\bm{W}}_k,\, {\bm{X}}{\bm{W}}_v$ —— 形状 $(N,H_q,d_h),(N,H_{kv},d_h),(N,H_{kv},d_h)$
- ${\bm{Q}}^{\rm idx}, {\bm{K}}^{\rm idx} \leftarrow \mathrm{stopgrad}({\bm{X}}){\bm{W}}^{\rm idx}_q,\, \mathrm{stopgrad}({\bm{X}}){\bm{W}}^{\rm idx}_k$ —— 形状 $(N,H_{kv},d_{\rm idx}),(N,1,d_{\rm idx})$;已断开梯度
- $M^{\rm idx} \leftarrow \mathrm{BlockMaxPool}({\bm{Q}}^{\rm idx}, {\bm{K}}^{\rm idx}, B_k)$ —— 形状 $(N,H_{kv},B)$;逐 group、因果
- ${\mathcal{I}} \leftarrow \mathrm{TopK}(M^{\rm idx},\, k)$ —— 被选块索引;已包含 local 块
- ${\bm{O}} \leftarrow \mathrm{TopKAttn}({\bm{Q}}, {\bm{K}}, {\bm{V}}, {\mathcal{I}})$ —— 形状 $(N,H_q,d_h)$;只 attend 被选中的块
- $\mathrm{output} \leftarrow {\bm{O}}{\bm{W}}_o$ —— 形状 $(N,d_{\rm model})$
- $\mathcal{L}_{\rm KL} \leftarrow \mathrm{KLdiv}({\bm{Q}}^{\rm idx}, {\bm{K}}^{\rm idx},\, \mathrm{stopgrad}({\bm{Q}}), \mathrm{stopgrad}({\bm{K}}),\, {\mathcal{I}})$ —— 在 ${\mathcal{I}}$ 诱导出的 token 上计算
- 返回 $\mathrm{output},\ \mathcal{L}_{\rm KL}$
3.3 计算复杂度
在相同的 $H_q$、$H_{kv}$、$d_h$ 和序列长度 $N$ 下,GQA 与 MSA 的 causal attention FLOPs 为
$$ F_{\rm GQA}(N)= 2 H_q d_h N^2, \qquad F_{\rm MSA}(N)= \underbrace{H_{kv} d_{\rm idx} N^2}_{\text{Index Branch}} + \underbrace{4 H_q d_h Nk B_k}_{\text{Main Branch}} . \tag{12} $$GQA 的主 attention 路径随完整上下文长度增长,而 MSA 用的是固定的选择预算 $kB_k$ 加上一个轻量的 index 计算;因此当 $kB_k \ll N$ 且 $H_{kv}d_{\rm idx} \ll H_qd_h$ 时,FLOPs 的差距随 $N$ 增大。
4 Kernel 设计
本节描述我们稀疏 prefill 实现中所用的 GPU kernel,包括 index TopK kernel、KV-outer 的 sparse attention 前向,以及 sparse KL loss 的反向。
4.1 Index 与 TopK
Exp-free selection。为了高效选出 top-$k$ 个 KV 块,index 模块直接对 index 分数 $s$ 排序。由于 softmax 保序($s_i \le s_j \iff \mathrm{softmax}(s)_i \le \mathrm{softmax}(s)_j$),分数之间的相对顺序不变,top-$k$ 索引也就不变。因此前向绕开了 softmax 的 max/exp/sum 三步,把原始分数直接送去做选择。
Per-thread register top-$k$。块大小 $B_k$ 和选择规模 $k$ 是与 top-$k$ kernel 协同设计的:较大的 $B_k$ 提高 attention 的算术强度(4.2 节),而在这个 $B_k$ 下取较小的 $k$,能让每行的候选块数 $B$ 和 $k$ 都低于通用 top-$k$ kernel 的甜点区——那些 kernel 要么靠大 $B$ 来摊薄多趟分桶的成本(radix selection),要么复杂度是 $O(B \log^2 B)$(bitonic sort)。我们采用 $B_k = 128$、$k = 16$。warp 的 32 个 lane 每个流式处理输入行的 1/32 跨步,并在 shared memory 里维护一个 $k$ 元素的最小堆。堆顶缓存在寄存器里,插入采用延迟写。最后用一轮 $k$ 次的 shuffle merge 把 32 个局部 TopK 结果合并。shared memory 的布局把每个 lane 映射到固定的 bank,避免冲突。
Benchmark。我们在 H800 GPU 上,以 fp32 输入、不排序输出的设定,与 torch.topk 以及 TileLang[Wang 等,2025]的 radix-select top-$k$ 做对比;延迟取 warmup 之后 50 次迭代的中位数。表 1 显示我们这个专用 kernel 在所有测试设定下都是最快的,在部署所用的 $k = 16$ 处增益最大。
表 1:形状为 $(N, B)$ 的 fp32 输入的 top-$k$ 延迟($\mu$s),各行独立处理。部署设定用 $B_k = 128$、$k = 16$,作为参照我们也报告了 $B_k = 64$ 下的 $k = 32$。所有实现产生的索引集合完全一致。
| 序列长度 $N$ | 块数 $B$ | $k$ | torch |
TileLang | 本文 | vs. torch |
vs. TileLang |
|---|---|---|---|---|---|---|---|
| $128$K | $1024$ | $16$ | $3970$ | $2864$ | $779$ | $5.1\times$ | $3.7\times$ |
| $128$K | $2048$ | $32$ | $5378$ | $3630$ | $1991$ | $2.7\times$ | $1.8\times$ |
| $512$K | $4096$ | $16$ | $33810$ | $17779$ | $7880$ | $4.3\times$ | $2.3\times$ |
| $512$K | $8192$ | $32$ | $57659$ | $26100$ | $21326$ | $2.7\times$ | $1.2\times$ |
4.2 Sparse Attention
我们重新审视在 query 与 key/value 长度相等的稀疏 prefill 下,迭代顺序该怎么选。令 $H_q$、$H_{kv}$、$G = H_q / H_{kv}$、$d_h$、$N$、$B_k$、$k$ 分别表示 query head 数、key-value head 数、GQA 比、head 维度、序列长度、KV 块大小,以及每个 query 选中的块数。为简化,下面的 IO 估算假定元素为 2 字节(bfloat16 量级的流量)。我们的 kernel 也支持 fp8;用 fp8 会等比缩放 IO 的绝对量,但不改变 Q-outer 与 KV-outer 迭代之间的比较结论。
把 query 放在外层循环,得到
$$ \begin{aligned} \mathrm{FLOPs} &= 4\, H_q\, N\, d_h\, k\, B_k, \\ \mathrm{IO} &= \underbrace{2 \cdot 2 \cdot H_q\, N\, d_h}_{\text{read}({\bm{Q}})+\text{write}({\bm{O}})} + \underbrace{2 \cdot 2 \cdot H_{kv}\, N\, k\, B_k\, d_h}_{\text{read}({\bm{K}}+{\bm{V}})}, \end{aligned} \tag{13, 14} $$于是 $\mathrm{FLOPs}/\mathrm{IO} \approx G$。
把 KV 块放在外层循环、并把选中了该块的那些 query 收集起来,则需要一个中间输出缓冲:
$$ \begin{aligned} \mathrm{FLOPs} &= 4\, H_q\, N\, d_h\, k\, B_k, \\ \mathrm{IO} &= \underbrace{2 \cdot 2 \cdot H_{kv}\, N\, d_h}_{\text{read}({\bm{K}}+{\bm{V}})} + \underbrace{2 \cdot 2 \cdot H_q\, N\, k\, d_h}_{\text{read}({\bm{Q}})+\text{write}({\bm{O}}_\text{buf})} + \underbrace{2 \cdot H_q\, N\, (k{+}1)\, d_h}_{\text{read}({\bm{O}}_\text{buf})+\text{write}({\bm{O}})}, \end{aligned} \tag{15, 16} $$于是 $\mathrm{FLOPs}/\mathrm{IO} \approx \tfrac{2}{3} B_k$。
由于实践中 $\tfrac{2}{3} B_k \gg G$,我们选择 KV-outer 迭代加 Q gather,以最大化算术强度。kernel 以一个 persistent grid 在 $(\textit{kv\_block}, \textit{kv\_head})$ tile 上执行。对每个 tile,由 TopK 选择结果构造的反向稀疏索引标出相关的 query 位置。这些 query 通过 TMA 拷贝载入 shared memory,每个 query token 一次,由一个 warp 的 32 个 lane 并行派发。
Pre-scheduled tile chunking。直接一个 CTA 对一个 tile 的映射,会被 sink 行主导——某个靠前的 KV 块几乎被每个 query 都选中——而同样的热点模式可以出现在任何热门 KV 块上。因此一个 GPU 调度 kernel 把每个 KV tile 沿它的 query 维切成每块最多约 $2 k B_k$ 个 query 的 chunk,让热 tile 扇出到许多共享同一份 $\mathbf{K}/\mathbf{V}$ 载入的 CTA 上。由于每个 query 的 $k$ 个部分结果现在由 $k$ 个 CTA 产生,调度器还给每个(query, chunk)对预先分配 $\mathbf{O}_\text{buf}$ 中的一个槽位 $s \in [0, k)$——与 query 索引 $i$ 一起打包成一个 32 位 handle——这样 attention kernel 就能把自己的部分结果写到预分配的偏移上,不需要原子操作。combine kernel 读取每个 query 的槽位计数,从而知道要合并多少个部分结果。
Two-phase forward。KV-outer 的切分不允许内联的 softmax 归一化,因为每个 query 的 $k$ 个部分结果是由 $k$ 个不同的 CTA 产生的。因此前向被拆成两个 kernel,中间以 HBM 缓冲隔开:$\mathbf{O}_\text{buf} \in \mathbb{R}^{k \times n \times H_q \times d}$(局部归一化后的部分输出)和 $\mathrm{LSE}_\text{buf} \in \mathbb{R}^{k \times n \times H_q}$(每个部分结果的 logsumexp)。attention kernel 跑上面那个工作列表,把每个部分结果写到它预分配的槽位。combine kernel 读取每个 query 的有效槽位,计算 $a = \max_s \mathrm{LSE}_s$ 和 $\mathrm{LSE}[i, h] = a + \log \sum_s \exp(\mathrm{LSE}_s - a)$,再构造归一化的 split-K 权重 $w_s = \exp(\mathrm{LSE}_s - \mathrm{LSE}[i, h])$。它输出 $\mathbf{O}[i, h] = \sum_s w_s\, \mathbf{O}_\text{buf}[s, i, h]$ 以及最终的 $\mathrm{LSE}[i, h]$。这两个 kernel 使用 Programmatic Dependent Launch 来隐藏 kernel 之间的启动延迟。
Query concatenation。KV-outer 迭代下,每个 KV tile 往往只关联几个到几十个 query 位置。一次只处理一个位置会让 score MMA 填不满:在 $G = 16$ 时,单个 query 位置只贡献 $G$ 个 query head,MMA 的 $M$ 维只有 16。在 Q-outer 迭代下,query 位置无法沿序列维拼接,因为它们通常选中不同的 KV 子集。但在 KV-outer 迭代下,某个 tile 收集到的所有位置共享同一批 KV 操作数。因此 kernel 把 $\lceil 128/G \rceil$ 个 query 位置连同它们各自的 $G$ 个 query head 打包在一起——全都在同一个 KV head 之下——凑成一个 $128 \times 128$ 的 score MMA。
4.3 Sparse KL Loss
LSE fusion。在最初的实现里,我们用一个专门的 kernel 计算 KL 散度的前向,并存下 $\mathrm{LSE}_{\rm main}$ 和 $\mathrm{LSE}_{\rm idx}$ 以便反向传播。但由于 KL loss 只影响反向梯度,我们把这一步优化成:在主 pass 期间直接把这些 LSE 值发到 global memory,从而完全跳过 KL loss 的前向。此外,在 index 分支计算期间,我们保存每个块的 LSE,并在 top-$k$ 个块上做一次 reduction 得到 $\mathrm{LSE}_{\rm idx}$。反向 kernel 随后把这些标量直接载入 softmax,消除了冗余的前向计算。
Dynamic load balancing。在变长序列和数据相关的稀疏性之下,每个 tile 的工作量相差几个数量级。kernel 以 persistent grid 运行,CTA 通过一个全局原子计数器领取工作;每个 tile 沿它收集到的 query 维被切成若干 sub-tile,sub-tile 的数量随该 tile 的 query 数缩放,同时受一个最小 sub-tile 粒度的约束,以摊薄每个 sub-tile 的固定开销。
5 实验
本节报告两个 109B 规模的实验,用来在一个以文本与图像/视频混合数据训练的原生多模态模型上验证 MSA 的最终设计。第一个从零训练一个原生 MSA 模型,记为 MSA-PT。第二个从一个 Full-Attention checkpoint 出发,把 dense attention 换成 MSA 后继续预训练,记为 MSA-CPT。两个模型都与 Full-Attention 基线使用同一架构族,只是把 dense attention 换成了 MSA 层。
5.1 设置
Model Structure。所有模型都使用同一个 41 层 MoE 骨干,总参数约 109B,每 token 激活 6B。前三层是 dense 层,其余 38 层是 MoE 层。模型使用 200K token 的词表,hidden size $d_{\rm model}=3072$。每个 attention 模块使用 MSA,64 个 query head、4 个 KV head、head 维度 128、RoPE 维度 64。每个 MoE 层有 128 个路由专家、1 个共享专家,top-4 路由专家选择。在稀疏训练与评测期间,两个 MSA 模型都使用块大小 $B_k=128$,每个 query、每个 GQA group 保留 $k=16$ 个 key-value 块。
Training Budget。所有模型都在总预算 3T token 下训练。MSA-PT 从零训练:在 40B token 的 indexer warmup 之后,余下的预训练全程保持稀疏训练。MSA-CPT 从一个在 2.6T token 上训练的 GQA Full-Attention checkpoint 出发,我们把 dense attention 换成 MSA 后继续训练 400B token:前 40B token 用于 indexer warmup,之后是稀疏继续预训练。
Evaluations。我们在同一套预训练评测套件上、用同等训练预算下的对应 checkpoint,评测 Full、MSA-PT 和 MSA-CPT。通用推理与问答方面用 MMLU[Hendrycks 等,2021]、MMLU-Pro[Wang 等,2024a]、BBH[Suzgun 等,2022]、GPQA Hard[Rein 等,2023]、ARC Challenge[Clark 等,2018]、TriviaQA[Joshi 等,2017]和 WinoGrande[Sakaguchi 等,2020]。数学与代码方面用 GSM8K[Cobbe 等,2021]、MGSM[Shi 等,2022]、MathVista[Lu 等,2024]、OlymMATH[Sun 等,2025]、HumanEval[Chen 等,2021]、EvalPlus[Liu 等,2023]、BigCodeBench[Zhuo 等,2025]和 MultiPL-E MBPP[Cassano 等,2023]。我们也评测多模态能力:图像 benchmark 包括 AI2D[Kembhavi 等,2016]、ChartQA[Masry 等,2022]、MMMU[Yue 等,2024]、OCRBench v2[Fu 等,2025]、CharXiv[Wang 等,2024b]、VisualWebBench[Liu 等,2024]和 CVBench[Tong 等,2024],视频 benchmark 包括 EgoSchema[Mangalam 等,2023]、LongVideoBench[Wu 等,2024]、MLVU[Zhou 等,2025]、MMVU[Zhao 等,2025b]、VideoMME[Fu 等,2024]和 TemporalBench[Cai 等,2024]。长上下文评测用 RULER[Hsieh 等,2024]和 HELMET[Yen 等,2025]。我们还额外报告了下游 agent 任务上的 perplexity,包括 $\tau^{2}$-bench[Barres 等,2025]、TheAgentCompany[Xu 等,2024]、Humanity’s Last Exam[Phan 等,2025]和 SWE-bench[Jimenez 等,2024]。
5.2 训练动态
图 2 把原生稀疏预训练与对应的 full-attention 运行做了对比。在 3T token 的整个训练过程中,两条 LM loss 曲线几乎无法区分,说明相对 full attention,MSA 没有引入可察觉的优化退化。梯度范数曲线在整个训练过程中也保持在同一区间,说明 MSA 不会导致异常的梯度波动或训练不稳定。这些结果表明,在大规模下训练一个 sparse attention 模型和训练 full-attention 基线一样稳定。
图 3 展示了从一个训练好的 full-attention checkpoint 过渡到稀疏继续预训练的过程。indexer warmup 阶段在稀疏 attention 启用之前,迅速把 KL loss 降下来。切换到稀疏 CPT 之后,KL loss 保持在低位。对每个 query 和 GQA head,令 ${\mathcal{I}}^{\star}$ 为由 Main Branch 分数诱导出的对应 Top-$k$ 块集合,$\widehat{{\mathcal{I}}}$ 为 Index Branch 的选择。block recall 是 $|{\mathcal{I}}^{\star}\cap\widehat{{\mathcal{I}}}|/|{\mathcal{I}}^{\star}|$,而 score recall 是 $\sum_{b\in{\mathcal{I}}^{\star}\cap\widehat{{\mathcal{I}}}}P_b/\sum_{b\in{\mathcal{I}}^{\star}}P_b$,其中 $P_b$ 是块 $b$ 内各 token 的 Main Branch attention 概率之和。block recall 保持在不错的水平,说明重要块被可靠地找回来了。更高的 score recall 进一步说明被检索到的块占据了 Main Branch 大部分的 attention 质量。综合来看,这些动态说明 warmup 提供了一个干净的转换阶段,而 CPT 的 indexer 在稀疏继续预训练期间保持了良好的对齐。
5.3 主要结果
表 2 在一组有代表性的预训练评测上对比了 Full、MSA-PT 和 MSA-CPT。两个稀疏模型都与 Full-Attention 基线大体相当,说明把 dense attention 换成 MSA 不会实质性地损害模型的通用语言、推理、多模态或 agent 导向的 perplexity 画像。两条训练路线体现出不同的强项。MSA-PT 在整个预训练过程中学习稀疏 pattern,在许多数学、图像、视频和长上下文检索 benchmark 上取得最强结果,说明原生稀疏预训练能让模型表示适配到稀疏 attention pattern 上。MSA-CPT 更保守:它保留了 Full-Attention checkpoint 的大部分行为,在多数文本、代码和 PPL 评测上都很接近,当已经有一个训练好的 dense checkpoint 时,它是一条实用的转换路线。剩下的差距是随 benchmark 而异的,并没有集中在某一个能力面上。
表 2:3T token 训练预算下有代表性的评测结果。Full 表示 Full-Attention 基线,MSA-PT 表示从零稀疏预训练,MSA-CPT 表示稀疏继续预训练。每行最优结果加粗;PPL 越低越好,其余越高越好。
| 分组 | Benchmark | Full | MSA-PT | MSA-CPT |
|---|---|---|---|---|
| 通用 | MMLU | 67.0 | 67.2 | 66.8 |
| 通用 | MMLU-Pro | 38.5 | 38.8 | 39.1 |
| 通用 | BBH | 67.7 | 66.6 | 66.1 |
| 通用 | GPQA Hard | 25.9 | 26.3 | 26.3 |
| 通用 | ARC Challenge | 82.7 | 82.5 | 82.9 |
| 通用 | TriviaQA | 66.0 | 65.5 | 67.7 |
| 通用 | WinoGrande | 58.3 | 60.9 | 62.0 |
| 数学 | GSM8K | 76.2 | 77.7 | 73.7 |
| 数学 | MGSM | 44.1 | 46.0 | 44.2 |
| 数学 | MathVista | 43.8 | 46.8 | 44.5 |
| 数学 | OlymMATH Easy P@100 | 23.0 | 26.0 | 22.0 |
| 代码 | HumanEval | 61.0 | 64.0 | 57.9 |
| 代码 | EvalPlus | 59.4 | 61.8 | 60.0 |
| 代码 | BigCodeBench | 44.8 | 44.0 | 45.7 |
| 代码 | MultiPL-E MBPP P@10 | 82.1 | 81.6 | 81.1 |
| 检索 | RULER-8K | 79.8 | 84.2 | 77.2 |
| 检索 | RULER-32K | 75.0 | 77.5 | 75.7 |
| 图像 | AI2D | 68.3 | 70.6 | 67.3 |
| 图像 | ChartQA | 75.0 | 75.4 | 71.4 |
| 图像 | MMMU | 46.8 | 45.9 | 44.5 |
| 图像 | OCRBench v2 | 55.0 | 55.7 | 54.3 |
| 图像 | CharXiv | 37.55 | 41.55 | 37.15 |
| 图像 | VisualWebBench | 55.6 | 68.4 | 59.4 |
| 图像 | CVBench | 57.0 | 59.7 | 58.8 |
| 视频 | EgoSchema | 29.6 | 37.6 | 25.8 |
| 视频 | LongVideoBench | 38.5 | 41.8 | 38.9 |
| 视频 | MLVU | 44.14 | 46.94 | 43.68 |
| 视频 | MMVU | 45.8 | 47.5 | 45.8 |
| 视频 | VideoMME | 41.11 | 45.48 | 39.65 |
| 视频 | TemporalBench | 49.4 | 53.4 | 50.6 |
| PPL $\downarrow$ | TAU2 | 1.155 | 1.148 | 1.150 |
| PPL $\downarrow$ | AgentCompany | 1.248 | 1.249 | 1.247 |
| PPL $\downarrow$ | HLE | 1.275 | 1.278 | 1.275 |
| PPL $\downarrow$ | SWE | 1.216 | 1.218 | 1.216 |
为了评估 MSA 在长上下文扩展之后是否仍然有效,我们在 MSA-CPT 模型上又做了一个扩展实验。从稀疏继续预训练的 checkpoint 出发,我们跑了约 140B token 的长上下文训练,然后在 HELMET 和 RULER 上评测。结果报告在表 3 中。经过扩展阶段后,MSA-CPT 仍然接近 Full-Attention 基线。由于每个 query 和 GQA group 仍然只 attend $kB_k=16\times 128=2{,}048$ 个 key-value token,这些结果说明 MSA 能在极紧的 attention 预算下保住长上下文能力。
表 3:MSA-CPT 在 HELMET 和 RULER 上的长上下文扩展结果。$\Delta$ 报告 MSA-CPT 与 Full-Attention 基线之差。「Overall」分数是各细粒度子任务的平均。所有指标越高越好。
| Benchmark | 子集 | Full | MSA-CPT | $\Delta$ |
|---|---|---|---|---|
| HELMET-128K | Overall | 46.53 | 45.93 | -0.60 |
| HELMET-128K | ICL | 70.40 | 72.80 | +2.40 |
| HELMET-128K | Rerank/RAG | 34.60 | 32.50 | -2.10 |
| RULER-128K | Overall | 72.00 | 72.12 | +0.12 |
| RULER-128K | CWE/FWE | 46.35 | 45.00 | -1.35 |
| RULER-128K | MK/MQ/MV | 96.63 | 98.87 | +2.24 |
| RULER-128K | QA1/QA2 | 47.80 | 46.80 | -1.00 |
| RULER-128K | VT | 97.80 | 96.80 | -1.00 |
支撑这些设计选择的额外消融实验放在附录里。具体来说,附录 B 研究了 Index Branch 的训练配方,包括梯度来源、KL 梯度的截断、warmup,以及与 sliding-window 稀疏基线的对比。附录 C 进一步考察了架构选择,例如块大小、强制 sink、局部选择,以及 Index Branch 的 value head。这些消融为主实验所用的最终 MSA 设计提供了经验依据。
5.4 效率
我们把 3.3 节的复杂度分析实例化到实验模型的配置上,同时报告理论上的 attention FLOPs 降低与实测的运行时加速。dense GQA 与 MSA 使用相同的 query head 数、key-value head 数、head 维度和上下文长度;唯一的区别是 dense GQA attend 完整上下文,而 MSA 先做索引选择、再在固定的 KV 预算上做 sparse attention。在我们的设定里,MSA 用 $B_k=128$、$k=16$,对应每个 query 选中 $kB_k=2{,}048$ 个 token 的预算。
如图 4 所示,在我们的设定下 MSA 相对 GQA 大幅降低了单 token attention FLOPs,且上下文越长降幅越大。在 $1\mathrm{M}$ token 处,同样的 head 配置下 FLOPs 降低达到 28.4 倍。实测的运行时加速遵循同样的 scaling 趋势,但预期不会与 FLOPs 降幅精确一致。sparse attention 引入了索引构造、top-$k$ 选择、反向索引物化、query 收集和负载均衡等开销,而且它的访存 pattern 不如 dense attention 规整。因此运行时加速小于理论 FLOPs 降幅,但它随上下文长度增长——dense 基线继续随完整序列长度扩张,而 MSA 把主 attention 预算固定住了。
6 相关工作
长上下文效率催生了大量高效 attention 的工作,大体可以分成两个方向:把 dense softmax attention 换成更便宜的线性或递归替代品,以及保留 softmax attention 但限制它的感受野。linear attention[Katharopoulos 等,2020;Choromanski 等,2021]把 softmax 核换成线性复杂度的替代物,而像 Mamba[Gu 和 Dao,2023]这样的状态空间模型则把 attention 换成在隐状态上的选择性递归。混合堆叠[MiniMax,2025a、b]把线性块与 full-attention 块交错排列,减少平方级层数的同时保留一部分精确 softmax 的容量。固定 pattern 的 attention 保留 softmax attention 但施加一个预定义的支撑集,包括局部窗口、全局 token[Beltagy 等,2020;Zaheer 等,2020],以及带 sliding window 的 attention sink[Xiao 等,2024b]。这些方法降低长上下文成本的方式,要么是部分或全部替换 dense attention,要么是使用一个与内容无关的 attention pattern。
在固定稀疏 pattern 之外,自适应 sparse attention 让被 attend 的支撑集取决于输入。已有方法的主要差别在于这个支撑集是何时构造的,以及选择器是否作为模型的一部分被训练。推理期稀疏化作用在一个预训练好的 Full-Attention 骨干上,只在服务期间构造稀疏支撑集。H2O[Zhang 等,2023]和 SnapKV[Li 等,2024]在解码期用累积的 attention 统计量裁剪 KV cache,Quest[Tang 等,2024]为每个 query 做 page 级的重要性估计,MInference[Jiang 等,2024]和 FlexPrefill[Lai 等,2025]在 prefill 时按 head 派发稀疏 kernel,InfLLM[Xiao 等,2024a]维护 attention sink、一个局部上下文窗口和可检索的 chunk。这些方法继承了 Full Attention 的训练成本,并且至少留下一个推理阶段仍处于接近 Full-Attention 的速度。原生训练的 sparse attention 设计在预训练期间就训练 indexer,是与 MSA 最接近的先前工作。NSA[Yuan 等,2025]面向 MQA/MHA 骨干,用三条并行分支:压缩 attention 提供粗粒度全局上下文、选择 attention 覆盖细粒度块、sliding window 负责局部上下文。InfLLM-V2[Zhao 等,2025a]通过把无参数的块选择与局部 sliding window 统一起来,实现零样本的 dense 到 sparse 切换。MoBA[Lu 等,2025]同样建立在 GQA 上,但使用非常大的 KV 块、以块平均后的 key 来打分,并且只通过语言建模梯度训练它的 indexer。DSA[DeepSeek-AI 等,2025]坐在 MLA 的 MQA 模式之上:一个多 head 的 ReLU-based lightning indexer 逐 token 打分,所有 query head 共享单一的 Top-$k$ 索引,选择是 token 级的。MSA 与这一邻域的区别在于两条被同时采纳的轴:per-GQA-group 的 Top-$k$ 共享结合块级选择,这带来了多 group、块粒度的检索,同时让 KV 读取保持连续。
高效 kernel 是 sparse attention 把理论 FLOP 降幅转化为墙钟加速的关键。FlashAttention[Dao 等,2022]和 FlashAttention-2[Dao,2024]引入了 IO 感知的分块 softmax attention,FlashDecoding[Dao 等,2023]把这套做法扩展到访存受限的解码。像 Flash-Sparse-Attention[Yan 等,2025]和 FlashMoBA[Xiao 等,2025]这样的开源块稀疏 kernel,则让这套递推的块稀疏变体变得可用。MSA 的 kernel(见第 4 节)复用了 FlashAttention 的算法骨架,但把循环顺序调整成适配 MSA 产生的那种 GQA 原生、块粒度的访存 pattern。
7 结论
我们提出了 MSA,一种与 Grouped-Query Attention 协同设计的 sparse attention 机制。这个架构在标准 GQA 层上挂一个轻量的 Index Branch:每个 GQA group 通过一个块级点积 indexer 独立选出一小组 key-value 块,Main Branch 则在被选中的块上做受限的 softmax attention。Index Branch 是一个纯粹的选择器,用一个针对 Main Branch 的 KL 对齐 loss 来训练,配合两阶段 warmup 调度,以及在 index 输入上的 stop-gradient——后者把辅助 loss 限制在 index 投影内部。在 109B-MoE 规模上,MSA 在绝大多数预训练与 agentic benchmark 上保住了 GQA Full-Attention 基线的能力,同时把 $1\mathrm{M}$ 上下文下的单 token attention 计算降低 28.4 倍,而这正是长上下文推理成为部署硬约束的区间。
Outlook。MSA 的几个核心决策——per-GQA-group 的独立选择、块级粒度、用 KL 对齐目标训练的 indexer——与当前多数开源前沿模型共用的 GQA 骨干是相容的,所以这套配方应该能几乎不加修改地迁移过去。两个方向是自然的下一步:一是补上长上下文检索上残留的那点差距,途径可以是更长的稀疏训练、推理时更大的选择预算,或者更丰富的 indexer 打分函数;二是把这套只有选择器的设计推广到预训练之外的场景,包括强化学习后训练和 agentic 部署——在那些场景里,长上下文成本是主导性的运营约束。
附录 A 可视化
为了更好地理解学出来的 indexer 到底选了什么,我们在图 5 中可视化了所有 query 块与 key 块配对上的 per-head Index Branch 选择概率。我们展示了来自一个靠前层(Layer 1)和一个靠后层(Layer 18)的四个 head,对应四个不同的 GQA group。跨层来看,学出来的稀疏 pattern 复现了 dense attention 中预期的主要结构:所有 head 都把高概率放在局部对角线上、一致地选中 sink 列,并把剩余预算留给少数几个长程相对位置。同时,非局部的选择在各 GQA group 之间并不相同。不同 group attend 不同的长程条带,同时共享共同的局部与 sink pattern,说明学出来的 indexer 捕捉到了 group 特有的稀疏 attention pattern,而不是塌缩成单一的全局选择 pattern。
我们进一步考察 MSA 模型中的 attention sink 现象。即使不显式强制 indexer 选择第一个 key-value 块,我们也观察到学出来的 Index Branch 在所有层和所有 head 上都自然地给最初那个块分配了很高的选择概率。图 6 展示了两个代表性层(Layer 4 和 Layer 24)的结果,每层采样八个 head。在两个层上,每个 head 都把相当一部分 attention 质量指向第一个 token。这证实了即使在我们的 sparse attention 机制里,attention 焦点也会自然涌现,并且在不同 head 和不同层上普遍存在。
附录 B 前期实验
本节给出在一个 pilot 模型上做的小规模消融研究。我们的目标是找出对稳定优化和强下游表现真正必要的那些训练设计选择。这些结果构成了第 3 节所述最终配方的经验依据。
B.1 设置
本节所有消融都使用一个 10B 参数的 pilot Transformer,架构族与正文的 MSA 模型相同,但只有 16 层。模型使用 200K token 的词表,hidden size $d_{\rm model}=2048$。每个 attention 模块使用 GQA,32 个 query head、4 个 KV head、head 维度 128、RoPE 维度 64。MoE 包含 64 个专家、top-4 专家路由,专家内部维度 1536。模型总参数 10.53B,每 token 激活 1.47B。优化器、学习率调度和 tokenizer 都与全规模配置一致。每次运行都在与全规模相同的预训练语料的一个子集上训练。
B.2 Index Branch 的梯度来源
训练 Index Branch 的一个核心难点是式 (7) 里的 top-$k$ 选择不可微。在朴素的 sparse attention 前向下,被选中的块索引只用作一个离散的路由决策。于是 index 投影 ${\bm{W}}^{\rm idx}_q$ 和 ${\bm{W}}^{\rm idx}_k$ 从语言建模目标那里拿不到有用的梯度,indexer 也就学不到该选哪些块。给 indexer 引入训练信号有若干可能的方式。我们考察两种既保留 sparse attention 结构、又能给 Index Branch 提供梯度的机制。
Index Branch output。第一种机制让 Index Branch 额外贡献一路 attention 输出。具体地,我们给 Index Branch 挂一个 value 投影,计算 ${\bm{O}}^{\rm idx}=\mathrm{Attn}({\bm{Q}}^{\rm idx},{\bm{K}}^{\rm idx},{\bm{V}}^{\rm idx})\in\mathbb{R}^{N\times H_q\times d_h}$。这一路输出经由一个单独的输出投影加到该层输出上,${\bm{O}}'={\bm{W}}_o{\bm{O}}+{\bm{W}}^{\rm idx}_o{\bm{O}}^{\rm idx}$。这个设计通过 Index Branch 对下一 token 预测的贡献来训练它。
KL loss。第二种机制直接监督 Index Branch,把它的选择分布匹配到被选支撑集上的 Main Branch。我们使用式 (10) 定义的辅助 loss $\mathcal{L}_{\rm KL}$。这个 loss 作用在 ${\bm{W}}^{\rm idx}_q$ 和 ${\bm{W}}^{\rm idx}_k$ 上,为 index 选择提供显式的训练信号。
为了分离这两种梯度来源的效果,我们用三种配置从零训练模型,且从第一步就使用 sparse attention:
- LM Loss only:Index Branch 的输出加到层输出上,模型只用语言建模 loss 训练,
- KL Loss only:丢弃 Index Branch 的输出,indexer 只通过辅助 KL loss 训练,
- LM Loss + KL Loss:两种机制都启用,
图 7 报告了每种配置相对于在同一数据上训练的 Full-Attention GQA 基线的逐 benchmark 差值。两个单信号配置表现出互补的弱点。LM Loss only 保住了短上下文能力,但在长上下文检索上表现很差:没有一个直接作用在 top-$k$ 选择本身上的目标,indexer 就几乎感受不到去选相关块的直接压力。KL Loss only 改善了检索,但削弱了短上下文能力:把 ${\bm{O}}_{\rm idx}$ 从层输出里去掉,减少了语言模型可用的 attention 容量。LM Loss + KL Loss 在这两个轴上取得了最好的平衡,也是本节其余消融所采用的配置。
基于这些结果,我们在本节其余消融中采用 LM Loss + KL Loss 配置。我们后面会在 C.3 节说明,一旦在全规模设定下用上 B.4 节引入的 indexer warmup,Index Branch 的输出就不再必要了。因此最终配方保留 KL 监督,但去掉了 Index Branch 的 value head 及其加性输出通路。
B.3 把 KL 梯度限制在 Index Branch 内
辅助 KL loss 的本意是训练 Index Branch 去匹配 Main Branch 的选择分布。在默认的 autograd 图下,KL 梯度会穿过 Index Branch 的 query 和 key 投影回到 hidden state,再经由残差流进入骨干。这时 KL loss 就变成了作用在骨干上的一个额外目标,而不是给 indexer 的局部监督信号。
我们从这种梯度路由中观察到两种失效模式。KL 系数较大时,偶发的 KL 梯度尖峰会传播进骨干,在几百步内造成梯度范数尖峰和 LM loss 发散(图 8)。即使在稳定的系数下,标准短上下文 benchmark 也会在训练过程中逐渐回退(图 9)。我们把这种回退归因于一种自蒸馏效应:骨干可以通过简化 Main Branch 的 attention 分布来降低 KL loss,而不是通过改进 Index Branch。
我们通过在 Index Branch 的输入处截断 KL 梯度来解决这两种失效模式(3.2 节)。这样每一层的 KL loss 就成了只作用于它自己那个 indexer 的局部监督信号。有了这个 detach,在不 detach 时会导致发散的同样 KL 系数下,LM loss 和梯度范数都保持稳定(图 8),短上下文的回退也消失了(图 9)。我们在后续所有运行中都使用这个 detach。
B.4 Indexer Warmup
我们观察到 Main Branch 的 attention 分布在训练最早期变化非常快。如图 10 所示,attention 熵从初始的平滑分布迅速跌到一个尖锐得多的分布,然后才进入表示学习的较慢阶段。这让初始化阶段的稀疏选择变得脆弱。如果从第 0 步就启用 top-$k$ 选择,Index Branch 必须去追一个快速移动的目标,而它自己的选择还几乎是随机的。糟糕的早期选择会把 Main Branch 路由到没有信息量的 token 上,这同时削弱了骨干的学习和 indexer 收到的 KL 监督。
我们用一个简短的 indexer warmup 来解决这个问题。warmup 期间,Main Branch 使用 full attention,而 Index Branch 由针对全序列 Main Branch 分布的 KL loss 来训练。这让骨干能在没有稀疏路由错误的情况下走过早期的锐化阶段,同时在 indexer 开始控制 token 选择之前给它一个有意义的初始化。$T_{\rm warm}$ 步之后,我们启用 top-$k$ 稀疏选择,并继续用限制在被选支撑集上的 KL loss 训练。
图 11 对比了有无这个 warmup 的预训练运行。warmup 过的那次在短上下文表现和长上下文检索上都更好。这些结果说明一段简短的 full-attention warmup 为稀疏训练提供了更好的初始化。因此在把 Full-Attention checkpoint 通过继续预训练转成 sparse attention 时,我们也采用这个 warmup。
B.5 可学习的 attention sink
图 6 的可视化显示,第一个 token 常常充当 attention sink:许多 head 会把不小的 attention 质量分给序列前缀,即使稀疏选择器并未被显式强制包含它。这就引出一个问题:这种 sink 行为是否应该由一个显式的可学习机制来表示,而不是被序列里第一个真实 token 吸收掉。我们因此测试了一种 GPT-OSS 风格的可学习 attention sink。具体地,每个 attention head 被赋予一个额外的可学习 sink logit,它在 attention softmax 里与正常的 key 位置竞争。
图 12 可视化了由此产生的 attention pattern。可学习 sink 在一些 head 上吸收了相当多的 attention 质量,但它并没有完全消除原来的第一 token sink。在若干 head 上,尤其是那些学出来的 sink 拿到质量很少的 head,第一个 token 仍然收到大量 attention,继续充当隐式 sink。
我们也在图 13 中对比了有无可学习 sink 的下游 perplexity。可学习 sink 变体相对默认设计没有带来清晰或一致的改善。考虑到它额外的参数、实现复杂度,以及它并不能完全抑制第一 token 的 sink 行为,我们没有把可学习 attention sink 放进最终配方。
B.6 动态稀疏选择 vs. 滑动窗口
为了评估动态选择的价值,我们把 MSA 与一个 FLOP 对齐的 sliding-window 基线做对比。这个基线去掉 Index Branch,改用一个固定的稀疏 pattern:每个 query attend 第一个 key 块,以及一个以该 query 结尾、token 预算相同的局部窗口。因此两种方法的选择预算相同,唯一的区别在于被选中的 token 是由位置固定下来的,还是动态挑出来的。
图 14 报告了下游 agent 任务上的 perplexity。在相同的稀疏选择预算下,sliding-window 模型在整条训练轨迹上的 perplexity 都高于 MSA。虽然两个模型都从更多训练 token 中获益,但固定的局部窗口 pattern 达不到动态稀疏选择的 perplexity。这说明对这些 agent 任务而言,位置固定的稀疏 pattern 不如内容相关的 token 选择合适。
附录 C 补充消融实验
C.1 块大小
MSA 的 Main Branch 里的 sparse attention 计算以连续 $B_k$ 个 token 的块为单位处理 key-value 对,这既影响模型表现也影响效率。更大的块能提升 kernel 效率,但可能因为选择粒度变粗而降低检索质量。我们在保持被选 token 总数不变的前提下调整 $B_k$,来考察这个权衡。相比主实验,这些运行使用了更少的训练迭代和评测套件的一个子集。
如表 4 所示,在这个设定下改变块大小对模型质量的影响有限。不同 $B_k$ 下的 PPL 结果几乎不变,把块大小从 32 增加到 64 或 128 时,RULER 分数也没有明显退化。这说明 MSA 可以在这些消融中以有限的质量损失,用更大的 key-value 块来提升 kernel 效率。
表 4:不同 key-value 块大小下的 perplexity 与长上下文检索分数。perplexity 越低越好,RULER 分数越高越好。
| Benchmark | Block 32 | Block 64 | Block 128 |
|---|---|---|---|
| PPL $\downarrow$ | |||
| TAU2 | 1.176 | 1.176 | 1.176 |
| AgentCompany | 1.266 | 1.276 | 1.266 |
| HLE | 1.299 | 1.299 | 1.300 |
| SWE | 1.233 | 1.233 | 1.233 |
| 长上下文检索 | |||
| RULER-8K | 72.5 | 72.8 | 73.8 |
| RULER-32K | 66.1 | 65.3 | 64.6 |
C.2 强制 sink 与局部选择
在早期的稀疏训练实验里,我们显式强制选择器包含两类块:序列里的第一个块,以及围绕 query 位置的一个固定局部窗口。第一个块对应常见的 attention sink pattern,而局部窗口保留了对短程建模很重要的邻近上下文,并给 indexer 提供了稠密的监督。这个设计主要是作为一种稳定化机制引入的:在 indexer 变得可靠之前,强制这些块能降低稀疏分支在早期训练中错过基本上下文的概率。
我们后来发现这些先验并不需要硬编码。当去掉对第一个块和固定局部窗口的强制选择后,训练出来的模型依然表现出这两种结构:attention 在有用的时候会集中到序列前缀,邻近 token 也依然被频繁选中。如表 5 所示,去掉强制 sink 和固定局部选择对标准模型质量影响很小:推理、代码和 PPL 指标几乎不变,长上下文检索也相当。这些结果说明稀疏模型能够在没有硬编码选择规则的情况下学到 sink 和局部选择 pattern。因此最终配方不强制第一个块或一个大的局部窗口,只强制那个特殊的、不完整的自身块。
表 5:强制 sink 与局部窗口选择的消融。除标了 $\downarrow$ 的以外,越高越好。
| Benchmark | No Forced | Forced |
|---|---|---|
| 通用知识与推理 | ||
| MMLU | 60.5 | 60.5 |
| MMLU-Pro | 32.5 | 33.4 |
| BBH | 58.2 | 58.2 |
| ARC Challenge | 78.1 | 77.9 |
| 数学 | ||
| GSM8K | 66.0 | 66.9 |
| MGSM | 35.8 | 36.3 |
| 代码 | ||
| EvalPlus | 54.0 | 53.6 |
| BigCodeBench | 35.6 | 35.7 |
| MultiPL-E MBPP P@10 | 80.1 | 79.5 |
| 图像 | ||
| ChartQA | 73.5 | 73.7 |
| MMMU | 43.6 | 42.9 |
| 视频 | ||
| VideoMMMU | 32.1 | 32.0 |
| PPL $\downarrow$ | ||
| TAU2 | 1.175 | 1.175 |
| AgentCompany | 1.268 | 1.266 |
| HLE | 1.301 | 1.300 |
| SWE | 1.235 | 1.233 |
| 长上下文检索 | ||
| RULER-8K | 71.6 | 71.7 |
| RULER-32K | 61.5 | 65.8 |
表 6:Index Branch value head 的继续预训练消融。
| Benchmark | With-value | No-value |
|---|---|---|
| 通用知识与推理 | ||
| MMLU | 66.4 | 67.3 |
| MMLU-Pro | 39.0 | 39.1 |
| BBH | 65.3 | 65.9 |
| ARC Challenge | 82.2 | 82.4 |
| 数学 | ||
| GSM8K | 77.6 | 76.4 |
| MathVista | 45.2 | 43.6 |
| MGSM | 48.4 | 47.6 |
| 代码 | ||
| HumanEval | 60.4 | 59.1 |
| EvalPlus | 57.7 | 58.7 |
| BigCodeBench | 46.0 | 44.0 |
| 图像 | ||
| AI2D | 69.3 | 70.4 |
| ChartQA | 75.3 | 74.9 |
| MMMU | 44.9 | 43.4 |
| OCRBench v2 | 53.2 | 53.9 |
| 视频 | ||
| MLVU | 42.4 | 43.9 |
| MMVU | 44.9 | 43.7 |
| PerceptionTest | 45.0 | 47.3 |
| 长上下文检索 | ||
| RULER-8K | 84.1 | 83.0 |
| RULER-32K | 79.7 | 80.4 |
C.3 Index Branch 的 value head
我们的前期实验(B.2 节)表明,让 Index Branch 额外提供一路 attention 输出,有助于模型从第 0 步就开始稀疏训练。但这个 index value head 引入了额外的计算和复杂度。既然 B.4 节的 indexer warmup 已经改善了稀疏训练的初始化,我们进一步消融这个 value head 是否仍然必要。
我们把原来的 with-value 设计与一个只用 KL 对齐信号训练 indexer 的 no-value 变体做对比。如表 6 所示,去掉 index value head 并没有在评测套件上导致系统性的退化。no-value 变体在一些通用推理 benchmark 上略好,而 with-value 变体在一些数学和代码任务上保有小幅优势。在多模态 benchmark 和长上下文检索上,差异同样是混合的。
总体来看,结果表明一旦用上 Index Branch warmup,index value head 就不是关键的了。它对下游质量的影响很小且随 benchmark 而异,两个变体都没有一致地压过对方。这说明在早先那套配方里 ${\bm{O}}_{\rm idx}$ 的主要作用是提供一个额外的早期训练信号,而不是在收敛时提供必要的容量。因此最终设计出于效率考虑去掉了 index value head。在推理时,top-$k$ indexer 只需要 ${\bm{Q}}_{\rm idx}{\bm{K}}_{\rm idx}^{\top}$ 的块级最大值,完全避开了 value 聚合通路和指数运算。
-
MiniMax Sparse Attention, https://arxiv.org/abs/2606.13392 ↩︎
-
MSA 推理 kernel, https://github.com/MiniMax-AI/MSA ↩︎
-
MiniMax-M3, https://huggingface.co/MiniMaxAI/MiniMax-M3 ↩︎
Author houmin
Publish July 30, 2026
LastMod July 31, 2026
License 本作品采用 CC BY-NC-ND 4.0 许可协议进行许可,转载时请注明原文链接
如果你在浏览博客的过程中发现了任何问题,欢迎在对应文章下评论。如果你有其他事情想要咨询,可以通过邮件联系我。