[Feature] Enable PTPC FP8 for compressed tensors moe (aiter kernel) (#12181)

This commit is contained in:
qichu-yun
2025-11-30 21:54:28 -08:00
committed by GitHub
parent 0825d7f4c6
commit bd0e690857
3 changed files with 73 additions and 18 deletions
@@ -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_kernel import is_fp8_fnuz
from sglang.srt.layers.quantization.fp8_utils import ( from sglang.srt.layers.quantization.fp8_utils import (
apply_fp8_linear, apply_fp8_linear,
apply_fp8_ptpc_linear,
normalize_e4m3fn_to_e4m3fnuz, normalize_e4m3fn_to_e4m3fnuz,
) )
from sglang.srt.layers.quantization.utils import requantize_with_max_scale from sglang.srt.layers.quantization.utils import requantize_with_max_scale
@@ -61,7 +62,6 @@ class CompressedTensorsW8A8Fp8(CompressedTensorsScheme):
) )
if input_scale is not None: if input_scale is not None:
layer.input_scale = Parameter(input_scale, requires_grad=False) layer.input_scale = Parameter(input_scale, requires_grad=False)
layer.weight = Parameter(weight.t(), requires_grad=False) layer.weight = Parameter(weight.t(), requires_grad=False)
layer.weight_scale = Parameter(max_w_scale, 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 weight_scale = layer.weight_scale.data
if _use_aiter: if _use_aiter:
# keep the weight as (N, K)
layer.weight = Parameter( layer.weight = Parameter(
shuffle_weight(weight, (16, 16)).t(), requires_grad=False shuffle_weight(weight, (16, 16)), requires_grad=False
) )
else: else:
layer.weight = Parameter(weight.t(), requires_grad=False) layer.weight = Parameter(weight.t(), requires_grad=False)
@@ -161,6 +162,17 @@ class CompressedTensorsW8A8Fp8(CompressedTensorsScheme):
x: torch.Tensor, x: torch.Tensor,
bias: Optional[torch.Tensor] = None, bias: Optional[torch.Tensor] = None,
) -> torch.Tensor: ) -> torch.Tensor:
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( return apply_fp8_linear(
input=x, input=x,
weight=layer.weight, weight=layer.weight,
@@ -898,3 +898,38 @@ def can_auto_enable_marlin_fp8() -> bool:
return 80 <= sm < 89 return 80 <= sm < 89
except Exception: except Exception:
return False 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)
+8
View File
@@ -707,6 +707,14 @@ class DeepseekV2MoE(nn.Module):
and self.shared_experts.gate_up_proj.weight.dtype == torch.float8_e4m3fn and self.shared_experts.gate_up_proj.weight.dtype == torch.float8_e4m3fn
) )
if self.shared_experts_is_fp8: if self.shared_experts_is_fp8:
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 ( assert (
self.shared_experts.gate_up_proj.quant_method.quant_config.weight_block_size 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.down_proj.quant_method.quant_config.weight_block_size