[AMD] Fix weight checking for AITER-shuffled block FP8 weights (#34330)

This commit is contained in:
Xinyu Kang
2026-09-11 15:49:56 -07:00
committed by GitHub
parent f963d7a27c
commit 6671cfc775
6 changed files with 211 additions and 17 deletions
@@ -0,0 +1,50 @@
# Copyright 2023-2024 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
import unittest
import torch
from sglang.srt.layers.quantization.fp8_utils import unshuffle_aiter_fp8_weight
from sglang.srt.utils import is_hip
from sglang.test.ci.ci_register import register_amd_ci
from sglang.test.test_utils import CustomTestCase
register_amd_ci(est_time=5, stage="jit-kernel-unit", runner_config="amd")
@unittest.skipUnless(is_hip(), "requires ROCm AITER")
class TestAiterFp8Utils(CustomTestCase):
def test_unshuffle_weight_round_trip(self):
from aiter.ops.shuffle import shuffle_weight
for shape in ((32, 64), (2, 32, 64)):
with self.subTest(shape=shape):
logical = (
torch.arange(
torch.Size(shape).numel(), device="cuda", dtype=torch.float32
)
.remainder(7)
.to(torch.float8_e4m3fn)
.reshape(shape)
)
shuffled = shuffle_weight(logical, layout=(16, 16))
torch.testing.assert_close(
unshuffle_aiter_fp8_weight(shuffled), logical
)
if __name__ == "__main__":
unittest.main()
@@ -25,6 +25,7 @@ from sglang.srt.layers.quantization.fp8_utils import (
quant_weight_ue8m0,
transform_scale_ue8m0,
)
from sglang.srt.utils import is_hip
from sglang.srt.utils.weight_checker import (
CheckEntry,
ChecksumInfo,
@@ -42,9 +43,10 @@ from sglang.srt.utils.weight_checker_comparator import (
Fp8BlockComparable,
RawComparable,
)
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_amd_ci(est_time=30, suite="stage-b-test-1-gpu-small-amd")
register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-small")
@@ -75,6 +77,7 @@ def _assert_entries_close(
torch.testing.assert_close(
a_ref.w_s, e_ref.w_s, msg=f"[{i}] w_s {a_name!r}"
)
assert a_ref.is_shuffled == e_ref.is_shuffled
else:
torch.testing.assert_close(
a_ref.tensor, e_ref.tensor, msg=f"[{i}] tensor {a_name!r}"
@@ -82,18 +85,61 @@ def _assert_entries_close(
def _build_fp8_quant_pair(device: str = "cuda"):
"""Construct a real fp8-quantized weight + matching fp32 + ue8m0-packed scales.
"""Construct a real fp8-quantized weight and matching fp32 scales.
Returns (qweight, sf_fp32, sf_packed_int32) so callers can pick which scale dtype
drives the _build_check_entries branch under test.
Returns (qweight, sf_fp32).
"""
weight_bf16 = torch.randn((256, 128), dtype=torch.bfloat16, device=device)
block_size = [128, 128]
qweight, sf_fp32 = quant_weight_ue8m0(
weight_dequant=weight_bf16, weight_block_size=block_size
)
sf_packed_int32 = transform_scale_ue8m0(sf_fp32, mn=qweight.shape[-2])
return qweight, sf_fp32, sf_packed_int32
return qweight, sf_fp32
# ---------------------------------------------------------------------------
# Shuffled FP8 integration
# ---------------------------------------------------------------------------
class TestShuffledFp8Comparable(CustomTestCase):
def test_iter_chunks_unshuffles_before_dequantization(self):
shuffled = torch.zeros((32, 64), dtype=torch.float8_e4m3fn)
scale = torch.ones((2, 4), dtype=torch.float32)
comparable = Fp8BlockComparable(shuffled, scale, is_shuffled=True)
with (
patch(
"sglang.srt.utils.weight_checker_comparator.unshuffle_fp8_weight",
side_effect=lambda weight: weight,
) as unshuffle,
patch(
"sglang.srt.utils.weight_checker_comparator.block_quant_dequant",
side_effect=lambda weight, *_args, **_kwargs: weight,
),
):
next(iter(comparable.iter_chunks()))
unshuffle.assert_called_once()
def test_dequantize_unshuffles_before_checksum(self):
shuffled = torch.zeros((32, 64), dtype=torch.float8_e4m3fn)
scale = torch.ones((2, 4), dtype=torch.float32)
comparable = Fp8BlockComparable(shuffled, scale, is_shuffled=True)
with (
patch(
"sglang.srt.utils.weight_checker_comparator.unshuffle_fp8_weight",
side_effect=lambda weight: weight,
) as unshuffle,
patch(
"sglang.srt.utils.weight_checker_comparator.block_quant_dequant",
side_effect=lambda weight, *_args, **_kwargs: weight,
),
):
comparable.dequantize()
unshuffle.assert_called_once_with(shuffled)
# ---------------------------------------------------------------------------
@@ -113,6 +159,8 @@ class _TinyModel(nn.Module):
self.register_buffer("running_mean", torch.zeros(4))
# Buffer names used to exercise weight checker's hard-coded filters.
self.register_buffer("rotary_emb_cos_sin_cache", torch.full((8,), 3.14))
self.register_buffer("rotary_emb_cos_cache", torch.full((8,), 1.62))
self.register_buffer("rotary_emb_sin_cache", torch.full((8,), 0.58))
self.register_buffer("rotary_emb_freqs_cis", torch.full((8,), 2.71))
self.register_buffer("gate_proj_weight_fp32_cache", torch.full((8,), 1.41))
@@ -260,8 +308,10 @@ class TestPostprocessTensors(CustomTestCase):
# --- fp8 quant pair (real dequant on real fp8 tensors) ---
@unittest.skipIf(is_hip(), "DeepGEMM is not supported on ROCm")
def test_fp8_quant_pair_yields_lazy_pair(self):
qweight, sf_fp32, sf_packed_int32 = _build_fp8_quant_pair()
qweight, sf_fp32 = _build_fp8_quant_pair()
sf_packed_int32 = transform_scale_ue8m0(sf_fp32, mn=qweight.shape[-2])
raw = {"x.weight": qweight, "x.weight_scale_inv": sf_packed_int32}
ref = Fp8BlockComparable(qweight, sf_packed_int32)
@@ -273,8 +323,24 @@ class TestPostprocessTensors(CustomTestCase):
[("x.weight", True, ref)],
)
def test_fp8_quant_pair_preserves_shuffled_flag(self):
qweight = torch.zeros((128, 128), dtype=torch.float8_e4m3fn)
scale = torch.ones((1, 1), dtype=torch.float32)
raw = {"x.weight": qweight, "x.weight_scale_inv": scale}
quantized_set = {
"x.weight": QuantizedWeight(
Fp8BlockComparable,
"x.weight_scale_inv",
is_shuffled=True,
)
}
_assert_entries_close(
_build_check_entries(raw, set(), quantized_set),
[("x.weight", True, Fp8BlockComparable(qweight, scale, True))],
)
def test_fp8_quant_pair_yield_order_alongside_other_entries(self):
qweight, sf_fp32, _ = _build_fp8_quant_pair()
qweight, sf_fp32 = _build_fp8_quant_pair()
bias = torch.ones(4, device="cuda")
raw = {
"x.weight": qweight,
@@ -453,11 +519,12 @@ class TestBuildQuantizedSet(CustomTestCase):
model.proj.register_parameter(
"weight_scale_inv", nn.Parameter(torch.zeros(1, 1), requires_grad=False)
)
model.proj.weight.is_shuffled = True
self.assertEqual(
_build_quantized_set(model),
{
"proj.weight": QuantizedWeight(
Fp8BlockComparable, "proj.weight_scale_inv"
Fp8BlockComparable, "proj.weight_scale_inv", is_shuffled=True
)
},
)
@@ -496,6 +563,8 @@ class TestSnapshot(_WeightCheckerTestBase):
"b",
"running_mean",
"rotary_emb_cos_sin_cache",
"rotary_emb_cos_cache",
"rotary_emb_sin_cache",
"rotary_emb_freqs_cis",
"gate_proj_weight_fp32_cache",
}
@@ -526,6 +595,16 @@ class TestResetTensors(_WeightCheckerTestBase):
self.checker._reset_tensors()
torch.testing.assert_close(self.model.rotary_emb_cos_sin_cache, before)
def test_skips_cos_cache(self):
before = self.model.rotary_emb_cos_cache.clone()
self.checker._reset_tensors()
torch.testing.assert_close(self.model.rotary_emb_cos_cache, before)
def test_skips_sin_cache(self):
before = self.model.rotary_emb_sin_cache.clone()
self.checker._reset_tensors()
torch.testing.assert_close(self.model.rotary_emb_sin_cache, before)
def test_skips_freqs_cis(self):
before = self.model.rotary_emb_freqs_cis.clone()
self.checker._reset_tensors()