💡 核心创新点
1.1 破除传统稀疏假设
在大模型长文本建模中,传统的稀疏自注意力(Sparse Attention)方案往往伴随着表达能力的严重牺牲,或者因为过于复杂的逐头(Per-Head)非正则化计算而导致吞吐极低。MiniMax 稀疏注意力(MSA) 破除了传统方法“精确表征必须依赖全量自注意力”的过往假设,开创性地将稀疏注意力建模拆解为双阶段过程。通过共享的分组查询注意力(GQA)特征,MSA 在保留原本高精度表征的同时,通过大幅度削减冗余上下文的检索,实现了计算复杂度的根本性重构。这一改变直接让百万级别超长文本的计算开销从二次方降低到线性级别,在实际应用部署中具备极强的可落地性。这是相较于传统方法的本质性机制跃迁。
1.2 分组块级检索机制
相比于传统的全自注意力机制,MSA 的核心差异在于其独特的 GQA 组级块检索模式。MSA 并没有使用传统的、对硬件极不友好的 Token 级稀疏检索,而是采用块级(Blockwise)自适应筛选,并使同一 GQA 分组内的所有 Query 头共享相同的检索结果。这种设计通过极轻量化的 索引分支(Index Branch) 动态预测与当前 Query 最相关的关键块(Key Blocks),而在 主分支(Main Branch) 中仅针对这些被选中的块进行高精度自注意力计算。这种粗细粒度的有机结合,不仅避免了无谓的全局矩阵乘法,而且由于规整的内存对齐机制,能够完美适配现代 GPU 的张量核心(Tensor Core)执行管线,使算法理论上的稀疏计算削减能够无损转化为实打实的吞吐量提升。
🔍 背景与动机
2.1 长文本时代的算力屏障
随着大语言模型(LLM)朝着智能体(Agent)工作流、超长代码仓库级推理、以及长效持久化记忆等场景演进,模型对数十万甚至数百万 Token 超长上下文的联合处理能力已经从“加分项”演变为“刚需”。然而,传统自注意力机制所带来的 二次方计算复杂度(Quadratic Complexity),如同一道难以逾越的算力大山,使长文本推理在算力和显存(Memory Footprint)两端都面临严峻的工程瓶颈。在实际部署中,巨大的 键值缓存(KV-Cache) 密集产生,常常直接导致系统显存溢出,或迫使系统由于高昂的推理延迟而放弃高并发吞吐。
2.2 现有解决方案的痛点
为了拓宽长文本能力的帕累托前沿,学术界和工业界此前曾尝试了多种替代架构。例如,线性注意力(Linear Attention)和状态空间模型(SSM,如 Mamba 等)虽然能够实现线性级的计算复杂度,但在需要精确位置匹配和高保真检索的“大海捞针”(Needle-in-a-Haystack)任务中,性能往往出现明显下滑;滑动窗口自注意力(Sliding Window)虽然计算简单,但由于其缺乏对长程依赖的感知能力,在涉及跨文档长链条推理时力不从心。
🎯 核心结论:寻找一个既能保持自注意力机制超凡检索精度,又能规避其二次方计算开销,同时还能保证极高 GPU 算子执行效率的稀疏自注意力框架,是打破这一瓶颈的关键。
为了通俗地理解这一问题,我们可以把长文本检索比作“在巨型档案馆中寻找历史档案”。传统的全局自注意力如公式
所示:
(1)
这里公式中的变量 , , 分别表示时刻 的 Query、Key、Value 向量。计算逻辑就像是派遣了整个团队去从头到尾阅读档案馆里的每一页文件,这在大规模长序列(大 )下,无疑效率极低且代价高昂;而 MSA 的设计逻辑则是先派出一名动作敏捷、记忆力极强的“轻量级图书管理员”(Index Branch)快速翻阅每个文件箱(Block)的侧边标签并挑出最相关的 个箱子(Top-k Blocks),随后专业的研究员们(Main Branch)只需要对这几个被挑出来的箱子进行深度精读。这种双阶段协同配合,既保留了极高的准确度,又大大加快了归档与检索效率。
📊 数据集与实验环境
3.1 109B MoE 模型配置
为了验证 MiniMax 稀疏注意力机制(MSA)在真实、大规模产业落地场景下的有效性,作者团队采用了一个参数量高达 109B 的混合专家模型(MoE) 进行从头预训练与微调测试。该模型的具体配置与训练环境如下:
- 模型架构:包含 41 层 Transformer,其中前 3 层为密集层(Dense Layers),后 38 层为混合专家层(MoE Layers)。每个 MoE 层包含 128 个路由专家(Routed Experts),1 个共享专家(Shared Expert),并采用 Top-4 的路由筛选策略。每次推理中,单 Token 激活的参数量约为 6B。
- 注意力规格:主模型配置 64 个 Query 头,4 个 KV 头,头维度 ,旋转位置编码(RoPE)维度为 64。这种参数比例使得 GQA 的分组比达到了 16。
- 稀疏参数:块大小 ,每个 GQA 分组保留的关键块数 ,即每个 Query 头最多与 个 Token 进行注意力计算,内存路由设计高度规则化。
- 词表大小:200K 词元(Tokens)。
3.2 训练超参与双轨路线
在硬件方面,该实验完全在高性能 NVIDIA H800 GPU 平台上运行,支持原生多模态联合训练,涵盖大规模文本和图像/视频混合数据集(如 AI2D、MMMU、EgoSchema 等)。模型共进行了两套方案的实验:
- 从头稀疏预训练 (MSA-PT):模型在总计 3T Tokens 预算下进行从头预训练。其中,最开始的 40B Tokens 采用全注意力机制进行索引器预热(Indexer Warmup),后续阶段全部转入块级稀疏训练模式。
- 渐进式继续训练 (MSA-CPT):首先在全注意力模式下预训练 2.6T Tokens,随后用 MSA 替换原先的密集注意力,继续训练 400B Tokens。同样在替换初期安排 40B Tokens 进行预热对齐。
- 优化器与超参:学习率、动态衰减策略等完全匹配标准 109B 模型的经典设置。通过在不同数据流中注入 MSA,验证了其对于多模态长文本输入的高吞吐承载能力。
🛠️ 核心机制与研究方法
4.1 双分支处理管线
MiniMax 稀疏注意力(MSA)的核心架构如图 1 所示,采用了极其精简的 Index-Main 双分支管线:

