From e1ca92a7fd6bbbd5e20d88100784844dc7d4b5d8 Mon Sep 17 00:00:00 2001 From: Yueming Yuan Date: Sun, 28 Jun 2026 18:44:05 -0700 Subject: [PATCH] [weight checker] refactor: add precision branch; allow ULP quant err; used chunked compare (#28974) --- python/sglang/srt/managers/io_struct.py | 1 + .../scheduler_components/weight_updater.py | 9 +- .../sglang/srt/model_executor/model_runner.py | 6 +- python/sglang/srt/utils/weight_checker.py | 210 +++++++------ .../srt/utils/weight_checker_comparator.py | 166 ++++++++++ .../unit/utils/test_weight_checker.py | 297 ++++++++++++------ .../utils/test_weight_checker_comparator.py | 206 ++++++++++++ 7 files changed, 706 insertions(+), 189 deletions(-) create mode 100644 python/sglang/srt/utils/weight_checker_comparator.py create mode 100644 test/registered/unit/utils/test_weight_checker_comparator.py diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 3d97dc487..842f8faeb 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -1712,6 +1712,7 @@ class ResumeMemoryOccupationReqOutput(BaseReq, kw_only=True): class CheckWeightsReqInput(BaseReq, kw_only=True): action: str = "checksum" + allow_quant_error: bool = False class CheckWeightsReqOutput(BaseReq, kw_only=True): diff --git a/python/sglang/srt/managers/scheduler_components/weight_updater.py b/python/sglang/srt/managers/scheduler_components/weight_updater.py index 1abfeef97..8c135d12f 100644 --- a/python/sglang/srt/managers/scheduler_components/weight_updater.py +++ b/python/sglang/srt/managers/scheduler_components/weight_updater.py @@ -268,12 +268,17 @@ class SchedulerWeightUpdaterManager: def check_weights(self, recv_req: CheckWeightsReqInput): try: - payload = self.tp_worker.model_runner.check_weights(action=recv_req.action) + payload = self.tp_worker.model_runner.check_weights( + action=recv_req.action, allow_quant_error=recv_req.allow_quant_error + ) if self.draft_worker is not None: draft_runner = _get_draft_model_runner(self.draft_worker) if draft_runner is not None: - draft_payload = draft_runner.check_weights(action=recv_req.action) + draft_payload = draft_runner.check_weights( + action=recv_req.action, + allow_quant_error=recv_req.allow_quant_error, + ) if payload is not None and draft_payload is not None: payload = _merge_checksum_payloads(payload, draft_payload) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 1afa7cfe6..08b6a550f 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -3205,8 +3205,10 @@ class ModelRunner(ModelRunnerKVCacheMixin): ) ShardedStateLoader.save_model(self.model, path, pattern, max_size) - def check_weights(self, action: str): - return self._weight_checker.handle(action=action) + def check_weights(self, action: str, allow_quant_error: bool = False): + return self._weight_checker.handle( + action=action, allow_quant_error=allow_quant_error + ) def update_weights_from_ipc(self, recv_req): """Update weights from IPC for checkpoint-engine integration.""" diff --git a/python/sglang/srt/utils/weight_checker.py b/python/sglang/srt/utils/weight_checker.py index 8e6459a37..5d44d05c1 100644 --- a/python/sglang/srt/utils/weight_checker.py +++ b/python/sglang/srt/utils/weight_checker.py @@ -1,17 +1,20 @@ import hashlib import logging import time -from typing import Dict, Iterable, Optional, Set, Tuple +from typing import Dict, Iterable, NamedTuple, Optional, Set import torch import torch.distributed as dist from pydantic import BaseModel, ConfigDict -from sglang.srt.layers.quantization.fp8_utils import ( - block_quant_dequant, - inverse_transform_scale_ue8m0, -) from sglang.srt.managers.mm_utils import tensor_hash +from sglang.srt.utils.weight_checker_comparator import ( + CHUNK_NUMEL, + ComparableWeight, + RawComparable, + compare_weights, + select_comparable_weight, +) logger = logging.getLogger(__name__) @@ -37,6 +40,17 @@ class ChecksumInfo(_StrictBaseModel): parallelism_info: ParallelismInfo +class CheckEntry(NamedTuple): + name: str + should_compare: bool + comparable: ComparableWeight + + +class QuantizedWeight(NamedTuple): + comparable_cls: type[ComparableWeight] + scale_name: str + + _NON_PERSISTENT_BUFFER_PATTERNS = ( "cos_sin_cache", "inv_freq", @@ -54,14 +68,16 @@ class WeightChecker: self._model_runner = model_runner self._snapshot_tensors = None - def handle(self, action: str) -> Optional[Dict]: - logger.info(f"[WeightChecker] handle action={action}") + def handle(self, action: str, allow_quant_error: bool = False) -> Optional[Dict]: + logger.info( + f"[WeightChecker] handle action={action} allow_quant_error={allow_quant_error}" + ) if action == "snapshot": return self._snapshot() elif action == "reset_tensors": return self._reset_tensors() elif action == "compare": - return self._compare() + return self._compare(allow_quant_error=allow_quant_error) elif action == "checksum": return self._compute_checksum() else: @@ -82,43 +98,44 @@ class WeightChecker: continue param.copy_(_random_like(param)) - def _compare(self): + def _compare(self, allow_quant_error: bool = False): assert self._snapshot_tensors is not None + quantized_set = _build_quantized_set(self._model_runner.model) skip_compare_names = { name for name, param in self._model_state() if getattr(param, "_skip_weight_check", False) } _check_tensors( - expect_tensors=_postprocess_tensors( - self._snapshot_tensors, skip_compare_names + expect_tensors=_build_check_entries( + self._snapshot_tensors, skip_compare_names, quantized_set ), - actual_tensors=_postprocess_tensors( - dict(self._model_state()), skip_compare_names + actual_tensors=_build_check_entries( + dict(self._model_state()), skip_compare_names, quantized_set ), + allow_quant_error=allow_quant_error, ) def _compute_checksum(self) -> Dict: torch.cuda.synchronize() start = time.perf_counter() + quantized_set = _build_quantized_set(self._model_runner.model) skip_compare_names = { name for name, param in self._model_state() if getattr(param, "_skip_weight_check", False) } - # Reuse the snapshot/compare postprocess pipeline so fp8 weights are - # dequantized to bf16 before hashing — two (qweight, scale) pairs that - # produce the same bf16 must produce the same checksum. - checksums = { - name: _hash_tensor(tensor.data) - for name, should_compare, tensor in _postprocess_tensors( - dict(self._model_state()), skip_compare_names - ) - if should_compare - } + # Hash the dequantized weight so two (qweight, scale) pairs with the same + # bf16 hash equal. + checksums = {} + for name, should_compare, comparable in _build_check_entries( + dict(self._model_state()), skip_compare_names, quantized_set + ): + if should_compare: + checksums[name] = _hash_tensor(comparable.dequantize().data) h = hashlib.sha256() for name in sorted(checksums): @@ -162,42 +179,50 @@ def _hash_tensor(t: torch.Tensor) -> str: def _check_tensors( - expect_tensors: Iterable[Tuple[str, bool, torch.Tensor]], - actual_tensors: Iterable[Tuple[str, bool, torch.Tensor]], + expect_tensors: Iterable[CheckEntry], + actual_tensors: Iterable[CheckEntry], + allow_quant_error: bool = False, ): - from sglang.srt.debug_utils.dumper import get_tensor_info - good_names = [] error_messages = [] info_messages = [] - for (expect_name, expect_should_compare, expect), ( + for (expect_name, should_compare, expect_comparable), ( actual_name, actual_should_compare, - actual, + actual_comparable, ) in zip(expect_tensors, actual_tensors, strict=True): assert expect_name == actual_name, f"{expect_name=} {actual_name=}" assert ( - expect_should_compare == actual_should_compare - ), f"{expect_should_compare=} {actual_should_compare=}" + should_compare == actual_should_compare + ), f"{should_compare=} {actual_should_compare=}" name = expect_name - should_compare = expect_should_compare - expect = expect.cuda() - actual = actual.cuda() - - if torch.all(expect == actual): - good_names.append(name) - else: - abs_diff = (actual.float() - expect.float()).abs() - msg = ( - f"name={name} " - f"max_abs_err={abs_diff.max()} " - f"mean_abs_err={abs_diff.mean()} " - f"{get_tensor_info(expect)=} " - f"{get_tensor_info(actual)=} " + try: + equal, max_abs_err, mean_abs_err, num_exceed = compare_weights( + expect_comparable, actual_comparable ) - (error_messages if should_compare else info_messages).append(msg) + except Exception as e: + e.add_note( + f"when handling {name=} expect={expect_comparable!r} actual={actual_comparable!r}" + ) + raise + if equal: + good_names.append(name) + continue + msg = ( + f"name={name} " + f"max_abs_err={max_abs_err} " + f"mean_abs_err={mean_abs_err} " + f"num_exceed={num_exceed} " + f"expect={expect_comparable!r} actual={actual_comparable!r} " + ) + if not should_compare: + info_messages.append(msg) + elif allow_quant_error and num_exceed == 0: + info_messages.append(msg + "(within quantization ULP tolerance)") + else: + error_messages.append(msg) logger.info(f"[check_tensors] equal tensors: {good_names}") if len(info_messages) > 0: @@ -212,7 +237,12 @@ def _random_like(t: torch.Tensor): dtype = t.dtype if dtype.is_floating_point: - return torch.rand(shape, device=device, dtype=torch.float32).to(dtype) + out = torch.empty(shape, device=device, dtype=dtype) + for chunk in out.view(-1).split(CHUNK_NUMEL): + chunk.copy_( + torch.rand(chunk.shape, device=device, dtype=torch.float32).to(dtype) + ) + return out if dtype == torch.bool: return torch.rand(shape, device=device) > 0.5 @@ -223,58 +253,44 @@ def _random_like(t: torch.Tensor): ) -def _postprocess_tensors( +def _build_quantized_set(model) -> Dict[str, QuantizedWeight]: + """Run the router over the model: {weight_name: QuantizedWeight} for each + quantized weight; weights absent from the set compare raw.""" + quantized_set = {} + for module_name, module in model.named_modules(): + comparable_cls = select_comparable_weight(getattr(module, "quant_method", None)) + if comparable_cls is None: + continue + prefix = f"{module_name}." if module_name else "" + own = {name for name, _ in module.named_parameters(recurse=False)} + for name in own: + scale = name.replace("weight", "weight_scale_inv") + if name.endswith("weight") and scale in own: + quantized_set[prefix + name] = QuantizedWeight( + comparable_cls, prefix + scale + ) + return quantized_set + + +def _build_check_entries( raw: Dict[str, torch.Tensor], skip_compare_names: Set[str], -) -> Iterable[Tuple[str, bool, torch.Tensor]]: - from sglang.srt.debug_utils.dumper import get_tensor_info - + quantized_set: Optional[Dict[str, QuantizedWeight]] = None, +) -> Iterable[CheckEntry]: + """Yields a CheckEntry per weight; quantized weights consume their scale, everything + else is raw.""" skip_compare_names = set(skip_compare_names) + quantized_set = quantized_set or {} + scale_names = {qw.scale_name for qw in quantized_set.values()} - # Skip non-persistent buffers (registered with persistent=False; recomputed - # after weight load and not part of the synced payload). - for name in raw: - if _is_non_persistent_buffer_name(name): - skip_compare_names.add(name) - logger.info(f"[check_tensors] Skipping non-persistent buffer: {name}") - - # dequant fp8 - quant_names = [ - name - for name in raw - # Match: `something.weight`, `something.experts.w2_weight` - if name.endswith("weight") and name.replace("weight", "weight_scale_inv") in raw - ] - quant_scale_names = [ - name.replace("weight", "weight_scale_inv") for name in quant_names - ] - skip_compare_names.update(quant_names) - skip_compare_names.update(quant_scale_names) - for name in quant_names: - w_q = raw[name] - w_s = raw[name.replace("weight", "weight_scale_inv")] - - try: - if w_s.dtype == torch.int32: - # UE8M0 packed format (Blackwell DeepGEMM) - w_s_for_dequant = inverse_transform_scale_ue8m0(w_s, mn=w_q.shape[-2]) - else: - w_s_for_dequant = w_s - - w_dequant = block_quant_dequant( - w_q, - w_s_for_dequant, - # TODO do not hardcode - block_size=[128, 128], - dtype=torch.bfloat16, + for name, tensor in raw.items(): + if name in scale_names: + continue # compared via its weight's comparable + if name in quantized_set: + qw = quantized_set[name] + yield CheckEntry(name, True, qw.comparable_cls(tensor, raw[qw.scale_name])) + else: + should_compare = name not in skip_compare_names and ( + not _is_non_persistent_buffer_name(name) ) - yield name, True, w_dequant - except Exception as e: - e.add_note( - f"when handling {name=} {get_tensor_info(w_q)=} {get_tensor_info(w_s)=}" - ) - raise - - for name in raw: - should_compare = name not in skip_compare_names - yield name, should_compare, raw[name] + yield CheckEntry(name, should_compare, RawComparable(tensor)) diff --git a/python/sglang/srt/utils/weight_checker_comparator.py b/python/sglang/srt/utils/weight_checker_comparator.py new file mode 100644 index 000000000..48f28d606 --- /dev/null +++ b/python/sglang/srt/utils/weight_checker_comparator.py @@ -0,0 +1,166 @@ +from typing import Iterable, NamedTuple, Optional, Tuple + +import torch + +from sglang.srt.layers.quantization.fp8 import Fp8LinearMethod, Fp8MoEMethod +from sglang.srt.layers.quantization.fp8_utils import ( + block_quant_dequant, + inverse_transform_scale_ue8m0, +) +from sglang.srt.layers.quantization.modelopt_quant import ( + ModelOptFp4LinearMethod, + ModelOptNvFp4FusedMoEMethod, +) + +# chunk to avoid too high GPU memory peak +CHUNK_NUMEL = 64 * 1024 * 1024 + + +class CompareResult(NamedTuple): + equal: bool + max_abs_err: float + mean_abs_err: float + num_exceed: int # elements past the combined per-side tolerance + + +class ComparableWeight: + """Base comparable-weight class; one subclass per precision or raw tensor.""" + + @staticmethod + def _quant_ulp(w_q: torch.Tensor) -> torch.Tensor: + """Per-element ULP of w_q in its own dtype.""" + finfo = torch.finfo(w_q.dtype) + x = w_q.to(torch.float32).abs() + # frexp: x = m * 2^e, m in [0.5, 1), so 2^(e-1) is x's binade base. + _, exponent = torch.frexp(x) + binade = torch.exp2((exponent - 1).to(torch.float32)) + # Zeros and subnormals share the spacing of the smallest normal binade. + binade = binade.masked_fill(x < finfo.smallest_normal, finfo.smallest_normal) + return binade * finfo.eps + + def iter_chunks(self) -> Iterable[Tuple[torch.Tensor, Optional[torch.Tensor]]]: + raise NotImplementedError + + def dequantize(self, dtype: torch.dtype = torch.bfloat16) -> torch.Tensor: + raise NotImplementedError + + +class Fp8BlockComparable(ComparableWeight): + """Deepseek-style FP8 quantization.""" + + def __init__(self, w_q: torch.Tensor, w_s: torch.Tensor): + self.w_q = w_q + self.w_s = w_s + + def __repr__(self) -> str: + return f"fp8_block(shape={tuple(self.w_q.shape)} dtype={self.w_q.dtype})" + + @staticmethod + def _normalize_scale(w_q: torch.Tensor, w_s: torch.Tensor) -> torch.Tensor: + if w_s.dtype == torch.int32: + w_s = inverse_transform_scale_ue8m0(w_s, mn=w_q.shape[-2]) + return w_s.to(torch.float32) + + @staticmethod + def _infer_block_size(w_q: torch.Tensor, w_s: torch.Tensor) -> list: + k, s_k = w_q.shape[-1], w_s.shape[-1] + assert k % s_k == 0, f"cannot infer block size from {w_q.shape=} {w_s.shape=}" + block = k // s_k + return [block, block] + + @staticmethod + def _iter_quant_chunks(w_q: torch.Tensor, w_s: torch.Tensor, block_n: int): + """Yields block-row-aligned (q_slice, s_slice) pairs of bounded size.""" + q3 = w_q.reshape(-1, *w_q.shape[-2:]) + s3 = w_s.reshape(-1, *w_s.shape[-2:]) + n, k = q3.shape[-2:] + rows = max(block_n, CHUNK_NUMEL // k // block_n * block_n) + for b in range(q3.shape[0]): + for r0 in range(0, n, rows): + r1 = min(r0 + rows, n) + yield q3[b, r0:r1], s3[b, r0 // block_n : -(-r1 // block_n)] + + def _scale_and_block_size(self): + s = self._normalize_scale(self.w_q, self.w_s) + return s, self._infer_block_size(self.w_q, s) + + def iter_chunks(self): + s, block_size = self._scale_and_block_size() + for q, s_chunk in self._iter_quant_chunks(self.w_q, s, block_size[0]): + q, s_chunk = q.cuda(), s_chunk.cuda() + yield ( + block_quant_dequant(q, s_chunk, block_size, dtype=torch.bfloat16), + block_quant_dequant( + self._quant_ulp(q), s_chunk, block_size, dtype=torch.float32 + ), + ) + + def dequantize(self, dtype: torch.dtype = torch.bfloat16) -> torch.Tensor: + s, block_size = self._scale_and_block_size() + return block_quant_dequant(self.w_q, s, block_size, dtype=dtype) + + +class RawComparable(ComparableWeight): + """Bitwise equal compare on raw tensor.""" + + def __init__(self, tensor: torch.Tensor): + self.tensor = tensor + + def __repr__(self) -> str: + return f"raw(shape={tuple(self.tensor.shape)} dtype={self.tensor.dtype})" + + def iter_chunks(self): + flat = self.tensor.reshape(-1) + for start in range(0, flat.numel(), CHUNK_NUMEL): + yield flat[start : start + CHUNK_NUMEL].cuda(), None + + def dequantize(self, dtype: torch.dtype = torch.bfloat16) -> torch.Tensor: + return self.tensor + + +def compare_weights( + expect: ComparableWeight, actual: ComparableWeight +) -> CompareResult: + """Chunked element-wise compare in ComparableWeight space.""" + equal = True + max_abs_err = torch.zeros((), dtype=torch.float32) + sum_abs_err = 0.0 + num_exceed = 0 + numel = 0 + for (expect_dq, expect_tol), (actual_dq, actual_tol) in zip( + expect.iter_chunks(), actual.iter_chunks(), strict=True + ): + assert ( + expect_dq.shape == actual_dq.shape + ), f"{expect_dq.shape=} {actual_dq.shape=}" + numel += expect_dq.numel() + abs_diff = (actual_dq.float() - expect_dq.float()).abs() + if torch.all(abs_diff == 0): + continue + equal = False + # |actual_dq - expect_dq| ≤ |actual_dq - w| + |expect_dq - w| ≤ actual_tol + expect_tol + tol = ( + 0.0 if expect_tol is None or actual_tol is None else expect_tol + actual_tol + ) + max_abs_err = torch.maximum(max_abs_err, abs_diff.max().cpu()) + sum_abs_err += abs_diff.sum().item() + # `~(diff <= tol)` instead of `diff > tol` so NaN counts as exceeding. + num_exceed += int((~(abs_diff <= tol)).sum()) + return CompareResult( + equal, max_abs_err.item(), sum_abs_err / max(numel, 1), num_exceed + ) + + +def select_comparable_weight(quant_method) -> Optional[type]: + """Map a module's quant_method to its ComparableWeight. None means raw (bitwise equal) compare.""" + if ( + isinstance(quant_method, (Fp8LinearMethod, Fp8MoEMethod)) + and quant_method.block_quant + and not quant_method.use_mxfp8 + ): + return Fp8BlockComparable + if isinstance(quant_method, (ModelOptFp4LinearMethod, ModelOptNvFp4FusedMoEMethod)): + raise NotImplementedError( + f"weight checker has no ComparableWeight for {type(quant_method).__name__}" + ) + return None diff --git a/test/registered/unit/utils/test_weight_checker.py b/test/registered/unit/utils/test_weight_checker.py index e54da8ab1..8ffae6b11 100644 --- a/test/registered/unit/utils/test_weight_checker.py +++ b/test/registered/unit/utils/test_weight_checker.py @@ -14,27 +14,33 @@ """Unit tests for sglang/srt/utils/weight_checker.py.""" import unittest -from typing import Iterable, List, Tuple +from typing import Iterable, List 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 ( + CheckEntry, ChecksumInfo, ParallelismInfo, + QuantizedWeight, WeightChecker, + _build_check_entries, + _build_quantized_set, _check_tensors, _hash_tensor, _is_non_persistent_buffer_name, - _postprocess_tensors, _random_like, ) +from sglang.srt.utils.weight_checker_comparator import ( + Fp8BlockComparable, + RawComparable, +) from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import CustomTestCase @@ -46,31 +52,39 @@ register_cuda_ci(est_time=30, stage="base-b", runner_config="1-gpu-small") # --------------------------------------------------------------------------- -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) +def _assert_entries_close( + actual: Iterable[CheckEntry], expected: Iterable[CheckEntry] +) -> None: + """Compare two streams of (name, should_compare, ComparableWeight).""" + actual_list: List[CheckEntry] = list(actual) + expected_list: List[CheckEntry] = 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( + for i, ((a_name, a_flag, a_ref), (e_name, e_flag, e_ref)) 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}" - ) + assert type(a_ref) is type(e_ref), f"[{i}] kind mismatch for {a_name!r}" + if isinstance(a_ref, Fp8BlockComparable): + torch.testing.assert_close( + a_ref.w_q, e_ref.w_q, msg=f"[{i}] w_q {a_name!r}" + ) + torch.testing.assert_close( + a_ref.w_s, e_ref.w_s, msg=f"[{i}] w_s {a_name!r}" + ) + else: + torch.testing.assert_close( + a_ref.tensor, e_ref.tensor, msg=f"[{i}] tensor {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. + drives the _build_check_entries branch under test. """ weight_bf16 = torch.randn((256, 128), dtype=torch.bfloat16, device=device) block_size = [128, 128] @@ -87,7 +101,7 @@ def _build_fp8_quant_pair(device: str = "cuda"): class _TinyModel(nn.Module): - """Mimics the buffer naming patterns _reset_tensors / _postprocess_tensors care about.""" + """Mimics the buffer naming patterns _reset_tensors / _build_check_entries care about.""" def __init__(self): super().__init__() @@ -173,9 +187,19 @@ class TestRandomLike(CustomTestCase): _random_like(t) torch.testing.assert_close(t, before) + def test_floating_point_chunked_generation(self): + with patch("sglang.srt.utils.weight_checker.CHUNK_NUMEL", 8): + out = _random_like(torch.zeros(64, dtype=torch.bfloat16)) + self.assertEqual(out.dtype, torch.bfloat16) + self.assertEqual(out.shape, (64,)) + self.assertGreater(out.unique().numel(), 8) + self.assertGreaterEqual(out.float().min().item(), 0.0) + # bf16 rounding may carry values just below 1.0 up to exactly 1.0 + self.assertLessEqual(out.float().max().item(), 1.0) + # --------------------------------------------------------------------------- -# _postprocess_tensors +# _build_check_entries # --------------------------------------------------------------------------- @@ -187,15 +211,17 @@ class TestPostprocessTensors(CustomTestCase): a = torch.randn(4) b = torch.randn(4) raw = {"a.weight": a, "b.bias": b} - _assert_triples_close( - _postprocess_tensors(raw, set()), - [("a.weight", True, a), ("b.bias", True, b)], + _assert_entries_close( + _build_check_entries(raw, set()), + [("a.weight", True, RawComparable(a)), ("b.bias", True, RawComparable(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, set()), [("x.weight", True, w)]) + _assert_entries_close( + _build_check_entries(raw, set()), [("x.weight", True, RawComparable(w))] + ) # --- non-persistent buffer skip --- @@ -206,70 +232,49 @@ class TestPostprocessTensors(CustomTestCase): "model.rotary_emb.cos_sin_cache": cache, "model.layers.0.weight": plain, } - _assert_triples_close( - _postprocess_tensors(raw, set()), + _assert_entries_close( + _build_check_entries(raw, set()), [ - ("model.rotary_emb.cos_sin_cache", False, cache), - ("model.layers.0.weight", True, plain), + ("model.rotary_emb.cos_sin_cache", False, RawComparable(cache)), + ("model.layers.0.weight", True, RawComparable(plain)), ], ) def test_skips_inv_freq_substring(self): t = torch.randn(4) - _assert_triples_close( - _postprocess_tensors({"model.rotary_emb.inv_freq": t}, set()), - [("model.rotary_emb.inv_freq", False, t)], + _assert_entries_close( + _build_check_entries({"model.rotary_emb.inv_freq": t}, set()), + [("model.rotary_emb.inv_freq", False, RawComparable(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}, set()), - [("model.layers.0.mlp.gate._weight_fp32", False, t)], + _assert_entries_close( + _build_check_entries({"model.layers.0.mlp.gate._weight_fp32": t}, set()), + [("model.layers.0.mlp.gate._weight_fp32", False, RawComparable(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}, set()), - [("weird.cos_sin_cache.foo.bar", False, t)], + _assert_entries_close( + _build_check_entries({"weird.cos_sin_cache.foo.bar": t}, set()), + [("weird.cos_sin_cache.foo.bar", False, RawComparable(t))], ) # --- fp8 quant pair (real dequant on real fp8 tensors) --- - def test_fp8_quant_pair_with_int32_scale_dequants_via_ue8m0(self): + def test_fp8_quant_pair_yields_lazy_pair(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, set()), - [ - ("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, set()), - [ - ("x.weight", True, expected_dequant), - ("x.weight", False, qweight), - ("x.weight_scale_inv", False, sf_fp32), - ], + ref = Fp8BlockComparable(qweight, sf_packed_int32) + quantized_set = { + "x.weight": QuantizedWeight(Fp8BlockComparable, "x.weight_scale_inv") + } + _assert_entries_close( + _build_check_entries(raw, set(), quantized_set), + [("x.weight", True, ref)], ) def test_fp8_quant_pair_yield_order_alongside_other_entries(self): @@ -280,17 +285,16 @@ class TestPostprocessTensors(CustomTestCase): "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, set()), + # scale_inv is consumed by its weight's comparable; y.bias stays raw. + ref = Fp8BlockComparable(qweight, sf_fp32) + quantized_set = { + "x.weight": QuantizedWeight(Fp8BlockComparable, "x.weight_scale_inv") + } + _assert_entries_close( + _build_check_entries(raw, set(), quantized_set), [ - ("x.weight", True, expected_dequant), - ("x.weight", False, qweight), - ("x.weight_scale_inv", False, sf_fp32), - ("y.bias", True, bias), + ("x.weight", True, ref), + ("y.bias", True, RawComparable(bias)), ], ) @@ -298,9 +302,9 @@ class TestPostprocessTensors(CustomTestCase): # 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}, set()), - [("x.weight_scale_inv", True, s)], + _assert_entries_close( + _build_check_entries({"x.weight_scale_inv": s}, set()), + [("x.weight_scale_inv", True, RawComparable(s))], ) @@ -313,13 +317,19 @@ 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())] + expect = [ + ("a", True, RawComparable(t.clone())), + ("b", True, RawComparable(t.clone())), + ] + actual = [ + ("a", True, RawComparable(t.clone())), + ("b", True, RawComparable(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))] + expect = [("a", True, RawComparable(torch.ones(2, 2)))] + actual = [("a", True, RawComparable(torch.zeros(2, 2)))] with self.assertRaises(Exception) as ctx: _check_tensors(expect_tensors=expect, actual_tensors=actual) msg = str(ctx.exception) @@ -328,30 +338,141 @@ class TestCheckTensors(CustomTestCase): 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))] + expect = [("a", False, RawComparable(torch.ones(2, 2)))] + actual = [("a", False, RawComparable(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))] + expect = [("a", True, RawComparable(torch.ones(2, 2)))] + actual = [("b", True, RawComparable(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))] + expect = [("a", True, RawComparable(torch.ones(2, 2)))] + actual = [("a", False, RawComparable(torch.ones(2, 2)))] with self.assertRaises(AssertionError): _check_tensors(expect_tensors=expect, actual_tensors=actual) + def test_chunked_raw_stats_match_unchunked(self): + expect = [("a", True, RawComparable(torch.zeros(10)))] + actual = [("a", True, RawComparable(torch.arange(10.0)))] + with patch("sglang.srt.utils.weight_checker_comparator.CHUNK_NUMEL", 3): + with self.assertRaises(Exception) as ctx: + _check_tensors(expect_tensors=expect, actual_tensors=actual) + self.assertIn("max_abs_err=9.0", str(ctx.exception)) + self.assertIn("mean_abs_err=4.5", str(ctx.exception)) + 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())] + expect = [ + ("a", True, RawComparable(t.clone())), + ("b", True, RawComparable(t.clone())), + ] + actual = [("a", True, RawComparable(t.clone()))] with self.assertRaises(ValueError): _check_tensors(expect_tensors=expect, actual_tensors=actual) +# --------------------------------------------------------------------------- +# _check_tensors + allow_quant_error +# --------------------------------------------------------------------------- + + +def _quantize_block_fp8(weight: torch.Tensor, scale_margin: float): + """Blockwise 128x128 fp8 quantization with a tweakable scale convention.""" + n, k = weight.shape + blocks = weight.float().view(n // 128, 128, k // 128, 128).permute(0, 2, 1, 3) + scale = blocks.abs().amax(dim=(-1, -2)) / 448.0 * scale_margin + q = (blocks / scale[:, :, None, None]).to(torch.float8_e4m3fn) + q = q.permute(0, 2, 1, 3).reshape(n, k) + return q, scale + + +class TestCheckTensorsAllowQuantError(CustomTestCase): + + def setUp(self): + torch.manual_seed(0) + weight = torch.randn(256, 256, device="cuda") * 0.02 + self.e_raw = self._as_raw(*_quantize_block_fp8(weight, 1.0)) + self.a_raw = self._as_raw(*_quantize_block_fp8(weight, 1.001)) + + @staticmethod + def _as_raw(q, s): + return {"x.weight": q, "x.weight_scale_inv": s} + + def _check(self, expect_raw, actual_raw, **kwargs): + quantized_set = { + "x.weight": QuantizedWeight(Fp8BlockComparable, "x.weight_scale_inv") + } + _check_tensors( + expect_tensors=_build_check_entries(expect_raw, set(), quantized_set), + actual_tensors=_build_check_entries(actual_raw, set(), quantized_set), + **kwargs, + ) + + def test_within_tolerance_passes_with_flag(self): + self._check(self.e_raw, self.a_raw, allow_quant_error=True) + + def test_within_tolerance_fails_without_flag(self): + with self.assertRaises(Exception) as ctx: + self._check(self.e_raw, self.a_raw) + self.assertIn("name=x.weight", str(ctx.exception)) + + def test_exceeding_tolerance_fails_with_flag(self): + bad_q = self.a_raw["x.weight"].clone().view(torch.uint8) + bad_q[::50] += 8 + bad = self._as_raw( + bad_q.view(torch.float8_e4m3fn), self.a_raw["x.weight_scale_inv"] + ) + with self.assertRaises(Exception) as ctx: + self._check(self.e_raw, bad, allow_quant_error=True) + self.assertIn("num_exceed", str(ctx.exception)) + + def test_flag_does_not_relax_non_quant_tensors(self): + expect = [("a", True, RawComparable(torch.ones(2, 2)))] + actual = [("a", True, RawComparable(torch.ones(2, 2) + 0.5))] + with self.assertRaises(Exception): + _check_tensors( + expect_tensors=expect, actual_tensors=actual, allow_quant_error=True + ) + + +# --------------------------------------------------------------------------- +# _build_quantized_set +# --------------------------------------------------------------------------- + + +class TestBuildQuantizedSet(CustomTestCase): + + def test_fp8_block_module_pairs_weight_and_scale(self): + from sglang.srt.layers.quantization.fp8 import Fp8LinearMethod + + method = Fp8LinearMethod.__new__(Fp8LinearMethod) + method.block_quant = True + method.use_mxfp8 = False + model = nn.Module() + model.proj = nn.Module() + model.proj.quant_method = method + model.proj.register_parameter( + "weight", nn.Parameter(torch.zeros(4, 4), requires_grad=False) + ) + model.proj.register_parameter( + "weight_scale_inv", nn.Parameter(torch.zeros(1, 1), requires_grad=False) + ) + self.assertEqual( + _build_quantized_set(model), + { + "proj.weight": QuantizedWeight( + Fp8BlockComparable, "proj.weight_scale_inv" + ) + }, + ) + + def test_no_quant_method_yields_empty_plan(self): + self.assertEqual(_build_quantized_set(_TinyModel()), {}) + + # --------------------------------------------------------------------------- # WeightChecker class # --------------------------------------------------------------------------- diff --git a/test/registered/unit/utils/test_weight_checker_comparator.py b/test/registered/unit/utils/test_weight_checker_comparator.py new file mode 100644 index 000000000..c5b677ff1 --- /dev/null +++ b/test/registered/unit/utils/test_weight_checker_comparator.py @@ -0,0 +1,206 @@ +# 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_comparator.py.""" + +import unittest +from unittest.mock import patch + +import torch + +from sglang.srt.layers.quantization.fp8_utils import ( + quant_weight_ue8m0, + transform_scale_ue8m0, +) +from sglang.srt.utils.weight_checker_comparator import ( + ComparableWeight, + Fp8BlockComparable, + compare_weights, + select_comparable_weight, +) +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=15, stage="base-b", runner_config="1-gpu-small") + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _compare_quant_pair(expect_q, expect_s, actual_q, actual_s): + return compare_weights( + Fp8BlockComparable(expect_q, expect_s), Fp8BlockComparable(actual_q, actual_s) + ) + + +def _build_fp8_quant_pair(device: str = "cuda"): + """Returns (qweight, fp32 scale, ue8m0-packed int32 scale) for one random weight.""" + 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 + + +# --------------------------------------------------------------------------- +# _quant_ulp +# --------------------------------------------------------------------------- + + +class TestQuantUlp(CustomTestCase): + + def test_matches_bruteforce_spacing_for_fp8(self): + for dtype in (torch.float8_e4m3fn, torch.float8_e5m2): + all_bits = torch.arange(256, dtype=torch.uint8).view(dtype) + vals = all_bits.to(torch.float32) + magnitudes = torch.unique(vals[torch.isfinite(vals) & (vals >= 0)]) + # Brute-force ULP: spacing to the next representable magnitude + # (the largest magnitude reuses the spacing below it). + spacing = magnitudes[1:] - magnitudes[:-1] + expected = torch.cat([spacing, spacing[-1:]]) + got = ComparableWeight._quant_ulp(magnitudes.to(dtype)) + torch.testing.assert_close(got, expected, rtol=0, atol=0) + + +# --------------------------------------------------------------------------- +# compare_weights +# --------------------------------------------------------------------------- + + +class TestCompareQuantPair(CustomTestCase): + """Chunked dequantized-space comparison of block-quantized pairs.""" + + @staticmethod + def _quantize(weight: torch.Tensor, scale_margin: float): + """Blockwise 128x128 fp8 quantization with a tweakable scale convention.""" + n, k = weight.shape + blocks = weight.float().view(n // 128, 128, k // 128, 128).permute(0, 2, 1, 3) + scale = blocks.abs().amax(dim=(-1, -2)) / 448.0 * scale_margin + q = (blocks / scale[:, :, None, None]).to(torch.float8_e4m3fn) + q = q.permute(0, 2, 1, 3).reshape(n, k) + return q, scale + + def setUp(self): + torch.manual_seed(0) + self.weight = torch.randn(256, 256, device="cuda") * 0.02 + self.e_q, self.e_s = self._quantize(self.weight, 1.0) + self.a_q, self.a_s = self._quantize(self.weight, 1.001) + + def test_identical_pair_is_equal(self): + equal, max_err, mean_err, num_exceed = _compare_quant_pair( + self.e_q, self.e_s, self.e_q.clone(), self.e_s.clone() + ) + self.assertTrue(equal) + self.assertEqual((max_err, mean_err, num_exceed), (0.0, 0.0, 0)) + + def test_ue8m0_packed_scale_equals_unpacked_scale(self): + qweight, sf_fp32, sf_packed_int32 = _build_fp8_quant_pair() + equal, *_ = _compare_quant_pair(qweight, sf_packed_int32, qweight, sf_fp32) + self.assertTrue(equal) + + def test_two_quantizations_stay_within_ulp_tolerance(self): + equal, max_err, mean_err, num_exceed = _compare_quant_pair( + self.e_q, self.e_s, self.a_q, self.a_s + ) + self.assertFalse(equal) + self.assertGreater(max_err, 0.0) + self.assertEqual(num_exceed, 0) + + def test_corruption_and_fp8_nan_exceed_tolerance(self): + bad_q = self.a_q.clone().view(torch.uint8) + bad_q[::50] += 8 # jumps a full binade; some bytes become fp8 NaN + equal, max_err, mean_err, num_exceed = _compare_quant_pair( + self.e_q, self.e_s, bad_q.view(torch.float8_e4m3fn), self.a_s + ) + self.assertFalse(equal) + self.assertGreater(num_exceed, 0) + + def test_chunked_result_matches_unchunked(self): + reference = _compare_quant_pair(self.e_q, self.e_s, self.a_q, self.a_s) + with patch("sglang.srt.utils.weight_checker_comparator.CHUNK_NUMEL", 128 * 128): + chunked = _compare_quant_pair(self.e_q, self.e_s, self.a_q, self.a_s) + self.assertEqual(chunked, reference) + + @staticmethod + def _quantize_partial(weight: torch.Tensor, scale_margin: float): + """128x128 block quant where the last block per dim may be partial.""" + n, k = weight.shape + s_n, s_k = -(-n // 128), -(-k // 128) + q = torch.empty(n, k, dtype=torch.float8_e4m3fn, device=weight.device) + scale = torch.empty(s_n, s_k, device=weight.device) + for i in range(s_n): + for j in range(s_k): + blk = weight[i * 128 : (i + 1) * 128, j * 128 : (j + 1) * 128].float() + s = blk.abs().amax() / 448.0 * scale_margin + s = s if s > 0 else weight.new_ones(()) + scale[i, j] = s + q[i * 128 : (i + 1) * 128, j * 128 : (j + 1) * 128] = (blk / s).to( + torch.float8_e4m3fn + ) + return q, scale + + def test_partial_last_block_infers_true_block_size(self): + # fused_qkv_a_proj_with_mqa out-dim is not a multiple of 128 (e.g. 2112 = + # 16*128 + 64), so the last row-block is partial. ceil(dim/num_blocks) + # would infer 125, misaligning scales; the true block size is 128. + n, k = 3 * 128 + 64, 256 + weight = torch.randn(n, k, device="cuda") * 0.02 + e_q, e_s = self._quantize_partial(weight, 1.0) + a_q, a_s = self._quantize_partial(weight, 1.001) + self.assertEqual(list(e_s.shape), [4, 2]) # ceil(448/128)=4, 256/128=2 + self.assertEqual(Fp8BlockComparable._infer_block_size(e_q, e_s), [128, 128]) + equal, _, _, num_exceed = _compare_quant_pair(e_q, e_s, a_q, a_s) + self.assertFalse(equal) + self.assertEqual(num_exceed, 0) + + def test_3d_expert_tensor(self): + q3 = self.e_q.reshape(2, 128, 256).contiguous() + s3 = self.e_s.reshape(2, 1, 2) + equal, *_ = _compare_quant_pair(q3, s3, q3.clone(), s3.clone()) + self.assertTrue(equal) + + +# --------------------------------------------------------------------------- +# select_comparable_weight +# --------------------------------------------------------------------------- + + +class TestSelectComparableWeight(CustomTestCase): + + def test_returns_none_when_not_a_quant_method(self): + self.assertIsNone(select_comparable_weight(None)) + + def test_returns_none_for_raw_safe_method(self): + from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod + + # unquantized / int4 / mxfp8 all route to raw (None). + fake = UnquantizedLinearMethod.__new__(UnquantizedLinearMethod) + self.assertIsNone(select_comparable_weight(fake)) + + def test_raises_on_nvfp4(self): + from sglang.srt.layers.quantization.modelopt_quant import ( + ModelOptFp4LinearMethod, + ) + + # nvfp4 has no ComparableWeight yet -> must raise, not silently raw-compare. + fake = ModelOptFp4LinearMethod.__new__(ModelOptFp4LinearMethod) + with self.assertRaises(NotImplementedError): + select_comparable_weight(fake) + + +if __name__ == "__main__": + unittest.main()