From 800deaaefab76c03716e5c4c3f24c5c00f5c063f Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Wed, 6 May 2026 22:48:48 +0800 Subject: [PATCH] Add unit and end-to-end tests for weight checker (#24536) --- test/registered/rl/test_weight_checker_e2e.py | 154 ++++++ .../unit/utils/test_weight_checker.py | 467 ++++++++++++++++++ 2 files changed, 621 insertions(+) create mode 100644 test/registered/rl/test_weight_checker_e2e.py create mode 100644 test/registered/unit/utils/test_weight_checker.py diff --git a/test/registered/rl/test_weight_checker_e2e.py b/test/registered/rl/test_weight_checker_e2e.py new file mode 100644 index 000000000..85df5a606 --- /dev/null +++ b/test/registered/rl/test_weight_checker_e2e.py @@ -0,0 +1,154 @@ +# 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. +# ============================================================================== +"""End-to-end test for the /weights_checker HTTP endpoint. + +Exercises the full HTTP -> tokenizer_manager -> scheduler -> model_runner -> +WeightChecker chain on a real engine. Unit tests in +test/registered/unit/utils/test_weight_checker.py cover the in-module +logic; this file is the thin integration cover plus interaction with +update_weights_from_tensor.""" + +import unittest +from typing import List, Tuple + +import requests +import torch + +from sglang.srt.utils import MultiprocessingSerializer, kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=150, suite="stage-b-test-1-gpu-small") + +_MODEL_NAME = "Qwen/Qwen3-0.6B" +# We address the up half via the HF-style unfused name "up_proj.weight". sglang's +# stacked_params_mapping rewrites this to "gate_up_proj.weight" with shard_id=1, +# so the upload writes only the up half of the fused tensor. Sending the fused +# name directly hits a name.replace() collision (gate_up_proj contains up_proj), +# producing a malformed key like "gate_gate_up_proj.weight" and crashing load. +_UP_PROJ_SHAPE = (3072, 1024) # intermediate_size, hidden_size for Qwen3-0.6B + + +class TestWeightCheckerE2E(CustomTestCase): + """All cases share one launched server (setUpClass). + + The reset case mutates weights to random; it is named to sort last so any + case that needs intact weights runs first. The server is torn down right + after, so leaving the engine in a corrupted state is harmless.""" + + @classmethod + def setUpClass(cls): + cls.url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + _MODEL_NAME, + cls.url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def _post(self, action: str) -> requests.Response: + return requests.post( + f"{self.url}/weights_checker", json={"action": action}, timeout=120 + ) + + def _update_weights( + self, named_tensors: List[Tuple[str, torch.Tensor]] + ) -> requests.Response: + return requests.post( + f"{self.url}/update_weights_from_tensor", + json={ + "serialized_named_tensors": [ + MultiprocessingSerializer.serialize(named_tensors, output_str=True) + ], + "flush_cache": True, + }, + timeout=120, + ) + + def test_a_snapshot_then_compare_unchanged_succeeds(self): + resp = self._post("snapshot") + self.assertEqual(resp.status_code, 200) + self.assertTrue(resp.json()["success"]) + + resp = self._post("compare") + self.assertEqual(resp.status_code, 200) + self.assertTrue(resp.json()["success"]) + + def test_b_unknown_action_returns_400(self): + resp = self._post("nonsense_action") + self.assertEqual(resp.status_code, 400) + self.assertIn("Unsupported", resp.json()["message"]) + + def test_c_update_with_diff_tensor_makes_compare_fail(self): + """A snapshot then an update with new bytes must make compare fail.""" + self.assertEqual(self._post("snapshot").status_code, 200) + + # The unfused HF name "up_proj" is what update_weights_from_tensor accepts; + # sglang's loader rewrites it onto the fused gate_up_proj tensor. + upload_name = "model.layers.5.mlp.up_proj.weight" + new_tensor = torch.full(_UP_PROJ_SHAPE, 1.5, device="cuda") + update_resp = self._update_weights([(upload_name, new_tensor)]) + self.assertEqual(update_resp.status_code, 200) + self.assertTrue(update_resp.json()["success"]) + + resp = self._post("compare") + self.assertEqual(resp.status_code, 400) + body = resp.json() + self.assertFalse(body["success"]) + # The error references the fused on-device parameter name, not the upload alias. + self.assertIn("model.layers.5.mlp.gate_up_proj.weight", body["message"]) + self.assertIn("max_abs_err", body["message"]) + + def test_d_update_with_same_tensor_keeps_compare_passing(self): + """Prime a param, snapshot, push the same bytes again, compare must pass.""" + param_name = "model.layers.6.mlp.up_proj.weight" + same_tensor = torch.full(_UP_PROJ_SHAPE, 0.25, device="cuda") + + # Step 1: prime the param to a known value. + self.assertTrue( + self._update_weights([(param_name, same_tensor)]).json()["success"] + ) + # Step 2: snapshot the now-primed state. + self.assertEqual(self._post("snapshot").status_code, 200) + # Step 3: push the exact same bytes again — should be a byte-perfect no-op. + self.assertTrue( + self._update_weights([(param_name, same_tensor)]).json()["success"] + ) + # Step 4: compare passes. + resp = self._post("compare") + self.assertEqual(resp.status_code, 200) + self.assertTrue(resp.json()["success"]) + + def test_z_snapshot_reset_compare_detects_diff(self): + """Destructive: leaves weights randomized. Named test_z_* so it runs last.""" + self.assertEqual(self._post("snapshot").status_code, 200) + self.assertEqual(self._post("reset_tensors").status_code, 200) + + resp = self._post("compare") + self.assertEqual(resp.status_code, 400) + body = resp.json() + self.assertFalse(body["success"]) + self.assertIn("max_abs_err", body["message"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/utils/test_weight_checker.py b/test/registered/unit/utils/test_weight_checker.py new file mode 100644 index 000000000..ca7b75102 --- /dev/null +++ b/test/registered/unit/utils/test_weight_checker.py @@ -0,0 +1,467 @@ +# 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. +# ============================================================================== +"""Unit tests for sglang/srt/utils/weight_checker.py.""" + +import unittest +from typing import Iterable, List, Tuple +from unittest.mock import patch + +import torch +from torch import nn + +from sglang.srt.layers.quantization.fp8_utils import ( + block_quant_dequant, + quant_weight_ue8m0, + transform_scale_ue8m0, +) +from sglang.srt.utils.weight_checker import ( + WeightChecker, + _check_tensors, + _postprocess_tensors, + _random_like, +) +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=30, suite="stage-b-test-1-gpu-small") + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +Triple = Tuple[str, bool, torch.Tensor] + + +def _assert_triples_close(actual: Iterable[Triple], expected: Iterable[Triple]) -> None: + """Compare two streams of (name, should_compare, tensor); element-wise tensor close.""" + actual_list: List[Triple] = list(actual) + expected_list: List[Triple] = list(expected) + assert len(actual_list) == len( + expected_list + ), f"length mismatch: actual={len(actual_list)} expected={len(expected_list)}" + for i, ((a_name, a_flag, a_t), (e_name, e_flag, e_t)) in enumerate( + zip(actual_list, expected_list) + ): + assert a_name == e_name, f"[{i}] name: {a_name!r} != {e_name!r}" + assert a_flag == e_flag, f"[{i}] should_compare: {a_flag} != {e_flag}" + torch.testing.assert_close( + a_t, e_t, msg=f"[{i}] tensor mismatch for {a_name!r}" + ) + + +def _build_fp8_quant_pair(device: str = "cuda"): + """Construct a real fp8-quantized weight + matching fp32 + ue8m0-packed scales. + + Returns (qweight, sf_fp32, sf_packed_int32) so callers can pick which scale dtype + drives the _postprocess_tensors branch under test. + """ + 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 + + +# --------------------------------------------------------------------------- +# Test fixtures +# --------------------------------------------------------------------------- + + +class _TinyModel(nn.Module): + """Mimics the buffer naming patterns _reset_tensors / _postprocess_tensors care about.""" + + def __init__(self): + super().__init__() + # requires_grad=False matches sglang's inference-time params, so _reset_tensors + # can do in-place copy_ on them (autograd would otherwise reject it). + self.w = nn.Parameter(torch.randn(4, 4), requires_grad=False) + self.b = nn.Parameter(torch.zeros(4), requires_grad=False) + self.register_buffer("running_mean", torch.zeros(4)) + # Buffer names that match weight_checker's hard-coded skip patterns. + self.register_buffer("rotary_emb_cos_sin_cache", torch.full((8,), 3.14)) + 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)) + + +class _FakeModelRunner: + """Minimal stand-in: WeightChecker only touches `.model.named_parameters()` and + `.model.named_buffers()`, nothing else.""" + + def __init__(self, model: nn.Module): + self.model = model + + +# --------------------------------------------------------------------------- +# _random_like +# --------------------------------------------------------------------------- + + +class TestRandomLike(CustomTestCase): + + def test_floating_point_preserves_dtype_shape_device(self): + for dtype in (torch.float32, torch.float16, torch.bfloat16): + t = torch.zeros(8, 4, dtype=dtype) + out = _random_like(t) + self.assertEqual(out.dtype, dtype) + self.assertEqual(out.shape, t.shape) + self.assertEqual(out.device, t.device) + self.assertGreater(out.float().abs().sum().item(), 0) + + def test_bool_returns_bool_with_both_values(self): + t = torch.zeros(1024, dtype=torch.bool) + out = _random_like(t) + self.assertEqual(out.dtype, torch.bool) + self.assertEqual(out.shape, t.shape) + self.assertEqual(out.device, t.device) + self.assertTrue(out.any().item()) + self.assertFalse(out.all().item()) + + def test_int_returns_correct_dtype_in_range(self): + for dtype in (torch.int8, torch.int32, torch.int64): + t = torch.zeros(256, dtype=dtype) + out = _random_like(t) + self.assertEqual(out.dtype, dtype) + self.assertEqual(out.shape, t.shape) + info = torch.iinfo(dtype) + self.assertGreaterEqual(out.min().item(), info.min) + self.assertLessEqual(out.max().item(), info.max) + self.assertGreater(out.unique().numel(), 1) + + def test_floating_point_values_in_unit_range(self): + t = torch.zeros(1024, dtype=torch.float32) + out = _random_like(t) + self.assertGreaterEqual(out.min().item(), 0.0) + self.assertLess(out.max().item(), 1.0) + + def test_does_not_mutate_input(self): + t = torch.full((16,), 5.0) + before = t.clone() + _random_like(t) + torch.testing.assert_close(t, before) + + +# --------------------------------------------------------------------------- +# _postprocess_tensors +# --------------------------------------------------------------------------- + + +class TestPostprocessTensors(CustomTestCase): + + # --- non-quant / non-skip --- + + def test_no_quant_yields_raw_with_should_compare_true(self): + a = torch.randn(4) + b = torch.randn(4) + raw = {"a.weight": a, "b.bias": b} + _assert_triples_close( + _postprocess_tensors(raw), + [("a.weight", True, a), ("b.bias", True, b)], + ) + + def test_weight_alone_without_scale_inv_does_not_trigger_dequant(self): + w = torch.randn(4) + raw = {"x.weight": w} + _assert_triples_close(_postprocess_tensors(raw), [("x.weight", True, w)]) + + # --- non-persistent buffer skip --- + + def test_skips_cos_sin_cache_substring(self): + cache = torch.randn(8) + plain = torch.randn(4) + raw = { + "model.rotary_emb.cos_sin_cache": cache, + "model.layers.0.weight": plain, + } + _assert_triples_close( + _postprocess_tensors(raw), + [ + ("model.rotary_emb.cos_sin_cache", False, cache), + ("model.layers.0.weight", True, plain), + ], + ) + + def test_skips_inv_freq_substring(self): + t = torch.randn(4) + _assert_triples_close( + _postprocess_tensors({"model.rotary_emb.inv_freq": t}), + [("model.rotary_emb.inv_freq", False, t)], + ) + + def test_skips_weight_fp32_substring(self): + t = torch.randn(4) + _assert_triples_close( + _postprocess_tensors({"model.layers.0.mlp.gate._weight_fp32": t}), + [("model.layers.0.mlp.gate._weight_fp32", False, t)], + ) + + def test_substring_match_not_endswith(self): + # Pattern can appear anywhere in the name, not just at the end. + t = torch.randn(4) + _assert_triples_close( + _postprocess_tensors({"weird.cos_sin_cache.foo.bar": t}), + [("weird.cos_sin_cache.foo.bar", False, t)], + ) + + # --- fp8 quant pair (real dequant on real fp8 tensors) --- + + def test_fp8_quant_pair_with_int32_scale_dequants_via_ue8m0(self): + qweight, sf_fp32, sf_packed_int32 = _build_fp8_quant_pair() + raw = {"x.weight": qweight, "x.weight_scale_inv": sf_packed_int32} + + # Reference: ue8m0 path inside _postprocess_tensors should eventually + # call block_quant_dequant with the unpacked fp32 scale. + expected_dequant = block_quant_dequant( + qweight, sf_fp32, block_size=[128, 128], dtype=torch.bfloat16 + ) + _assert_triples_close( + _postprocess_tensors(raw), + [ + ("x.weight", True, expected_dequant), + ("x.weight", False, qweight), + ("x.weight_scale_inv", False, sf_packed_int32), + ], + ) + + def test_fp8_quant_pair_with_fp32_scale_dequants_directly(self): + qweight, sf_fp32, _ = _build_fp8_quant_pair() + raw = {"x.weight": qweight, "x.weight_scale_inv": sf_fp32} + + expected_dequant = block_quant_dequant( + qweight, sf_fp32, block_size=[128, 128], dtype=torch.bfloat16 + ) + _assert_triples_close( + _postprocess_tensors(raw), + [ + ("x.weight", True, expected_dequant), + ("x.weight", False, qweight), + ("x.weight_scale_inv", False, sf_fp32), + ], + ) + + def test_fp8_quant_pair_yield_order_alongside_other_entries(self): + qweight, sf_fp32, _ = _build_fp8_quant_pair() + bias = torch.ones(4, device="cuda") + raw = { + "x.weight": qweight, + "x.weight_scale_inv": sf_fp32, + "y.bias": bias, + } + expected_dequant = block_quant_dequant( + qweight, sf_fp32, block_size=[128, 128], dtype=torch.bfloat16 + ) + # All dequant entries come first, then a raw pass over every key. + _assert_triples_close( + _postprocess_tensors(raw), + [ + ("x.weight", True, expected_dequant), + ("x.weight", False, qweight), + ("x.weight_scale_inv", False, sf_fp32), + ("y.bias", True, bias), + ], + ) + + def test_only_scale_without_weight_does_not_trigger_dequant(self): + # Without the matching `.weight`, no quant pair forms; the scale_inv flows + # through as a normal entry with should_compare=True. + s = torch.zeros(1, 1, dtype=torch.int32) + _assert_triples_close( + _postprocess_tensors({"x.weight_scale_inv": s}), + [("x.weight_scale_inv", True, s)], + ) + + +# --------------------------------------------------------------------------- +# _check_tensors (implementation moves both sides via .cuda()) +# --------------------------------------------------------------------------- + + +class TestCheckTensors(CustomTestCase): + + def test_passes_when_all_equal(self): + t = torch.ones(2, 2) + expect = [("a", True, t.clone()), ("b", True, t.clone())] + actual = [("a", True, t.clone()), ("b", True, t.clone())] + _check_tensors(expect_tensors=expect, actual_tensors=actual) + + def test_raises_when_should_compare_true_and_diff(self): + expect = [("a", True, torch.ones(2, 2))] + actual = [("a", True, torch.zeros(2, 2))] + with self.assertRaises(Exception) as ctx: + _check_tensors(expect_tensors=expect, actual_tensors=actual) + msg = str(ctx.exception) + self.assertIn("name=a", msg) + self.assertIn("max_abs_err", msg) + + def test_passes_when_should_compare_false_even_if_diff(self): + # should_compare=False -> diff is logged, not raised. + expect = [("a", False, torch.ones(2, 2))] + actual = [("a", False, torch.zeros(2, 2))] + _check_tensors(expect_tensors=expect, actual_tensors=actual) + + def test_asserts_on_name_mismatch(self): + expect = [("a", True, torch.ones(2, 2))] + actual = [("b", True, torch.ones(2, 2))] + with self.assertRaises(AssertionError): + _check_tensors(expect_tensors=expect, actual_tensors=actual) + + def test_asserts_on_should_compare_mismatch(self): + expect = [("a", True, torch.ones(2, 2))] + actual = [("a", False, torch.ones(2, 2))] + with self.assertRaises(AssertionError): + _check_tensors(expect_tensors=expect, actual_tensors=actual) + + def test_zip_strict_raises_on_length_mismatch(self): + t = torch.ones(2, 2) + expect = [("a", True, t.clone()), ("b", True, t.clone())] + actual = [("a", True, t.clone())] + with self.assertRaises(ValueError): + _check_tensors(expect_tensors=expect, actual_tensors=actual) + + +# --------------------------------------------------------------------------- +# WeightChecker class +# --------------------------------------------------------------------------- + + +class _WeightCheckerTestBase(CustomTestCase): + """Shared fixture: fresh _TinyModel + WeightChecker per test, on CUDA. + + The model lives on CUDA so that _snapshot's `.detach().cpu()` produces + an independent CPU copy. On a CPU model `.cpu()` is a no-op and the + snapshot would alias the live storage, which masks reset-then-compare + divergence. + """ + + def setUp(self): + torch.manual_seed(0) + self.model = _TinyModel().cuda() + self.checker = WeightChecker(model_runner=_FakeModelRunner(self.model)) + + +class TestSnapshot(_WeightCheckerTestBase): + + def test_captures_params_and_buffers(self): + self.checker._snapshot() + keys = set(self.checker._snapshot_tensors.keys()) + expected = { + "w", + "b", + "running_mean", + "rotary_emb_cos_sin_cache", + "rotary_emb_freqs_cis", + "gate_proj_weight_fp32_cache", + } + self.assertEqual(keys, expected) + + def test_detaches_and_moves_to_cpu(self): + self.checker._snapshot() + for tensor in self.checker._snapshot_tensors.values(): + self.assertEqual(tensor.device.type, "cpu") + # Mutating the live model must not affect the snapshot copy. + original_w = self.checker._snapshot_tensors["w"].clone() + with torch.no_grad(): + self.model.w.data.fill_(99.0) + torch.testing.assert_close(self.checker._snapshot_tensors["w"], original_w) + + +class TestResetTensors(_WeightCheckerTestBase): + + def test_changes_normal_params_in_place(self): + before_w = self.model.w.clone() + before_w_ptr = self.model.w.data_ptr() + self.checker._reset_tensors() + # In-place: storage pointer unchanged. + self.assertEqual(self.model.w.data_ptr(), before_w_ptr) + self.assertFalse(torch.equal(self.model.w, before_w)) + + def test_skips_cos_sin_cache(self): + before = self.model.rotary_emb_cos_sin_cache.clone() + self.checker._reset_tensors() + torch.testing.assert_close(self.model.rotary_emb_cos_sin_cache, before) + + def test_skips_freqs_cis(self): + before = self.model.rotary_emb_freqs_cis.clone() + self.checker._reset_tensors() + torch.testing.assert_close(self.model.rotary_emb_freqs_cis, before) + + def test_skips_weight_fp32(self): + before = self.model.gate_proj_weight_fp32_cache.clone() + self.checker._reset_tensors() + torch.testing.assert_close(self.model.gate_proj_weight_fp32_cache, before) + + +class TestCompare(_WeightCheckerTestBase): + + def test_without_snapshot_raises(self): + with self.assertRaises(AssertionError): + self.checker._compare() + + def test_passes_when_unchanged(self): + self.checker._snapshot() + self.checker._compare() # no exception + + def test_fails_after_reset_on_normal_param(self): + self.checker._snapshot() + self.checker._reset_tensors() + with self.assertRaises(Exception) as ctx: + self.checker._compare() + msg = str(ctx.exception) + self.assertTrue(("name=w" in msg) or ("name=b" in msg)) + + def test_passes_when_only_skipped_buffer_diverges(self): + self.checker._snapshot() + # Mutate a non-persistent skip-pattern buffer; compare must still pass. + with torch.no_grad(): + self.model.rotary_emb_cos_sin_cache.fill_(99.0) + self.checker._compare() + + def test_passes_after_reset_then_restoring_normal_params(self): + # Full lifecycle: reset (skips cos_sin_cache et al.), then restore non-skip + # params by hand. Compare must pass — proving reset+postprocess skip lists agree. + self.checker._snapshot() + snapshot = {k: v.clone() for k, v in self.checker._snapshot_tensors.items()} + self.checker._reset_tensors() + with torch.no_grad(): + for name, tensor in self.model.named_parameters(): + tensor.data.copy_(snapshot[name].to(tensor.device)) + for name, tensor in self.model.named_buffers(): + tensor.data.copy_(snapshot[name].to(tensor.device)) + self.checker._compare() + + +class TestHandle(_WeightCheckerTestBase): + + def test_routes_to_actions(self): + with patch.object(self.checker, "_snapshot") as m_snap, patch.object( + self.checker, "_reset_tensors" + ) as m_reset, patch.object(self.checker, "_compare") as m_compare: + self.checker.handle("snapshot") + self.checker.handle("reset_tensors") + self.checker.handle("compare") + m_snap.assert_called_once() + m_reset.assert_called_once() + m_compare.assert_called_once() + + def test_unknown_action_raises(self): + with self.assertRaises(Exception) as ctx: + self.checker.handle("nonsense_action") + self.assertIn("Unsupported", str(ctx.exception)) + + +if __name__ == "__main__": + unittest.main()