MiniMax 发布了 MSA(MiniMax 稀疏注意力),这是一种直接构建在分组查询注意力(GQA)之上的稀疏注意力方法。它针对一个瓶颈:长上下文下 softmax 注意力的二次方成本。MiniMax 研究团队在一个使用原生多模态数据训练的 109B 参数混合专家模型中对其进行了测试。他们还开源了一个推理内核,并发布了一个生产模型 MiniMax-M3。
什么是 MSA(MiniMax 稀疏注意力)
MSA(MiniMax 稀疏注意力)将注意力分解为两个阶段:一个索引分支和一个主分支。索引分支决定每个查询应该读取哪些键值块。然后,主分支仅对这些块执行精确的 softmax 注意力。
选择是在块粒度上进行的,而不是按每个 token。默认块大小为 Bk = 128 个 token。每个查询和 GQA 组保留 k = 16 个块。这将每个查询的预算固定为 kBk = 2,048 个键值 token。
这两种成本结构不同。密集 GQA 注意力的每个查询成本随 O(N) 缩放,即整个上下文。MSA 的成本随 O(kBk) 缩放,随着 N 的增长而保持固定。因此,随着上下文长度的增加,计算差距会扩大。
选择在每个 GQA 组内共享,但在组之间独立。一个键值头服务于多个查询头,它们共享一个块集合。不同的组可以关注不同的长距离区域。
两个分支如何工作
索引分支仅向标准 GQA 层添加两个投影矩阵。它为每个 GQA 组定义一个索引查询头和一个共享的索引键头。它对可见的键 token 进行评分,然后将这些评分通过最大池化聚合到块级别。
然后,一个 Top-k 算子为每个查询和组选择得分最高的块。包含查询的本地块始终被包含在内。这可以防止选择器丢弃查询的紧邻区域。
主分支从选定的块中收集因果可见的 token。它应用仅限于这些 token 的缩放点积 softmax 注意力。每个查询头保留自己的查询投影,但共享该组的块集合。
报告中的可视化展示了学习型索引器选择的内容。注意力头集中在局部对角线和第一个块上。它们将剩余预算保留给少数长距离条带。


