MarkTechPost(RSS)
精选
70AI 编辑部评分,满分 100

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

2026-06-17 15:44· 49天前· Asif Razzaq
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 稀疏注意力),这是一种直接构建在分组查询注意力(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 注意力。每个查询头保留自己的查询投影,但共享该组的块集合。

报告中的可视化展示了学习型索引器选择的内容。注意力头集中在局部对角线和第一个块上。它们将剩余预算保留给少数长距离条带。

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

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 与四种原生训练的稀疏设计进行了对比。

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

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

MSA 的独特组合是每 GQA 组的 Top-k 共享与块级选择相结合。这使得 KV 读取保持连续,同时为每个组提供独立的检索。

质量方面表现良好。两种稀疏模型总体上与全注意力基线保持竞争力。

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

基准测试全注意力MSA-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 上仍与全注意力保持接近。每个查询仍然只关注 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 自身的评估套件,而非第三方复现。

来源:MarkTechPost(RSS) · marktechpost.com

同一事件 · 1

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

MarkTechPost(RSS)·2026-06-17 15:44·49天前·Asif Razzaq
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 倍。

正文 · AI 翻译

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 注意力。每个查询头保留自己的查询投影,但共享该组的块集合。

报告中的可视化展示了学习型索引器选择的内容。注意力头集中在局部对角线和第一个块上。它们将剩余预算保留给少数长距离条带。

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

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 与四种原生训练的稀疏设计进行了对比。

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

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

MSA 的独特组合是每 GQA 组的 Top-k 共享与块级选择相结合。这使得 KV 读取保持连续,同时为每个组提供独立的检索。

质量方面表现良好。两种稀疏模型总体上与全注意力基线保持竞争力。

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

基准测试全注意力MSA-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 上仍与全注意力保持接近。每个查询仍然只关注 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 自身的评估套件,而非第三方复现。

来源:MarkTechPost(RSS)· marktechpost.com

同一事件 · 1