Better MoE Model Inference with Warp Decode
Better MoE Model Inference with Warp Decode
发表时间: 2026-04 · Blog post by Cursor (cursor.com)
原文: https://cursor.com/blog/warp-decode
作者/机构:Less Wright, Federico Cassano & Zhiyuan Zhang
速读
一句话结论 本文提出了一种名为 Warp Decode 的混合专家模型(MoE)推理内核优化方法,通过将并行化维度从“按专家分配”翻转为“按输出神经元分配”,在 Blackwell GPU 上实现了 1.84 倍的解码吞吐量提升和 1.4 倍的计算精度提升。
要解决什么问题 在混合专家模型(MoE)的推理系统中,传统的标准做法是围绕“专家”来组织 token 的生成路径。这种以专家为中心的数据流在 Prefill(预填充)阶段或大批量处理时效率很高,因为多个 token 共享同一个专家,足以摊销组织数据的开销。然而,在自回归解码阶段,系统每次只生成一个 token(即小批量解码),此时传统路径会暴露出严重的机制卡点:为了把 token 路由给对应的专家并收集结果,系统必须执行大量的数据搬运和布局管理工作。在传统 MoE 路径的 8 个执行阶段中,有 5 个阶段(如对齐内存边界的填充、分散 token、组合中间结果等)纯粹是为了维护专家维度的数据结构而存在的簿记操作,完全不包含实际的数学计算。这种管理与调度开销不仅浪费了计算资源,还强制要求分配额外的中间内存缓冲区,导致在 Blackwell 等先进 GPU 上进行单 token 解码时,内存带宽的利用率极低,严重拖慢了推理速度。
怎么做的 Warp Decode 的核心思路是彻底翻转并行化轴,不再将 GPU 的计算资源(warp,即 32 个并行处理通道的集合)分配给各个专家,而是直接将每个 warp 精确分配给单个输出值(神经元)。这种围绕输出而非专家组织内核的设计,使得每个 warp 在整个生命周期内完全独立,只需负责产生一个输出标量,从而直接绕过了传统方法中繁琐的阶段交接和跨 warp 同步卡点。该方法将整个 MoE 计算层压缩为两个高度融合的内核。首先是 Gate/Up 内核,每个协作线程阵列包含 8 个 warp,每个 warp 负责计算一个 token 与一个路由专家配对的中间神经元。warp 会读取该神经元的 gate 和 up 权重行,流式传输输入激活向量,并在私有寄存器中将 MXFP8 格式的权重动态转换为 FP32 进行点积累加。由于激活向量只需读取一次即可复用于两个投影,系统直接计算 $SiLU(gate) \times up$ 并写入中间值,全程无需共享内存暂存。其次是 Down 内核,每个 warp 负责一个 token 的单一输出维度。它会循环遍历所有 top-k 路由专家,加载向下投影权重,并将每个专家的路由权重直接折叠到一个 FP32 寄存器累加器中。处理完所有专家后,系统调用 PTX 级别的 `shfl.sync.bfly` 指令,通过 warp 内的蝶形归约交换寄存器数据,直接得出最终的加权组合结果。这种设计的关键在于实现了极端的 warp 独立性,在机制上消除了三大开销:一是消除了为满足分组内核要求而对专家 token 列表进行的内存填充;二是消除了分散和组合步骤,八个专家的中间结果直接在寄存器中折叠,永远不会在全局内存中具体化;三是移除了传统路径必需的激活收集缓冲区和每个专家的输出缓冲区,每个 token 节省了超过 32 KB 的中间流量,从而将宝贵的 L2 缓存容量全部释放给真正决定性能的权重读取。由于 warp 之间没有共享的可变状态,GPU 调度程序可以无视正确性约束,随时切换停顿的 warp,完美隐藏了内存延迟。
效果如何 实验在 NVIDIA B200 GPU(具备 148 个流式多处理器)上进行,模型采用 Qwen-3 风格的 MoE 架构(隐藏维度 2048,从多专家中路由 top-8),输入激活保持 BF16 格式,权重使用 MXFP8。对比的基线方法是传统的“以专家为中心”的 MoE 推理路径,该路线代表了当前主流的通过收集 token、按专家分配计算并重新组装结果的标准流水线实现。在端到端解码任务中,Warp Decode 实现了 1.84 倍的吞吐量提升,且该增益在所有上下文长度下保持平坦,证明其是纯粹的生成时优化。在硬件效率方面,当批处理大小为 32 时,该方法维持了 3.95 TB/s 的吞吐量,达到了 B200 硬件连续内存读取峰值的 58%,剩余差距主要源于专家路由带来的随机内存访问延迟。同时,Warp Decode 带来了 1.4 倍的准确度提升,因为传统基线在计算时会将 BF16 激活转换为 MXFP8 再转回,导致层间舍入误差累积;而新方法全程保持激活为 BF16、累加器为 FP32,输出结果更接近完整的 FP32 参考值,最小余弦相似度大于 0.999996。作者也明确指出了该方法的局限性:Warp Decode 并非通用替代方案,在 Prefill 阶段或大批量推理等高容量场景下,由于大量 token 共享同一专家足以摊销管理开销,传统的以专家为中心的方法依然具有优势。
主要贡献
核心问题:大多数混合专家(MoE)推理系统围绕专家(experts)来组织token的生成路径。虽然这种做法在规模化应用中是标准方法,但在Blackwell GPU上进行小批量解码(small-batch decode)时,会导致大量的管理与调度开销。
研究目标:探索在Blackwell GPU上MoE解码所能达到的最大实际内存带宽,并以此为基础优化推理架构。
创新点:提出了一种称为“Warp Decode”的方法。该方法彻底翻转了并行化轴(parallelism axis),不再将warp分配给专家,而是将每个warp分配给单个输出值(神经元),从而围绕输出而非专家来组织内核。
核心效果:Warp decode是一种罕见的能同时提高性能和准确性的内核设计。在Blackwell架构上,它实现了1.84倍的吞吐量提升,同时输出结果比传统方法更接近完整的FP32参考值(准确度提升1.4倍)。这极大地加速了Composer模型的研究和训练流程。
背景知识/关键Observation/设计原则
传统MoE路径的局限性:现代MoE模型将每个token路由到专门的专家网络子集中(例如在某一层从128个专家中选择8个)。标准的实现方式是围绕这些专家来组织所有的计算,包括收集每个专家所需的token、执行数学运算并重新组装结果。这种方法在Prefill阶段和大批量(large batches)处理时表现良好,因为每个专家分担的工作量足以抵消组织数据的开销。然而,在自回归解码(autoregressive decode)步骤中,每次只生成一个token,缺乏足够的共享工作量来证明这种开销的合理性。在传统路径的8个阶段中,有5个阶段纯粹是为了管理以专家为中心的数据布局而存在的,并不执行任何实际计算。
方法细节
更改并行化维度以消除开销:Warp decode通过将并行化核心从专家重新组织为输出,彻底消除了传统方法中的五个“簿记(bookkeeping)”步骤。现代GPU以32个并行处理通道(称为warp)为一组执行指令。在Warp decode方法中,每个warp被精确分配计算一个输出值。该warp直接从内存中流式传输其所需的权重数据,在单个运行总和中聚合所有8个路由专家的结果,并最终写入一个结果。这种warp的独立性使得Warp decode在运行时无需任何分段(staging)、交接(handoffs)、跨warp同步点或中间缓冲区。整个MoE计算层被压缩为两个内核:moe_gate_up_3d_batched和moe_down_3d_batched。
Gate/Up内核的工作机制:在gate/up内核中,每个协作线程阵列(CTA)由8个warp组成,每个warp负责一个中间神经元,对应一个token和一个路由专家的配对。warp加载路由专家ID,读取该神经元的gate和up权重行,并流式传输输入激活向量。在此过程中,MXFP8格式的权重被动态转换为FP32,并且两个点积都在私有寄存器中累加。由于这两个内核被融合为一次传递,激活向量只需读取一次即可立即重新用于两个投影,而无需任何共享内存进行数据暂存。在完成warp级别的归约(reduction)后,warp应用$SiLU(gate) \times up$并写入一个中间值。
Down内核的工作机制:在down内核中,每个warp负责一个token的一个输出维度。它循环遍历所有top-k路由专家,加载相关的向下投影(down-projection)权重行,并流式传输中间激活,同时将每个专家的路由权重折叠到一个单一的运行FP32累加器中。在处理完所有专家后,系统使用__shfl_xor_sync通过warp级别的蝶形归约(butterfly reduction)来归约32个通道本地的部分和。该操作直接编译为PTX的shfl.sync.bfly指令,这是一种单一的硬件原语,可在warp内的通道之间交换寄存器,从而完全绕过共享内存。这种设计的优势在于不需要L1缓存往返、不存在bank冲突、也不需要显式的屏障(barriers),因为同步已经通过通道掩码(lane mask)内置于指令中。最终的加权top-k组合不是作为一个单独的epilogue,而是直接成为投影本身的一部分。Warp decode中的每个warp都是独立的,并在其整个生命周期内获得一个单一且稳定的分配:产生一个输出标量。正是这种warp独立性消除了传统路径所需的共享内存暂存、跨warp同步和中间缓冲区。
通过消除阶段提升吞吐量:Warp decode通过消除传统路径所需的阶段和缓冲区,以及创建允许更好调度和延迟隐藏的warp独立性来实现性能提升。阶段消除提供了大部分吞吐量增益,具体包括消除填充(padding)、分散(scattering)和组合(combine)步骤。这种消除需要从根本上重新组织并行性,而不仅仅是融合传统流水线的各个阶段。
* 消除填充:传统路径将每个专家的token列表填充到2的幂或128字节边界,以符合分组内核的要求。在单token解码时,这是无法摊销的开销。Warp decode路径通过从不形成每个专家的批次,完全避免了这种开销。
* 消除分散和组合:传统路径在每个专家完成后,将八个中间结果写入GPU内存,然后运行单独的归约步骤来组合它们。Warp decode路径将每个专家的路由权重折叠到warp内的运行累加器中,这八个中间结果永远不会在内存中具体化,从而节省了后续归约通道的写入和读取成本。
通过消除缓冲区释放缓存容量:这种重组还移除了传统路径因其以专家为中心的布局而需要的两个中间内存缓冲区。
* 第一个是激活收集缓冲区(activation gather buffer),它是输入激活向量的副本,并重新排列为以专家为主的布局。在批处理大小为1时,这是对已存在数据的完整复制。
* 第二个是每个专家的输出缓冲区(per-expert output buffer)。假设有8个专家和2048的隐藏维度,每个token在BF16格式下会分配、写入、立即读取一次并丢弃 $8 \times 2048 \times 2$ 字节 = 32 KB的数据。
Warp decode通过将8个专家的贡献折叠到跨32个warp通道的寄存器累加器中来消除这两个缓冲区,在最终写入单标量之前没有任何数据到达全局内存。每个token移除超过32 KB的中间缓冲区流量,为真正决定性能的权重行释放了L2缓存容量。
Warp独立性优化硬件调度:重组使得保留的计算速度更快,因为内核在设计上是极度并行(embarrassingly parallel)的:每个warp完全独立于其他所有warp。由于每个warp拥有正好一个输出标量并仅读取其所需的权重行,因此warp之间没有共享的可变状态。在单个warp层面,这种独立性是完全的:输入激活是只读的,累加器存在于私有寄存器中,输出写入到一个唯一的地址。从硬件调度程序的角度来看,整个输出维度是一个扁平的独立工作项池。GPU的warp调度程序可以随时以任何顺序发出任何warp,而没有任何正确性约束。当一个warp因等待内存加载而停顿时,调度程序会立即切换到另一个warp。在B200的148个流式多处理器(SM)中,有数千个warp在运行,内存延迟几乎完全被其他warp的有用计算所隐藏。此外,该内核呈线性扩展,使得输出维度翻倍就会使独立的warp数量翻倍,而无需增加同步。在token批处理维度上也是如此,因此调度程序看到的是一个扁平的工作命名空间,节点之间没有边缘。这与传统路径形成鲜明对比,传统路径中专家级别的GEMM内核需要块内(intra-block)协调。
实验环境
- 模型架构:Qwen-3 风格的 MoE 模型(典型参数为从多专家中路由 top-8,隐藏维度为 2048)。
- 数据类型:输入激活值保持 BF16 格式;权重使用 MXFP8 并在计算时动态转为 FP32;累加器使用 FP32。
- 硬件配置:NVIDIA B200 GPUs(具备 148 个流式多处理器 / SMs)。
- 软件配置:内部推理系统,结合底层 PTX 指令优化。
实验结果
端到端解码吞吐量:
* 实验内容:在 NVIDIA B200 GPU 上运行 Qwen-3 风格模型,测试内部推理系统在不同上下文长度下的吞吐量表现。
* 实验结果与分析:Warp decode 实现了 1.84 倍的吞吐量提升。吞吐量增益在所有上下文长度桶(context-length buckets)中保持平坦(flat),证实这是一种纯粹的生成时(generation-time)性能改进,不依赖于提示词(prompt)的长度。
准确度提升:
* 实验内容:对比 Warp decode 与传统路径在数据量化转换上的误差积累情况。
* 实验结果与分析:Warp decode 的输出比传统路径更接近完整的 32 位基准(ground truth),准确度提升了 1.4 倍。传统方法将 BF16 激活转换为 MXFP8 再转换回来,引入了在模型层间累积的舍入误差底线;而 Warp decode 始终保持激活为 BF16,累加器为 FP32,归约操作从未在降级的输入上进行,从而显著提升了计算质量。
硬件效率:
* 实验内容:测试 Warp decode 接近 B200 硬件最大吞吐量(连续内存读取峰值为 6.8 TB/s)的程度,并验证结果正确性。
* 实验结果与分析:在批处理大小 B=32 时,Warp decode 维持了 3.95 TB/s 的吞吐量,达到硬件峰值的 58%。剩余的差距主要反映了专家路由产生的随机访问模式(如单个token可能路由到非相邻的专家 5, 8, 14, 19 等)带来的内存延迟成本。在所有批处理大小下,算法与参考实现的正确性保持高度一致:最小余弦相似度 > 0.999996,最大绝对差值为 0.001953。
结论
Warp decode 并不是对以专家为中心的执行方式的通用替代方案。对于 Prefill 和大批量推理等高容量工作负载,由于许多 token 共享同一个专家,组织数据的开销可以被有效摊销,以专家为中心的方法仍然具有优势。然而,在缺乏足够共享工作的场景(如 MoE 自回归解码)中,Warp decode 通过消除大量管理开销展现出巨大优势。作为持续改进 Composer 模型的重要组成部分,Warp decode 这种推理端的优化决定了模型输出到达开发者的速度和准确性,其价值与预训练数据和强化学习(RL)的投入同等重要。
💬 评论讨论
欢迎在这里分享您的想法和见解!