该架构最大程度地保留了现有的 GQA 软硬件生态。它不改变 backbone 的层级传输逻辑,而是通过在输入阶段注入极轻量化的自适应变换,完成注意力的动态路由。
4.2 索引分支 (Index Branch) 的数学原理
4.3.1 块级最大池化打分
索引分支为每个 GQA 分组引入了独立的索引 Query 投影矩阵,并共享一个通用的 K 投影矩阵,以维持低开销。对于每个 Query 向量,我们首先计算 Token 级别的点积相似度评分,然后再在块(Block)的层级上利用最大池化(Max-Pooling)将其聚合成块级评分。其数学公式表达如下:
在公式 (6) 中, 和 代表投影至 维度的低秩索引向量。 计算了 Query 与 Key 之间的 Token 级关联评分,而 则通过在划分好的块 内取最大值,实现了对整个块表现力的鲁棒评估。在获取到块级打分矩阵 后,我们采用 Top-k 算法筛选出得分最高的 个候选块索引集 。特别地,局部物理块(Local Block) 默认无条件包含在集合中,以增强近距离上下文和因果逻辑的建模稳定性。
4.3 主分支 (Main Branch) 块稀疏计算
4.3.1 自适应注意力聚合
一旦索引分支计算出块级索引集合 ,主分支则使用正常的注意力权重投影进行实际注意力计算。这一步完全规避了传统稀疏注意力所采用的重构技巧,而是执行局部的、因果可见的密集 Softmax 注意力。计算公式表示如下:
在公式 (8) 中, 为主分支中第 个 Query 头的特征向量,而 与 分别表示根据索引分支输出,对该 GQA 分组下的 Key-Value 缓存进行块级 Gather 操作后的矩阵。通过这种动态重组方式,原本需要处理全量上下文的自注意力成本被彻底压缩为仅处理 规模的局部稠密运算。
4.4 GPU 内核深度协同设计
4.4.1 免指数 Top-K 快速筛选
为了把 MSA 理论上的 FLOPs 计算量削减切实转化为端到端的吞吐飞跃,作者对 GPU 算子执行路径进行了极致的协同设计(Co-Design)。
- 免指数排序优化 (Exp-Free Selection):由于 Softmax 变换本身具有单调递增的保序性,在计算 Top-k 分数时,系统可以完全跳过指数求和、极大值减法等昂贵的标准 Softmax 算子操作,直接将原始点积得分送入 Top-k 筛选。这规避了冗余的计算瓶颈。
- 寄存器内极速 Min-Heap 维护:针对 的特化场景,内核设计了基于 Warp 内 32 个 Lane 独立流式读取并维护一个寄存器小顶堆(Min-Heap)的高效实现。这种设计避免了频繁的共享显存冲突,相比
torch.topk和 TileLang 实现了数倍的延迟压缩(数据详见表 1)。

