diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py index 80735f3da..f04479dbb 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -1376,18 +1376,18 @@ class FlashInferFP4MoE(FusedMoE): hidden_states_scale=hs_scale_linear.view(torch.float8_e4m3fn).reshape( *hs_scale_linear.shape[:-1], -1 ), - gemm1_weights=self.gemm1_weights_fp4_shuffled.data, - gemm1_weights_scale=self.gemm1_scales_fp4_shuffled.data.view( - torch.float8_e4m3fn - ), + # B300 port: HEAD's align_fp4_moe_weights_for_flashinfer_trtllm now + # stores the prepared (shuffled) fp4 weights/scales into the standard + # param names instead of gemm*_fp4_shuffled (weights=uint8, scales + # already fp8_e4m3fn — the .view below is a no-op kept for clarity). + gemm1_weights=self.w13_weight.data, + gemm1_weights_scale=self.w13_weight_scale.data.view(torch.float8_e4m3fn), gemm1_bias=None, gemm1_alpha=None, gemm1_beta=None, gemm1_clamp_limit=None, - gemm2_weights=self.gemm2_weights_fp4_shuffled.data, - gemm2_weights_scale=self.gemm2_scales_fp4_shuffled.data.view( - torch.float8_e4m3fn - ), + gemm2_weights=self.w2_weight.data, + gemm2_weights_scale=self.w2_weight_scale.data.view(torch.float8_e4m3fn), gemm2_bias=None, output1_scale_scalar=self.g1_scale_c.data, output1_scale_gate_scalar=self.g1_alphas.data,