[Fix] compressed-tensors block FP8: requantize weight scales to UE8M0 for DeepGEMM on Blackwell (#28662)

This commit is contained in:
Jimmy Shong
2026-06-26 21:41:18 +00:00
committed by GitHub
parent 7b02eab7a6
commit e745b3af22
5 changed files with 364 additions and 60 deletions
@@ -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}")
+29 -54
View File
@@ -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,