Fix wrongly large dumped file and handle non intrusive hook reset in dumper (#19124)
This commit is contained in:
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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__]))
|
||||
|
||||
Reference in New Issue
Block a user