跳到正文
北京时间
原文
MarkTechPost(RSS)· Asif Razzaq·· 2026-06-17精选AI 评分70

MiniMax 发布 MSA 稀疏注意力方法,开源推理内核并推出 MiniMax-M3 模型

MiniMax Sparse Attention (MSA): a Two-Branch Block-Sparse Attention Trained on a 109B-Parameter MoE With a 3T-Token Budget

AI 导读

MiniMax 发布 MSA(MiniMax Sparse Attention),一种构建在 Grouped Query Attention 上的稀疏注意力方法。它将注意力分解为索引分支与主分支:索引分支以块粒度(默认 128 token)为每个 GQA 组选择 16 个 token 块(固定预算 2048 个键值 token),主分支仅在这些块上执行精确 softmax 注意力。MSA 在 109B 参数 MoE 模型上训练,开源了面向 NVIDIA SM100 GPU 的推理内核 `fmha_sm100`(MIT 许可,支持 BF16/FP8/NVFP4/FP4),并发布生产模型 MiniMax-M3。MSA-PT 在 MMLU、GSM8K、HumanEval、RULER-8K、RULER-32K 上分别达 67.2、77.7、64.0、84.2、77.5,与全注意力基线持平。128K 上下文下,其 exp-free Top-k 选择比 `torch.topk` 快 5.1 倍。

推荐理由

MiniMax 把长上下文注意力从 O(N) 压到固定每查询 2048 token,还同时开源高效内核与生产模型,对做长上下文 agent 的团队是即时可用的方法,遗憾是只限 SM100 GPU。

正文 · AI 翻译

MiniMax 发布了 MSA(MiniMax Sparse Attention),一种直接构建在分组查询注意力(GQA)之上的稀疏注意力方法。它针对的是一个瓶颈:长上下文下 softmax 注意力的二次方开销。MiniMax 研究团队在一个使用原生多模态数据训练的 109B 参数混合专家模型中对其进行了测试。他们还开源了一个推理内核,并发布了一款生产级模型 MiniMax-M3。

什么是 MSA(MiniMax Sparse Attention)

MSA(MiniMax Sparse Attention)将注意力分解为两个阶段:一个索引分支和一个主分支。索引分支决定每个查询应读取哪些键值块。主分支随后仅对这些块执行精确的 softmax 注意力。

选择以块为粒度进行,而非按 token。默认块大小为 Bk = 128 个 token。每个查询和 GQA 组保留 k = 16 个块。这将每个查询的预算固定为 kBk = 2,048 个键值 token。

两种开销结构不同。稠密 GQA 注意力对每个查询的扩展为 O(N),即完整上下文。MSA 的扩展为 O(kBk),随着 N 增长保持固定。因此,随着上下文长度增加,计算差距会不断扩大。

选择在每个 GQA 组内部共享,但在组与组之间相互独立。一个键值头服务多个查询头,它们共享一个块集合。不同的组可以关注不同的长距离区域。

两个分支如何工作

Index Branch 只向标准 GQA 层添加两个投影矩阵。它为每个 GQA 组定义一个索引查询头,并定义一个共享的索引键头。它对可见的 key token 打分,然后将这些分数最大池化到块级别。

随后,一个 Top-k 算子为每个查询和组选出得分最高的块。包含该查询的本地块始终被纳入。这防止选择器丢弃该查询的紧邻邻域。

Main Branch 从选定的块中收集因果可见的 token。它对这些 token 施加受限的缩放点积 softmax 注意力。每个查询头保留自己的查询投影,但共享该组的块集合。

报告中有一张可视化图展示了学习到的索引器所选择的内容。各头集中在本地对角线和第一个块上。它们将剩余预算留给少数几条长程条纹。

https://arxiv.org/pdf/2606.13392v1
https://arxiv.org/pdf/2606.13392v1

MSA 如何训练

Top-k 选择不可微分,因此语言建模损失无法训练索引投影。MSA 通过 KL 对齐损失解决了这一问题。该损失将 Index Branch 的分布与 Main Branch 的注意力模式相匹配。教师是 Main Branch 在选定 token 上的组平均分布。

三种机制共同稳定了稀疏训练。梯度分离(Gradient Detach)对 Index Branch 的输入施加 stop-gradient。这将 KL 损失限制在索引投影上,而非主干网络。若没有这一机制,更大的 KL 系数会导致梯度尖峰和损失发散。

索引器预热(Indexer Warmup)在最初的迭代中让两个分支都运行完整注意力。索引器在控制路由之前先从 KL 损失中学习。强制局部块(Forced Local Block)为邻近上下文保留一个槽位。

消融实验塑造了最终方案。早期的一个变体添加了一个带有独立输出的 Index Branch 值头。一旦使用了预热,该值头就不再必要。最终设计出于效率考量将其移除。

