diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py index b9ca03d8c..8db9ac7f9 100644 --- a/python/sglang/srt/layers/quantization/fp8.py +++ b/python/sglang/srt/layers/quantization/fp8.py @@ -66,6 +66,7 @@ from sglang.srt.layers.quantization.fp8_utils import ( normalize_e4m3fn_to_e4m3fnuz, requant_block_scale_ue8m0_for_deepgemm, resolve_mxfp8_dense_gemm_backend, + torch_w8a8_block_fp8_linear, unshuffle_aiter_fp8_weight, use_aiter_bpreshuffle_gemm, ) @@ -799,6 +800,29 @@ class Fp8LinearMethod(LinearMethodBase): layer.aiter_bpreshuffled = True layer.weight.is_shuffled = True + if ( + is_xpu() + and self.w8a8_block_fp8_linear is torch_w8a8_block_fp8_linear + and self.weight_block_size in ([1, 128], [128, 128]) + and layer.weight_scale_inv.ndim == 2 + ): + # Keep the checkpoint's logical [N-blocks, K-blocks] shape, but use + # transpose-contiguous storage. For [1, 128], scaled_mm transposes + # scale_b internally; for [128, 128], the wrapper passes scale_b.t(). + # This avoids a per-forward contiguous/copy in either path. + scale = layer.weight_scale_inv.data + scale_b_is_contiguous = scale.t().is_contiguous() + if not scale_b_is_contiguous: + scale_reordered = torch.empty_strided( + scale.shape, + (1, scale.shape[0]), + dtype=scale.dtype, + device=scale.device, + ) + scale_reordered.copy_(scale) + with torch.no_grad(): + layer.weight_scale_inv.set_(scale_reordered) + def _process_mxfp8_linear_weight_scale(self, layer: Module) -> None: if not self.use_mxfp8: return diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index d4ae58608..361899487 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -592,6 +592,58 @@ def dispatch_w8a8_block_fp8_linear() -> Callable: return _dispatch_auto_backend() +def torch_w8a8_block_fp8_linear( + input: torch.Tensor, + weight: torch.Tensor, + block_size: List[int], + weight_scale: torch.Tensor, + input_scale: Optional[torch.Tensor] = None, + bias: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Run block-FP8 linear with Torch's scaled_mm implementation.""" + if not isinstance(block_size, (list, tuple)) or len(block_size) != 2: + raise ValueError( + f"XPU block-FP8 scaled_mm expects a two-dimensional weight_block_size, " + f"but got {block_size}" + ) + block_n, block_k = block_size + if block_k != 128 or block_n not in (1, 128): + raise ValueError( + "XPU block-FP8 scaled_mm supports weight_block_size [1, 128] or " + f"[128, 128], but got {block_size}" + ) + + scale_b_recipe = ( + torch.nn.functional.ScalingType.BlockWise1x128 + if block_n == 1 + else torch.nn.functional.ScalingType.BlockWise128x128 + ) + input_2d = input.reshape(-1, input.shape[-1]) + if input_scale is None: + q_input, activation_scale = per_token_group_quant_fp8(input_2d, block_k) + else: + q_input = input_2d + activation_scale = input_scale.reshape(-1, input_scale.shape[-1]) + + if q_input.stride(-1) != 1: + q_input = q_input.contiguous() + if weight.stride(-1) != 1: + weight = weight.contiguous() + weight_t = weight.t() + scale_b = weight_scale if block_n == 1 else weight_scale.t() + output = torch.nn.functional.scaled_mm( + q_input, + weight_t, + activation_scale, + torch.nn.functional.ScalingType.BlockWise1x128, + scale_b, + scale_b_recipe, + bias=bias, + output_dtype=torch.bfloat16 if input_scale is not None else input.dtype, + ) + return output.view(*input.shape[:-1], weight.shape[0]) + + def resolve_mxfp8_dense_gemm_backend() -> Mxfp8DenseGemmBackend: """Pick the MXFP8 dense linear backend, honoring `--fp8-gemm-backend` only when it names a backend that owns an MXFP8 dense kernel.""" @@ -804,7 +856,8 @@ def _dispatch_auto_backend() -> Callable: # 3. CUTLASS (if SM120 GPU and CUDA 12.8+) # 4. AITER (if AMD GPU with AITER enabled) # 5. NPU (Ascend) - # 6. Triton (fallback) + # 6. XPU (Intel GPU, PyTorch torch._scaled_mm) + # 7. Triton (fallback) if deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM: return deepgemm_w8a8_block_fp8_linear_with_fallback @@ -820,6 +873,8 @@ def _dispatch_auto_backend() -> Callable: ) return npu_w8a8_mxfp8_linear + elif _is_xpu: + return torch_w8a8_block_fp8_linear else: return triton_w8a8_block_fp8_linear diff --git a/test/registered/e2e/xpu/test_xpu_fp8_linear.py b/test/registered/e2e/xpu/test_xpu_fp8_linear.py new file mode 100644 index 000000000..f0bb30a8d --- /dev/null +++ b/test/registered/e2e/xpu/test_xpu_fp8_linear.py @@ -0,0 +1,336 @@ +""" +SGLang XPU backend integration tests for FP8 scaled_mm linear paths. + +Usage: + python3 -m unittest test.registered.e2e.xpu.test_xpu_fp8_linear + pytest test/registered/e2e/xpu/test_xpu_fp8_linear.py +""" + +import unittest +from typing import List + +import torch + +from sglang.srt.layers.quantization.fp8_utils import ( + dispatch_w8a8_block_fp8_linear, + per_token_group_quant_fp8, + torch_w8a8_block_fp8_linear, +) +from sglang.test.ci.ci_register import register_xpu_ci +from sglang.test.test_utils import CustomTestCase + +register_xpu_ci(est_time=20, suite="stage-b-test-1-gpu-xpu") + + +def reference_block_fp8_matmul( + q_input: torch.Tensor, + activation_scale: torch.Tensor, + weight: torch.Tensor, + block_size: List[int], + weight_scale: torch.Tensor, + bias: torch.Tensor = None, +) -> torch.Tensor: + """Reference using the exact quantized operands passed to scaled_mm.""" + block_n, block_k = block_size + + # Dequantize A: q_input [M, K], activation_scale [M, K // block_k] + M, K = q_input.shape + scale_a_expanded = ( + activation_scale.unsqueeze(-1).expand(M, K // block_k, block_k).reshape(M, K) + ) + a_dequant = q_input.float() * scale_a_expanded.float() + + # Dequantize B: weight [N, K], weight_scale [N // block_n, K // block_k] + N = weight.shape[0] + scale_b_expanded = ( + weight_scale.unsqueeze(1) + .unsqueeze(-1) + .expand(N // block_n, block_n, K // block_k, block_k) + .reshape(N, K) + ) + b_dequant = weight.float() * scale_b_expanded.float() + + output = torch.matmul(a_dequant, b_dequant.t()) + if bias is not None: + output = output + bias.float() + return output + + +class TestXPUFP8Linear(CustomTestCase): + def setUp(self): + if not torch.xpu.is_available(): + self.skipTest("XPU is not available") + self.device = "xpu" + + def test_w8a8_block_fp8_dispatch_on_xpu(self): + """Verify that dispatch_w8a8_block_fp8_linear cleanly routes to torch_w8a8_block_fp8_linear on XPU.""" + dispatched_fn = dispatch_w8a8_block_fp8_linear() + self.assertIs(dispatched_fn, torch_w8a8_block_fp8_linear) + + def test_torch_w8a8_block_fp8_linear_shapes_and_dims(self): + """Test small sizes with varied M, 2D/3D shapes, and bias to keep memory minimal.""" + K, N = 256, 256 + block_size = [128, 128] + + # Weights: [N, K], weight_scale: [N // 128, K // 128] + weight = torch.randn(N, K, dtype=torch.bfloat16, device=self.device).to( + torch.float8_e4m3fn + ) + weight_scale = ( + torch.rand(N // 128, K // 128, dtype=torch.float32, device=self.device) + + 0.1 + ) + bias = torch.randn(N, dtype=torch.bfloat16, device=self.device) + + # 1. 2D inputs with different M (decode M=1, small batch M=8, unaligned M=17) + for M in [1, 8, 17]: + x_2d = torch.randn(M, K, dtype=torch.bfloat16, device=self.device) + out_no_bias = torch_w8a8_block_fp8_linear( + x_2d, weight, block_size, weight_scale + ) + self.assertEqual(out_no_bias.shape, (M, N)) + self.assertEqual(out_no_bias.dtype, torch.bfloat16) + + out_bias = torch_w8a8_block_fp8_linear( + x_2d, weight, block_size, weight_scale, bias=bias + ) + self.assertEqual(out_bias.shape, (M, N)) + torch.testing.assert_close( + out_bias, out_no_bias + bias, atol=0.05, rtol=0.01 + ) + + # 2. 3D input: [Batch, SeqLen, Hidden] + x_3d = torch.randn(2, 4, K, dtype=torch.bfloat16, device=self.device) + out_3d = torch_w8a8_block_fp8_linear(x_3d, weight, block_size, weight_scale) + self.assertEqual(out_3d.shape, (2, 4, N)) + + x_3d_noncontiguous = x_3d.transpose(0, 1) + out_3d_noncontiguous = torch_w8a8_block_fp8_linear( + x_3d_noncontiguous, weight, block_size, weight_scale + ) + self.assertEqual(out_3d_noncontiguous.shape, (4, 2, N)) + + def test_torch_w8a8_block_fp8_linear_prequantized(self): + """Test pre-quantized input branch (input_scale is not None).""" + M, K, N = 16, 256, 256 + block_size = [128, 128] + block_k = block_size[1] + + x_bf16 = torch.randn(M, K, dtype=torch.bfloat16, device=self.device) + q_input, input_scale = per_token_group_quant_fp8(x_bf16, block_k) + + weight = torch.randn(N, K, dtype=torch.bfloat16, device=self.device).to( + torch.float8_e4m3fn + ) + weight_scale = ( + torch.rand(N // 128, K // 128, dtype=torch.float32, device=self.device) + + 0.1 + ) + + out = torch_w8a8_block_fp8_linear( + q_input, weight, block_size, weight_scale, input_scale=input_scale + ) + self.assertEqual(out.shape, (M, N)) + self.assertEqual(out.dtype, torch.bfloat16) + + scale_a = ( + input_scale.unsqueeze(-1).expand(M, K // block_k, block_k).reshape(M, K) + ) + scale_b = ( + weight_scale.unsqueeze(1) + .unsqueeze(-1) + .expand(N // 128, 128, K // 128, 128) + .reshape(N, K) + ) + ref = torch.matmul( + q_input.float() * scale_a.float(), + (weight.float() * scale_b.float()).t(), + ).to(torch.bfloat16) + torch.testing.assert_close(out, ref, rtol=0.05, atol=0.1) + + def test_torch_w8a8_block_fp8_linear_non_contiguous_views(self): + """Test that pre-quantized non-contiguous slice views are handled safely without crash.""" + M, K, N = 16, 256, 256 + block_size = [128, 128] + + x = torch.randn(M * 2, K, dtype=torch.bfloat16, device=self.device) + qx, scale = per_token_group_quant_fp8(x, block_size[1]) + qx_strided = qx[::2, :] + scale_strided = scale[::2, :] + self.assertFalse(qx_strided.is_contiguous()) + self.assertFalse(scale_strided.is_contiguous()) + + weight = torch.randn(N, K, dtype=torch.bfloat16, device=self.device).to( + torch.float8_e4m3fn + ) + weight_scale = ( + torch.rand(N // 128, K // 128, dtype=torch.float32, device=self.device) + + 0.1 + ) + + out_q_strided = torch_w8a8_block_fp8_linear( + qx_strided, + weight, + block_size, + weight_scale, + input_scale=scale_strided.contiguous(), + ) + self.assertEqual(out_q_strided.shape, (M, N)) + + out_scale_strided = torch_w8a8_block_fp8_linear( + qx_strided.contiguous(), + weight, + block_size, + weight_scale, + input_scale=scale_strided, + ) + torch.testing.assert_close(out_scale_strided, out_q_strided) + + scale_transpose_contiguous = scale_strided.t().contiguous().t() + self.assertFalse(scale_transpose_contiguous.is_contiguous()) + self.assertTrue(scale_transpose_contiguous.t().is_contiguous()) + out_scale_transpose_contiguous = torch_w8a8_block_fp8_linear( + qx_strided.contiguous(), + weight, + block_size, + weight_scale, + input_scale=scale_transpose_contiguous, + ) + torch.testing.assert_close(out_scale_transpose_contiguous, out_q_strided) + + weight_noncontiguous = torch.empty( + N * 2, K, dtype=torch.float8_e4m3fn, device=self.device + ) + weight_noncontiguous[::2] = weight + weight_noncontiguous = weight_noncontiguous[::2] + self.assertFalse(weight_noncontiguous.is_contiguous()) + self.assertEqual(weight_noncontiguous.stride(-1), 1) + weight_scale_noncontiguous = weight_scale.t().contiguous().t() + out_weight_views = torch_w8a8_block_fp8_linear( + qx_strided.contiguous(), + weight_noncontiguous, + block_size, + weight_scale_noncontiguous, + input_scale=scale_strided.contiguous(), + ) + self.assertEqual(out_weight_views.shape, (M, N)) + torch.testing.assert_close(out_weight_views, out_q_strided) + + weight_last_dim_strided = torch.empty( + N, K, 2, dtype=torch.float8_e4m3fn, device=self.device + ) + weight_last_dim_strided[..., 0] = weight + weight_last_dim_strided = weight_last_dim_strided[..., 0] + self.assertFalse(weight_last_dim_strided.is_contiguous()) + self.assertNotEqual(weight_last_dim_strided.stride(-1), 1) + out_weight_strided = torch_w8a8_block_fp8_linear( + qx_strided.contiguous(), + weight_last_dim_strided, + block_size, + weight_scale, + input_scale=scale_strided.contiguous(), + ) + torch.testing.assert_close(out_weight_strided, out_q_strided) + + def test_torch_w8a8_block_fp8_linear_rejects_invalid_block_size(self): + x = torch.randn(8, 256, dtype=torch.bfloat16, device=self.device) + weight = torch.randn(256, 256, dtype=torch.bfloat16, device=self.device).to( + torch.float8_e4m3fn + ) + weight_scale = torch.ones(2, 2, dtype=torch.float32, device=self.device) + for block_size in ([], [128], [128, 128, 128], [1, 32], [64, 64]): + with self.subTest(block_size=block_size), self.assertRaises(ValueError): + torch_w8a8_block_fp8_linear(x, weight, block_size, weight_scale) + + def test_torch_w8a8_block_fp8_linear_non_square_weight_scale(self): + """Keep the public v2 weight-scale orientation correct when N != K.""" + M, K, N = 8, 256, 384 + block_size = [128, 128] + x = torch.randn(M, K, dtype=torch.bfloat16, device=self.device) + weight = torch.randn(N, K, dtype=torch.bfloat16, device=self.device).to( + torch.float8_e4m3fn + ) + weight_scale = ( + torch.rand(N // 128, K // 128, dtype=torch.float32, device=self.device) + + 0.1 + ) + out = torch_w8a8_block_fp8_linear(x, weight, block_size, weight_scale) + q_input, input_scale = per_token_group_quant_fp8(x, block_size[1]) + ref = reference_block_fp8_matmul( + q_input, input_scale, weight, block_size, weight_scale + ).to(torch.bfloat16) + self.assertEqual(out.shape, (M, N)) + torch.testing.assert_close(out, ref, rtol=0.05, atol=0.1) + + def test_torch_w8a8_block_fp8_linear_numerical_accuracy(self): + """Compare torch_w8a8_block_fp8_linear with dequantized baseline on XPU.""" + M, K, N = 16, 256, 256 + block_size = [128, 128] + + x = torch.randn(M, K, dtype=torch.bfloat16, device=self.device) + weight = torch.randn(N, K, dtype=torch.bfloat16, device=self.device).to( + torch.float8_e4m3fn + ) + weight_scale = ( + torch.rand(N // 128, K // 128, dtype=torch.float32, device=self.device) + + 0.1 + ) + bias = torch.randn(N, dtype=torch.bfloat16, device=self.device) + + out_torch = torch_w8a8_block_fp8_linear( + x, weight, block_size, weight_scale, bias=bias + ) + weight_scale_post_load = torch.empty_strided( + weight_scale.shape, + (1, weight_scale.shape[0]), + dtype=weight_scale.dtype, + device=weight_scale.device, + ) + weight_scale_post_load.copy_(weight_scale) + out_post_load = torch_w8a8_block_fp8_linear( + x, weight, block_size, weight_scale_post_load, bias=bias + ) + torch.testing.assert_close(out_post_load, out_torch, rtol=0, atol=0) + q_input, input_scale = per_token_group_quant_fp8(x, block_size[1]) + out_ref = reference_block_fp8_matmul( + q_input, input_scale, weight, block_size, weight_scale, bias=bias + ).to(torch.bfloat16) + + torch.testing.assert_close(out_torch, out_ref, rtol=0.05, atol=0.1) + + def test_torch_w8a8_block_fp8_linear_1x128_recipe(self): + """Validate the alternate [1, 128] weight-scale recipe with N != K.""" + M, K, N = 8, 256, 384 + block_size = [1, 128] + x = torch.randn(M, K, dtype=torch.bfloat16, device=self.device) + weight = torch.randn(N, K, dtype=torch.bfloat16, device=self.device).to( + torch.float8_e4m3fn + ) + weight_scale = ( + torch.rand(N, K // 128, dtype=torch.float32, device=self.device) + 0.1 + ) + bias = torch.randn(N, dtype=torch.bfloat16, device=self.device) + + out = torch_w8a8_block_fp8_linear( + x, weight, block_size, weight_scale, bias=bias + ) + weight_scale_post_load = torch.empty_strided( + weight_scale.shape, + (1, weight_scale.shape[0]), + dtype=weight_scale.dtype, + device=weight_scale.device, + ) + weight_scale_post_load.copy_(weight_scale) + out_post_load = torch_w8a8_block_fp8_linear( + x, weight, block_size, weight_scale_post_load, bias=bias + ) + torch.testing.assert_close(out_post_load, out, rtol=0, atol=0) + q_input, input_scale = per_token_group_quant_fp8(x, block_size[1]) + ref = reference_block_fp8_matmul( + q_input, input_scale, weight, block_size, weight_scale, bias + ).to(torch.bfloat16) + torch.testing.assert_close(out, ref, rtol=0.05, atol=0.1) + + +if __name__ == "__main__": + unittest.main()