From 1685d29f21279349907d8ebbb417a7c339af07a8 Mon Sep 17 00:00:00 2001 From: sonle5 <51179712+hdt98@users.noreply.github.com> Date: Sun, 2 Aug 2026 14:06:09 +0700 Subject: [PATCH] [AMD] Fix DeepSeek-V4 fused-RMS FP8 scale metadata on gfx950 (#31727) --- .../srt/layers/quantization/fp8_utils.py | 11 ++++ python/sglang/srt/models/deepseek_v4.py | 8 +++ .../quant/test_fused_rms_fp8_group_quant.py | 42 ++++++++++++++ .../unit/layers/test_fp8_bpreshuffle_scale.py | 55 +++++++++++++++++++ 4 files changed, 116 insertions(+) diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index 037ae9412..38a524bea 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -112,6 +112,17 @@ def materialize_bpreshuffle_fp8_scale(scale: torch.Tensor) -> torch.Tensor: return scale.t().contiguous().t() if scale.dim() == 2 else scale +def view_aiter_fused_rms_transposed_fp8_scale(scale: torch.Tensor) -> torch.Tensor: + """Expose AITER fused-RMS ``transpose_scale=True`` storage logically. + + The fused-RMS op returns transposed physical bytes through a row-major-looking + view. Restore logical ``[M, G]`` indexing without copying those bytes. + """ + if scale.dim() != 2: + return scale + return torch.as_strided(scale, scale.shape, (1, scale.shape[0])) + + def materialize_bpreshuffle_fp8_scale_tuple( value: Tuple[torch.Tensor, ...], ) -> Tuple[torch.Tensor, ...]: diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 52fb8c098..0c59f2c6c 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -84,6 +84,9 @@ from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.moe import get_moe_a2a_backend, should_use_dp_reduce_scatterv from sglang.srt.layers.moe.fused_moe_triton import FusedMoE +from sglang.srt.layers.quantization.fp8_utils import ( + view_aiter_fused_rms_transposed_fp8_scale, +) from sglang.srt.layers.rotary_embedding import get_rope_wrapper from sglang.srt.layers.utils import PPMissingLayer, get_layer_id from sglang.srt.layers.utils.cp_utils import ( @@ -243,6 +246,11 @@ def _fused_rmsnorm_fp8_quant(hidden_states, weight, eps): output_unquantized_inp1=True, transpose_scale=_use_aiter_bpreshuffle_gfx95, ) + if _use_aiter_bpreshuffle_gfx95: + x_quant = ( + x_quant[0], + view_aiter_fused_rms_transposed_fp8_scale(x_quant[1]), + ) return x_quant, x_bf16 diff --git a/test/registered/quant/test_fused_rms_fp8_group_quant.py b/test/registered/quant/test_fused_rms_fp8_group_quant.py index 02dbfbbaf..793c6a131 100644 --- a/test/registered/quant/test_fused_rms_fp8_group_quant.py +++ b/test/registered/quant/test_fused_rms_fp8_group_quant.py @@ -4,6 +4,10 @@ import unittest import torch import torch.nn.functional as F +from sglang.srt.layers.quantization.fp8_utils import ( + materialize_bpreshuffle_fp8_scale, + view_aiter_fused_rms_transposed_fp8_scale, +) from sglang.test.ci.ci_register import register_amd_ci from sglang.test.test_utils import CustomTestCase @@ -145,6 +149,44 @@ class TestFusedRMSFP8GroupQuant(CustomTestCase): with self.subTest(M=M, N1=N1, N2=N2, group_size=g, dtype=dtype, seed=seed): self._case(M, N1, N2, g, dtype, seed) + def test_transposed_scale_matches_bpreshuffle_layout_contract(self): + from aiter.ops.triton.fused_fp8_quant import fused_rms_fp8_group_quant + + common_kwargs = dict( + inp2=None, + inp2_weight=None, + inp2_epsilon=None, + group_size=128, + dtype_quant=torch.float8_e4m3fn, + res1=None, + output_unquantized_inp1=False, + ) + + for m, k in ((1, 1024), (64, 1024), (1, 4096), (64, 4096)): + with self.subTest(m=m, k=k): + torch.manual_seed(0) + x = torch.randn(m, k, dtype=torch.bfloat16, device="cuda") + weight = torch.ones(k, dtype=torch.float32, device="cuda") + + (q_row_major, scale_row_major), *_ = fused_rms_fp8_group_quant( + x, weight, 1e-6, transpose_scale=False, **common_kwargs + ) + (q_transposed, scale_transposed), *_ = fused_rms_fp8_group_quant( + x, weight, 1e-6, transpose_scale=True, **common_kwargs + ) + + repaired = view_aiter_fused_rms_transposed_fp8_scale(scale_transposed) + materialized = materialize_bpreshuffle_fp8_scale(repaired) + + torch.testing.assert_close(q_transposed, q_row_major, rtol=0, atol=0) + torch.testing.assert_close(repaired, scale_row_major, rtol=0, atol=0) + torch.testing.assert_close( + materialized, scale_row_major, rtol=0, atol=0 + ) + self.assertEqual(repaired.stride(), (1, repaired.shape[0])) + self.assertEqual(repaired.data_ptr(), scale_transposed.data_ptr()) + self.assertEqual(materialized.data_ptr(), scale_transposed.data_ptr()) + if __name__ == "__main__": unittest.main(verbosity=2) diff --git a/test/registered/unit/layers/test_fp8_bpreshuffle_scale.py b/test/registered/unit/layers/test_fp8_bpreshuffle_scale.py index 34b49c527..b9e12d884 100644 --- a/test/registered/unit/layers/test_fp8_bpreshuffle_scale.py +++ b/test/registered/unit/layers/test_fp8_bpreshuffle_scale.py @@ -1,10 +1,12 @@ import unittest +from unittest.mock import patch import torch from sglang.srt.layers.quantization.fp8_utils import ( materialize_bpreshuffle_fp8_scale, materialize_bpreshuffle_fp8_scale_tuple, + view_aiter_fused_rms_transposed_fp8_scale, ) from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -31,6 +33,59 @@ class TestBpreshuffleScaleMaterialization(CustomTestCase): self.assertTrue(torch.equal(rematerialized, scale)) self.assertEqual(rematerialized.stride(), materialized.stride()) + self.assertEqual(rematerialized.data_ptr(), materialized.data_ptr()) + + def test_repairs_aiter_scale_before_downstream_layout_handling(self): + """AITER-transposed scale bytes must retain their logical indexing. + + AITER ``transpose_scale=True`` returns transposed physical storage with + row-major-looking metadata. Treating that metadata as logical layout + permutes the scales during CK materialization. + """ + logical_scale = torch.arange(12, dtype=torch.float32).reshape(3, 4) + aiter_scale = logical_scale.t().contiguous().view(logical_scale.shape) + + repaired = view_aiter_fused_rms_transposed_fp8_scale(aiter_scale) + materialized = materialize_bpreshuffle_fp8_scale(repaired) + renormalized = view_aiter_fused_rms_transposed_fp8_scale(repaired) + + self.assertTrue(torch.equal(repaired, logical_scale)) + self.assertTrue(torch.equal(materialized, logical_scale)) + self.assertTrue(torch.equal(renormalized, logical_scale)) + self.assertEqual(repaired.stride(), (1, logical_scale.shape[0])) + self.assertEqual(repaired.data_ptr(), aiter_scale.data_ptr()) + self.assertEqual(materialized.data_ptr(), aiter_scale.data_ptr()) + self.assertEqual(renormalized.stride(), repaired.stride()) + self.assertEqual(renormalized.data_ptr(), aiter_scale.data_ptr()) + + def test_deepseek_v4_repairs_fused_rms_scale_at_producer(self): + """DeepSeek-V4 must repair fused-RMS scale metadata before CK consumes it.""" + from sglang.srt.models import deepseek_v4 + + q_input = torch.ones((3, 1024), dtype=torch.float32) + x_bf16 = torch.ones((3, 1024), dtype=torch.bfloat16) + logical_scale = torch.arange(24, dtype=torch.float32).reshape(3, 8) + aiter_scale = logical_scale.t().contiguous().view(logical_scale.shape) + fused_output = ((q_input, aiter_scale), x_bf16, None, None) + + with ( + patch.object( + deepseek_v4, + "fused_rms_fp8_group_quant", + return_value=fused_output, + create=True, + ), + patch.object(deepseek_v4, "_use_aiter_bpreshuffle_gfx95", True), + ): + x_quant, x_unquantized = deepseek_v4._fused_rmsnorm_fp8_quant( + q_input, torch.ones(1024), 1e-6 + ) + + self.assertIs(x_quant[0], q_input) + self.assertIs(x_unquantized, x_bf16) + self.assertTrue(torch.equal(x_quant[1], logical_scale)) + self.assertEqual(x_quant[1].stride(), (1, logical_scale.shape[0])) + self.assertEqual(x_quant[1].data_ptr(), aiter_scale.data_ptr()) def test_tuple_helper_keeps_extra_tuple_payload(self): q_input = torch.ones((3, 8), dtype=torch.float32)