diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py index 10991637f..4d2e9a110 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py @@ -21,25 +21,49 @@ from sglang.srt.layers.quantization.fp8_kernel import ( per_token_group_quant_fp8, scaled_fp8_quant, ) -from sglang.srt.utils.common import next_power_of_2 +from sglang.srt.utils.common import ( + is_flashinfer_available, + is_sm120_supported, + next_power_of_2, +) if TYPE_CHECKING: from sglang.srt.layers.moe.token_dispatcher import ( StandardCombineInput, StandardDispatchOutput, ) +if is_flashinfer_available() and is_sm120_supported(): + from flashinfer import fp4_quantize +else: + from sgl_kernel import scaled_fp4_quant as fp4_quantize -def align_fp8_moe_weights_for_flashinfer_trtllm(layer: Module) -> None: - """Prepare FP8 MoE weights/scales for FlashInfer TRT-LLM kernels.""" +def align_fp8_moe_weights_for_flashinfer_trtllm( + layer: Module, swap_w13_halves: bool = False +) -> None: + """Prepare FP8 MoE weights/scales for FlashInfer TRT-LLM kernels. + + Args: + layer: The MoE layer to process. + swap_w13_halves: If True, swap W13 halves from [Up, Gate] to [Gate, Up]. + This is needed for ModelOpt FP8 checkpoints which store weights in + [Up, Gate] order, while regular FP8 checkpoints store them in [Gate, Up]. + """ from flashinfer import reorder_rows_for_gated_act_gemm, shuffle_matrix_a - # Note: No need to swap W13 halves, they are already in the correct order: - # [Gate, Up] w13_weight = cast(torch.Tensor, layer.w13_weight) w2_weight = cast(torch.Tensor, layer.w2_weight) num_experts, two_n, hidden = w13_weight.shape + # Optionally swap W13 halves: [Up, Gate] -> [Gate, Up] + if swap_w13_halves: + inter = two_n // 2 + w13_weight = ( + w13_weight.reshape(num_experts, 2, inter, hidden) + .flip(dims=[1]) + .reshape(num_experts, two_n, hidden) + ) + w13_interleaved_list = [ reorder_rows_for_gated_act_gemm(w13_weight[i]) for i in range(num_experts) ] @@ -90,6 +114,69 @@ def align_fp8_moe_weights_for_flashinfer_trtllm(layer: Module) -> None: layer.output2_scales_scalar = Parameter(output2_scales_scalar, requires_grad=False) +def align_fp4_moe_weights_for_flashinfer_trtllm(layer: Module) -> None: + """Prepare FP4 MoE weights/scales for FlashInfer TRT-LLM kernels. + + This function handles the weight transformation needed for FP4 TRTLLM MoE: + - Reorders weights for gated activation GEMM + - Shuffles weights and scales for transposed MMA output + - Computes the output scale factors + """ + from sglang.srt.layers.quantization.utils import ( + prepare_static_weights_for_trtllm_fp4_moe, + ) + + w13_weight = cast(torch.Tensor, layer.w13_weight) + w2_weight = cast(torch.Tensor, layer.w2_weight) + w13_weight_scale = cast(torch.Tensor, layer.w13_weight_scale) + w2_weight_scale = cast(torch.Tensor, layer.w2_weight_scale) + + ( + gemm1_weights_fp4_shuffled, + gemm1_scales_fp4_shuffled, + gemm2_weights_fp4_shuffled, + gemm2_scales_fp4_shuffled, + ) = prepare_static_weights_for_trtllm_fp4_moe( + w13_weight, + w2_weight, + w13_weight_scale, + w2_weight_scale, + w2_weight.size(-2), # hidden_size + w13_weight.size(-2) // 2, # intermediate_size + w13_weight.size(0), # num_experts + ) + + # Set flashinfer parameters + layer.gemm1_weights_fp4_shuffled = Parameter( + gemm1_weights_fp4_shuffled, requires_grad=False + ) + layer.gemm2_weights_fp4_shuffled = Parameter( + gemm2_weights_fp4_shuffled, requires_grad=False + ) + layer.gemm1_scales_fp4_shuffled = Parameter( + gemm1_scales_fp4_shuffled, requires_grad=False + ) + layer.gemm2_scales_fp4_shuffled = Parameter( + gemm2_scales_fp4_shuffled, requires_grad=False + ) + + # Compute additional scaling factor needed for TRT-LLM + w2_input_scale_quant = cast(torch.Tensor, layer.w2_input_scale_quant) + g1_alphas = cast(torch.Tensor, layer.g1_alphas) + layer.g1_scale_c = Parameter( + (w2_input_scale_quant * g1_alphas).to(torch.float32), + requires_grad=False, + ) + + # Clean up weights that won't be used by TRT-LLM + del ( + layer.w2_weight, + layer.w2_weight_scale, + layer.w13_weight, + layer.w13_weight_scale, + ) + + @dataclass class FlashInferTrtllmFp8MoeQuantInfo(MoeQuantInfo): """Quantization payload consumed by FlashInfer TRT-LLM FP8 MoE kernels.""" @@ -102,6 +189,7 @@ class FlashInferTrtllmFp8MoeQuantInfo(MoeQuantInfo): global_num_experts: int local_expert_offset: int local_num_experts: int + intermediate_size: int routing_method_type: int @@ -116,9 +204,9 @@ class FlashInferTrtllmFp8MoeQuantInfo(MoeQuantInfo): output1_scales_scalar: torch.Tensor | None = None output1_scales_gate_scalar: torch.Tensor | None = None output2_scales_scalar: torch.Tensor | None = None + use_routing_scales_on_input: bool = False -@register_fused_func("none", "flashinfer_trtllm") def fused_experts_none_to_flashinfer_trtllm_fp8( dispatch_output: StandardDispatchOutput, quant_info: FlashInferTrtllmFp8MoeQuantInfo, @@ -180,9 +268,11 @@ def fused_experts_none_to_flashinfer_trtllm_fp8( gemm2_weights_scale=quant_info.w2_weight_scale_inv, num_experts=quant_info.global_num_experts, top_k=topk_config.top_k, - n_group=topk_config.num_expert_group, - topk_group=topk_config.topk_group, - intermediate_size=int(quant_info.w2_weight.shape[2]), + n_group=( + topk_config.num_expert_group if topk_config.num_expert_group else 0 + ), + topk_group=topk_config.topk_group if topk_config.topk_group else 0, + intermediate_size=quant_info.intermediate_size, local_expert_offset=quant_info.local_expert_offset, local_num_experts=quant_info.local_num_experts, routed_scaling_factor=( @@ -219,9 +309,11 @@ def fused_experts_none_to_flashinfer_trtllm_fp8( output2_scales_scalar=quant_info.output2_scales_scalar, num_experts=quant_info.global_num_experts, top_k=topk_config.top_k, - n_group=topk_config.num_expert_group, - topk_group=topk_config.topk_group, - intermediate_size=int(quant_info.w2_weight.shape[2]), + n_group=( + topk_config.num_expert_group if topk_config.num_expert_group else 0 + ), + topk_group=topk_config.topk_group if topk_config.topk_group else 0, + intermediate_size=quant_info.intermediate_size, local_expert_offset=quant_info.local_expert_offset, local_num_experts=quant_info.local_num_experts, routed_scaling_factor=( @@ -229,9 +321,179 @@ def fused_experts_none_to_flashinfer_trtllm_fp8( if runner_config.routed_scaling_factor is not None else 1.0 ), - use_routing_scales_on_input=False, + use_routing_scales_on_input=quant_info.use_routing_scales_on_input, routing_method_type=routing_method_type, tune_max_num_tokens=next_power_of_2(a_q.shape[0]), ) return StandardCombineInput(hidden_states=output) + + +@dataclass +class FlashInferTrtllmFp4MoeQuantInfo(MoeQuantInfo): + """Quantization payload consumed by FlashInfer TRT-LLM FP4 MoE kernels.""" + + # Shuffled FP4 weights (processed by align_fp4_moe_weights_for_flashinfer_trtllm) + gemm1_weights_fp4_shuffled: torch.Tensor + gemm2_weights_fp4_shuffled: torch.Tensor + gemm1_scales_fp4_shuffled: torch.Tensor + gemm2_scales_fp4_shuffled: torch.Tensor + + # Scaling factors + g1_scale_c: torch.Tensor + g1_alphas: torch.Tensor + g2_alphas: torch.Tensor + w13_input_scale_quant: torch.Tensor + + # Expert-parallel metadata + global_num_experts: int + local_expert_offset: int + local_num_experts: int + intermediate_size_per_partition: int + + routing_method_type: int + + +def quantize_hidden_states_fp4( + hidden_states: torch.Tensor, + input_scale_quant: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + """ + Quantize hidden states to FP4 for TRTLLM MoE. + + Global scale factor is set by ModelOptNvFp4FusedMoEMethod during weight loading. + Only block scales are computed at runtime for efficiency. + + Returns (packed_fp4_uint8, scale_float8_e4m3fn_runtime) + """ + + # flashinfer.fp4_quantize returns (packed_uint8, scale_fp8) + # Only the block scales are computed at runtime + hs_fp4_bytes, hs_sf_bytes = fp4_quantize( + hidden_states, + input_scale_quant, + 16, # sf_vec_size + False, # use_ue8m0 + False, # is_sf_swizzled_layout + ) + + seq_len, hidden_size = hidden_states.shape + hs_fp4 = hs_fp4_bytes.reshape(seq_len, hidden_size // 2) + # TRT-LLM expects hidden state scales shaped as [seq_len, hidden_size // 16] + hs_sf = hs_sf_bytes.view(torch.float8_e4m3fn).reshape(seq_len, hidden_size // 16) + + return hs_fp4, hs_sf + + +def fused_experts_none_to_flashinfer_trtllm_fp4( + dispatch_output: StandardDispatchOutput, + quant_info: FlashInferTrtllmFp4MoeQuantInfo, + runner_config: MoeRunnerConfig, +) -> StandardCombineInput: + """FlashInfer TRTLLM FP4 MoE forward pass. + + This function handles the FP4 TRTLLM MoE path that was previously in + FlashInferFP4MoE.forward_impl and ModelOptNvFp4FusedMoEMethod.apply. + """ + from flashinfer.fused_moe import trtllm_fp4_block_scale_moe + + from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput + from sglang.srt.layers.moe.topk import TopKOutputChecker + from sglang.srt.layers.moe.utils import RoutingMethodType + + assert runner_config.activation == "silu", "Only silu is supported for FP4 MoE." + assert runner_config.is_gated, "Only gated MoEs are supported for FP4 MoE." + + hidden_states = dispatch_output.hidden_states + topk_output = dispatch_output.topk_output + assert TopKOutputChecker.format_is_bypassed(topk_output) + + router_logits = topk_output.router_logits + topk_config = topk_output.topk_config + routing_method_type = quant_info.routing_method_type + + # Quantize hidden states to FP4 + hs_fp4, hs_scale_linear = quantize_hidden_states_fp4( + hidden_states, quant_info.w13_input_scale_quant + ) + + # DeepSeekV3 style routing requires float32 router logits + if routing_method_type == RoutingMethodType.DeepSeekV3: + router_logits = router_logits.to(torch.float32) + + correction_bias = ( + None + if topk_config.correction_bias is None + else topk_config.correction_bias.to(hidden_states.dtype) + ) + + with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + num_tokens = hs_fp4.shape[0] + hidden_size = ( + hs_fp4.shape[-1] * 2 if hs_fp4.dtype == torch.uint8 else hs_fp4.shape[-1] + ) + symm_output = torch.empty( + num_tokens, hidden_size, dtype=torch.bfloat16, device=hs_fp4.device + ) + + result = trtllm_fp4_block_scale_moe( + routing_logits=router_logits, + routing_bias=correction_bias, + hidden_states=hs_fp4, + hidden_states_scale=hs_scale_linear.view(torch.float8_e4m3fn).flatten(), + gemm1_weights=quant_info.gemm1_weights_fp4_shuffled, + gemm1_weights_scale=quant_info.gemm1_scales_fp4_shuffled.view( + torch.float8_e4m3fn + ), + gemm1_bias=None, + gemm1_alpha=None, + gemm1_beta=None, + gemm1_clamp_limit=None, + gemm2_weights=quant_info.gemm2_weights_fp4_shuffled, + gemm2_weights_scale=quant_info.gemm2_scales_fp4_shuffled.view( + torch.float8_e4m3fn + ), + gemm2_bias=None, + output1_scale_scalar=quant_info.g1_scale_c, + output1_scale_gate_scalar=quant_info.g1_alphas, + output2_scale_scalar=quant_info.g2_alphas, + num_experts=quant_info.global_num_experts, + top_k=topk_config.top_k, + n_group=topk_config.num_expert_group, + topk_group=topk_config.topk_group, + intermediate_size=quant_info.intermediate_size_per_partition, + local_expert_offset=quant_info.local_expert_offset, + local_num_experts=quant_info.local_num_experts, + routed_scaling_factor=runner_config.routed_scaling_factor, + tile_tokens_dim=None, + routing_method_type=( + routing_method_type + if routing_method_type is not None + else RoutingMethodType.Default + ), + do_finalize=True, + tune_max_num_tokens=next_power_of_2(hs_fp4.shape[0]), + output=symm_output, + )[0] + + return StandardCombineInput(hidden_states=result) + + +@register_fused_func("none", "flashinfer_trtllm") +def fused_experts_none_to_flashinfer_trtllm( + dispatch_output: StandardDispatchOutput, + quant_info: MoeQuantInfo, + runner_config: MoeRunnerConfig, +) -> StandardCombineInput: + """Dispatch to FP8 or FP4 FlashInfer TRT-LLM MoE based on quant_info type.""" + if isinstance(quant_info, FlashInferTrtllmFp4MoeQuantInfo): + return fused_experts_none_to_flashinfer_trtllm_fp4( + dispatch_output, quant_info, runner_config + ) + if isinstance(quant_info, FlashInferTrtllmFp8MoeQuantInfo): + return fused_experts_none_to_flashinfer_trtllm_fp8( + dispatch_output, quant_info, runner_config + ) + raise TypeError( + f"Unexpected quant_info type for flashinfer_trtllm: {type(quant_info)}" + ) diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py index 8b5d2ea0a..cb3ca7e0d 100644 --- a/python/sglang/srt/layers/quantization/fp8.py +++ b/python/sglang/srt/layers/quantization/fp8.py @@ -1467,6 +1467,7 @@ class Fp8MoEMethod(FusedMoEMethodBase): global_num_experts=global_num_experts, local_expert_offset=moe_ep_rank * num_local_experts, local_num_experts=num_local_experts, + intermediate_size=layer.w2_weight.shape[2], routing_method_type=int( getattr(layer, "routing_method_type", RoutingMethodType.DeepSeekV3) ), diff --git a/python/sglang/srt/layers/quantization/modelopt_quant.py b/python/sglang/srt/layers/quantization/modelopt_quant.py index 7a0dfed08..6c3b1288e 100755 --- a/python/sglang/srt/layers/quantization/modelopt_quant.py +++ b/python/sglang/srt/layers/quantization/modelopt_quant.py @@ -44,7 +44,6 @@ from sglang.srt.layers.quantization.utils import ( convert_to_channelwise, is_layer_skipped, per_tensor_dequantize, - prepare_static_weights_for_trtllm_fp4_moe, requantize_with_max_scale, swizzle_blockscale, ) @@ -600,81 +599,12 @@ class ModelOptFp8MoEMethod(FusedMoEMethodBase): # Align FP8 weights to FlashInfer per-tensor kernel layout if enabled if get_moe_runner_backend().is_flashinfer_trtllm(): - from flashinfer import reorder_rows_for_gated_act_gemm, shuffle_matrix_a - - # 1) Swap W13 halves: [Up, Gate] -> [Gate, Up] expected by FI - num_experts, two_n, hidden = layer.w13_weight.shape - inter = two_n // 2 - w13_swapped = ( - layer.w13_weight.reshape(num_experts, 2, inter, hidden) - .flip(dims=[1]) - .reshape(num_experts, two_n, hidden) + from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import ( + align_fp8_moe_weights_for_flashinfer_trtllm, ) - # 2) Reorder rows for fused gated activation (W13) - w13_interleaved = [ - reorder_rows_for_gated_act_gemm(w13_swapped[i]) - for i in range(num_experts) - ] - w13_interleaved = torch.stack(w13_interleaved).reshape( - num_experts, two_n, hidden - ) - - # 3) Shuffle weights for transposed MMA output (both W13, W2) - epilogue_tile_m = 128 - w13_shuffled = [ - shuffle_matrix_a(w13_interleaved[i].view(torch.uint8), epilogue_tile_m) - for i in range(num_experts) - ] - w2_shuffled = [ - shuffle_matrix_a(layer.w2_weight[i].view(torch.uint8), epilogue_tile_m) - for i in range(num_experts) - ] - - layer.w13_weight = Parameter( - torch.stack(w13_shuffled).view(torch.float8_e4m3fn), - requires_grad=False, - ) - layer.w2_weight = Parameter( - torch.stack(w2_shuffled).view(torch.float8_e4m3fn), - requires_grad=False, - ) - - # Precompute and register per-expert output scaling factors for FI MoE - if get_moe_runner_backend().is_flashinfer_trtllm(): - # Note: w13_input_scale and w2_input_scale are scalar Parameters post-reduction - assert ( - hasattr(layer, "w13_input_scale") and layer.w13_input_scale is not None - ) - assert hasattr(layer, "w2_input_scale") and layer.w2_input_scale is not None - assert ( - hasattr(layer, "w13_weight_scale") - and layer.w13_weight_scale is not None - ) - assert ( - hasattr(layer, "w2_weight_scale") and layer.w2_weight_scale is not None - ) - - input_scale = layer.w13_input_scale.to(torch.float32) - activation_scale = layer.w2_input_scale.to(torch.float32) - w13_weight_scale = layer.w13_weight_scale.to(torch.float32) - w2_weight_scale = layer.w2_weight_scale.to(torch.float32) - - output1_scales_scalar = ( - w13_weight_scale * input_scale * (1.0 / activation_scale) - ) - output1_scales_gate_scalar = w13_weight_scale * input_scale - output2_scales_scalar = activation_scale * w2_weight_scale - - layer.output1_scales_scalar = Parameter( - output1_scales_scalar, requires_grad=False - ) - layer.output1_scales_gate_scalar = Parameter( - output1_scales_gate_scalar, requires_grad=False - ) - layer.output2_scales_scalar = Parameter( - output2_scales_scalar, requires_grad=False - ) + # ModelOpt FP8 stores weights in [Up, Gate] order, so we need to swap + align_fp8_moe_weights_for_flashinfer_trtllm(layer, swap_w13_halves=True) elif get_moe_runner_backend().is_flashinfer_cutlass(): assert ( hasattr(layer, "w13_input_scale") and layer.w13_input_scale is not None @@ -725,28 +655,19 @@ class ModelOptFp8MoEMethod(FusedMoEMethodBase): get_moe_runner_backend().is_flashinfer_trtllm() and TopKOutputChecker.format_is_bypassed(topk_output) ): - router_logits = topk_output.router_logits + from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import ( + FlashInferTrtllmFp8MoeQuantInfo, + fused_experts_none_to_flashinfer_trtllm_fp8, + ) + from sglang.srt.layers.moe.utils import RoutingMethodType + topk_config = topk_output.topk_config - # Constraints + # Constraints for ModelOpt FP8 MoE assert ( self.moe_runner_config.activation == "silu" ), "Only silu is supported for flashinfer fp8 moe" - from flashinfer import RoutingMethodType - from flashinfer.fused_moe import trtllm_fp8_per_tensor_scale_moe - - correction_bias = ( - None - if topk_config.correction_bias is None - else topk_config.correction_bias - ) - # Pre-quantize activations to FP8 per-tensor using provided input scale - x_fp8, _ = scaled_fp8_quant(x, layer.w13_input_scale) - - use_routing_scales_on_input = True - routed_scaling_factor = self.moe_runner_config.routed_scaling_factor - # Enforce Llama4 routing for ModelOpt FP8 MoE for now. # TODO(brayden): support other routing methods assert topk_config.top_k == 1, "ModelOpt FP8 MoE requires top_k==1" @@ -756,50 +677,26 @@ class ModelOptFp8MoEMethod(FusedMoEMethodBase): assert ( not topk_config.topk_group ), "ModelOpt FP8 MoE does not support grouped top-k" - routing_method_type = RoutingMethodType.Llama4 - # FlashInfer TRTLLM requires routing_logits (and bias) to be bfloat16 - routing_logits_cast = router_logits.to(torch.bfloat16) - routing_bias_cast = ( - None if correction_bias is None else correction_bias.to(torch.bfloat16) + quant_info = FlashInferTrtllmFp8MoeQuantInfo( + w13_weight=layer.w13_weight, + w2_weight=layer.w2_weight, + global_num_experts=layer.num_experts, + local_expert_offset=layer.moe_ep_rank * layer.num_local_experts, + local_num_experts=layer.num_local_experts, + intermediate_size=layer.w2_weight.shape[2], + routing_method_type=RoutingMethodType.Llama4, + block_quant=False, + w13_input_scale=layer.w13_input_scale, + output1_scales_scalar=layer.output1_scales_scalar, + output1_scales_gate_scalar=layer.output1_scales_gate_scalar, + output2_scales_scalar=layer.output2_scales_scalar, + use_routing_scales_on_input=True, ) - with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() - ): - # FIXME: there is a bug in the trtllm_fp8_block_scale_moe. - # It ignored the `output`` argument. https://github.com/flashinfer-ai/flashinfer/blob/da01b1bd8f9f22aec8c0eea189ad54860b034947/flashinfer/fused_moe/core.py#L1323-L1325 - # so we put the whole function under the ``use_symmetric_memory`` context manager. - # If the bug is fixed, we can only put the output tensor allocation under the context manager. - output = trtllm_fp8_per_tensor_scale_moe( - routing_logits=routing_logits_cast, - routing_bias=routing_bias_cast, - hidden_states=x_fp8, - gemm1_weights=layer.w13_weight, - output1_scales_scalar=layer.output1_scales_scalar, - output1_scales_gate_scalar=layer.output1_scales_gate_scalar, - gemm2_weights=layer.w2_weight, - output2_scales_scalar=layer.output2_scales_scalar, - num_experts=layer.num_experts, - top_k=topk_config.top_k, - n_group=0, - topk_group=0, - intermediate_size=layer.w2_weight.shape[2], - local_expert_offset=layer.moe_ep_rank * layer.num_local_experts, - local_num_experts=layer.num_local_experts, - routed_scaling_factor=( - routed_scaling_factor - if routed_scaling_factor is not None - else 1.0 - ), - use_routing_scales_on_input=use_routing_scales_on_input, - routing_method_type=routing_method_type, - tune_max_num_tokens=next_power_of_2(x.shape[0]), - ) - - from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput - - return StandardCombineInput(hidden_states=output) + return fused_experts_none_to_flashinfer_trtllm_fp8( + dispatch_output, quant_info, self.moe_runner_config + ) if get_moe_runner_backend().is_flashinfer_cutlass(): activation = ACT_STR_TO_TYPE_MAP[self.moe_runner_config.activation] @@ -1532,49 +1429,12 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): and reorder_rows_for_gated_act_gemm is not None and shuffle_matrix_sf_a is not None ): + from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import ( + align_fp4_moe_weights_for_flashinfer_trtllm, + ) + # FlashInfer TRTLLM processing - handles both w13 and w2 - ( - gemm1_weights_fp4_shuffled, - gemm1_scales_fp4_shuffled, - gemm2_weights_fp4_shuffled, - gemm2_scales_fp4_shuffled, - ) = prepare_static_weights_for_trtllm_fp4_moe( - layer.w13_weight, - layer.w2_weight, - layer.w13_weight_scale, - layer.w2_weight_scale, - layer.w2_weight.size(-2), # hidden_size - layer.w13_weight.size(-2) // 2, # intermediate_size - layer.w13_weight.size(0), # num_experts - ) - - # Set flashinfer parameters - layer.gemm1_weights_fp4_shuffled = Parameter( - gemm1_weights_fp4_shuffled, requires_grad=False - ) - layer.gemm2_weights_fp4_shuffled = Parameter( - gemm2_weights_fp4_shuffled, requires_grad=False - ) - layer.gemm1_scales_fp4_shuffled = Parameter( - gemm1_scales_fp4_shuffled, requires_grad=False - ) - layer.gemm2_scales_fp4_shuffled = Parameter( - gemm2_scales_fp4_shuffled, requires_grad=False - ) - - # Additional parameter needed for TRT-LLM - layer.g1_scale_c = Parameter( - (layer.w2_input_scale_quant * layer.g1_alphas).to(torch.float32), - requires_grad=False, - ) - - # Clean up weights that won't be used by TRT-LLM - del ( - layer.w2_weight, - layer.w2_weight_scale, - layer.w13_weight, - layer.w13_weight_scale, - ) + align_fp4_moe_weights_for_flashinfer_trtllm(layer) else: # CUTLASS processing - handle w13 and w2 separately @@ -1644,6 +1504,10 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig ): self.moe_runner_config = moe_runner_config + if get_moe_runner_backend().is_flashinfer_trtllm(): + self.runner = MoeRunner( + MoeRunnerBackend.FLASHINFER_TRTLLM, moe_runner_config + ) def apply( self, @@ -1662,12 +1526,36 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): ), f"{activation=} missing from {ACT_STR_TO_TYPE_MAP.keys()=}" moe_runner_config = self.moe_runner_config - # Check if this is a FlashInferFP4MoE layer that should handle its own forward + # FlashInfer TRTLLM FP4 path - layer has shuffled weights only when + # backend is flashinfer_trtllm if hasattr(layer, "gemm1_weights_fp4_shuffled"): - # This layer was processed with flashinfer TRTLLM - delegate to its own forward - from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput + from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import ( + FlashInferTrtllmFp4MoeQuantInfo, + ) + from sglang.srt.layers.moe.utils import RoutingMethodType - return StandardCombineInput(hidden_states=layer.forward(x, topk_output)) + # Determine routing method type based on layer configuration + routing_method_type = getattr( + layer, "routing_method_type", RoutingMethodType.Default + ) + + quant_info = FlashInferTrtllmFp4MoeQuantInfo( + gemm1_weights_fp4_shuffled=layer.gemm1_weights_fp4_shuffled.data, + gemm2_weights_fp4_shuffled=layer.gemm2_weights_fp4_shuffled.data, + gemm1_scales_fp4_shuffled=layer.gemm1_scales_fp4_shuffled.data, + gemm2_scales_fp4_shuffled=layer.gemm2_scales_fp4_shuffled.data, + g1_scale_c=layer.g1_scale_c.data, + g1_alphas=layer.g1_alphas.data, + g2_alphas=layer.g2_alphas.data, + w13_input_scale_quant=layer.w13_input_scale_quant, + global_num_experts=layer.num_experts, + local_expert_offset=layer.moe_ep_rank * layer.num_local_experts, + local_num_experts=layer.num_local_experts, + intermediate_size_per_partition=layer.intermediate_size_per_partition, + routing_method_type=routing_method_type, + ) + + return self.runner.run(dispatch_output, quant_info) if self.enable_flashinfer_cutlass_moe: from sglang.srt.layers.moe.token_dispatcher import DispatchOutputChecker