- 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
5.9 KiB
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 测试覆盖