Flash-MSA: 利用稀疏注意力内核加速百万token训练

Hacker News Top 工具

摘要

介绍Flash-MSA,首个针对MiniMax稀疏注意力在Hopper和Blackwell GPU上的高性能开源训练内核,实现高效的百万token训练。

暂无内容
查看原文
查看缓存全文

缓存时间: 2026/07/12 22:52

# Flash-MSA: 利用稀疏注意力内核加速百万Token训练 来源: https://nanduruganesh.github.io/flash-msa/ [\[Github\]](https://github.com/nanduruganesh/flash-msa)[\[MiniMax论文\]](https://arxiv.org/abs/2606.13392)[\[训练器\]](https://github.com/nanduruganesh/Megatron-LM) plot*Flash-MSA 与 Flash-Attention 独立训练步骤对比。*1 (https://nanduruganesh.github.io/flash-msa/#fn:1) 几款前沿模型 \[1 (https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro),2 (https://huggingface.co/zai-org/GLM-5.2),3 (https://huggingface.co/deepseek-ai/DeepSeek-V3.2),4 (https://huggingface.co/meituan-longcat/LongCat-2.0),5 (https://huggingface.co/MiniMaxAI/MiniMax-M3)\] 使用了稀疏注意力来大幅加速推理,但尚未有人发布能够高效训练的代码。今天我首次介绍针对 Hopper 和 Blackwell GPU 的、基于 CuTeDSL 的、开源且性能优异的 MiniMax 稀疏注意力训练内核。所有开发工作均在 Spheron (https://www.spheron.network/) 的 H100 和 B200 租赁实例上完成,并参考了 FA4 (https://github.com/Dao-AILab/flash-attention/tree/main/flash_attn)、MSA 推理 (https://github.com/MiniMax-AI/MSA) 和 Codex。 **免责声明:这不是官方实现,我与 MiniMax 无关联。** ## 关于 MSA MSA 与 Deepseek 稀疏注意力 (https://arxiv.org/abs/2512.02556) 类似,但有一些核心改动 fig_1*图1 来自 MSA 论文* ### 1. 分块稀疏性 与代理注意力选择单个 KV 用于主注意力不同,它通过代理分数上的最大池化,以块(大小为128)为单位进行选择。这为内核带来了一些良好的缓存特性。 ### 2. 主注意力使用 GQA 而非 MLA 这一点尤其重要,因为据我所知,目前西方实验室尚未将 MLA 纳入其训练中,这使得前沿模型(如 GLM-5.2、DSv4)中流行、适配 MLA 的稀疏注意力公式对这里的模型不可用。 ### 3. 代理头的分组专业化 用 GQA 替代 MLA 在每个层内引入了独立的查询组,使得我们可以选择每个代理头对应不同的 KV 子集,而不是像 DSA 那样对整个注意力层求和评分。有证据 (https://arxiv.org/abs/2410.10819) 表明注意力头会自然关注不同的 token,因此这一变化应能提升主注意力的表达能力。 ## 内核设计 fig_2*内核序列的高级概述* 为了高效运行 MSA,我们需要尽量减少重复工作,并避免寄存器/共享内存过载。除了常规的 Flash 寄存器(Q 块、KV 块、O 累加器、LSE 累加器)外,在前向过程中我们还需要考虑流式 top-k 累加器。在反向过程中,我们必须为双重注意力联合传递腾出空间,以便同时计算主注意力和代理注意力的梯度,因为代理梯度需要同时访问代理注意力和主注意力概率。分块稀疏性的一个好处是,由于我们只需要缓存块的索引,而不是像 DSA 那样缓存单个 token,因此可以一直保存块索引直到反向传播 —— 这意味着在整个训练步骤中,只有代理前向过程与上下文长度呈平方关系,其他所有部分都使用代理前向过程中缓存的稀疏块。 ### 前向 操作顺序:代理注意力 -> 稀疏主注意力 -> 将主注意力输出传递到下一层,保存主 LSE 用于反向。 #### 代理注意力 代理点积与常规 Flash Attention 略有不同,因为我们不再需要累加输出,但在遍历键时,我们必须跟踪 top-k 注意力分数及其对应索引。与 Flash 不同,我也不为反向过程累加 LSE,而是通过在反向过程中对稀疏激活执行非常廉价的重计算来获得代理点积的 LSE。在实践中,这比在前向过程中融合 LSE+top-k 更快。当计算每个 QK^T 块时,我获取每个块的因果局部最大分数,并对当前每个查询行的 top-k 值(保存在寄存器中)执行插入排序。我必须将键块一分为二来为这些 top-k 寄存器腾出空间。此外,MSA 规定每个 token 的局部块必须采用非掩码的滑动窗口方式,因此我将每个查询的局部 KV 块的注意力分数设为无穷大。 #### 主注意力 主注意力只是一个块稀疏的 Flash Attention 前向过程。这已在 MoBA (https://arxiv.org/abs/2502.13189) 中被实现,所以我复制了他们的巧妙技巧 (https://github.com/MoonshotAI/MoBA/blob/master/moba/moba_efficient.py),将块稀疏注意力重新参数化为变长 Flash。 ### 反向 为了计算代理头的梯度,我们需要融合代理注意力和主注意力的反向过程,因为代理训练信号需要同时访问代理注意力概率和主注意力概率。 由于我们从正向过程保存了块索引,并且只在稀疏 KV 激活上训练两种注意力,因此反向过程可以在线性时间内运行。首先,我们提取缓存的 block_indices,并反转映射 $B$[batch, proxy head, query, top_k_slot] -> [key block] 为 $B^{\ast}$[batch, proxy head, key block] -> [使用该块的查询]。我们使用 $B^{\ast}$ 来调度查询块,以优化共享稀疏 KV 块的复用。 然后,我们在选定的块上运行一个快速的稀疏代理注意力前向过程(再次使用 MoBA 变长技巧)以获取代理 LSE,接着流式执行融合的代理-主反向任务,加载 QKV、Q_proxy、K_proxy 和 main_lse 的块。为了容纳如此多头到寄存器中,我们必须减少一次使用的 Q 块和 KV 块的大小。在每个流中,我们计算主注意力和代理注意力概率,计算 dQ、dK、dV,然后根据 KL 训练项计算代理 dQ、dK: #### KL 散度损失 回忆 DSA 中的原始 KL 损失项:$L^\iota = \sum_t D_{KL}(p_t, s_t \Vert \text{Softmax}(I_t, s_t))$ 将索引器和主注意力概率分布具体化以累积 KL 散度需要大量的共享内存读写和额外的寄存器使用,这将显著降低训练速度。幸运的是,有一个技巧可以让我们以原子方式进行反向传播,同时在数学上等价于完整的 KL 损失。 将代理注意力概率设为 $p_{px}$,主注意力概率设为 $p$,展开 KL 项: \[L^\iota = \sum_t D_{KL}(p_t \Vert p_{px,t}) = \sum_t p_t \cdot \log\left(\frac{p_t}{p_{px,t}}\right)\] 使用对数规则再次展开: \[L^\iota = \sum_t p (\log(p_t) - \log(p_{px,t})) = \sum_t (p_t \cdot \log(p_t) - p_t \cdot \log(p_{px,t}))\] 我们想要计算传入下一隐层的梯度,即位置 i(不是 t)的 pre-softmax 注意力分数 $z_{px,i}$: \[\frac{\partial L^\iota}{\partial z_{px,i}} = \frac{\partial}{\partial z_{px,i}} \sum_t (p_t \cdot \log(p_t) - p_t \cdot \log(p_{px,t}))\] 主概率 $p_t$ 与代理/KL 损失图无关,因此在此偏导中视为常数。 \[\frac{\partial L^\iota}{\partial z_{px,i}} = \frac{\partial}{\partial z_{px,i}} \sum_t - p_t \cdot \log(p_{px,t})\] 我们知道 softmax 对数概率偏导 (https://math.stackexchange.com/a/2340848):位置 t 的 softmax 输出 ($p_{px,t}$) 相对于 pre-softmax logit ($z_{px,i}$) 的导数:$\frac{\partial \log(p_{px,t})}{\partial {z_{px,i}}} = \delta_{it} - p_{px,i}$ \[\frac{\partial L^\iota}{\partial z_{px,i}} = - \sum_t p_t (\delta_{it} - p_{px,i}) = - \sum_t p_t \delta_{it} + \sum_t p_t p_{px,i}\] 克罗内克函数 $\delta_{it}$ 仅在 i==t 时非零,所以 $\sum_t p_t \delta_{it} = p_i$。 \[\frac{\partial L^\iota}{\partial z_{px,i}} = - p_i + \sum_t p_t p_{px,i} = - p_i + p_{px,i} \sum_t p_t\] 因为 $p_t$ 是一个概率分布,$\sum_t p_t = 1$。 \[\frac{\partial L^\iota}{\partial z_{px,i}} = - p_i + p_{px,i}\] 换句话说,KL 损失对代理分数的梯度 = 代理概率 - 主概率。这就是我在内核中使用的项,无需完全具体化 KL 即可计算代理梯度。 ### 预热内核 在预热模式下,主注意力的前向过程是密集的,不使用 block_indices,因此可以完全跳过代理前向过程,我们可以在反向过程中完全训练它。对于主注意力预热前向内核,我直接调用 Flash 并保存其返回的输出和 LSE,并返回一个占位 KL。在反向过程中,我在索引器上调用密集 Flash 仅获取 LSE,然后复用来自稀疏 MSA 内核的融合代理+主注意力反向过程。 ### 正确性 为了验证内核前向和反向的正确性,我在 Eager PyTorch 中实现了 MSA,并在多种配置下扫描了两个实现的前向输出和反向梯度的余弦相似度。扫描使用 bf16 精度,反向过程包括目标输出损失和内部 KL 损失。通常 bf16 的精度容忍度为 0.01。 Eager vs. 内核余弦相似度 | 批次 | 序列 | Q 头数 | 前向 | 反向投影梯度 | 输出 | QKV | 代理 Q | 代理 K --- | --- | --- | --- | --- | --- | --- | --- | --- | --- | | 1 | 4096 | 8 | 0.9996 | 0.9996 | 0.9996 | 0.9999 | 1.0000 | 1.0000 | | 2 | 4096 | 8 | 0.9996 | 0.9996 | 0.9996 | 0.9999 | 1.0000 | 1.0000 | | 4 | 4096 | 8 | 0.9996 | 0.9996 | 0.9996 | 0.9999 | 1.0000 | 1.0000 | | 1 | 8192 | 8 | 0.9985 | 0.9985 | 0.9985 | 0.9995 | 0.9999 | 1.0000 | | 2 | 8192 | 8 | 0.9985 | 0.9985 | 0.9985 | 0.9995 | 1.0000 | 1.0000 | | 4 | 8192 | 8 | 0.9985 | 0.9985 | 0.9985 | 0.9995 | 1.0000 | 1.0000 | | 1 | 4096 | 16 | 0.9996 | 0.9995 | 0.9995 | 0.9999 | 0.9999 | 1.0000 | | 2 | 4096 | 16 | 0.9996 | 0.9996 | 0.9996 | 0.9999 | 1.0000 | 1.0000 | | 4 | 4096 | 16 | 0.9996 | 0.9996 | 0.9996 | 0.9999 | 1.0000 | 1.0000 | | 1 | 8192 | 16 | 0.9985 | 0.9984 | 0.9984 | 0.9997 | 0.9999 | 1.0000 | | 2 | 8192 | 16 | 0.9985 | 0.9984 | 0.9985 | 0.9997 | 1.0000 | 1.0000 | | 4 | 8192 | 16 | 0.9985 | 0.9984 | 0.9985 | 0.9997 | 1.0000 | 1.0000 | | 1 | 4096 | 32 | 0.9996 | 0.9996 | 0.9996 | 1.0000 | 0.9999 | 1.0000 | | 2 | 4096 | 32 | 0.9996 | 0.9995 | 0.9996 | 1.0000 | 0.9999 | 1.0000 | | 4 | 4096 | 32 | 0.9996 | 0.9995 | 0.9996 | 1.0000 | 1.0000 | 1.0000 | | 1 | 8192 | 32 | 0.9985 | 0.9983 | 0.9984 | 0.9998 | 0.9999 | 1.0000 | | 2 | 8192 | 32 | 0.9985 | 0.9984 | 0.9984 | 0.9998 | 0.9999 | 1.0000 ## 下一步计划 ### 提升融合反向过程的并行性 目前反向过程受限于低张量管线利用率和低占用率,这是因为融合反向过程需要大量寄存器/共享内存来同时运行两个注意力。例如,在本文的吞吐量扫描中,Flash-MSA 反向过程使用了 138 个寄存器/线程,105 KB 共享内存/CTA,而 H100/B200 支持最多 255 个寄存器/线程和 228 KB 共享内存/CTA,因此我只能达到 1 个 CTA/SM。为了增加符合条件的 warp,我尝试了不同的 Q/KV 分块配置,更窄的分块以实现 2 个 CTA/SM,但由此带来的总 CTA 增加反而导致了净减速。Flash-MSA 反向过程的理论占用率为 12.5%,而 Flash-Attention 为 18.75%。 ### 路由器架构加速 GLM 已经证明,使用 IndexShare (https://arxiv.org/abs/2603.12201) 跨层共享代理头是稳定且更快的。此外,索引器在推理时似乎总是以低精度服务,因此如果训练时也使用低精度,可能有助于解决训练-推理匹配问题,并在稳定的情况下大幅加快训练速度。 ### 上下文并行 为了在长上下文场景下扩展 LLM 训练,某种形式的 CP 是必须的,否则训练器会很快耗尽内存。这里有几种选择。 #### 1. 逐头 All-Gather 为清楚起见,我将注意力头的“MSA 组”定义为主注意力查询头分配给每个代理查询头的子集。由于 MSA 组在前向/反向过程中相互独立,因此可以轻松使用 TP 风格的 CP 折叠,CP rank 最高可达 `num_proxy_heads`,每个设备持有(序列长度 / CP rank)个 token 及其对应的 MSA 组,然后在调用 MSA 之前对整个序列执行 A2A 通信。这不需要对内核进行任何修改,事实上,你目前可能就可以在当前 Megatron MSA 分支 (https://github.com/nanduruganesh/Megatron-LM) 中实现。 #### 2. Ring 在这里实现 Ring 风格并行 (https://christianjmills.com/posts/cuda-mode-notes/lecture-013/) 更具挑战性,但很可能是 MSA 实现 CP 的最优方式,因为它允许更好的重叠和比 a2a 更高的 CP rank。Ring 需要在代理前向过程中实现类似重叠交换的方式,能够跨设备流式传输 top-k 值和索引,然后广播回所有设备,接着在跨设备的主注意力中进行巧妙的稀疏键主机设备查找和检索。我不知道如何将其集成到内核中,因此在发布这个核心功能之前,我需要联系本地的 CP 专家。无论如何,这将是下一篇博客文章的主题。 ### 关于联合索引器-主注意力训练的说明 尽管论文中确实包含一种 MSA 公式,其中在代理头上添加了值和输出投影,然后将代理输出添加到主注意力的输出中,以便将 CE 损失的梯度引入代理权重并据称提升知识型评估,但实现这一方案会因添加新的头和累加器而减慢前向/反向过程。MSA 本身在论文的 Table 6 中也指出,适当的预热可以弥补不使用 CE 损失训练索引器的不足。 ### 关于调度器要求代理头数 ≥ KV 头数的说明 由于我通过将每个代理头映射到其对应的主注意力 GQA 组来进行调度,当前内核要求 MSA 组数 ≥ GQA 组数。可能存在一种简单的方法来反转调度器/融合反向过程的映射,但这只在模型具有比代理头更多的 KV 头时才需要。然而,我预计稳定训练至少需要 4 个 MSA 组,因此要使 KV 头数 > 代理头数,你需要超过 4 个 GQA 组,而在当前的 Transformer 体系下,我强烈憎恨那些每层 KV 缓存超过 1024 的模型,所以我不打算实现这条路径。 ### 跨 Top-k 的扫描 我认为随着块数量的减少,可视化稀疏性优势会很有趣。当我有计算资源时,我想对 GQA 和代理 MQA 配置进行更多扫描。我还想使用 MSA 对现有的 GQA 基础模型(如 Qwen3)进行持续的预训练以测试转换,这需要更多的计算资源。 topk_sweep *Top-k 扫描,配置与 MSA vs. FA 扫描相同。*

相似文章

MiniMax 稀疏注意力

Hugging Face Daily Papers

MiniMax 稀疏注意力 引入了一种分块稀疏注意力机制,针对超长上下文的大语言模型实现了显著的加速。在1M上下文长度下,每个token的注意力计算减少28.4倍,在H800 GPU上预填充阶段实际速度提升14.2倍,解码阶段提升7.6倍。该方法附带了一个开源推理内核以及一个公开发布的多模态模型。