From d7a03c7ebfd7cdd3c299d8e4ff82764434c36f7a Mon Sep 17 00:00:00 2001 From: Jonah Bernard <96398205+Jonahcb@users.noreply.github.com> Date: Tue, 24 Feb 2026 19:47:53 -0500 Subject: [PATCH] [MoE Refactor] Refactor FlashInferFusedMoE into FusedMoE and flashinfer_trtllm.py (#19266) --- python/sglang/srt/layers/moe/ep_moe/layer.py | 3 +- .../srt/layers/moe/fused_moe_triton/layer.py | 98 ------------------- .../moe/moe_runner/flashinfer_trtllm.py | 74 ++++++++++++++ .../sglang/srt/layers/quantization/unquant.py | 23 ++++- 4 files changed, 93 insertions(+), 105 deletions(-) diff --git a/python/sglang/srt/layers/moe/ep_moe/layer.py b/python/sglang/srt/layers/moe/ep_moe/layer.py index ebcc696ec..ce44c0b70 100644 --- a/python/sglang/srt/layers/moe/ep_moe/layer.py +++ b/python/sglang/srt/layers/moe/ep_moe/layer.py @@ -15,7 +15,6 @@ from sglang.srt.layers.moe import ( get_moe_runner_backend, ) from sglang.srt.layers.moe.fused_moe_triton.layer import ( - FlashInferFusedMoE, FusedMoE, moe_forward_piecewise_cuda_graph_impl, ) @@ -707,7 +706,7 @@ def get_moe_impl_class(quant_config: Optional[QuantizationConfig]): or quant_config.get_name() == "compressed_tensors" ): # FlashInferFusedMoE support bf16, fp8 and compressed_tensors - return FlashInferFusedMoE + return FusedMoE if get_moe_runner_backend().is_flashinfer_cutlass(): return FusedMoE 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 de8a07ab3..cb25bdc8b 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -40,7 +40,6 @@ from sglang.srt.layers.moe.token_dispatcher.base import BaseDispatcher from sglang.srt.layers.moe.token_dispatcher.flashinfer import FlashinferDispatcher from sglang.srt.layers.moe.token_dispatcher.standard import ( StandardDispatcher, - StandardDispatchOutput, ) from sglang.srt.layers.moe.topk import ( BypassedTopKOutput, @@ -1129,103 +1128,6 @@ class FusedMoE(torch.nn.Module): self.meta_overlap_args = None -class FlashInferFusedMoE(FusedMoE): - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - - def forward(self, hidden_states: torch.Tensor, topk_output: TopKOutput): - if is_in_piecewise_cuda_graph(): - if not TopKOutputChecker.format_is_standard(topk_output): - # Make sure there is torch lib op registration for the whole moe layer - return self.forward_impl(hidden_states, topk_output) - else: - return moe_forward_piecewise_cuda_graph_impl( - hidden_states, - topk_output.topk_weights, - topk_output.topk_ids, - topk_output.router_logits, - self.layer_id, - ) - else: - return self.forward_impl(hidden_states, topk_output) - - def forward_impl(self, hidden_states: torch.Tensor, topk_output: TopKOutput): - assert ( - self.moe_runner_config.activation == "silu" - ), "Only silu is supported for flashinfer trtllm moe" - assert self.quant_method is not None - assert ( - topk_output.topk_config.renormalize - ), "Renormalize is required for flashinfer trtllm moe" - assert ( - self.num_fused_shared_experts == 0 - ), "Fused shared experts are not supported for flashinfer trtllm moe" - assert ( - self.moe_runner_config.is_gated - ), "Only gated MoEs are supported for flashinfer trtllm moe" - - assert TopKOutputChecker.format_is_bypassed(topk_output) - - router_logits = topk_output.router_logits - topk_config = topk_output.topk_config - correction_bias = topk_config.correction_bias - routed_scaling_factor = self.moe_runner_config.routed_scaling_factor - - if isinstance(self.quant_method, UnquantizedFusedMoEMethod): - # lazy import - try: - from flashinfer.fused_moe import trtllm_bf16_moe - except ImportError as e: - raise ImportError( - "Can't import trtllm_bf16_moe from flashinfer. " - "Please check flashinfer version to use bf16 with flashinfer_trtllm backend." - ) from e - - with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() - ): - # TODO: Now trtllm_bf16_moe doesn't support inplace output, - # we can move this out when it support that. - final_hidden_states = trtllm_bf16_moe( - routing_logits=router_logits, - routing_bias=correction_bias, - hidden_states=hidden_states, - gemm1_weights=self.w13_weight, - gemm2_weights=self.w2_weight, - num_experts=self.num_experts, - top_k=topk_config.top_k, - n_group=topk_config.num_expert_group, - topk_group=topk_config.topk_group, - intermediate_size=self.intermediate_size_per_partition, - local_expert_offset=self.moe_ep_rank * self.num_local_experts, - local_num_experts=self.num_local_experts, - routing_method_type=self.routing_method_type, - routed_scaling_factor=routed_scaling_factor, - tune_max_num_tokens=next_power_of_2(hidden_states.shape[0]), - ) - - else: - - final_hidden_states = self.quant_method.apply( - layer=self, - dispatch_output=StandardDispatchOutput( - hidden_states=hidden_states, - hidden_states_scale=None, - topk_output=topk_output, - ), - ).hidden_states - - # NOTE for symmetric memory tagging: - # We do not create the context in this function. - # Instead, we create the context and tagging inside each FusedMoEMethodBase - # This can allow fine-grained tagging. - - if self.reduce_results and (self.moe_tp_size > 1 or self.moe_ep_size > 1): - final_hidden_states = tensor_model_parallel_all_reduce(final_hidden_states) - - return final_hidden_states - - class FlashInferFP4MoE(FusedMoE): """FP4 TRTLLM MoE implementation using FlashInfer.""" 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 3aa20b17a..b76e37b0b 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py @@ -483,6 +483,76 @@ def fused_experts_none_to_flashinfer_trtllm_fp4( return StandardCombineInput(hidden_states=result) +@dataclass +class FlashInferTrtllmBf16MoeQuantInfo(MoeQuantInfo): + """Quantization payload consumed by FlashInfer TRT-LLM BF16 MoE kernels.""" + + gemm1_weights: torch.Tensor + gemm2_weights: torch.Tensor + + # Expert-parallel metadata + global_num_experts: int + local_expert_offset: int + + +def fused_experts_none_to_flashinfer_trtllm_bf16( + dispatch_output: StandardDispatchOutput, + quant_info: FlashInferTrtllmBf16MoeQuantInfo, + runner_config: MoeRunnerConfig, +) -> StandardCombineInput: + # lazy import + try: + from flashinfer.fused_moe import trtllm_bf16_moe + except ImportError as e: + raise ImportError( + "Can't import trtllm_bf16_moe from flashinfer. " + "Please check flashinfer version to use bf16 with flashinfer_trtllm backend." + ) from e + + assert ( + runner_config.activation == "silu" + ), "Only silu is supported for flashinfer trtllm moe" + assert ( + dispatch_output.topk_output.topk_config.renormalize + ), "Renormalize is required for flashinfer trtllm moe" + assert ( + runner_config.num_fused_shared_experts == 0 + ), "Fused shared experts are not supported for flashinfer trtllm moe" + assert ( + runner_config.is_gated + ), "Only gated MoEs are supported for flashinfer trtllm moe" + from sglang.srt.layers.moe.topk import TopKOutputChecker + + assert TopKOutputChecker.format_is_bypassed(dispatch_output.topk_output) + + hidden_states = dispatch_output.hidden_states + topk_output = dispatch_output.topk_output + topk_config = topk_output.topk_config + + with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + + # Call the fused kernel + final_hidden_states = trtllm_bf16_moe( + routing_logits=topk_output.router_logits, + routing_bias=topk_config.correction_bias, + hidden_states=hidden_states, + gemm1_weights=quant_info.gemm1_weights, + gemm2_weights=quant_info.gemm2_weights, + 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=runner_config.intermediate_size_per_partition, + local_expert_offset=quant_info.local_expert_offset, + local_num_experts=runner_config.num_local_experts, + routing_method_type=runner_config.routing_method_type, + routed_scaling_factor=runner_config.routed_scaling_factor, + tune_max_num_tokens=next_power_of_2(hidden_states.shape[0]), + ) + + return StandardCombineInput(hidden_states=final_hidden_states) + + @register_fused_func("none", "flashinfer_trtllm") def fused_experts_none_to_flashinfer_trtllm( dispatch_output: StandardDispatchOutput, @@ -498,6 +568,10 @@ def fused_experts_none_to_flashinfer_trtllm( return fused_experts_none_to_flashinfer_trtllm_fp8( dispatch_output, quant_info, runner_config ) + if isinstance(quant_info, FlashInferTrtllmBf16MoeQuantInfo): + return fused_experts_none_to_flashinfer_trtllm_bf16( + 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/unquant.py b/python/sglang/srt/layers/quantization/unquant.py index c77e7d21f..83c5cc1ae 100644 --- a/python/sglang/srt/layers/quantization/unquant.py +++ b/python/sglang/srt/layers/quantization/unquant.py @@ -322,11 +322,12 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp): self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig ): self.moe_runner_config = moe_runner_config - backend = ( - MoeRunnerBackend.TRITON_KERNELS - if self.use_triton_kernels - else MoeRunnerBackend.TRITON - ) + if self.use_flashinfer_trtllm_moe: + backend = MoeRunnerBackend.FLASHINFER_TRTLLM + elif self.use_triton_kernels: + backend = MoeRunnerBackend.TRITON_KERNELS + else: + backend = MoeRunnerBackend.TRITON self.runner = MoeRunner(backend, moe_runner_config) @property @@ -385,6 +386,18 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp): tune_max_num_tokens=next_power_of_2(x.shape[0]), )[0] return StandardCombineInput(hidden_states=output) + elif self.use_flashinfer_trtllm_moe: + from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import ( + FlashInferTrtllmBf16MoeQuantInfo, + ) + + quant_info = FlashInferTrtllmBf16MoeQuantInfo( + gemm1_weights=layer.w13_weight, + gemm2_weights=layer.w2_weight, + global_num_experts=layer.num_experts, + local_expert_offset=layer.moe_ep_rank * layer.num_local_experts, + ) + return self.runner.run(dispatch_output, quant_info) else: # Skip aiter fused_moe when using non-auto MoE backend (e.g., triton, triton_kernels) # because aiter CK kernels don't support all GEMM dimensions