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 d4dafa224..59530433c 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 @@ -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}") diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py index 5250003e9..c11e7e0a5 100644 --- a/python/sglang/srt/layers/quantization/fp8.py +++ b/python/sglang/srt/layers/quantization/fp8.py @@ -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: diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index 132ed9399..ab7d91b27 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -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, diff --git a/test/manual/quant/test_compressed_tensors_fp8_block_ue8m0.py b/test/manual/quant/test_compressed_tensors_fp8_block_ue8m0.py new file mode 100644 index 000000000..285ea114d --- /dev/null +++ b/test/manual/quant/test_compressed_tensors_fp8_block_ue8m0.py @@ -0,0 +1,108 @@ +"""Regression test for the Blackwell DeepGEMM UE8M0 weight-scale bug in the +compressed-tensors block-FP8 scheme (sgl-project/sglang#28662). + +On Blackwell (SM100) block-wise FP8 dispatches to DeepGEMM, which quantizes +activations to UE8M0 scales and expects the weight scales to be UE8M0-packed +as well. If ``CompressedTensorsW8A8Fp8`` leaves the raw float32 block scales, +DeepGEMM combines UE8M0 activation scales with float32 weight scales and emits +NaN logits. ``process_weights_after_loading`` must requantize the weight scales +to UE8M0 when DeepGEMM is the active runner. + +Manual / Blackwell-only: requires SM100 + DeepGEMM (DEEPGEMM_SCALE_UE8M0). +""" + +import unittest + +import torch + +from sglang.test.test_utils import CustomTestCase + + +class TestCompressedTensorsFp8BlockUE8M0(CustomTestCase): + @unittest.skipUnless(torch.cuda.is_available(), "needs CUDA") + def test_block_fp8_weight_scales_requantized_to_ue8m0(self): + from sglang.srt.layers import deep_gemm_wrapper + + if not ( + deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM + and deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0 + ): + self.skipTest("requires Blackwell DeepGEMM (DEEPGEMM_SCALE_UE8M0)") + + from compressed_tensors.quantization import ( + QuantizationArgs, + QuantizationStrategy, + ) + + from sglang.srt.layers.quantization.compressed_tensors.schemes.compressed_tensors_w8a8_fp8 import ( + CompressedTensorsW8A8Fp8, + ) + + torch.manual_seed(0) + device = "cuda" + block = [128, 128] + N, K = 512, 1024 # N % 64 == 0 and K % 128 == 0 -> DeepGEMM-supported shape + + # A bf16 reference weight, block-quantized to fp8 + float32 block scales, + # exactly as a compressed-tensors `float-quantized` checkpoint stores it. + w_ref = torch.randn(N, K, device=device, dtype=torch.bfloat16) * 0.1 + w_view = w_ref.float().view(N // 128, 128, K // 128, 128) + scale = w_view.abs().amax(dim=(1, 3)).clamp(min=1e-4) / 448.0 # (N/128, K/128) + w_fp8 = (w_view / scale[:, None, :, None]).to(torch.float8_e4m3fn).view(N, K) + w_deq = ( + (w_fp8.float().view(N // 128, 128, K // 128, 128) * scale[:, None, :, None]) + .view(N, K) + .to(torch.bfloat16) + ) + + weight_quant = QuantizationArgs( + num_bits=8, + type="float", + strategy=QuantizationStrategy.BLOCK, + symmetric=True, + dynamic=False, + block_structure=block, + ) + scheme = CompressedTensorsW8A8Fp8( + weight_quant=weight_quant, is_static_input_scheme=False + ) + + from sglang.srt.layers.quantization.fp8_utils import ( + deepgemm_w8a8_block_fp8_linear_with_fallback, + ) + + if ( + scheme.w8a8_block_fp8_linear + is not deepgemm_w8a8_block_fp8_linear_with_fallback + ): + self.skipTest("DeepGEMM is not the active block-FP8 runner") + + # Build the layer directly (create_weights would require an initialized + # tensor-parallel group); the fix lives in process_weights_after_loading. + layer = torch.nn.Module() + layer.weight = torch.nn.Parameter(w_fp8, requires_grad=False) + layer.weight_scale = torch.nn.Parameter( + scale.to(torch.float32), requires_grad=False + ) + layer.orig_dtype = torch.bfloat16 + + scheme.process_weights_after_loading(layer) + + # The fix: weight scales must be requantized to UE8M0 for the DeepGEMM runner. + self.assertTrue( + getattr(layer.weight_scale, "format_ue8m0", False), + "block-FP8 weight scales were not requantized to UE8M0", + ) + + x = torch.randn(8, K, device=device, dtype=torch.bfloat16) * 0.1 + y = scheme.apply_weights(layer, x) + + # Pre-fix this path produced NaN; post-fix it matches the bf16 dequant ref. + self.assertTrue(torch.isfinite(y).all(), "output contains NaN/Inf") + y_ref = x.float() @ w_deq.float().t() + rel_err = ((y.float() - y_ref).norm() / y_ref.norm()).item() + self.assertLess(rel_err, 0.05, f"relative error too high: {rel_err}") + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/layers/quantization/test_deepgemm_ue8m0_requant.py b/test/registered/unit/layers/quantization/test_deepgemm_ue8m0_requant.py new file mode 100644 index 000000000..412dbad5b --- /dev/null +++ b/test/registered/unit/layers/quantization/test_deepgemm_ue8m0_requant.py @@ -0,0 +1,167 @@ +"""CPU guards for DeepGEMM UE8M0 weight-scale requantization decisions.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +import unittest +from unittest.mock import patch + +import torch +from compressed_tensors.quantization import QuantizationStrategy + +import sglang.srt.layers.quantization.fp8_utils as fp8_utils +from sglang.srt.layers import deep_gemm_wrapper +from sglang.srt.layers.quantization.compressed_tensors.schemes.compressed_tensors_w8a8_fp8 import ( + CompressedTensorsW8A8Fp8, +) +from sglang.test.test_utils import CustomTestCase + +BLOCK_SIZE = [128, 128] + + +def _make_params(n: int = 64, k: int = 128): + weight = torch.nn.Parameter(torch.zeros((n, k)), requires_grad=False) + weight_scale = torch.nn.Parameter(torch.ones((1, 1)), requires_grad=False) + weight_scale.format_ue8m0 = False + return weight, weight_scale + + +class TestDeepGemmUE8M0Requant(CustomTestCase): + def _enabled_deepgemm_ue8m0(self): + return patch.multiple( + deep_gemm_wrapper, + ENABLE_JIT_DEEPGEMM=True, + DEEPGEMM_SCALE_UE8M0=True, + ) + + def test_helper_requants_supported_deepgemm_bf16_once(self): + weight, weight_scale = _make_params() + + with self._enabled_deepgemm_ue8m0(), patch.object( + fp8_utils, "requant_weight_ue8m0_inplace" + ) as requant: + fired = fp8_utils.requant_block_scale_ue8m0_for_deepgemm( + weight, + weight_scale, + BLOCK_SIZE, + use_deepgemm_runner=True, + output_dtype=torch.bfloat16, + weight_shape=weight.shape, + ) + fired_again = fp8_utils.requant_block_scale_ue8m0_for_deepgemm( + weight, + weight_scale, + BLOCK_SIZE, + use_deepgemm_runner=True, + output_dtype=torch.bfloat16, + weight_shape=weight.shape, + ) + + self.assertTrue(fired) + self.assertFalse(fired_again) + self.assertTrue(weight_scale.format_ue8m0) + requant.assert_called_once_with(weight, weight_scale, BLOCK_SIZE) + + def test_helper_skips_non_bf16_output(self): + weight, weight_scale = _make_params() + + with self._enabled_deepgemm_ue8m0(), patch.object( + fp8_utils, "requant_weight_ue8m0_inplace" + ) as requant: + fired = fp8_utils.requant_block_scale_ue8m0_for_deepgemm( + weight, + weight_scale, + BLOCK_SIZE, + use_deepgemm_runner=True, + output_dtype=torch.float16, + weight_shape=weight.shape, + ) + + self.assertFalse(fired) + self.assertFalse(weight_scale.format_ue8m0) + requant.assert_not_called() + + def test_helper_skips_shape_deepgemm_will_not_run(self): + weight, weight_scale = _make_params(n=96, k=128) + + with self._enabled_deepgemm_ue8m0(), patch.object( + fp8_utils, "requant_weight_ue8m0_inplace" + ) as requant: + fired = fp8_utils.requant_block_scale_ue8m0_for_deepgemm( + weight, + weight_scale, + BLOCK_SIZE, + use_deepgemm_runner=True, + output_dtype=torch.bfloat16, + weight_shape=weight.shape, + ) + + self.assertFalse(fired) + self.assertFalse(weight_scale.format_ue8m0) + requant.assert_not_called() + + def test_helper_skips_non_deepgemm_runner(self): + weight, weight_scale = _make_params() + + with self._enabled_deepgemm_ue8m0(), patch.object( + fp8_utils, "requant_weight_ue8m0_inplace" + ) as requant: + fired = fp8_utils.requant_block_scale_ue8m0_for_deepgemm( + weight, + weight_scale, + BLOCK_SIZE, + use_deepgemm_runner=False, + output_dtype=torch.bfloat16, + weight_shape=weight.shape, + ) + + self.assertFalse(fired) + self.assertFalse(weight_scale.format_ue8m0) + requant.assert_not_called() + + def test_helper_skips_unsupported_block_size(self): + weight, weight_scale = _make_params() + unsupported_block_size = [128, 256] + + with self._enabled_deepgemm_ue8m0(), patch.object( + fp8_utils, "requant_weight_ue8m0_inplace" + ) as requant: + fired = fp8_utils.requant_block_scale_ue8m0_for_deepgemm( + weight, + weight_scale, + unsupported_block_size, + use_deepgemm_runner=True, + output_dtype=torch.bfloat16, + weight_shape=weight.shape, + ) + + self.assertFalse(fired) + self.assertFalse(weight_scale.format_ue8m0) + requant.assert_not_called() + + def test_compressed_tensors_block_processing_preserves_ue8m0_marker(self): + scheme = CompressedTensorsW8A8Fp8.__new__(CompressedTensorsW8A8Fp8) + scheme.strategy = QuantizationStrategy.BLOCK + scheme.is_static_input_scheme = False + scheme.weight_block_size = BLOCK_SIZE + scheme.w8a8_block_fp8_linear = ( + fp8_utils.deepgemm_w8a8_block_fp8_linear_with_fallback + ) + + layer = torch.nn.Module() + layer.weight, layer.weight_scale = _make_params() + layer.orig_dtype = torch.bfloat16 + + with self._enabled_deepgemm_ue8m0(), patch.object( + fp8_utils, "requant_weight_ue8m0_inplace" + ) as requant: + scheme.process_weights_after_loading(layer) + scheme.process_weights_after_loading(layer) + + self.assertTrue(layer.weight_scale.format_ue8m0) + requant.assert_called_once() + + +if __name__ == "__main__": + unittest.main(verbosity=3)