From bd0e69085760fc0c6defc5092bff9d9b08ab4627 Mon Sep 17 00:00:00 2001 From: qichu-yun Date: Mon, 1 Dec 2025 13:54:28 +0800 Subject: [PATCH] [Feature] Enable PTPC FP8 for compressed tensors moe (aiter kernel) (#12181) --- .../schemes/compressed_tensors_w8a8_fp8.py | 34 ++++++++++++------ .../srt/layers/quantization/fp8_utils.py | 35 +++++++++++++++++++ python/sglang/srt/models/deepseek_v2.py | 22 ++++++++---- 3 files changed, 73 insertions(+), 18 deletions(-) diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w8a8_fp8.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w8a8_fp8.py index 7ea1545a0..ef136a782 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w8a8_fp8.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w8a8_fp8.py @@ -18,6 +18,7 @@ from sglang.srt.layers.quantization.compressed_tensors.schemes import ( from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz from sglang.srt.layers.quantization.fp8_utils import ( apply_fp8_linear, + apply_fp8_ptpc_linear, normalize_e4m3fn_to_e4m3fnuz, ) from sglang.srt.layers.quantization.utils import requantize_with_max_scale @@ -61,7 +62,6 @@ class CompressedTensorsW8A8Fp8(CompressedTensorsScheme): ) if input_scale is not None: layer.input_scale = Parameter(input_scale, requires_grad=False) - layer.weight = Parameter(weight.t(), requires_grad=False) layer.weight_scale = Parameter(max_w_scale, requires_grad=False) @@ -83,8 +83,9 @@ class CompressedTensorsW8A8Fp8(CompressedTensorsScheme): weight_scale = layer.weight_scale.data if _use_aiter: + # keep the weight as (N, K) layer.weight = Parameter( - shuffle_weight(weight, (16, 16)).t(), requires_grad=False + shuffle_weight(weight, (16, 16)), requires_grad=False ) else: layer.weight = Parameter(weight.t(), requires_grad=False) @@ -161,12 +162,23 @@ class CompressedTensorsW8A8Fp8(CompressedTensorsScheme): x: torch.Tensor, bias: Optional[torch.Tensor] = None, ) -> torch.Tensor: - return apply_fp8_linear( - input=x, - weight=layer.weight, - weight_scale=layer.weight_scale, - input_scale=layer.input_scale, - bias=bias, - use_per_token_if_dynamic=True, - compressed_tensor_quant=True, - ) + if _use_aiter and self.strategy == QuantizationStrategy.CHANNEL: + return apply_fp8_ptpc_linear( + input=x, + weight=layer.weight, + weight_scale=layer.weight_scale, + input_scale=layer.input_scale, + bias=bias, + use_per_token_if_dynamic=True, + compressed_tensor_quant=True, + ) + else: + return apply_fp8_linear( + input=x, + weight=layer.weight, + weight_scale=layer.weight_scale, + input_scale=layer.input_scale, + bias=bias, + use_per_token_if_dynamic=True, + compressed_tensor_quant=True, + ) diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index 1005e9139..5ad8b502d 100644 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -898,3 +898,38 @@ def can_auto_enable_marlin_fp8() -> bool: return 80 <= sm < 89 except Exception: return False + + +def apply_fp8_ptpc_linear( + input: torch.Tensor, + weight: torch.Tensor, + weight_scale: torch.Tensor, + input_scale: Optional[torch.Tensor] = None, + input_scale_ub: Optional[torch.Tensor] = None, + bias: Optional[torch.Tensor] = None, + cutlass_fp8_supported: bool = cutlass_fp8_supported(), + use_per_token_if_dynamic: bool = False, + pad_output: Optional[bool] = None, + compressed_tensor_quant: bool = False, +) -> torch.Tensor: + # View input as 2D matrix for fp8 methods + input_2d = input.view(-1, input.shape[-1]) + + # weight is transposed (K, N) + output_shape = [*input.shape[:-1], weight.shape[1]] + + q_input, x_scale = aiter.per_token_quant_hip(input_2d, quant_dtype=aiter.dtypes.fp8) + + per_tensor_weights = (weight_scale.numel() == 1) and weight_scale.dim() < 2 + per_tensor_activations = (x_scale.numel() == 1) and x_scale.dim() < 2 + + if not (per_tensor_weights and per_tensor_activations): + # weight is in (N, K) + output_shape = [*input.shape[:-1], weight.shape[0]] + + output = aiter.gemm_a8w8_bpreshuffle( + q_input, weight, x_scale, weight_scale, None, input.dtype + ) + if bias is not None: + output = output + bias + return output.view(*output_shape) diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 163d12845..c74eef2f1 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -707,13 +707,21 @@ class DeepseekV2MoE(nn.Module): and self.shared_experts.gate_up_proj.weight.dtype == torch.float8_e4m3fn ) if self.shared_experts_is_fp8: - assert ( - self.shared_experts.gate_up_proj.quant_method.quant_config.weight_block_size - == self.shared_experts.down_proj.quant_method.quant_config.weight_block_size - ) - self.shared_experts_weight_block_size = ( - self.shared_experts.gate_up_proj.quant_method.quant_config.weight_block_size - ) + if ( + _use_aiter + and config.quantization_config.get("quant_method") + == "compressed-tensors" + ): + # For compressed-tensors ptpc model, don't need to check the weight_block_size + pass + else: + assert ( + self.shared_experts.gate_up_proj.quant_method.quant_config.weight_block_size + == self.shared_experts.down_proj.quant_method.quant_config.weight_block_size + ) + self.shared_experts_weight_block_size = ( + self.shared_experts.gate_up_proj.quant_method.quant_config.weight_block_size + ) self.top_k = config.num_experts_per_tok