diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py index 26aee883b..1aa43e3e3 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py @@ -160,6 +160,36 @@ def _maybe_apply_routed_scaling_factor( return output +def _prescale_router_weight_on_input( + dispatch_output: StandardDispatchOutput | FlashinferDispatchOutput, + runner_config: MoeRunnerConfig, +) -> StandardDispatchOutput | FlashinferDispatchOutput: + if not runner_config.apply_router_weight_on_input: + return dispatch_output + + topk_output = dispatch_output.topk_output + topk_weights = topk_output.topk_weights + + if dispatch_output.hidden_states_scale is not None: + raise NotImplementedError( + "apply_router_weight_on_input is not supported when activations are " + "quantized before dispatch (flashinfer_cutlass fp4 all-gather path)." + ) + + assert topk_weights.dim() == 2 and topk_weights.shape[-1] == 1, ( + "apply_router_weight_on_input requires topk=1" + ) + + hidden_states = dispatch_output.hidden_states * topk_weights.to( + dispatch_output.hidden_states.dtype + ) + unit_scales = torch.ones_like(topk_weights, dtype=torch.float32) + return dispatch_output._replace( + hidden_states=hidden_states, + topk_output=topk_output._replace(topk_weights=unit_scales), + ) + + def _prepare_input( dispatch_output, quant_info: FlashInferCutlassMoeQuantInfo, @@ -197,8 +227,9 @@ def _run_flashinfer_cutlass( ) -> torch.Tensor: flashinfer_cutlass_fused_moe, _ = _flashinfer_cutlass_fused_moe() + dispatch_output = _prescale_router_weight_on_input(dispatch_output, runner_config) topk_output = dispatch_output.topk_output - topk_weights = topk_output.topk_weights + topk_weights = topk_output.topk_weights.to(torch.float32) topk_ids = topk_output.topk_ids x, x_sf, output_dtype, output_col = _prepare_input( dispatch_output, quant_info, runner_config @@ -267,9 +298,6 @@ def fused_experts_none_to_flashinfer_cutlass( assert isinstance(quant_info, FlashInferCutlassMoeQuantInfo), ( f"Unexpected quant_info type for flashinfer_cutlass: {type(quant_info)}" ) - assert not runner_config.apply_router_weight_on_input, ( - "apply_router_weight_on_input is not supported for FlashInfer CUTLASS" - ) output = _run_flashinfer_cutlass( dispatch_output=dispatch_output, @@ -292,9 +320,6 @@ def fused_experts_flashinfer_to_flashinfer_cutlass( assert isinstance(quant_info, FlashInferCutlassMoeQuantInfo), ( f"Unexpected quant_info type for flashinfer_cutlass: {type(quant_info)}" ) - assert not runner_config.apply_router_weight_on_input, ( - "apply_router_weight_on_input is not supported for FlashInfer CUTLASS" - ) output = _run_flashinfer_cutlass( dispatch_output=dispatch_output, diff --git a/python/sglang/srt/layers/quantization/modelopt_quant.py b/python/sglang/srt/layers/quantization/modelopt_quant.py index 5a15eeecf..129b17b5a 100755 --- a/python/sglang/srt/layers/quantization/modelopt_quant.py +++ b/python/sglang/srt/layers/quantization/modelopt_quant.py @@ -3053,7 +3053,6 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): assert activation in _SUPPORTED_ACT_STRS or ( activation == "situ" and moe_runner_backend.is_flashinfer_trtllm() ), f"{activation=} is unsupported by {moe_runner_backend}" - moe_runner_config = self.moe_runner_config if moe_runner_backend.is_flashinfer_megamoe(): from sglang.srt.layers.moe.flashinfer_megamoe import ( @@ -3179,9 +3178,6 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): FlashInferCutlassMoeQuantInfo, ) - assert not moe_runner_config.apply_router_weight_on_input, ( - "apply_router_weight_on_input is not supported for Flashinfer" - ) quant_info = FlashInferCutlassMoeQuantInfo( quant_type="fp4", w13_weight=layer.w13_weight,