From 5e670760bd18206cdea5df45e129b22e762b1ee6 Mon Sep 17 00:00:00 2001 From: leavelet Date: Sun, 21 Jun 2026 16:24:17 +0000 Subject: [PATCH] B300 NVFP4: fix FlashInferFP4MoE.forward attr names to match HEAD weight-prep MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The team-custom FlashInferFP4MoE.forward_impl read gemm{1,2}_{weights,scales}_fp4_shuffled, but the HEAD-advanced align_fp4_moe_weights_for_flashinfer_trtllm now stores into w13_weight/w2_weight/ w13_weight_scale/w2_weight_scale (weights uint8, scales already fp8). Rename the 4 reads (Option a: preserves the team forward incl. its DeepSeekV3 fp32 router_logits cast — GLM-5.2 uses DeepSeekV3 routing, so option-b delegation that drops the cast was avoided). g1_scale_c/g1_alphas/g2_alphas/ w13_input_scale_quant already match HEAD writes (opus-verified 1:1 format/semantics, no silent-garbage). Cannot wholesale-port layer.py/ep_moe (team EP/CP infra). Watch at launch: draft accept-len + sane output. Co-Authored-By: Claude Opus 4.8 (1M context) --- .../srt/layers/moe/fused_moe_triton/layer.py | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) 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,