Files
DeepGEMM/megamoe-research-reports/pr352_sm90_split_l1l2_megamoe_code_review.md
T
Xinyi Liu 062cb160cf Phase 0: SM90 MegaMoE design doc, reference baseline, nsys script
- MEGAMOE_SM90_DESIGN.md: complete design document with finalized decisions
  (fused single kernel, cooperative + single-WG, dynamic BLOCK_M, etc.)
- tests/test_mega_moe_sm90.py: PyTorch FP32/BF16 reference implementation
  for dispatch → L1 GEMM → SwiGLU → L2 GEMM → combine pipeline
- scripts/run_nsys_mega_moe_sm90.sh: nsys profiling wrapper script
- megamoe-research-reports/: research analysis of PR304/323/347/352/357/360
2026-06-16 18:01:12 +08:00

5.9 KiB
Raw Blame History

PR352 SM90 Split L1/L2 MegaMoE 代码review报告

范围

  • Worktree: pr-352
  • HEAD: 655075ef3
  • 审查方式: 代码review
  • 主题: SM90 MegaMoE 增强 — phase 分离、compact frontend、多 epilogue 策略

实现概述

SM90 FP8 MegaMoE Enhanced Kernel

deep_gemm/include/deep_gemm/impls/sm90_fp8_mega_moe.cuh (~2507 lines)

本 kernel 是 pr-323 的重度增强版。核心创新是使用 C++ 宏 将同一个 kernel body 参数化为多种配置变体。

warp 分工与线程布局(按 Phase 策略枚举):

Compact 模式 (kCompactFrontendWarpgroup, topk=2, cluster=1, default):

Warp Index 所属 Warpgroup Role
0–1 (2 warps, 64 threads) WG0 Dispatch
2–3 (2 warps, 64 threads) WG0 TMA A+SFA / B+SFB
4–11 (8 warps, 256 threads) WG1, WG2 Math WGMMA + epilogue + combine

精确 warp 统计: 12 warps = 384 threads = 3 warpgroups

Serial/Wide N 模式 (BLOCK_M=32, topk≥8, cluster=2):

Warp Index 所属 Warpgroup Role
0–3 (4 warps, 128 threads) WG0 Dispatch
4–7 (4 warps, 128 threads) WG1 TMA A+SFA / B+SFB / MMA issue / idle
8–11 (4 warps, 128 threads) WG2 Math WGMMA + epilogue + combine

精确 warp 统计: 12 warps = 384 threads = 3 warpgroups

Serial/Wide N 模式 (BLOCK_M=64, topk≥8, cluster=2):

Warp Index 所属 Warpgroup Role
0–3 (4 warps, 128 threads) WG0 Dispatch
4–7 (4 warps, 128 threads) WG1 TMA A+SFA / B+SFB / MMA issue / idle
8–15 (8 warps, 256 threads) WG2, WG3 Math WGMMA + epilogue + combine

精确 warp 统计: 16 warps = 512 threads = 4 warpgroups |

Compact Frontend 模式 (dispatch+TMA 共享 WG0):

当 kCompactFrontendWarpgroup 为 true 时,dispatch (2 warps) + TMA (2 warps) 共享一个 128-thread warpgroup:

Config Dispatch TMA Epilogue Total Threads Reg Budget
Compact (default) 64 64 256 384 48+48+208→59,392
Wide (≥128) 128 128 256 512 48+40+208=64,512

Phase 策略(通过宏模板参数选择):

模式 说明
kSerialNWarpgroups math warpgroups 串行处理 L1 再 L2
kWideNWarpgroups math warpgroups 使用更多 N block 并行
kFusedL1L2Warpgroups 同一 WG 同时处理 L1 和 L2(与 pr-323 类似)
kUseMMASync 使用 MMA sync 路径(BLOCK_M=32 时启用)
kCompactFrontendWarpgroup dispatch+TMA 共享 warpgroup

寄存器分配(宏推导):