MSA 的训练方式
Top-k 选择是不可微的,因此语言建模损失无法训练索引投影。MSA 通过 KL 对齐损失解决了这个问题。该损失使索引分支的分布与主分支的注意力模式相匹配。教师分布是主分支在所选 token 上的分组平均分布。
三种机制稳定了稀疏训练。梯度分离(Gradient Detach)对索引分支的输入应用了停止梯度。这使 KL 损失局限于索引投影,而非主干网络。如果没有它,较大的 KL 系数会导致梯度尖峰和损失发散。
索引器预热(Indexer Warmup)在最初的迭代中让两个分支都运行完整注意力。索引器在控制路由之前先从 KL 损失中学习。强制局部块(Forced Local Block)为邻近上下文预留了一个槽位。
消融实验塑造了最终的方案。早期的一个变体为索引分支添加了一个带有自身输出的值头。一旦使用了预热,该值头就不再必要。最终设计出于效率考虑将其移除。
MSA 支持两种训练路径。MSA-PT 在 400 亿 token 的索引器预热后从头开始训练。MSA-CPT 则转换一个在 2.6 万亿 token 上训练好的密集 GQA 检查点,然后继续训练 4000 亿 token,其中包括 400 亿 token 的预热。
内核协同设计
理论上的稀疏性如果没有匹配的 GPU 路径,就无法转化为速度提升。MSA 将该算法与两个内核设计思路相结合。
第一个是无指数(exp-free)的 Top-k 选择。Softmax 保留了顺序,因此对原始分数进行排序会产生相同的索引。该内核在选择之前跳过了 max、exp 和 sum 步骤。在上下文长度为 128K、k = 16 的情况下,它比 torch.topk 快 5.1 倍。它也比 TileLang 的基数选择内核快 3.7 倍。
第二种是带查询收集的 KV 外部稀疏注意力。遍历 KV 块相比遍历查询会提高算术强度。该内核将 ⌈128/G⌉ 个查询位置打包成一个 128×128 的分数 MMA。两阶段前向传播将注意力计算和合并步骤分散到多个 CTA 中。
开源内核 fmha_sm100 面向 NVIDIA SM100 GPU。它提供密集 FlashAttention 以及稀疏 Top-k 内核,采用 MIT 许可证发布。支持 BF16、FP8、NVFP4 和 FP4 精度。
MSA 与其他稀疏方法的比较
研究团队将 MSA 与四种原生训练的稀疏设计进行了对比。
下表总结了所描述的差异。
| 方法 | 骨干架构 | 选择粒度 | 索引器 / 选择信号 |
|---|---|---|---|
| MSA | GQA | 块级(B_k = 128),每 GQA 组的 Top-k | KL 对齐损失 |
| NSA | MQA / MHA | 压缩 + 选定块 + 滑动窗口 | 原生(端到端)训练 |
| InfLLM-V2 | 密集↔稀疏可切换 | 无参数块选择 + 滑动窗口 | 无参数(无训练索引器) |
| MoBA | GQA | 非常大的 KV 块(块平均键) | 仅 LM 梯度 |
| DSA | MLA(MQA 模式) | Token 级;跨头共享单个 Top-k | ReLU 闪电索引器 |
MSA 的独特组合是每 GQA 组的 Top-k 共享与块级选择相结合。这使得 KV 读取保持连续,同时为每个组提供独立的检索。
质量方面表现良好。两种稀疏模型总体上与全注意力基线保持竞争力。
下表展示了在 3T token 预算下的代表性结果。
| 基准测试 | 全注意力 | MSA-PT | MSA-CPT |
|---|---|---|---|
| MMLU | 67.0 | 67.2 | 66.8 |
| GSM8K | 76.2 | 77.7 | 73.7 |
| HumanEval | 61.0 | 64.0 | 57.9 |
| RULER-8K | 79.8 | 84.2 | 77.2 |
| RULER-32K | 75.0 | 77.5 | 75.7 |
| VideoMME | 41.11 | 45.48 | 39.65 |
经过长上下文扩展后,MSA-CPT 在 HELMET-128K 和 RULER-128K 上仍与全注意力保持接近。每个查询仍然只关注 2,048 个键值 token。
解释器演示
使用场景与示例
MSA 针对上下文长度是部署关键约束的工作负载。
- 长周期智能体:一个跨越数百个推理和行动步骤的智能体会积累大量记录。对此历史记录进行密集注意力计算会呈二次方增长。MSA 将每个查询的预算保持在 2,048 个 token,与长度无关。
- 仓库级代码推理:加载完整仓库的编码智能体可能超过数十万 token。索引器将每个查询路由到少数相关代码块,无关文件则被排除在选定集合之外。
- 持久化记忆:长期运行的助手会不断积累对话状态。MSA 每次查询仅读取固定大小的最相关块切片,随着记忆增长,解码成本基本保持平稳。
- 长视频理解:该模型原生支持多模态,并基于图像和视频数据训练。MSA-PT 在多个视频基准测试(包括 VideoMME 和 TemporalBench)中,在三轮运行中得分最高。稀疏选择机制可扩展到长视觉 token 序列。
运行内核
最快路径是使用 Hugging Face 的 kernels 库。
# pip install -U kernels
from kernels import get_kernel
kernel_module = get_kernel("MiniMaxAI/msa", version=0)
sparse_atten_func = kernel_module.sparse_atten_func
sparse_atten_func(...) 该仓库还直接展示了规划器、索引器和注意力调用的实现。
import torch
from fmha_sm100 import fmha_sm100, fmha_sm100_plan, sparse_topk_select
page_size, topk = 128, 16
# Dense proxy pass: per-block max score from a cheap Q slice.
proxy_plan = fmha_sm100_plan(
qo_lens, kv_lens, proxy_q.shape[1],
num_kv_heads=1, page_size=page_size, output_maxscore=True,
)
_, max_score = fmha_sm100(
proxy_q, proxy_k_pages, proxy_v_pages, proxy_plan,
kv_indices=kv_indices, output_o=False, output_maxscore=True,
)
# Block scores -> selected KV block indexes.
kv_block_indexes = sparse_topk_select(
max_score.contiguous(), topk, num_valid_pages=num_pages,
)
# Sparse attention over the selected blocks.
sparse_plan = fmha_sm100_plan(
qo_lens, kv_lens, q.shape[1],
num_kv_heads=k_pages.shape[1], page_size=page_size, kv_block_num=topk,
)
out, _ = fmha_sm100(
q, k_pages, v_pages, sparse_plan,
kv_indices=kv_indices, kv_block_indexes=kv_block_indexes,
) 这些是该仓库的官方使用示例。输入是由调用方准备的分页键值张量。首次运行会对索引器进行 JIT 编译,这可能需要几分钟。环境要求包括 SM100 GPU、CUDA 工具包和 Python 3.10 或更高版本。
优势与不足
优势
- 在报告设置下,每 token 注意力计算量在 1M 上下文时下降 28.4 倍。
- 在 H800 上,实测的墙钟加速比在 1M 上下文时达到 14.2 倍(预填充)和 7.6 倍(解码)。
- 该设计仅向标准 GQA 层添加了两个投影矩阵。
- 它既支持从头训练,也支持从密集检查点进行转换。
- 推理内核以 MIT 许可证发布。
不足与待解决问题
- 已发布的内核针对 NVIDIA SM100 架构;其他架构需要单独适配。
- 在某些子任务上,与全注意力机制相比,仍存在残余的长上下文检索差距。
- 报告中的加速比假设了特定的头配置和 H800 环境。
- 与普通密集层相比,KL 损失增加了训练时的复杂度。
- 结果来自 MiniMax 自身的评估套件,而非第三方复现。