Add unit and end-to-end tests for weight checker (#24536)
This commit is contained in:
@@ -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()
|
||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user