kNumDispatchRegisters = 48
kNumNonEpilogueRegisters = kCompactFrontendWarpgroup ? 48 : 40
kNumEpilogueRegisters = (kSerialNWarpgroups or kWideNWarpgroups) ? 256
                      : ((kUseMMASync and BLOCK_M==32) ? 240 : 208)

Compact mode 下 non-epilogue 必须与 dispatch 使用相同的 register 数 (48),因为它们共享同一个 warpgroup(WG0)。

megamoe_sm90 branch 具体配置(来自 csrc/jit_kernels/heuristics/mega_moe.hpp 的 1025 行扩展启发式):

Topk Tokens BLOCK_M BLOCK_N BLOCK_K Epilogue Threads Cluster Compact Phase 备注
2 所有 64 128 128 256 1 Yes Fused
8 ≤128 32 128 128 128 2 No Serial NW kSerialNWarpgroups 硬编码为 false,实际不可达
8 ≤576 64 128 128 256 2 No Serial NW 同上
8 >576 64 256 128 256 2 No Serial NW 同上
9+ ≤128 32 128 128 128 2 No Serial NW 同上
9+ ≤576 64 128 128 256 2 No Serial NW 同上
9+ >576 64 256 128 256 2 No Serial NW 同上

代码review发现

高: 宏驱动模板实例化导致编译膨胀

每次调用 sm90_fp8_mega_moe 会通过宏展开 4 个 kernel 实例(INSTANTIATE_KERNEL_WITH_PHASE_POLICY × 1 fuse + 1 serial_N + 1 wide_N + 1 mma_sync)。每个实例有独立的 __launch_bounds__ 和 register 分配,JIT 编译时间较长。

中: Compact frontend 下 dispatch 使用 48 reg/warp — 仍然偏紧

dispatch warp 需要 rank round-robin 选择、SF copy 等复杂操作。48 reg/warp 可能在某些 shape 下触发 register spilling,但由于与 TMA warp 共享 WG0,无法单独增加而不破坏 budget。

中: Scheduler 保留严格整除约束

PR352 分叉自 #316(早于 PR347),因此其 scheduler (mega_moe.cuh:38) 仍使用 kNumExpertsPerRank % kNumExpertsPerWave == 0 的严格约束。PR347 的放宽修复(> 0 && <=)未被合入。对非 2 的幂 per-rank expert 数的 shape,可能触发编译期断言失败。

低: Compact 模式 register 余量尚充足

Compact 模式实际 register 消耗为 48×64+48×64+208×256=59,392(文档之前误写为 64,512),占预算约 90.6%,有约 5K register 余量。与 "tight" 的描述不同,实际还有一定 headroom。

中: 启发式文件增长到 1025 行,可维护性下降

heuristics/mega_moe.hpp 混合了 SM100 和 SM90 路径,且 SM90 部分包含大量硬编码的 topk/token 分支表。建议拆分为 mega_moe_sm90.hpp + mega_moe_sm100.hpp(类似 pr-360 的做法)。

低: Serial NW 模式下 kNumEpilogueRegisters=256 可能溢出

256 reg/warp × 256 epilogue threads = 65536,超过 64K reg budget。需要确认是否有其他约束(如减少 dispatch 线程数)来保证不溢出。

正面评价

  • 宏驱动架构灵活,一个 kernel body 支持 4 种执行策略
  • Compact frontend 优化了资源利用率(H100 上 dispatch 不需要独占 warpgroup)
  • 多种 phase 策略覆盖了不同 token 量的最优执行路径
  • Serial NW 模式的 epilogue 获得 256 reg/warp,适合计算密集场景

建议检查清单

  • 验证 Serial NW 256 reg/warp 配置不超标
  • 实测 compact frontend 下 dispatch warp 的 spilling 情况
  • 考虑拆分 heuristics 文件降低维护成本
  • 确认所有 4 种 phase 策略的 correctness 测试覆盖