Prhub

#46915 [Bugfix][Test] Fix test_flashinfer_cutlass_mxfp4_fused_moe on sm90 (stale weight/scale interleave)

原始 PR 作者 wentian-byte 合并时间 2026-06-28 02:30 文件变更 1 提交数 4 评论 3 代码增减 +14 / -15

执行摘要

修复 SM90 MXFP4 MoE 测试中的权重 / 缩放交错问题

Issue #46585 报告:在 H20(sm90)上,test_flashinfer_cutlass_mxfp4_fused_moe 出现约 89% 的精度不匹配,远超 20% 的阈值。经诊断,测试为 SM90 混合输入 cutlass MoE 准备的权重/缩放布局与 kernel 期望不一致——权重完全没有交错,缩放使用了手写的 _interleave_scales_lastdim_by4 函数,但该函数在 K > 128 时无法正确布局,导致 FP4→BF16 LUT 读取错误的字节位置。

该 PR 值得快速合并和精读,因为它展示了一个常见的测试陷阱:手写数据布局函数与 kernel 期望不一致导致测试假阴性。学习点是优先使用 kernel 配套库提供的官方助手函数,而非手动实现。

讨论亮点

该 PR 仅涉及 1 个文件,review 评论较少。yzong-rh 和 yewentao256 均批准,yzong-rh 确认在 1×H100 上本地测试通过。pre-commit 检查因缺乏 verified 状态而标红,但已通过格式修复提交解决。

实现拆解

  1. 识别根因:分析发现 test_flashinfer_cutlass_mxfp4_fused_moe 在调用 cutlass_fused_moe(use_w4_group_scaling=True) 时,权重(packed 4-bit)未做交错,缩放使用了手写函数 _interleave_scales_lastdim_by4,与 SM90 混合 dtype GEMM 要求的交错布局不匹配。

  2. 移除手写函数:删除文件中不再使用的 _interleave_scales_lastdim_by4 函数(约 12 行),并调整导入顺序以满足 lint 要求。

  3. 改用 flashinfer 官方函数:在测试中导入 flashinfer.fused_moe 中的 interleave_moe_weights_for_sm90_mixed_gemminterleave_moe_scales_for_sm90_mixed_gemm,分别对权重和缩放进行正确的交错处理。对于权重调用时指定 quant_type="fp4",以匹配 FP4 量化格式。

  4. 验证:在 1×H20(sm90, flashinfer 0.6.12)上验证,测试结果从 1 个失败(89% 不匹配)变为全部 8 个参数化用例通过。

文件 模块 状态 重要度
tests/kernels/moe/test_ocp_mx_moe.py MoE 测试 modified 5.41

关键符号

_interleave_scales_lastdim_by4 interleave_moe_weights_for_sm90_mixed_gemm interleave_moe_scales_for_sm90_mixed_gemm

关键源码片段

tests/kernels/moe/test_ocp_mx_moe.py test-coverage

该文件包含所有变更:移除手写缩放交错函数 `_interleave_scales_lastdim_by4`,改用 flashinfer 官方交错函数 `interleave_moe_weights_for_sm90_mixed_gemm` 和 `interleave_moe_scales_for_sm90_mixed_gemm`,并调整导入顺序。

# 旧代码:手写缩放交错,未处理权重交错
# w13_s_inter = _interleave_scales_lastdim_by4(w13_s)
# w2_s_inter = _interleave_scales_lastdim_by4(w2_scale)# 新代码:使用 flashinfer 官方助手函数,同时处理权重和缩放
from flashinfer.fused_moe import (
    interleave_moe_scales_for_sm90_mixed_gemm,
    interleave_moe_weights_for_sm90_mixed_gemm,
)# 对权重应用交错:quant_type="fp4" 匹配 FP4 量化格式
w13_q_swapped = interleave_moe_weights_for_sm90_mixed_gemm(
    w13_q_swapped, quant_type="fp4"
)
w2_q = interleave_moe_weights_for_sm90_mixed_gemm(w2_q, quant_type="fp4")# 对缩放应用交错
w13_s_inter = interleave_moe_scales_for_sm90_mixed_gemm(w13_s)
w2_s_inter = interleave_moe_scales_for_sm90_mixed_gemm(w2_scale)

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

风险极低。变更局限在一个测试文件内,仅修改了权重和缩放的预处理方式,不涉及任何生产代码或 kernel 逻辑。测试本身在 sm90 GPU 上才执行,对其他架构无影响。手动函数 _interleave_scales_lastdim_by4 被删除,但该函数仅在此测试中使用,且已被官方函数替代。

  • 对用户:无直接影响,仅修复了特定硬件(Hopper GPU)上的测试失败。
  • 对系统:提高了测试覆盖率在 sm90 上的可靠性,使得后续开发能正确验证 MXFP4 MoE kernel。
  • 对团队:消除了一个已知的 CI 噪音,降低维护负担。

关联 Issue

#46585 [Bug]: test_flashinfer_cutlass_mxfp4_fused_moe accuracy mismatch on H20 (sm90) — 89% mismatch vs 20% threshold

完整报告

参与讨论