From c016e66e231ca50452d3875ddc1f40d4035425a0 Mon Sep 17 00:00:00 2001 From: biondizzle Date: Thu, 14 May 2026 11:27:58 +0000 Subject: [PATCH] Add CUDA sync + NaN/Inf check after each expert GEMM in grouped kernel --- src/nvfp4_megamoe_kernel/cutlass_nvfp4_gemm/kernel.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/src/nvfp4_megamoe_kernel/cutlass_nvfp4_gemm/kernel.py b/src/nvfp4_megamoe_kernel/cutlass_nvfp4_gemm/kernel.py index f62ae2eb..5f8315aa 100644 --- a/src/nvfp4_megamoe_kernel/cutlass_nvfp4_gemm/kernel.py +++ b/src/nvfp4_megamoe_kernel/cutlass_nvfp4_gemm/kernel.py @@ -79,6 +79,15 @@ def cutlass_grouped_nvfp4_gemm( M_expert, N, K, ) # (M_expert, N) bfloat16 + # Check for CUDA errors after each expert GEMM + err = torch.cuda.current_stream().synchronize() + + # Validate output + if torch.isnan(expert_out).any() or torch.isinf(expert_out).any(): + if MEGA_MOE_DEBUG: + print(f"[cutlass_grouped_gemm] WARNING: expert {e} produced NaN/Inf, skipping") + continue + # Scatter back with routing weights for t_idx, token_idx in enumerate(token_indices): for k_idx in range(num_topk):