From 326b788ab4bb0a422989005cdbd9f9cfaf0079e3 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Sun, 22 Feb 2026 16:20:08 +0800 Subject: [PATCH] Fix wrongly large dumped file and handle non intrusive hook reset in dumper (#19124) --- python/sglang/srt/debug_utils/dumper.py | 85 ++++++-- .../debug_utils/test_dump_comparator.py | 2 +- test/registered/debug_utils/test_dumper.py | 199 +++++++++++++++++- 3 files changed, 255 insertions(+), 31 deletions(-) diff --git a/python/sglang/srt/debug_utils/dumper.py b/python/sglang/srt/debug_utils/dumper.py index de8c19b75..7366bf64c 100644 --- a/python/sglang/srt/debug_utils/dumper.py +++ b/python/sglang/srt/debug_utils/dumper.py @@ -7,6 +7,7 @@ import socket import threading import time from abc import ABC, abstractmethod +from collections.abc import Callable from contextlib import contextmanager from copy import deepcopy from dataclasses import asdict, dataclass, field, fields, replace @@ -130,8 +131,8 @@ class DumperConfig(_BaseConfig): enable_output_console: bool = True enable_value: bool = True enable_grad: bool = False - enable_model_value: bool = True - enable_model_grad: bool = True + enable_model_value: bool = False + enable_model_grad: bool = False exp_name: Optional[str] = None enable_http_server: bool = True cleanup_previous: bool = False @@ -200,6 +201,7 @@ class _Dumper: def __init__(self, *, config: DumperConfig): self._config = config self._state = _DumperState() + self._non_intrusives: list["_NonIntrusiveDumper"] = [] # ------------------------------- public :: core --------------------------------- @@ -280,7 +282,9 @@ class _Dumper: mode = self._config.non_intrusive_mode if mode == "off": return None - return _NonIntrusiveDumper(dumper=self, model=model, mode=mode) + non_intrusive = _NonIntrusiveDumper(dumper=self, model=model, mode=mode) + self._non_intrusives.append(non_intrusive) + return non_intrusive # ------------------------------- public :: secondary --------------------------------- @@ -291,6 +295,9 @@ class _Dumper: self._config = self._config.with_defaults(**kwargs) def reset(self) -> None: + for non_intrusive in self._non_intrusives: + non_intrusive.remove() + self._non_intrusives.clear() self._state = _DumperState() @contextmanager @@ -488,6 +495,7 @@ class _NonIntrusiveDumper: ): self._dumper = dumper self._mode = mode + self._handles: list = [] for module_name, module in model.named_modules(): if ctx := self._detect_module_ctx(module_name, module): @@ -498,13 +506,18 @@ class _NonIntrusiveDumper: module_name=module_name, is_root=is_root ) hook = self._make_forward_hook(module_name=module_name, is_root=is_root) - _register_forward_hook_or_replace_fn( + self._handles += _register_forward_hook_or_replace_fn( module, pre_hook=pre_hook, hook=hook, mode="replace_fn" if is_root else "hook", ) + def remove(self) -> None: + for handle in self._handles: + handle.remove() + self._handles.clear() + @classmethod def _detect_module_ctx( cls, module_name: str, module: "torch.nn.Module" @@ -518,12 +531,16 @@ class _NonIntrusiveDumper: def _register_ctx_hooks(self, module: "torch.nn.Module", *, ctx: dict) -> None: clear_ctx = {k: None for k in ctx} - module.register_forward_pre_hook( - lambda _mod, _input, _ctx=ctx: self._dumper.set_ctx(**_ctx) + self._handles.append( + module.register_forward_pre_hook( + lambda _mod, _input, _ctx=ctx: self._dumper.set_ctx(**_ctx) + ) ) - module.register_forward_hook( - lambda _mod, _input, _output, _clear=clear_ctx: self._dumper.set_ctx( - **_clear + self._handles.append( + module.register_forward_hook( + lambda _mod, _input, _output, _clear=clear_ctx: self._dumper.set_ctx( + **_clear + ) ) ) @@ -578,7 +595,7 @@ def _register_forward_hook_or_replace_fn( pre_hook, hook, mode: str, -) -> None: +) -> list: """Attach pre/post forward hooks to *module*. mode="hook" — standard ``register_forward_pre_hook`` / ``register_forward_hook`` @@ -586,10 +603,15 @@ def _register_forward_hook_or_replace_fn( mode="replace_fn" — monkey-patch ``module.forward`` so hooks fire even when callers invoke ``.forward()`` directly (as sglang does for the root model). + + Returns a list of handle objects with a ``.remove()`` method that undoes + the registration. """ if mode == "hook": - module.register_forward_pre_hook(pre_hook) - module.register_forward_hook(hook) + return [ + module.register_forward_pre_hook(pre_hook), + module.register_forward_hook(hook), + ] elif mode == "replace_fn": original_forward = module.forward @@ -601,6 +623,13 @@ def _register_forward_hook_or_replace_fn( return output module.forward = _wrapped + + class _Handle: + def remove(self) -> None: + assert module.forward is _wrapped + module.forward = original_forward + + return [_Handle()] else: raise ValueError(f"Unknown mode {mode!r}") @@ -609,6 +638,7 @@ def _register_forward_hook_or_replace_fn( def _torch_save(value, path: str): + value = _clone_if_view(value) try: try: return torch.save(value, path) @@ -623,17 +653,32 @@ def _torch_save(value, path: str): print(f"[Dumper] Observe error={e} when saving data, skip the tensor") -def _strip_parameter(value): - """Strip nn.Parameter to plain Tensor so it can be pickled.""" - if isinstance(value, torch.nn.Parameter): - return value.data - if isinstance(value, dict) and isinstance( - value.get("value"), torch.nn.Parameter - ): - return {**value, "value": value["value"].data} +def _map_tensor(value, fn: Callable[[torch.Tensor], torch.Tensor]): + if isinstance(value, dict): + return {k: _map_tensor(v, fn) for k, v in value.items()} + if isinstance(value, torch.Tensor): + return fn(value) return value +def _clone_if_view(value): + def _fn(t: torch.Tensor) -> torch.Tensor: + if t.untyped_storage().nbytes() > t.nelement() * t.element_size(): + return t.clone() + return t + + return _map_tensor(value, _fn) + + +def _strip_parameter(value): + def _fn(t: torch.Tensor) -> torch.Tensor: + if isinstance(t, torch.nn.Parameter): + return t.data + return t + + return _map_tensor(value, _fn) + + def _collective_with_timeout(fn, operation_name: str, timeout_seconds: int = 60): completed = threading.Event() diff --git a/test/registered/debug_utils/test_dump_comparator.py b/test/registered/debug_utils/test_dump_comparator.py index acb4de147..3cab0eae7 100644 --- a/test/registered/debug_utils/test_dump_comparator.py +++ b/test/registered/debug_utils/test_dump_comparator.py @@ -89,7 +89,7 @@ class TestEndToEnd(CustomTestCase): from argparse import Namespace from sglang.srt.debug_utils.dump_comparator import main - from sglang.srt.debug_utils.dumper import _Dumper, DumperConfig + from sglang.srt.debug_utils.dumper import DumperConfig, _Dumper with tempfile.TemporaryDirectory() as d1, tempfile.TemporaryDirectory() as d2: baseline_tensor = torch.randn(10, 10) diff --git a/test/registered/debug_utils/test_dumper.py b/test/registered/debug_utils/test_dumper.py index 3398fb51d..5688ca17f 100644 --- a/test/registered/debug_utils/test_dumper.py +++ b/test/registered/debug_utils/test_dumper.py @@ -13,12 +13,13 @@ import torch import torch.distributed as dist from sglang.srt.debug_utils.dumper import ( + DumperConfig, _collective_with_timeout, _deepcopy_or_clone, _Dumper, - DumperConfig, _format_tags, _get_default_exp_name, + _map_tensor, _materialize_value, _MegatronPlugin, _obj_to_dict, @@ -287,6 +288,48 @@ class TestDumperPureFunctions: assert "min=None" in get_tensor_info(torch.tensor([])) +class TestMapTensor: + def test_bare_tensor(self): + t = torch.randn(4) + result = _map_tensor(t, lambda x: x * 2) + assert torch.equal(result, t * 2) + + def test_bare_tensor_no_change(self): + t = torch.randn(4) + result = _map_tensor(t, lambda x: x) + assert result is t + + def test_dict_with_tensor_values(self): + t1 = torch.randn(3) + t2 = torch.randn(5) + value = {"a": t1, "b": t2, "meta": "not a tensor"} + result = _map_tensor(value, lambda x: x.clone()) + assert torch.equal(result["a"], t1) + assert torch.equal(result["b"], t2) + assert result["a"] is not t1 + assert result["b"] is not t2 + assert result["meta"] == "not a tensor" + + def test_dict_no_tensors(self): + value = {"a": 1, "b": "hello"} + result = _map_tensor(value, lambda x: x.clone()) + assert result == value + + def test_nested_dict(self): + inner_t = torch.randn(3) + value = {"outer": {"inner": inner_t, "label": "ok"}, "top": torch.randn(2)} + result = _map_tensor(value, lambda x: x.clone()) + assert torch.equal(result["outer"]["inner"], inner_t) + assert result["outer"]["inner"] is not inner_t + assert result["outer"]["label"] == "ok" + assert result is not value + assert result["outer"] is not value["outer"] + + def test_non_tensor_non_dict(self): + result = _map_tensor(42, lambda x: x.clone()) + assert result == 42 + + class TestTorchSave: def test_normal(self, tmp_path): path = str(tmp_path / "a.pt") @@ -308,6 +351,22 @@ class TestTorchSave: assert torch.equal(torch.load(path, weights_only=True), param.data) + def test_shared_storage_not_bloated(self, tmp_path): + big = torch.randn(1000, 1000) + view = big[0] + path = str(tmp_path / "view.pt") + + _torch_save({"value": view, "meta": {}}, path) + + file_size = Path(path).stat().st_size + expected_max = view.nelement() * view.element_size() * 10 + assert file_size < expected_max, ( + f"File {file_size} bytes but view is only " + f"{view.nelement() * view.element_size()} bytes — " + f"torch.save likely serialized the full " + f"{big.nelement() * big.element_size()} byte storage" + ) + def test_silent_skip(self, tmp_path, capsys): path = str(tmp_path / "c.pt") @@ -852,7 +911,9 @@ class TestKvFilter: class TestDumpModel: def test_grad_basic(self, tmp_path): - d = _make_test_dumper(tmp_path, enable_model_value=False) + d = _make_test_dumper( + tmp_path, enable_model_grad=True, enable_model_value=False + ) model = torch.nn.Linear(4, 2) x = torch.randn(3, 4) y = model(x).sum() @@ -866,7 +927,9 @@ class TestDumpModel: ) def test_value_basic(self, tmp_path): - d = _make_test_dumper(tmp_path, enable_model_grad=False) + d = _make_test_dumper( + tmp_path, enable_model_value=True, enable_model_grad=False + ) model = torch.nn.Linear(4, 2, bias=False) d.dump_model(model, name_prefix="model") @@ -877,7 +940,9 @@ class TestDumpModel: ) def test_no_grad_skipped(self, tmp_path): - d = _make_test_dumper(tmp_path, enable_model_value=False) + d = _make_test_dumper( + tmp_path, enable_model_grad=True, enable_model_value=False + ) model = torch.nn.Linear(4, 2) d.dump_model(model, name_prefix="model") @@ -886,7 +951,9 @@ class TestDumpModel: assert len(filenames) == 0 def test_filter(self, tmp_path): - d = _make_test_dumper(tmp_path, filter="weight") + d = _make_test_dumper( + tmp_path, enable_model_value=True, enable_model_grad=True, filter="weight" + ) model = torch.nn.Linear(4, 2) x = torch.randn(3, 4) y = model(x).sum() @@ -901,7 +968,9 @@ class TestDumpModel: ) def test_grad_file_content(self, tmp_path): - d = _make_test_dumper(tmp_path, enable_model_value=False) + d = _make_test_dumper( + tmp_path, enable_model_grad=True, enable_model_value=False + ) model = torch.nn.Linear(4, 2, bias=False) x = torch.ones(1, 4) y = model(x).sum() @@ -913,7 +982,9 @@ class TestDumpModel: assert torch.equal(_load_dump(path)["value"], model.weight.grad) def test_disable_model_grad(self, tmp_path): - d = _make_test_dumper(tmp_path, enable_model_grad=False) + d = _make_test_dumper( + tmp_path, enable_model_value=True, enable_model_grad=False + ) model = torch.nn.Linear(4, 2) x = torch.randn(3, 4) y = model(x).sum() @@ -925,7 +996,9 @@ class TestDumpModel: assert all("grad" not in f for f in filenames) def test_parameter_saved_as_parameter(self, tmp_path): - d = _make_test_dumper(tmp_path, enable_model_grad=False) + d = _make_test_dumper( + tmp_path, enable_model_value=True, enable_model_grad=False + ) model = torch.nn.Linear(4, 2, bias=False) d.dump_model(model, name_prefix="p") @@ -940,7 +1013,9 @@ class TestDumpModel: def __reduce_ex__(self, protocol): raise RuntimeError("not pickleable") - d = _make_test_dumper(tmp_path, enable_model_grad=False) + d = _make_test_dumper( + tmp_path, enable_model_value=True, enable_model_grad=False + ) model = torch.nn.Linear(4, 2, bias=False) model.weight = BadParam(model.weight.data) @@ -953,7 +1028,9 @@ class TestDumpModel: assert torch.equal(loaded["value"], model.weight.data) def test_disable_model_value(self, tmp_path): - d = _make_test_dumper(tmp_path, enable_model_value=False) + d = _make_test_dumper( + tmp_path, enable_model_grad=True, enable_model_value=False + ) model = torch.nn.Linear(4, 2, bias=False) x = torch.ones(1, 4) y = model(x).sum() @@ -1097,6 +1174,57 @@ class TestReset: assert d._state.step == 1 assert d._state.dump_index == 1 + def test_reset_removes_non_intrusive_hooks(self, tmp_path): + model = torch.nn.Sequential( + torch.nn.Linear(4, 4), + torch.nn.ReLU(), + torch.nn.Linear(4, 4), + ) + d = _make_test_dumper(tmp_path, non_intrusive_mode="all") + d.register_non_intrusive_dumper(model) + + x = torch.randn(2, 4) + with d.capture_output() as captured: + model(x) + assert len(captured) > 0 + + d.reset() + d.configure(enable=True, dir=str(tmp_path), non_intrusive_mode="all") + + with d.capture_output() as captured_after: + model(x) + assert len(captured_after) == 0 + + def test_reset_removes_non_intrusive_hooks_multiple_models(self, tmp_path): + model_a = torch.nn.Sequential( + torch.nn.Linear(4, 4), + torch.nn.ReLU(), + ) + model_b = torch.nn.Sequential( + torch.nn.Linear(4, 4), + torch.nn.ReLU(), + ) + d = _make_test_dumper(tmp_path, non_intrusive_mode="all") + d.register_non_intrusive_dumper(model_a) + d.register_non_intrusive_dumper(model_b) + + x = torch.randn(2, 4) + with d.capture_output() as captured: + model_a(x) + model_b(x) + assert len(captured) > 0 + + d.reset() + d.configure(enable=True, dir=str(tmp_path), non_intrusive_mode="all") + + with d.capture_output() as captured_a: + model_a(x) + assert len(captured_a) == 0 + + with d.capture_output() as captured_b: + model_b(x) + assert len(captured_b) == 0 + def _dumper_worker(rank, http_port: int, stop_event): """Minimal distributed dumper worker: configure, step (triggers ZMQ setup), then wait.""" @@ -1966,5 +2094,56 @@ class TestDumperE2E: kill_process_tree(proc.pid) +class TestRegisterForwardHook: + @pytest.mark.parametrize("mode", ["hook", "replace_fn"]) + def test_handles_removable(self, mode): + call_log: list[str] = [] + + def pre_hook(_module, _input): + call_log.append("pre") + + def hook(_module, _input, _output): + call_log.append("post") + + module = torch.nn.Linear(4, 4) + handles = _register_forward_hook_or_replace_fn( + module, + pre_hook=pre_hook, + hook=hook, + mode=mode, + ) + + x = torch.randn(2, 4) + if mode == "hook": + module(x) + else: + module.forward(x) + assert call_log == ["pre", "post"] + + call_log.clear() + for h in handles: + h.remove() + + if mode == "hook": + module(x) + else: + module.forward(x) + assert call_log == [] + + def test_replace_fn_remove_asserts_on_rewrap(self): + module = torch.nn.Linear(4, 4) + handles = _register_forward_hook_or_replace_fn( + module, + pre_hook=lambda _m, _i: None, + hook=lambda _m, _i, _o: None, + mode="replace_fn", + ) + + module.forward = lambda *a, **kw: None + + with pytest.raises(AssertionError): + handles[0].remove() + + if __name__ == "__main__": sys.exit(pytest.main([__file__]))