4.4.2 KV 外循环算子优化
在实现主分支自注意力时,传统的注意力算子往往采用 Q-Outer(Query 外循环)策略,但这在块稀疏下会导致严重的 Tensor Core 硬件利用率不足。因为在稀疏场景下,各个 Query 激活的 Key 块是不均匀且高度发散的。
💡 关键技巧:MSA 设计并采用了 KV-Outer(KV 外循环) 算子实现。该实现将选定的 KV 块作为外层遍历核心,动态聚集选中该块的 Query 序列进行矩阵拼接,从而能充分填满 Tensor Core 的 MMA(矩阵乘累加)基本单元,极大释放了 Tensor Core 的算力。
此外,为了防止由于某些热门 KV 块被大量 Query 同时选中而导致的算力负载倾斜(Sink Rows),内核额外引入了 预调度瓦片分块技术(Pre-scheduled Tile Chunking),通过两阶段合并机制消除了原子操作锁竞争,保证了极其优异的算力负载均衡。
4.5 辅助对齐损失与训练稳定化
为了使非微分的 Top-k 块选择能够稳定且有意义地训练,MSA 在前向和后向传播上采用了如下两项工程创新:
-
KL 散度对齐损失 (KL-Alignment Loss):引入一个专门的辅助损失,强行将索引分支预测出的局部 Token 分布,去贴合主注意力分支在选定块内的实际 Softmax 概率均值,为其提供清晰合理的全局表征监督。
-
梯度阻断技术 (Gradient Detach):
🎯 核心结论:如果不做梯度控制,KL 损失产生的反向梯度会顺着投影矩阵渗透回主模型的 Backbone 骨干网络,导致严重的自 distill(自蒸馏)退化与训练发散。
因此,MSA 强制在索引分支的输入端添加
stop_gradient(梯度阻断),使辅助损失只能作用并优化索引分支自身的轻量投影权重,彻底隔绝了对基础特征层的影响。这也是保障 109B 模型训练极佳收敛稳定性的幕后功臣。
📈 结果与深度分析
5.1 全面评测与基线对比
为了全面分析 MSA 在下游任务中的表现,作者在包含通用语言、数学、代码、多模态以及长文本在内的数十项评测基准上,对全注意力基线模型(Full)、从头训练模型(MSA-PT)以及继续训练模型(MSA-CPT)进行了对比,部分代表性结果如表 2 所示:

从实验数据可以看出,基于 MSA 的稀疏模型在各项能力维度上均能高度对齐甚至超越 Full-Attention 基准线。尤其是在长文本检索任务(如 RULER-8K/32K)上,MSA-PT 分别达到了 84.2% 和 77.5% 的优异成绩,超越了 Full-Attention(分别为 79.8% 和 75.0%)。这证明了其在极低注意力算力预算(仅 激活 Token)下,依然能完美维系并增强长文本特征聚合能力,消融实验中移除索引分支会导致长文本检索几乎完全失效(RULER 下降超过 50%)。
5.2 极致的推理效率加速
通过利用 MSA 的复杂度公式
,我们可以清晰地量化这一理论收益:
(12)
其中由于 ,第二项主分支所消耗的 FLOPs 在长序列(大 )下远远低于 GQA 对应的二次方复杂度。在配备 64 个 Query 头、4 个 KV 头的 H800 GPU 环境下进行的高清效率基准测试中,MSA 显现出了统治级的推理效率表现:

🚀 重大突破:在 1M(100万)Tokens 的极端超长上下文输入下,MSA 相比 GQA 实现了高达 28.4 倍 的理论 attention 计算量削减。得益于协同设计的 GPU 内核,该框架在 H800 上达成了 14.2 倍的 Prefill 首字延迟加速,以及 7.6 倍的 Decoding 解码延迟加速,彻底颠覆了超长上下文的线上服务成本结构。
5.3 局限性与未来展望
⚠️ 局限性:尽管 MSA 通过自适应索引取得了极其显著的加速和对齐效果,但分析也显现了一些硬性局限。当前版本的索引打分系统依然存在注意力汇聚(Attention Sink)的固有倾向,即对序列第一个块的打分持久高企。虽然局部物理块的无条件保留机制一定程度上减缓了对这一块的选择压力,但如何在不施加固定规则的情况下,让索引器更细粒度地理解高频局部的逻辑转移,仍是一个具有挑战性的课题。
未来,作者团队建议沿着三个方向继续拓展 MSA 的研究边界:
- 自适应变长选择:探索根据序列特征或层深度,动态调整保留块数 的机制。
- 多模态特征融合:在多模态、超大图像高频注意力中,优化块级索引以适配空间位置对齐。
- 强化学习后训练 (RL Post-Training):将 MSA 的高效率进一步集成至推理时搜索、蒙特卡洛树搜索等计算密集型智能体逻辑中,充分发掘百万级长文本的极致吞吐潜力。
📋 参考文献
- Ainslie et al., "GQA: Training generalized multi-query transformer models from multi-head checkpoints", EMNLP 2023.
- Wang et al., "Tilelang: A composable tiled programming model for AI systems", arXiv 2025.
- Dao et al., "FlashAttention-2: Faster attention with better parallelism and work partitioning", ICLR 2024.