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):