diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_nvfp4_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_nvfp4_moe.py index 84bd4ae8f..001734c96 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_nvfp4_moe.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_nvfp4_moe.py @@ -5,13 +5,8 @@ from typing import TYPE_CHECKING import torch -from sglang.srt.distributed import get_tp_group -from sglang.srt.distributed.device_communicators.pynccl_allocator import ( - use_symmetric_memory, -) -from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig -from sglang.srt.layers.moe.utils import RoutingMethodType, get_moe_runner_backend +from sglang.srt.layers.moe.utils import get_moe_runner_backend from sglang.srt.layers.quantization.compressed_tensors.schemes import ( CompressedTensorsMoEScheme, ) @@ -22,7 +17,7 @@ from sglang.srt.layers.quantization.utils import ( replace_parameter, swizzle_blockscale, ) -from sglang.srt.utils import next_power_of_2, set_weight_attrs +from sglang.srt.utils import set_weight_attrs logger = logging.getLogger(__name__) @@ -282,7 +277,11 @@ class CompressedTensorsW4A4Nvfp4MoE(CompressedTensorsMoEScheme): ): self.moe_runner_config = moe_runner_config if self.use_flashinfer_trtllm: - self.runner = MoeRunner(MoeRunnerBackend.TRITON, moe_runner_config) + import sglang.srt.layers.moe.moe_runner.flashinfer_trtllm # noqa: F401 – triggers @register_fused_func + + self.runner = MoeRunner( + MoeRunnerBackend.FLASHINFER_TRTLLM, moe_runner_config + ) else: import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass # noqa: F401 – triggers @register_fused_func @@ -296,92 +295,38 @@ class CompressedTensorsW4A4Nvfp4MoE(CompressedTensorsMoEScheme): dispatch_output: StandardDispatchOutput, ) -> CombineInput: - from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput - x = dispatch_output.hidden_states - topk_output = dispatch_output.topk_output if self.use_flashinfer_trtllm: - from flashinfer import trtllm_fp4_block_scale_moe - - from sglang.srt.layers.quantization.fp4_utils import fp4_quantize - - router_logits = topk_output.router_logits - topk_config = topk_output.topk_config - - # global_scale must be shape [1] (strict in cute-dsl backend). - hs_fp4_bytes, hs_sf_bytes = fp4_quantize( - x, - layer.w13_input_scale_quant[:1], - self.group_size, # sf_vec_size - False, # use_ue8m0 - False, # is_sf_swizzled_layout - ) - hs_fp4 = hs_fp4_bytes.reshape(x.shape[0], x.shape[1] // 2) - hs_scale = hs_sf_bytes.view(torch.float8_e4m3fn).reshape( - *hs_sf_bytes.shape[:-1], -1 - ) - - correction_bias = ( - None - if topk_config.correction_bias is None - else topk_config.correction_bias.to(x.dtype) + from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import ( + FlashInferTrtllmFp4MoeQuantInfo, ) assert layer.routing_method_type is not None - # DeepSeekV3 style routing requires float32 router logits - if layer.routing_method_type == RoutingMethodType.DeepSeekV3: - router_logits = router_logits.to(torch.float32) - - routed_scaling_factor = self.moe_runner_config.routed_scaling_factor - routed_scaling_factor = ( - routed_scaling_factor if routed_scaling_factor is not None else 1.0 - ) - - 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 - ) - - output = trtllm_fp4_block_scale_moe( - routing_logits=router_logits, - routing_bias=correction_bias, - hidden_states=hs_fp4, - hidden_states_scale=hs_scale, - gemm1_weights=layer.w13_weight, - gemm1_weights_scale=layer.w13_weight_scale.view(torch.float8_e4m3fn), - gemm1_bias=None, - gemm1_alpha=None, - gemm1_beta=None, - gemm1_clamp_limit=None, - gemm2_weights=layer.w2_weight, - gemm2_weights_scale=layer.w2_weight_scale.view(torch.float8_e4m3fn), - gemm2_bias=None, - output1_scale_scalar=layer.g1_scale_c, - output1_scale_gate_scalar=layer.g1_alphas, - output2_scale_scalar=layer.g2_alphas, - num_experts=layer.num_experts, - top_k=topk_config.top_k, - n_group=topk_config.num_expert_group, - topk_group=topk_config.topk_group, - intermediate_size=layer.intermediate_size_per_partition, + quant_info = FlashInferTrtllmFp4MoeQuantInfo( + w13_weight=layer.w13_weight, + w2_weight=layer.w2_weight, + w13_weight_scale=layer.w13_weight_scale, + w2_weight_scale=layer.w2_weight_scale, + g1_scale_c=layer.g1_scale_c, + g1_alphas=layer.g1_alphas, + g2_alphas=layer.g2_alphas, + # global_scale must be shape [1]: the cute-dsl fp4_quantize + # backend reshapes it to [1]. process_weights_after_loading + # expands this to [num_local_experts] (all equal), so slice the + # single value out (matches ModelOptNvFp4FusedMoEMethod, which + # feeds a scalar). + w13_input_scale_quant=layer.w13_input_scale_quant[:1], + global_num_experts=layer.num_experts, local_expert_offset=layer.moe_ep_rank * layer.num_local_experts, local_num_experts=layer.num_local_experts, - routed_scaling_factor=routed_scaling_factor, + intermediate_size_per_partition=layer.intermediate_size_per_partition, routing_method_type=layer.routing_method_type, - do_finalize=True, - tune_max_num_tokens=next_power_of_2(hs_fp4.shape[0]), - output=symm_output, - )[0] + use_per_token_activation=False, + gemm1_clamp_limit=None, + ) + return self.runner.run(dispatch_output, quant_info) else: from sglang.srt.layers.moe.moe_runner.flashinfer_cutlass import ( FlashInferCutlassMoeQuantInfo, @@ -411,5 +356,3 @@ class CompressedTensorsW4A4Nvfp4MoE(CompressedTensorsMoEScheme): apply_routed_scaling_factor=False, ) return self.runner.run(dispatch_output, quant_info) - - return StandardCombineInput(hidden_states=output)