Prhub

#32668 Enable GPT-OSS FlashInfer MXFP4 on SM120

原始 PR 作者 mmangkad 合并时间 2026-07-30 08:04 文件变更 3 提交数 5 评论 8 代码增减 +329 / -17

执行摘要

在 SM120 上启用 GPT-OSS FlashInfer MXFP4 MoE

在SM120(Blackwell架构)上,FlashInfer CUTLASS MXFP4内核比现有Marlin内核提供显著性能优势(最高30%吞吐提升)。PR body中的基准测试显示从低并发到高并发全面超越Marlin,因此希望为GPT-OSS模型默认启用该内核。

值得精读。本PR展示了如何为特定GPU架构添加新的MoE内核后端,包括初始化检测、对齐约束处理、权重后处理以及自动切换。对于涉及多架构内核支持的开发者有很好的参考价值。

讨论亮点

在PR审核中,b8zhong询问FlashInfer是否已支持SM90,如果支持是否可以完全删除triton MoE内核。mmangkad回复SM90已支持,但由于minimax m3(MXFP8)和ROCm平台仍依赖triton内核的某些部分,目前无法完全删除。这一讨论凸显了内核维护的复杂性:多架构、多精度格式共存时,不能简单用单一后端替代所有场景。

实现拆解

  1. mxfp4.py中注册SM120内核路径:在Mxfp4MoEMethod.__init__()中新增is_sm120_supported()分支,设置_fi_kernel = 'cutlass_sm120';在create_weights()中将SM120与SM90的padding策略合并(要求dimension % 128 == 0),因为这些CUTLASS内核有相同的对齐约束。

  2. mxfp4.py中新增权重后处理函数:在process_weights_after_loading()中为cutlass_sm120分支调用新增的_process_weights_for_sm120_cutlass(),该函数包含_stack_up_gate_w13()_pad_w2_3d()_apply_sm120_cutlass()等辅助函数,负责将模型权重重新排列为FlashInfer CUTLASS内核所需的布局(halved [up; gate]布局),并执行padding。

  3. overrides.py中修改SM120自动选择:当检测到SM120且模型使用MXFP4量化格式时,将moe_runner_backend'marlin'改为'flashinfer_mxfp4',使得SM120用户无需手动指定即可启用新内核。

  4. 新增单元测试test_mxfp4_sm120_cutlass.py:添加test_gpt_oss_sm120_padding_layout_and_kernel函数,在SM120上构建模拟层并调用Mxfp4MoEMethod的权重处理,然后通过FlashInfer cutlass_fused_moe执行前向,验证结果与显式调用FlashInfer直接路径一致,确保padding和内核行为正确。

文件 模块 状态 重要度
python/sglang/srt/layers/quantization/mxfp4.py 量化层 modified 8.68
test/registered/unit/layers/quantization/test_mxfp4_sm120_cutlass.py 集成测试 modified 6.46
python/sglang/srt/arg_groups/overrides.py 参数覆盖 modified 5.02

关键符号

_process_weights_for_sm120_cutlass _stack_up_gate_w13 _pad_w2_3d _apply_sm120_cutlass

关键源码片段

python/sglang/srt/layers/quantization/mxfp4.py core-logic

核心变更,新增 SM120 内核路径和权重处理函数,是功能实现的主文件。

# Inside Mxfp4MoEMethod.__init__():
self._fi_kernel: Optional[str] = None
if self.use_flashinfer:
    if is_sm100_supported():
        self._fi_kernel = "trtllm_sm100"
    elif is_sm120_supported():
        # SM120 -> use FlashInfer CUTLASS MXFP8 x MXFP4 MoE kernel
        self._fi_kernel = "cutlass_sm120"
    elif is_sm90_supported():
        if not _FI_HAS_SM90_CUTLASS_MXFP4:
            raise RuntimeError(...)
        self._fi_kernel = "cutlass_sm90"
    else:
        raise NotImplementedError(
            "moe_runner_backend=flashinfer_mxfp4 requires SM90, SM100, or SM120."
        )# Later in process_weights_after_loading():
if self._fi_kernel == "cutlass_sm120":
    self._process_weights_for_sm120_cutlass(layer)
    return

评论区精华

FlashInfer 对 SM90 的支持与 triton 内核删除可行性 设计

b8zhong 询问 FlashInfer 是否已支持 SM90,如果支持是否可以完全删除 triton MoE 内核。

结论:mmangkad 回复 SM90 已支持,但由于 minimax m3(MXFP8)和 ROCm 平台仍依赖 triton 内核的某些部分,目前无法完全删除。 · 已解决

风险与影响

  1. 回归风险overrides.py中的自动选择修改可能影响其他架构或非GPT-OSS模型的默认后端选择,需确保条件判断互斥且完备。
  2. 正确性风险:新增的_process_weights_for_sm120_cutlass函数涉及复杂的权重重排和padding,如果输入维度不满足内核约束(如%128 != 0),可能导致静默错误或数值异常。单元测试覆盖了特定尺寸(hidden=160, intermediate=160),但真实模型维度可能触发未测试的边界条件。
  3. 依赖风险:依赖FlashInfer版本,需确保cutlass_fused_moe支持MXFP8×MXFP4的SM120入口,否则抛出NotImplementedError。
  4. 性能风险:N/A(基准测试已显示积极收益)。

对用户:SM120用户使用GPT-OSS模型时将自动获得FlashInfer MXFP4内核,无需任何配置更改,性能提升显著。对其他架构无影响。对系统:代码量增加约330行,主要集中在权重处理函数,未引入新的外部依赖。对团队:维护成本略有增加,但内核选择逻辑更清晰,且与SM90复用padding策略降低了长期维护负担。

量化层核心路径变更 SM120 依赖硬件可用性 权重处理逻辑新增 依赖 FlashInfer 版本

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论