[Fix] compressed-tensors block FP8: requantize weight scales to UE8M0 for DeepGEMM on Blackwell (#28662)
This commit is contained in:
+24
-6
@@ -21,8 +21,10 @@ 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,
|
||||
deepgemm_w8a8_block_fp8_linear_with_fallback,
|
||||
dispatch_w8a8_block_fp8_linear,
|
||||
normalize_e4m3fn_to_e4m3fnuz,
|
||||
requant_block_scale_ue8m0_for_deepgemm,
|
||||
validate_fp8_block_shape,
|
||||
)
|
||||
from sglang.srt.layers.quantization.utils import requantize_with_max_scale
|
||||
@@ -188,15 +190,31 @@ class CompressedTensorsW8A8Fp8(CompressedTensorsLinearScheme):
|
||||
|
||||
elif self.strategy == QuantizationStrategy.BLOCK:
|
||||
assert self.is_static_input_scheme is False
|
||||
weight = layer.weight
|
||||
weight_scale = layer.weight_scale
|
||||
|
||||
if is_fp8_fnuz():
|
||||
weight, weight_scale, _ = normalize_e4m3fn_to_e4m3fnuz(
|
||||
weight=weight, weight_scale=weight_scale
|
||||
weight=layer.weight, weight_scale=layer.weight_scale
|
||||
)
|
||||
layer.weight = Parameter(weight.data, requires_grad=False)
|
||||
layer.weight_scale = Parameter(weight_scale.data, requires_grad=False)
|
||||
layer.weight = Parameter(weight.data, requires_grad=False)
|
||||
layer.weight_scale = Parameter(weight_scale.data, requires_grad=False)
|
||||
layer.weight_scale.format_ue8m0 = False
|
||||
else:
|
||||
layer.weight.requires_grad_(False)
|
||||
layer.weight_scale.requires_grad_(False)
|
||||
|
||||
# On Blackwell, block-FP8 dispatches to DeepGEMM, which needs the
|
||||
# weight scales UE8M0-packed to match its UE8M0 activation scales.
|
||||
use_deepgemm_runner = (
|
||||
self.w8a8_block_fp8_linear
|
||||
is deepgemm_w8a8_block_fp8_linear_with_fallback
|
||||
)
|
||||
requant_block_scale_ue8m0_for_deepgemm(
|
||||
layer.weight,
|
||||
layer.weight_scale,
|
||||
self.weight_block_size,
|
||||
use_deepgemm_runner=use_deepgemm_runner,
|
||||
output_dtype=getattr(layer, "orig_dtype", None),
|
||||
weight_shape=layer.weight.shape,
|
||||
)
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unknown quantization strategy {self.strategy}")
|
||||
|
||||
@@ -57,13 +57,14 @@ from sglang.srt.layers.quantization.fp8_utils import (
|
||||
apply_fp8_linear,
|
||||
can_auto_enable_marlin_fp8,
|
||||
cutlass_fp8_supported,
|
||||
deepgemm_w8a8_block_fp8_linear_with_fallback,
|
||||
dispatch_w8a8_block_fp8_linear,
|
||||
dispatch_w8a8_mxfp8_linear,
|
||||
get_fp8_gemm_runner_backend,
|
||||
input_to_float8,
|
||||
mxfp8_group_quantize,
|
||||
normalize_e4m3fn_to_e4m3fnuz,
|
||||
requant_weight_ue8m0_inplace,
|
||||
requant_block_scale_ue8m0_for_deepgemm,
|
||||
)
|
||||
from sglang.srt.layers.quantization.kv_cache import BaseKVCacheMethod
|
||||
from sglang.srt.layers.quantization.marlin_utils_fp8 import prepare_fp8_layer_for_marlin
|
||||
@@ -535,37 +536,19 @@ class Fp8LinearMethod(LinearMethodBase):
|
||||
self._process_mxfp8_linear_weight_scale(layer)
|
||||
return
|
||||
else:
|
||||
# For fp8 linear weights run with deepgemm, the weights and scales need be requantized to ue8m0
|
||||
from sglang.srt.layers.quantization.fp8_utils import (
|
||||
deepgemm_w8a8_block_fp8_linear_with_fallback,
|
||||
# Requantize block scales to UE8M0 when DeepGEMM is the active runner.
|
||||
use_deepgemm_runner = (
|
||||
self.w8a8_block_fp8_linear
|
||||
is deepgemm_w8a8_block_fp8_linear_with_fallback
|
||||
)
|
||||
from sglang.srt.model_loader.utils import (
|
||||
should_deepgemm_weight_requant_ue8m0,
|
||||
requant_block_scale_ue8m0_for_deepgemm(
|
||||
layer.weight,
|
||||
layer.weight_scale_inv,
|
||||
getattr(self.quant_config, "weight_block_size", None),
|
||||
use_deepgemm_runner=use_deepgemm_runner,
|
||||
output_dtype=getattr(layer, "orig_dtype", None),
|
||||
weight_shape=layer.weight.shape,
|
||||
)
|
||||
|
||||
# Only requantize to UE8M0 if DeepGEMM can actually run
|
||||
# this layer. If the dtype or shape is unsupported, the GEMM
|
||||
# falls back to triton at runtime, which needs float32 scales.
|
||||
if (
|
||||
should_deepgemm_weight_requant_ue8m0(
|
||||
weight_block_size=getattr(
|
||||
self.quant_config, "weight_block_size", None
|
||||
),
|
||||
output_dtype=getattr(layer, "orig_dtype", None),
|
||||
weight_shape=layer.weight.shape,
|
||||
)
|
||||
and (
|
||||
self.w8a8_block_fp8_linear
|
||||
is deepgemm_w8a8_block_fp8_linear_with_fallback
|
||||
)
|
||||
and (not layer.weight_scale_inv.format_ue8m0)
|
||||
):
|
||||
requant_weight_ue8m0_inplace(
|
||||
layer.weight,
|
||||
layer.weight_scale_inv,
|
||||
self.quant_config.weight_block_size,
|
||||
)
|
||||
layer.weight_scale_inv.format_ue8m0 = True
|
||||
weight, weight_scale = layer.weight.data, layer.weight_scale_inv.data
|
||||
|
||||
layer.weight.data = weight.data
|
||||
@@ -1334,9 +1317,6 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
||||
# For fp8 moe run with deepgemm, the expert weights and scales need be requantized to ue8m0
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
from sglang.srt.layers.moe.ep_moe.layer import DeepEPMoE
|
||||
from sglang.srt.model_loader.utils import (
|
||||
should_deepgemm_weight_requant_ue8m0,
|
||||
)
|
||||
|
||||
# Check if MoE will actually use DeepGEMM runner
|
||||
will_use_deepgemm = self.is_deepgemm_moe_runner_backend_enabled()
|
||||
@@ -1379,28 +1359,23 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
||||
layer.w13_weight_scale_inv.format_ue8m0 = True
|
||||
layer.w2_weight_scale_inv.format_ue8m0 = True
|
||||
|
||||
if (
|
||||
not self.is_fp4_expert
|
||||
and should_deepgemm_weight_requant_ue8m0(
|
||||
weight_block_size=getattr(
|
||||
self.quant_config, "weight_block_size", None
|
||||
),
|
||||
)
|
||||
and will_use_deepgemm
|
||||
and not layer.w13_weight_scale_inv.format_ue8m0
|
||||
):
|
||||
assert isinstance(
|
||||
layer, DeepEPMoE
|
||||
), "DeepGemm MoE is only supported with DeepEPMoE"
|
||||
if not self.is_fp4_expert:
|
||||
weight_block_size = self.quant_config.weight_block_size
|
||||
requant_weight_ue8m0_inplace(
|
||||
layer.w13_weight, layer.w13_weight_scale_inv, weight_block_size
|
||||
)
|
||||
requant_weight_ue8m0_inplace(
|
||||
layer.w2_weight, layer.w2_weight_scale_inv, weight_block_size
|
||||
)
|
||||
layer.w13_weight_scale_inv.format_ue8m0 = True
|
||||
layer.w2_weight_scale_inv.format_ue8m0 = True
|
||||
if requant_block_scale_ue8m0_for_deepgemm(
|
||||
layer.w13_weight,
|
||||
layer.w13_weight_scale_inv,
|
||||
weight_block_size,
|
||||
use_deepgemm_runner=will_use_deepgemm,
|
||||
):
|
||||
assert isinstance(
|
||||
layer, DeepEPMoE
|
||||
), "DeepGemm MoE is only supported with DeepEPMoE"
|
||||
requant_block_scale_ue8m0_for_deepgemm(
|
||||
layer.w2_weight,
|
||||
layer.w2_weight_scale_inv,
|
||||
weight_block_size,
|
||||
use_deepgemm_runner=True,
|
||||
)
|
||||
|
||||
def _process_mxfp8_moe_weights(self, layer: Module, quantize: bool = True) -> None:
|
||||
|
||||
|
||||
@@ -1279,6 +1279,42 @@ def requant_weight_ue8m0_inplace(weight, weight_scale_inv, weight_block_size):
|
||||
weight_scale_inv.data = new_weight_scale_inv
|
||||
|
||||
|
||||
def requant_block_scale_ue8m0_for_deepgemm(
|
||||
weight: torch.nn.Parameter,
|
||||
weight_scale: torch.nn.Parameter,
|
||||
weight_block_size: Optional[List[int]],
|
||||
use_deepgemm_runner: bool,
|
||||
output_dtype: Optional[torch.dtype] = None,
|
||||
weight_shape=None,
|
||||
) -> bool:
|
||||
"""Requantize block-FP8 weight scales to UE8M0 in place for DeepGEMM.
|
||||
|
||||
No-op (returns False) unless the caller selected the DeepGEMM runner, the
|
||||
block size is 128x128 (the only layout the requant kernel supports), the
|
||||
scales are not already UE8M0, and DeepGEMM can run the layer (bf16 output,
|
||||
aligned shape). Returns True when it requantizes.
|
||||
"""
|
||||
from sglang.srt.model_loader.utils import (
|
||||
should_deepgemm_weight_requant_ue8m0,
|
||||
)
|
||||
|
||||
if (
|
||||
not use_deepgemm_runner
|
||||
or weight_block_size != [128, 128]
|
||||
or getattr(weight_scale, "format_ue8m0", False)
|
||||
or not should_deepgemm_weight_requant_ue8m0(
|
||||
weight_block_size=weight_block_size,
|
||||
output_dtype=output_dtype,
|
||||
weight_shape=weight_shape,
|
||||
)
|
||||
):
|
||||
return False
|
||||
|
||||
requant_weight_ue8m0_inplace(weight, weight_scale, weight_block_size)
|
||||
weight_scale.format_ue8m0 = True
|
||||
return True
|
||||
|
||||
|
||||
def requant_weight_ue8m0(
|
||||
weight: torch.Tensor,
|
||||
weight_scale_inv: torch.Tensor,
|
||||
|
||||
Reference in New Issue
Block a user