MSA 支持两条训练路线。MSA-PT 在 40B token 的索引器预热后从零开始训练。MSA-CPT 转换一个在 2.6T token 上训练的稠密 GQA 检查点。随后它继续训练 400B token,其中包括 40B token 的预热。

内核协同设计

理论上的稀疏性若没有匹配的 GPU 路径,就无法转化为速度。MSA 将算法与两个内核思路相结合。

第一个是无 exp 的 Top-k 选择。Softmax 保持顺序不变,因此对原始分数排序会得到相同的索引。该内核在选择之前跳过了 max、exp 和求和步骤。在 128K 上下文下,配合 k = 16,它的运行速度比 torch.topk 快 5.1×。它还比 TileLang 的基数选择内核快 3.7×。

第二种是带 query gather 的 KV-outer 稀疏注意力。相比遍历 query,遍历 KV 块能提高算术强度。该 kernel 将 ⌈128/G⌉ 个 query 位置打包进一个 128×128 的 score MMA。两阶段前向传播将注意力步骤与合并步骤拆分到不同 CTA 上执行。

这个开源 kernel fmha_sm100面向 NVIDIA SM100 GPU。它以 MIT 许可证发布,包含稠密 FlashAttention 以及稀疏 Top-k kernel。它支持 BF16、FP8、NVFP4 和 FP4 精度。

MSA 与其他稀疏方法的对比

研究团队将 MSA 与四种原生训练的稀疏设计进行了对比。

下表总结了它所描述的差异。

方法骨干架构选择粒度索引器 / 选择信号
MSAGQA块级(B_k = 128),按 GQA 组 Top-kKL 对齐损失
NSAMQA / MHA压缩 + 选定块 + 滑动窗口原生(端到端)训练
InfLLM-V2稠密↔稀疏可切换无参数块选择 + 滑动窗口无参数(无训练索引器)
MoBAGQA超大 KV 块(块平均键)仅 LM 梯度
DSAMLA(MQA 模式)Token 级;各头共享单一 Top-kReLU lightning 索引器

MSA 的独特之处在于按 GQA 分组进行 Top-k 共享,并结合块级选择。这样既保持了 KV 读取的连续性,又让每个分组拥有自己的检索。

质量方面也站得住脚。两个稀疏模型总体上与 Full-Attention 基线保持相当的竞争力。

下表展示了 3T-token 预算下的代表性结果。

基准FullMSA-PTMSA-CPT
MMLU67.067.266.8
GSM8K76.277.773.7
HumanEval61.064.057.9
RULER-8K79.884.277.2
RULER-32K75.077.575.7
VideoMME41.1145.4839.65

在长上下文扩展之后,MSA-CPT 在 HELMET-128K 和 RULER-128K 上仍与 Full 保持接近。每次查询仍然只关注 2,048 个 key-value token。

讲解演示场

用例与示例

MSA 面向的是上下文长度成为部署瓶颈的工作负载。

  • 长时程智能体:一个跨越数百个推理与行动步骤的智能体会累积出庞大的对话记录。对这种历史做稠密注意力会呈二次方增长。MSA 无论长度如何,都将每次查询的预算控制在 2,048 个 token。
  • 仓库级代码推理:一个加载完整仓库的编程智能体可能超过数十万个 token。索引器将每次查询路由到少数相关块。无关文件被排除在所选集合之外。
  • 持久记忆:一个长期运行的助手会不断增长其对话状态。MSA 每次查询读取最相关块的一个固定大小切片。随着记忆增长,解码成本大致保持平稳。
  • 长视频理解:该模型原生多模态,并在图像和视频数据上训练。MSA-PT 在多个视频基准上取得了三次运行中的最高分,包括 VideoMME 和 TemporalBench。稀疏选择可扩展到长视觉 token 序列。

运行 Kernel

最快的路径是使用 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(...)

该仓库还直接展示了 planner、indexer 和 attention 调用。

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 Toolkit 以及 Python 3.10 或更高版本。

优势与不足

优势

  • 在报告设定的场景下,1M 上下文时每 token 的注意力计算量下降 28.4×。
  • 在 H800 上,1M 上下文时实测的墙钟时间加速达到 prefill 14.2×、解码 7.6×。
  • 该设计仅向标准 GQA 层添加了两个投影矩阵。
  • 它同时支持从零开始训练以及从稠密 checkpoint 转换。
  • 推理 kernel 以 MIT 许可证发布。

不足与未解问题

  • 已发布的 kernel 面向 NVIDIA SM100;其他架构需要另行开发。
  • 在某些子任务上,与全注意力相比仍存在残余的长上下文检索差距。
  • 所报告的加速基于特定的 head 配置和 H800 环境。
  • KL 损失相比普通稠密层增加了训练时的复杂度。
  • 结果来自 MiniMax 自家的评测套件,而非第三方复现。

来源:MarkTechPost(RSS) · marktechpost.com