From f06e2d3d1ff011d4125f48e6dc84ddc608f85841 Mon Sep 17 00:00:00 2001 From: Zilin Zhu Date: Wed, 17 Jun 2026 09:53:33 +0700 Subject: [PATCH] Support asymmetric compressed-tensors MoE (#27690) --- .../srt/layers/moe/fused_moe_triton/layer.py | 2 + .../schemes/compressed_tensors_wNa16_moe.py | 57 ++++++++++++++++++- 2 files changed, 57 insertions(+), 2 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 a24dfeb04..923036a91 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -815,6 +815,7 @@ class FusedMoE(torch.nn.Module): "CompressedTensorsWNA16TritonMoE", ] ) + and "zero" not in weight_name else loaded_weight ) @@ -1034,6 +1035,7 @@ class FusedMoE(torch.nn.Module): "CompressedTensorsWNA16TritonMoE", ] ) + and "zero" not in weight_name else loaded_weight ) diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py index 0d6a9b896..5b0198712 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py @@ -22,6 +22,7 @@ from sglang.srt.layers.quantization.compressed_tensors.schemes import ( from sglang.srt.layers.quantization.marlin_utils import ( marlin_make_workspace, marlin_moe_permute_scales, + moe_awq_to_marlin_zero_points, ) from sglang.srt.layers.quantization.utils import replace_parameter from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip, set_weight_attrs @@ -69,7 +70,7 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): self.strategy = config.strategy self.group_size = config.group_size self.actorder = config.actorder - assert config.symmetric, "Only symmetric quantization is supported for MoE" + self.sym = config.symmetric if not ( self.quant_config.quant_format == CompressionFormat.pack_quantized.value @@ -129,7 +130,7 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): # In the case where we have actorder/g_idx, # we do not partition the w2 scales - load_full_w2 = self.actorder and self.group_size != -1 + load_full_w2 = (self.actorder != "static") and self.group_size != -1 if load_full_w2: w2_scales_size = intermediate_size_per_partition * layer.moe_tp_size @@ -177,6 +178,32 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): layer.register_parameter("w13_weight_shape", w13_weight_shape) set_weight_attrs(w13_weight_shape, extra_weight_attrs) + # add zero param + if not self.sym: + w13_qzeros = torch.nn.Parameter( + torch.empty( + num_experts, + num_groups_w13, + 2 * intermediate_size_per_partition // self.packed_factor, + dtype=torch.int32, + ), + requires_grad=False, + ) + layer.register_parameter("w13_weight_zero_point", w13_qzeros) + set_weight_attrs(w13_qzeros, extra_weight_attrs) + + w2_qzeros = torch.nn.Parameter( + torch.empty( + num_experts, + num_groups_w2, + hidden_size // self.packed_factor, + dtype=torch.int32, + ), + requires_grad=False, + ) + layer.register_parameter("w2_weight_zero_point", w2_qzeros) + set_weight_attrs(w2_qzeros, extra_weight_attrs) + w13_g_idx = torch.nn.Parameter( torch.empty( num_experts, @@ -236,6 +263,10 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): layer._original_shapes["w2_weight_scale"] = tuple(w2_scale.shape) layer._original_shapes["w13_weight_scale"] = tuple(w13_scale.shape) + if not self.sym: + layer._original_shapes["w13_weight_zero_point"] = w13_qzeros.shape + layer._original_shapes["w2_weight_zero_point"] = tuple(w2_qzeros.shape) + def process_weights_after_loading(self, layer: torch.nn.Module) -> None: # Skip if the layer is already converted to Marlin format to prevent double-packing. @@ -339,6 +370,24 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): ) replace_tensor("w2_weight_scale", marlin_w2_scales) + # Repack zero + if not self.sym: + marlin_w13_zp = moe_awq_to_marlin_zero_points( + layer.w13_weight_zero_point, + size_k=layer.w13_weight_zero_point.shape[1], + size_n=layer.w13_weight_zero_point.shape[2] * self.packed_factor, + num_bits=self.num_bits, + ) + replace_tensor("w13_weight_zero_point", marlin_w13_zp) + + marlin_w2_zp = moe_awq_to_marlin_zero_points( + layer.w2_weight_zero_point, + size_k=layer.w2_weight_zero_point.shape[1], + size_n=layer.w2_weight_zero_point.shape[2] * self.packed_factor, + num_bits=self.num_bits, + ) + replace_tensor("w2_weight_zero_point", marlin_w2_zp) + layer.workspace = marlin_make_workspace(layer.w13_weight_packed.device, 4) layer.is_marlin_converted = True @@ -376,6 +425,8 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): w13_g_idx=getattr(layer, "w13_weight_g_idx", None), w2_g_idx=getattr(layer, "w2_weight_g_idx", None), is_k_full=self.is_k_full, + w13_qzeros=layer.w13_weight_zero_point if not self.sym else None, + w2_qzeros=layer.w2_weight_zero_point if not self.sym else None, ) def apply_weights( @@ -422,6 +473,8 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): g_idx2=layer.w2_weight_g_idx, sort_indices1=layer.w13_g_idx_sort_indices, sort_indices2=layer.w2_g_idx_sort_indices, + w1_zeros=layer.w13_weight_zero_point if not self.sym else None, + w2_zeros=layer.w2_weight_zero_point if not self.sym else None, num_bits=self.num_bits, is_k_full=self.is_k_full, routed_scaling_factor=self.moe_runner_config.routed_scaling_factor,