From a883906a248a28981672018f76458943b0e63cac Mon Sep 17 00:00:00 2001 From: HandH1998 <1335248067@qq.com> Date: Mon, 26 Jan 2026 00:15:53 +0800 Subject: [PATCH] Support mxint4 flashinfer_trtllm moe gemm (#16892) --- .../srt/layers/moe/fused_moe_triton/layer.py | 7 + .../sglang/srt/layers/moe/moe_runner/base.py | 7 +- .../compressed_tensors/compressed_tensors.py | 13 + .../compressed_tensors_moe.py | 346 +++++++++++++++++- 4 files changed, 367 insertions(+), 6 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 139fd2fc2..019843ae0 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -54,6 +54,9 @@ from sglang.srt.layers.quantization.base_config import ( FusedMoEMethodBase, QuantizationConfig, ) +from sglang.srt.layers.quantization.compressed_tensors.compressed_tensors_moe import ( + CompressedTensorsMxInt4MoEMethod, +) from sglang.srt.layers.quantization.fp8 import Fp8MoEMethod from sglang.srt.layers.quantization.modelopt_quant import ModelOptNvFp4FusedMoEMethod from sglang.srt.layers.quantization.unquant import UnquantizedFusedMoEMethod @@ -253,6 +256,7 @@ class FusedMoE(torch.nn.Module): gemm1_alpha=gemm1_alpha, gemm1_clamp_limit=gemm1_clamp_limit, is_gated=is_gated, + routing_method_type=routing_method_type, ) self.quant_method: Optional[FusedMoEMethodBase] = None @@ -688,6 +692,7 @@ class FusedMoE(torch.nn.Module): isinstance(self.quant_method, ModelOptNvFp4FusedMoEMethod) or isinstance(self.quant_method, Fp8MoEMethod) or isinstance(self.quant_method, UnquantizedFusedMoEMethod) + or isinstance(self.quant_method, CompressedTensorsMxInt4MoEMethod) ): shard_id = {"w1": "w3", "w3": "w1", "w2": "w2"}[shard_id] @@ -1140,6 +1145,7 @@ class FlashInferFusedMoE(FusedMoE): 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 @@ -1170,6 +1176,7 @@ class FlashInferFusedMoE(FusedMoE): 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]), ) diff --git a/python/sglang/srt/layers/moe/moe_runner/base.py b/python/sglang/srt/layers/moe/moe_runner/base.py index 118206a90..12dd2ba6a 100644 --- a/python/sglang/srt/layers/moe/moe_runner/base.py +++ b/python/sglang/srt/layers/moe/moe_runner/base.py @@ -6,7 +6,11 @@ from typing import TYPE_CHECKING, Callable, Optional, Tuple, TypeGuard import torch -from sglang.srt.layers.moe.utils import MoeA2ABackend, MoeRunnerBackend +from sglang.srt.layers.moe.utils import ( + MoeA2ABackend, + MoeRunnerBackend, + RoutingMethodType, +) if TYPE_CHECKING: from sglang.srt.layers.moe.moe_runner.triton import ( @@ -33,6 +37,7 @@ class MoeRunnerConfig: top_k: Optional[int] = None num_fused_shared_experts: Optional[int] = None params_dtype: Optional[torch.dtype] = None + routing_method_type: Optional[RoutingMethodType] = None # Runner configuration activation: str = "silu" diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py index 8253f6ebe..7b0374851 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py @@ -471,6 +471,19 @@ class CompressedTensorsConfig(QuantizationConfig): return is_channel_group and input_quant_none and is_symmetric and is_static + def _is_mxint4a16(self, weight_quant: BaseModel, input_quant: BaseModel) -> bool: + input_quant_none = input_quant is None + is_symmetric = weight_quant.symmetric + is_mxint4 = ( + weight_quant.num_bits == 4 + and weight_quant.type == QuantizationType.INT + and weight_quant.strategy == QuantizationStrategy.GROUP.value + and weight_quant.group_size == 32 + ) + is_static = not weight_quant.dynamic + + return is_mxint4 and input_quant_none and is_symmetric and is_static + def _is_dynamic_token_w4( self, weight_quant: BaseModel, input_quant: BaseModel ) -> bool: diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors_moe.py index 9e10dd837..62cf492ed 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors_moe.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors_moe.py @@ -11,7 +11,11 @@ import torch from compressed_tensors import CompressionFormat from compressed_tensors.quantization import QuantizationStrategy -from sglang.srt.distributed import get_tensor_model_parallel_world_size, get_tp_group +from sglang.srt.distributed import ( + get_moe_expert_parallel_rank, + get_tensor_model_parallel_world_size, + get_tp_group, +) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -21,10 +25,15 @@ from sglang.srt.hardware_backend.npu.quantization.fused_moe_method_npu import ( NPUW8A8Int8DynamicMoEMethod, ) from sglang.srt.layers.dp_attention import is_allocation_symmetric -from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig +from sglang.srt.layers.moe import ( + MoeRunner, + MoeRunnerBackend, + MoeRunnerConfig, + get_moe_runner_backend, +) from sglang.srt.layers.moe.cutlass_moe_params import CutlassMoEParams, CutlassMoEType from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo -from sglang.srt.layers.moe.utils import get_moe_runner_backend +from sglang.srt.layers.moe.utils import RoutingMethodType from sglang.srt.layers.quantization.base_config import FusedMoEMethodBase from sglang.srt.layers.quantization.compressed_tensors.schemes import ( WNA16_SUPPORTED_BITS, @@ -47,6 +56,7 @@ from sglang.srt.layers.quantization.utils import ( from sglang.srt.utils import ( get_bool_env_var, is_cuda, + is_flashinfer_available, is_hip, is_npu, next_power_of_2, @@ -74,6 +84,17 @@ if _use_aiter: from aiter.fused_moe import fused_moe from aiter.ops.shuffle import shuffle_weight +if is_flashinfer_available(): + from flashinfer.fp4_quantization import block_scale_interleave + from flashinfer.fused_moe import ( + convert_to_block_layout, + trtllm_mxint4_block_scale_moe, + ) + from flashinfer.fused_moe.core import ( + _maybe_get_cached_w3_w1_permute_indices, + get_w2_permute_indices_with_cache, + ) + logger = logging.getLogger(__name__) @@ -90,6 +111,7 @@ __all__ = [ "CompressedTensorsW8A8Fp8MoEMethod", "NPUCompressedTensorsW8A8Int8MoEMethod", "CompressedTensorsWNA16MoEMethod", + "CompressedTensorsMxInt4MoEMethod", "NPUCompressedTensorsW4A16Int4DynamicMoEMethod", ] @@ -114,8 +136,17 @@ class CompressedTensorsMoEMethod(FusedMoEMethodBase): if quant_config._is_wNa16_group_channel(weight_quant, input_quant): if not _is_npu: - logger.info_once("Using CompressedTensorsWNA16MarlinMoEMethod") - return CompressedTensorsWNA16MoEMethod(quant_config) + if ( + quant_config._is_mxint4a16(weight_quant, input_quant) + and get_moe_runner_backend().is_flashinfer_trtllm() + ): + logger.info_once( + "Using CompressedTensorsMxInt4MoEMethod with flashinfer_trtllm backend" + ) + return CompressedTensorsMxInt4MoEMethod(quant_config) + else: + logger.info_once("Using CompressedTensorsWNA16MarlinMoEMethod") + return CompressedTensorsWNA16MoEMethod(quant_config) else: if ( quant_config._is_dynamic_token_w4(weight_quant, input_quant) @@ -1764,3 +1795,308 @@ class NPUCompressedTensorsW4A16Int4DynamicMoEMethod(CompressedTensorsMoEMethod): group_list, output_dtype, ) + + +class CompressedTensorsMxInt4MoEMethod(CompressedTensorsMoEMethod): + def __init__(self, quant_config: CompressedTensorsConfig): + self.quant_config = quant_config + config = self.quant_config.target_scheme_map["Linear"].get("weights") + self.num_bits = config.num_bits + self.packed_factor = 32 // config.num_bits + self.strategy = config.strategy + self.group_size = config.group_size + self.actorder = config.actorder + assert ( + config.strategy == "group" + and config.group_size == 32 + and config.num_bits == 4 + ), "MxInt4 only supports group strategy with group size 32" + assert config.symmetric, "Only symmetric quantization is supported for MoE" + assert ( + get_moe_runner_backend().is_flashinfer_trtllm() + ), "MxInt4 only supports flashinfer_trtllm backend" + assert ( + not config.actorder + ), "Actorder is not supported by flashinfer_trtllm backend" + self.moe_ep_rank = get_moe_expert_parallel_rank() + + if self.quant_config.quant_format != CompressionFormat.pack_quantized.value: + raise ValueError( + f"For Fused MoE layers, only {CompressionFormat.pack_quantized.value} " + "is supported for the mxint4" + ) + self._cache_permute_indices = {} + + def create_weights( + self, + layer: torch.nn.Module, + num_experts: int, + hidden_size: int, + intermediate_size_per_partition: int, + params_dtype: torch.dtype, + **extra_weight_attrs, + ): + extra_weight_attrs.update({"quant_method": self.strategy}) + w13_weight = torch.nn.Parameter( + torch.empty( + num_experts, + 2 * intermediate_size_per_partition, + hidden_size // self.packed_factor, + dtype=torch.int32, + ), + requires_grad=False, + ) + layer.register_parameter("w13_weight_packed", w13_weight) + set_weight_attrs(w13_weight, extra_weight_attrs) + + w2_weight = torch.nn.Parameter( + torch.empty( + num_experts, + hidden_size, + intermediate_size_per_partition // self.packed_factor, + dtype=torch.int32, + ), + requires_grad=False, + ) + layer.register_parameter("w2_weight_packed", w2_weight) + set_weight_attrs(w2_weight, extra_weight_attrs) + + w2_scales_size = intermediate_size_per_partition + num_groups_w2 = w2_scales_size // self.group_size + num_groups_w13 = hidden_size // self.group_size + + assert params_dtype == torch.bfloat16 + w13_scale = torch.nn.Parameter( + torch.ones( + num_experts, + 2 * intermediate_size_per_partition, + num_groups_w13, + dtype=params_dtype, + ), + requires_grad=False, + ) + layer.register_parameter("w13_weight_scale", w13_scale) + set_weight_attrs(w13_scale, extra_weight_attrs) + + w2_scale = torch.nn.Parameter( + torch.ones(num_experts, hidden_size, num_groups_w2, dtype=params_dtype), + requires_grad=False, + ) + layer.register_parameter("w2_weight_scale", w2_scale) + set_weight_attrs(w2_scale, extra_weight_attrs) + + w13_weight_shape = torch.nn.Parameter( + torch.empty(num_experts, 2), requires_grad=False + ) + + layer.register_parameter("w13_weight_shape", w13_weight_shape) + set_weight_attrs(w13_weight_shape, extra_weight_attrs) + + w2_weight_shape = torch.nn.Parameter( + torch.empty(num_experts, 2), requires_grad=False + ) + layer.register_parameter("w2_weight_shape", w2_weight_shape) + set_weight_attrs(w2_weight_shape, extra_weight_attrs) + + layer.a13_scale = None + layer.a2_scale = None + + # Adapted from https://github.com/flashinfer-ai/flashinfer/blob/main/tests/moe/test_trtllm_gen_fused_moe.py + def prepare_static_weights_for_kernel( + self, + gemm1_weights, + gemm2_weights, + gemm1_scales, + gemm2_scales, + num_experts, + ): + """Prepare quantized weights for kernel (done offline with weights).""" + + epilogue_tile_m = 128 + gemm1_weights_mxint4_shuffled = [] + gemm1_scales_shuffled = [] + gemm2_weights_mxint4_shuffled = [] + gemm2_scales_shuffled = [] + + def repack(w): + assert w.dim() == 2 and w.dtype == torch.int32 + shifts = torch.arange(0, 32, 4, dtype=torch.int32, device=w.device) + w = (w.unsqueeze(2) >> shifts) & 0x0F + w = (w - 8).to(torch.int8).reshape(w.shape[0], -1, 2) + w = (w[..., 0] & 0x0F) | ((w[..., 1] & 0x0F) << 4) + w = w.to(torch.uint8) + return w + + for i in range(num_experts): + # NOTE(HandH1998): + # the huggingface weight format follows (w/s + 8) to pack, + # however, trtllm requires (w/s) to pack + # we need to convert the weight to trtllm's format first + cur_expert_gemm1_weight = repack(gemm1_weights[i]) + cur_expert_gemm2_weight = repack(gemm2_weights[i]) + + # Calculate the permute indices for the following: + # 1. Reorder rows of W1 and scales for fused gated activation + # 2. Shuffle weights and scaling factors for transposed mma output + # for both w3_w1 and w2 weights and scale factors + permute_indices = _maybe_get_cached_w3_w1_permute_indices( + self._cache_permute_indices, + cur_expert_gemm1_weight, + epilogue_tile_m, + ) + gemm1_weights_shuffled = cur_expert_gemm1_weight[ + permute_indices.to(gemm1_weights.device) + ].contiguous() + permute_sf_indices = _maybe_get_cached_w3_w1_permute_indices( + self._cache_permute_indices, + gemm1_scales[i].to(torch.bfloat16), + epilogue_tile_m, + num_elts_per_sf=32, + ) + gemm1_scales_shuffled.append( + block_scale_interleave( + gemm1_scales[i] + .to(torch.bfloat16)[permute_sf_indices.to(gemm1_scales.device)] + .contiguous() + ) + ) + + permute_indices = get_w2_permute_indices_with_cache( + self._cache_permute_indices, + cur_expert_gemm2_weight, + epilogue_tile_m, + ) + gemm2_weights_shuffled = cur_expert_gemm2_weight[ + permute_indices.to(gemm2_weights.device) + ].contiguous() + + permute_sf_indices = get_w2_permute_indices_with_cache( + self._cache_permute_indices, + gemm2_scales[i].to(torch.bfloat16), + epilogue_tile_m, + num_elts_per_sf=16, + ) + gemm2_scales_shuffled.append( + block_scale_interleave( + gemm2_scales[i] + .to(torch.bfloat16)[permute_sf_indices.to(gemm2_scales.device)] + .contiguous() + ) + ) + + block_k = 128 + gemm1_weights_shuffled = convert_to_block_layout( + gemm1_weights_shuffled.view(torch.uint8), block_k + ) + gemm2_weights_shuffled = convert_to_block_layout( + gemm2_weights_shuffled.view(torch.uint8), block_k + ) + + gemm1_weights_mxint4_shuffled.append(gemm1_weights_shuffled) + gemm2_weights_mxint4_shuffled.append(gemm2_weights_shuffled) + + gemm1_weights_mxint4_shuffled = torch.stack(gemm1_weights_mxint4_shuffled) + gemm2_weights_mxint4_shuffled = torch.stack(gemm2_weights_mxint4_shuffled) + gemm1_scales_shuffled = torch.stack(gemm1_scales_shuffled).view(torch.bfloat16) + gemm2_scales_shuffled = torch.stack(gemm2_scales_shuffled).view(torch.bfloat16) + + return ( + gemm1_weights_mxint4_shuffled, + gemm1_scales_shuffled, + gemm2_weights_mxint4_shuffled, + gemm2_scales_shuffled, + ) + + def process_weights_after_loading(self, layer: torch.nn.Module) -> None: + + num_experts = layer.w13_weight_packed.shape[0] + ( + gemm1_weights_mxint4_shuffled, + gemm1_scales_shuffled, + gemm2_weights_mxint4_shuffled, + gemm2_scales_shuffled, + ) = self.prepare_static_weights_for_kernel( + layer.w13_weight_packed, + layer.w2_weight_packed, + layer.w13_weight_scale, + layer.w2_weight_scale, + num_experts=num_experts, + ) + replace_parameter(layer, "w13_weight_packed", gemm1_weights_mxint4_shuffled) + replace_parameter(layer, "w2_weight_packed", gemm2_weights_mxint4_shuffled) + replace_parameter(layer, "w13_weight_scale", gemm1_scales_shuffled) + replace_parameter(layer, "w2_weight_scale", gemm2_scales_shuffled) + + def create_moe_runner( + self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig + ): + self.moe_runner_config = moe_runner_config + + def apply( + self, + layer: torch.nn.Module, + dispatch_output: StandardDispatchOutput, + ) -> CombineInput: + from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput + + assert ( + self.moe_runner_config.is_gated + ), "Only gated MoEs are supported for flashinfer mxint4" + + x = dispatch_output.hidden_states + topk_output = dispatch_output.topk_output + + router_logits = topk_output.router_logits + topk_config = topk_output.topk_config + correction_bias = ( + None + if topk_config.correction_bias is None + else topk_config.correction_bias.to(x.dtype) + ) + + local_num_experts = self.moe_runner_config.num_local_experts + routing_method_type = layer.routing_method_type + assert routing_method_type is not None + # DeepSeekV3 style routing requires float32 router logits, + # see this PR for details: https://github.com/flashinfer-ai/flashinfer/commit/d84e1d560da0a27961c19ca788d96c19cb9dcfb6 + if 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 = x.shape[0] + hidden_size = x.shape[-1] + symm_output = torch.empty( + num_tokens, hidden_size, dtype=torch.bfloat16, device=x.device + ) + + output = trtllm_mxint4_block_scale_moe( + routing_logits=router_logits, # float + routing_bias=correction_bias, + hidden_states=x, + gemm1_weights=layer.w13_weight_packed, + gemm1_weights_scale=layer.w13_weight_scale, + gemm1_alpha=self.moe_runner_config.gemm1_alpha, + gemm1_beta=None, + gemm1_clamp_limit=self.moe_runner_config.gemm1_clamp_limit, + gemm2_weights=layer.w2_weight_packed, + gemm2_weights_scale=layer.w2_weight_scale, + num_experts=self.moe_runner_config.num_experts, + top_k=topk_config.top_k, + n_group=topk_config.num_expert_group, + topk_group=topk_config.topk_group, + intermediate_size=self.moe_runner_config.intermediate_size_per_partition, + local_expert_offset=self.moe_ep_rank * local_num_experts, + local_num_experts=local_num_experts, + routed_scaling_factor=routed_scaling_factor, + routing_method_type=routing_method_type, + tune_max_num_tokens=next_power_of_2(x.shape[0]), + output=symm_output, + ) + + return StandardCombineInput(hidden_states=output)