Enhance reset, states, http in dumper (#19095)

This commit is contained in:
fzyzcjy
2026-02-22 16:17:41 +08:00
committed by GitHub
parent 4091b720c5
commit c1f497e20e
5 changed files with 385 additions and 165 deletions
+1 -1
View File
@@ -25,7 +25,7 @@ class DumpLoader:
from sglang.srt.debug_utils.dumper import dumper from sglang.srt.debug_utils.dumper import dumper
step = dumper._step step = dumper._state.step
conditions = dict(name=name, step=step, **kwargs) conditions = dict(name=name, step=step, **kwargs)
row = find_row(self._df, conditions=conditions) row = find_row(self._df, conditions=conditions)
assert ( assert (
+111 -119
View File
@@ -9,7 +9,7 @@ import time
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from contextlib import contextmanager from contextlib import contextmanager
from copy import deepcopy from copy import deepcopy
from dataclasses import asdict, dataclass, fields, replace from dataclasses import asdict, dataclass, field, fields, replace
from functools import cached_property from functools import cached_property
from http.server import BaseHTTPRequestHandler, HTTPServer from http.server import BaseHTTPRequestHandler, HTTPServer
from pathlib import Path from pathlib import Path
@@ -122,7 +122,7 @@ _DEFAULT_EXP_NAME_PREFIX = "dump_"
@dataclass(frozen=True) @dataclass(frozen=True)
class _DumperConfig(_BaseConfig): class DumperConfig(_BaseConfig):
enable: bool = False enable: bool = False
filter: Optional[str] = None filter: Optional[str] = None
dir: str = "/tmp/dumper" dir: str = "/tmp/dumper"
@@ -158,6 +158,15 @@ class _DumperConfig(_BaseConfig):
# -------------------------------------- dumper core ------------------------------------------ # -------------------------------------- dumper core ------------------------------------------
@dataclass
class _DumperState:
dump_index: int = 0
step: int = 0
global_ctx: dict = field(default_factory=dict)
captured_output_data: Optional[dict] = None
cleanup_previous_handled: bool = False
class _Dumper: class _Dumper:
"""Utility to dump tensors, which can be useful when comparison checking models. """Utility to dump tensors, which can be useful when comparison checking models.
@@ -188,17 +197,9 @@ class _Dumper:
Related: `sglang.srt.debug_utils.dump_comparator` for dump comparison Related: `sglang.srt.debug_utils.dump_comparator` for dump comparison
""" """
def __init__(self, *, config: _DumperConfig): def __init__(self, *, config: DumperConfig):
self._config = config self._config = config
self._state = _DumperState()
self._http_server_handled = not config.enable_http_server
self._cleanup_previous_handled = not config.cleanup_previous
self._dump_index = 0
self._step = 0
self._global_ctx: dict = {}
self._captured_output_data: Optional[dict] = None
self._rpc_broadcast: "_RpcBroadcastBase" = _LocalOnlyBroadcast(self)
# ------------------------------- public :: core --------------------------------- # ------------------------------- public :: core ---------------------------------
@@ -209,7 +210,7 @@ class _Dumper:
def step(self): def step(self):
"""This should be called on all ranks at the end of each iteration.""" """This should be called on all ranks at the end of each iteration."""
self._ensure_http_server() self._http_manager # noqa: B018
if not self._config.enable: if not self._config.enable:
return return
@@ -217,8 +218,8 @@ class _Dumper:
# Users may want to `dump` only on some ranks, thus determine name here # Users may want to `dump` only on some ranks, thus determine name here
self._ensure_exp_name() self._ensure_exp_name()
self._step += 1 self._state.step += 1
print(f"[Dumper] [{time.time()}] step={self._step}") print(f"[Dumper] [{time.time()}] step={self._state.step}")
def dump(self, name: str, value, save: bool = True, **kwargs) -> None: def dump(self, name: str, value, save: bool = True, **kwargs) -> None:
self._dump_inner( self._dump_inner(
@@ -267,15 +268,15 @@ class _Dumper:
... ...
dumper.set_ctx(layer_id=None) dumper.set_ctx(layer_id=None)
""" """
self._global_ctx = { self._state.global_ctx = {
k: v for k, v in (self._global_ctx | kwargs).items() if v is not None k: v for k, v in (self._state.global_ctx | kwargs).items() if v is not None
} }
def register_non_intrusive_dumper( def register_non_intrusive_dumper(
self, self,
model: "torch.nn.Module", model: "torch.nn.Module",
) -> Optional["_NonIntrusiveDumper"]: ) -> Optional["_NonIntrusiveDumper"]:
self._ensure_http_server() self._http_manager # noqa: B018
mode = self._config.non_intrusive_mode mode = self._config.non_intrusive_mode
if mode == "off": if mode == "off":
return None return None
@@ -290,48 +291,29 @@ class _Dumper:
self._config = self._config.with_defaults(**kwargs) self._config = self._config.with_defaults(**kwargs)
def reset(self) -> None: def reset(self) -> None:
self._dump_index = 0 self._state = _DumperState()
self._step = 0
self._global_ctx = {}
@contextmanager @contextmanager
def capture_output(self): def capture_output(self):
assert self._captured_output_data is None assert self._state.captured_output_data is None
self._captured_output_data = {} self._state.captured_output_data = {}
try: try:
yield self._captured_output_data yield self._state.captured_output_data
finally: finally:
self._captured_output_data = None self._state.captured_output_data = None
def get_state(self) -> dict: def get_state(self) -> dict:
return { return {
"config": asdict(self._config), "config": asdict(self._config),
"dump_index": self._dump_index, "dump_index": self._state.dump_index,
"step": self._step, "step": self._state.step,
} }
# ------------------------- public :: only used internally ----------------------------- @cached_property
def _http_manager(self) -> Optional["_DumperHttpManager"]:
def _handle_http_control_request( if self._config.server_port_parsed is None:
self, *, method: str, body: dict[str, Any] return None
) -> list[dict]: return _DumperHttpManager(self)
return self._rpc_broadcast._handle_http_control_request_inner(
method=method, body=body
)
def _handle_http_control_request_inner(
self, *, method: str, body: dict[str, Any]
) -> dict:
if method == "get_state":
return self.get_state()
elif method == "configure":
self.configure(**body)
return {}
elif method == "reset":
self.reset()
return {}
else:
raise ValueError(f"Unknown dumper control method: {method!r}")
# ------------------------- private :: related to dump ----------------------------- # ------------------------- private :: related to dump -----------------------------
@@ -348,12 +330,12 @@ class _Dumper:
value_tag: str, value_tag: str,
grad_tag: str, grad_tag: str,
) -> None: ) -> None:
self._ensure_http_server() self._http_manager # noqa: B018
if not self._config.enable: if not self._config.enable:
return return
tags = dict(name=name, **extra_kwargs, **self._global_ctx) tags = dict(name=name, **extra_kwargs, **self._state.global_ctx)
if (f := self._config.filter) is not None and re.search( if (f := self._config.filter) is not None and re.search(
f, _format_tags(tags) f, _format_tags(tags)
) is None: ) is None:
@@ -405,7 +387,7 @@ class _Dumper:
if not tensor.requires_grad: if not tensor.requires_grad:
return return
captured_step = self._step captured_step = self._state.step
captured_tags = dict(name=f"grad__{name}", **deepcopy(extra_kwargs)) captured_tags = dict(name=f"grad__{name}", **deepcopy(extra_kwargs))
def grad_hook(grad: torch.Tensor) -> None: def grad_hook(grad: torch.Tensor) -> None:
@@ -429,13 +411,13 @@ class _Dumper:
step: Optional[int] = None, step: Optional[int] = None,
) -> None: ) -> None:
self._ensure_exp_name() self._ensure_exp_name()
self._dump_index += 1 self._state.dump_index += 1
rank = _get_rank() rank = _get_rank()
full_kwargs = dict( full_kwargs = dict(
step=(step if step is not None else self._step), step=(step if step is not None else self._state.step),
rank=rank, rank=rank,
dump_index=self._dump_index, dump_index=self._state.dump_index,
**tags, **tags,
) )
full_filename = _format_tags(full_kwargs) + ".pt" full_filename = _format_tags(full_kwargs) + ".pt"
@@ -452,19 +434,22 @@ class _Dumper:
f"sample_value={get_truncated_value(value)}" f"sample_value={get_truncated_value(value)}"
) )
capturing = self._captured_output_data is not None capturing = self._state.captured_output_data is not None
if save and (self._config.enable_output_file or capturing): if save and (self._config.enable_output_file or capturing):
output_data = { output_data = {
"value": value.data if isinstance(value, torch.nn.Parameter) else value, "value": value,
"meta": dict(**full_kwargs, **self._static_meta), "meta": dict(**full_kwargs, **self._static_meta),
} }
if capturing: if capturing:
output_data["value"] = _deepcopy_or_clone(output_data["value"]) output_data["value"] = _deepcopy_or_clone(output_data["value"])
self._captured_output_data[tags["name"]] = output_data self._state.captured_output_data[tags["name"]] = output_data
else: else:
if not self._cleanup_previous_handled: if (
self._cleanup_previous_handled = True not self._state.cleanup_previous_handled
and self._config.cleanup_previous
):
self._state.cleanup_previous_handled = True
_cleanup_old_dumps( _cleanup_old_dumps(
Path(self._config.dir), exp_name=self._config.exp_name Path(self._config.dir), exp_name=self._config.exp_name
) )
@@ -478,33 +463,6 @@ class _Dumper:
def _static_meta(self) -> dict: def _static_meta(self) -> dict:
return _compute_static_meta() return _compute_static_meta()
# Even if DUMPER_ENABLE=0, users may want to use HTTP endpoint to enable it
def _ensure_http_server(self):
if self._http_server_handled:
return
self._http_server_handled = True
http_port = self._config.server_port_parsed
if http_port is None:
return
rpc_broadcast = _create_zmq_rpc_broadcast(
self,
timeout_seconds=self._config.collective_timeout,
)
if _get_rank() == 0:
assert rpc_broadcast is not None
self._rpc_broadcast = rpc_broadcast
if http_port == "reuse":
print(
"[Dumper] Standalone HTTP server disabled, reusing existing ports"
)
else:
_start_http_server(prefix="/dumper/", target=self, http_port=http_port)
print(f"[Dumper] HTTP server started on port {http_port}")
def _ensure_exp_name(self): def _ensure_exp_name(self):
if self._config.exp_name is None: if self._config.exp_name is None:
name = _get_default_exp_name( name = _get_default_exp_name(
@@ -656,15 +614,26 @@ def _torch_save(value, path: str):
return torch.save(value, path) return torch.save(value, path)
except RuntimeError as e: except RuntimeError as e:
if "not pickleable" in str(e): if "not pickleable" in str(e):
# Some parameter subclasses with extra fields are not pickleable stripped = _strip_parameter(value)
if isinstance(value, torch.nn.Parameter): if stripped is not value:
print(f"[Dumper] Observe error={e} and try pickling value.data") print(f"[Dumper] Observe error={e} and try pickling .data")
return _torch_save(value.data, path) return _torch_save(stripped, path)
raise raise
except Exception as e: except Exception as e:
print(f"[Dumper] Observe error={e} when saving data, skip the tensor") 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}
return value
def _collective_with_timeout(fn, operation_name: str, timeout_seconds: int = 60): def _collective_with_timeout(fn, operation_name: str, timeout_seconds: int = 60):
completed = threading.Event() completed = threading.Event()
@@ -726,7 +695,10 @@ def _cleanup_old_dumps(base_dir: Path, exp_name: Optional[str] = None) -> None:
print(f"[Dumper] Cleaned up {entry}") print(f"[Dumper] Cleaned up {entry}")
if dist.is_initialized(): if dist.is_initialized():
dist.barrier() _collective_with_timeout(
dist.barrier,
operation_name="barrier in _cleanup_old_dumps",
)
def _get_rank(): def _get_rank():
@@ -792,6 +764,51 @@ def _compute_static_meta():
return result return result
# -------------------------------------- http manager ------------------------------------------
class _DumperHttpManager:
def __init__(self, dumper: "_Dumper"):
self._dumper = dumper
http_port = self._dumper._config.server_port_parsed
rpc_broadcast = _create_zmq_rpc_broadcast(
self,
timeout_seconds=self._dumper._config.collective_timeout,
)
if _get_rank() == 0:
assert rpc_broadcast is not None
self._rpc_broadcast = rpc_broadcast
if http_port == "reuse":
print(
"[Dumper] Standalone HTTP server disabled, reusing existing ports"
)
else:
_start_http_server(prefix="/dumper/", target=self, http_port=http_port)
print(f"[Dumper] HTTP server started on port {http_port}")
# ------------------------------- public ---------------------------------
def handle_request(self, *, method: str, body: dict[str, Any]) -> list[dict]:
return self._rpc_broadcast._handle_request_inner(method=method, body=body)
# ------------------------------- private ---------------------------------
def _handle_request_inner(self, *, method: str, body: dict[str, Any]) -> dict:
if method == "get_state":
return self._dumper.get_state()
elif method == "configure":
self._dumper.configure(**body)
return {}
elif method == "reset":
self._dumper.reset()
return {}
else:
raise ValueError(f"Unknown dumper control method: {method!r}")
# -------------------------------------- http control server ------------------------------------------ # -------------------------------------- http control server ------------------------------------------
@@ -812,9 +829,7 @@ def _make_http_handler(*, prefix: str, target):
try: try:
req_body = self._get_request_body() req_body = self._get_request_body()
print(f"[Dumper#{_get_rank()}] HTTP {self.path} {req_body=}") print(f"[Dumper#{_get_rank()}] HTTP {self.path} {req_body=}")
result = target._handle_http_control_request( result = target.handle_request(method=method, body=req_body)
method=method, body=req_body
)
resp_body = json.dumps(result).encode() resp_body = json.dumps(result).encode()
self.send_response(200) self.send_response(200)
self.send_header("Content-Type", "application/json") self.send_header("Content-Type", "application/json")
@@ -924,19 +939,6 @@ class _RpcBroadcastBase:
self._handles = handles self._handles = handles
class _LocalOnlyBroadcast(_RpcBroadcastBase):
"""Calls methods directly on the local dumper, wrapping the result in a list."""
def __init__(self, dumper: "_Dumper"):
self._dumper = dumper
def __getattr__(self, method_name: str):
def call(*args, **kwargs):
return [getattr(self._dumper, method_name)(*args, **kwargs)]
return call
class _ZmqRpcBroadcast(_RpcBroadcastBase): class _ZmqRpcBroadcast(_RpcBroadcastBase):
"""Broadcasts method calls to all ZMQ RPC handles. """Broadcasts method calls to all ZMQ RPC handles.
@@ -959,16 +961,6 @@ class _ZmqRpcBroadcast(_RpcBroadcastBase):
# --------------------------------- copied code (avoid dependency) -------------------------------------- # --------------------------------- copied code (avoid dependency) --------------------------------------
def get_int_env_var(name: str, default: int = 0) -> int:
value = os.getenv(name)
if value is None or not value.strip():
return default
try:
return int(value)
except ValueError:
return default
def _get_local_ip_by_remote() -> Optional[str]: def _get_local_ip_by_remote() -> Optional[str]:
# try ipv4 # try ipv4
s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
@@ -1163,7 +1155,7 @@ _plugins: list[_FrameworkPlugin] = [_SGLangPlugin(), _MegatronPlugin()]
# -------------------------------------- singleton ------------------------------------------ # -------------------------------------- singleton ------------------------------------------
dumper = _Dumper(config=_DumperConfig.from_env()) dumper = _Dumper(config=DumperConfig.from_env())
# -------------------------------------- other utility functions ------------------------------------------ # -------------------------------------- other utility functions ------------------------------------------
+1 -1
View File
@@ -2972,7 +2972,7 @@ class Scheduler(
not torch.distributed.is_initialized() not torch.distributed.is_initialized()
or torch.distributed.get_rank() == 0 or torch.distributed.get_rank() == 0
): ):
response = dumper._handle_http_control_request( response = dumper._http_manager.handle_request(
method=recv_req.method, body=recv_req.body method=recv_req.method, body=recv_req.body
) )
self.send_to_tokenizer.send_output( self.send_to_tokenizer.send_output(
@@ -89,7 +89,7 @@ class TestEndToEnd(CustomTestCase):
from argparse import Namespace from argparse import Namespace
from sglang.srt.debug_utils.dump_comparator import main 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 _Dumper, DumperConfig
with tempfile.TemporaryDirectory() as d1, tempfile.TemporaryDirectory() as d2: with tempfile.TemporaryDirectory() as d1, tempfile.TemporaryDirectory() as d2:
baseline_tensor = torch.randn(10, 10) baseline_tensor = torch.randn(10, 10)
@@ -98,7 +98,7 @@ class TestEndToEnd(CustomTestCase):
dump_dirs = [] dump_dirs = []
for d, tensor in [(d1, baseline_tensor), (d2, target_tensor)]: for d, tensor in [(d1, baseline_tensor), (d2, target_tensor)]:
dumper = _Dumper( dumper = _Dumper(
config=_DumperConfig( config=DumperConfig(
enable=True, enable=True,
dir=d, dir=d,
enable_http_server=False, enable_http_server=False,
+270 -42
View File
@@ -14,12 +14,15 @@ import torch.distributed as dist
from sglang.srt.debug_utils.dumper import ( from sglang.srt.debug_utils.dumper import (
_collective_with_timeout, _collective_with_timeout,
_deepcopy_or_clone,
_Dumper, _Dumper,
_DumperConfig, DumperConfig,
_format_tags, _format_tags,
_get_default_exp_name,
_materialize_value, _materialize_value,
_MegatronPlugin, _MegatronPlugin,
_obj_to_dict, _obj_to_dict,
_register_forward_hook_or_replace_fn,
_SGLangPlugin, _SGLangPlugin,
_torch_save, _torch_save,
dumper, dumper,
@@ -54,25 +57,25 @@ def _capture_stdout():
class TestDumperConfig: class TestDumperConfig:
def test_from_env_defaults_match_dataclass_defaults(self): def test_from_env_defaults_match_dataclass_defaults(self):
assert _DumperConfig.from_env() == _DumperConfig() assert DumperConfig.from_env() == DumperConfig()
def test_from_env_bool(self): def test_from_env_bool(self):
with temp_set_env(DUMPER_ENABLE="1"): with temp_set_env(DUMPER_ENABLE="1"):
assert _DumperConfig.from_env().enable is True assert DumperConfig.from_env().enable is True
with temp_set_env(DUMPER_ENABLE="false"): with temp_set_env(DUMPER_ENABLE="false"):
assert _DumperConfig.from_env().enable is False assert DumperConfig.from_env().enable is False
def test_from_env_str(self): def test_from_env_str(self):
with temp_set_env(DUMPER_FILTER="layer_id=0"): with temp_set_env(DUMPER_FILTER="layer_id=0"):
assert _DumperConfig.from_env().filter == "layer_id=0" assert DumperConfig.from_env().filter == "layer_id=0"
def test_from_env_dir(self): def test_from_env_dir(self):
with temp_set_env(DUMPER_DIR="/my/dir"): with temp_set_env(DUMPER_DIR="/my/dir"):
assert _DumperConfig.from_env().dir == "/my/dir" assert DumperConfig.from_env().dir == "/my/dir"
def test_from_env_int(self): def test_from_env_int(self):
with temp_set_env(DUMPER_COLLECTIVE_TIMEOUT="120"): with temp_set_env(DUMPER_COLLECTIVE_TIMEOUT="120"):
assert _DumperConfig.from_env().collective_timeout == 120 assert DumperConfig.from_env().collective_timeout == 120
def test_configure_overrides(self): def test_configure_overrides(self):
d = _make_test_dumper("/tmp") d = _make_test_dumper("/tmp")
@@ -83,87 +86,119 @@ class TestDumperConfig:
def test_type_validation(self): def test_type_validation(self):
with pytest.raises(TypeError, match="enable.*expected bool.*got str"): with pytest.raises(TypeError, match="enable.*expected bool.*got str"):
_DumperConfig(enable="yes") DumperConfig(enable="yes")
with pytest.raises( with pytest.raises(
TypeError, match="collective_timeout.*expected int.*got str" TypeError, match="collective_timeout.*expected int.*got str"
): ):
_DumperConfig(collective_timeout="abc") DumperConfig(collective_timeout="abc")
with pytest.raises(TypeError, match="filter.*expected str.*got int"): with pytest.raises(TypeError, match="filter.*expected str.*got int"):
_DumperConfig(filter=123) DumperConfig(filter=123)
def test_configure_default_skips_when_env_set(self): def test_configure_default_skips_when_env_set(self):
with temp_set_env(DUMPER_FILTER="from_env"): with temp_set_env(DUMPER_FILTER="from_env"):
d = _Dumper(config=_DumperConfig.from_env()) d = _Dumper(config=DumperConfig.from_env())
d.configure_default(filter="from_code") d.configure_default(filter="from_code")
assert d._config.filter == "from_env" assert d._config.filter == "from_env"
def test_configure_default_applies_when_no_env(self): def test_configure_default_applies_when_no_env(self):
d = _Dumper(config=_DumperConfig.from_env()) d = _Dumper(config=DumperConfig.from_env())
d.configure_default(filter="from_code") d.configure_default(filter="from_code")
assert d._config.filter == "from_code" assert d._config.filter == "from_code"
def test_from_env_whitespace_treated_as_unset(self):
with temp_set_env(DUMPER_FILTER=" "):
assert DumperConfig.from_env().filter is None
def test_may_enable_default_false(self): def test_may_enable_default_false(self):
d = _Dumper(config=_DumperConfig()) d = _Dumper(config=DumperConfig())
assert d.may_enable is False assert d.may_enable is False
def test_may_enable_true_when_enabled(self): def test_may_enable_true_when_enabled(self):
d = _Dumper(config=_DumperConfig(enable=True)) d = _Dumper(config=DumperConfig(enable=True))
assert d.may_enable is True assert d.may_enable is True
def test_may_enable_true_when_server_port_set(self): def test_may_enable_true_when_server_port_set(self):
d = _Dumper(config=_DumperConfig(server_port="40000")) d = _Dumper(config=DumperConfig(server_port="40000"))
assert d.may_enable is True assert d.may_enable is True
d2 = _Dumper(config=_DumperConfig(server_port="reuse")) d2 = _Dumper(config=DumperConfig(server_port="reuse"))
assert d2.may_enable is True assert d2.may_enable is True
class TestServerPortParsed:
def test_negative_returns_none(self):
assert DumperConfig(server_port="-1").server_port_parsed is None
def test_zero_returns_none(self):
assert DumperConfig(server_port="0").server_port_parsed is None
def test_positive_returns_int(self):
result = DumperConfig(server_port="40000").server_port_parsed
assert result == 40000
assert isinstance(result, int)
def test_reuse_returns_string(self):
assert DumperConfig(server_port="reuse").server_port_parsed == "reuse"
class TestDefaultExpName:
def test_starts_with_prefix(self):
name = _get_default_exp_name(timeout_seconds=5)
assert name.startswith("dump_")
def test_suffix_format(self):
name = _get_default_exp_name(timeout_seconds=5)
suffix = name[len("dump_") :]
assert len(suffix) == 22
assert suffix[8] == "_"
class TestKvPairsParsing: class TestKvPairsParsing:
def test_from_kv_pairs_none_returns_defaults(self): def test_from_kv_pairs_none_returns_defaults(self):
assert _DumperConfig.from_kv_pairs(None) == _DumperConfig() assert DumperConfig.from_kv_pairs(None) == DumperConfig()
def test_from_kv_pairs_empty_returns_defaults(self): def test_from_kv_pairs_empty_returns_defaults(self):
assert _DumperConfig.from_kv_pairs([]) == _DumperConfig() assert DumperConfig.from_kv_pairs([]) == DumperConfig()
def test_from_kv_pairs_bool_field(self): def test_from_kv_pairs_bool_field(self):
cfg = _DumperConfig.from_kv_pairs(["enable=true"]) cfg = DumperConfig.from_kv_pairs(["enable=true"])
assert cfg.enable is True assert cfg.enable is True
assert cfg.dir == "/tmp/dumper" assert cfg.dir == "/tmp/dumper"
def test_from_kv_pairs_bool_numeric(self): def test_from_kv_pairs_bool_numeric(self):
assert _DumperConfig.from_kv_pairs(["enable=1"]).enable is True assert DumperConfig.from_kv_pairs(["enable=1"]).enable is True
assert _DumperConfig.from_kv_pairs(["enable=0"]).enable is False assert DumperConfig.from_kv_pairs(["enable=0"]).enable is False
def test_from_kv_pairs_int_field(self): def test_from_kv_pairs_int_field(self):
cfg = _DumperConfig.from_kv_pairs(["collective_timeout=120"]) cfg = DumperConfig.from_kv_pairs(["collective_timeout=120"])
assert cfg.collective_timeout == 120 assert cfg.collective_timeout == 120
assert type(cfg.collective_timeout) is int assert type(cfg.collective_timeout) is int
def test_from_kv_pairs_int_field_zero_stays_int(self): def test_from_kv_pairs_int_field_zero_stays_int(self):
cfg = _DumperConfig.from_kv_pairs(["collective_timeout=0"]) cfg = DumperConfig.from_kv_pairs(["collective_timeout=0"])
assert cfg.collective_timeout == 0 assert cfg.collective_timeout == 0
assert type(cfg.collective_timeout) is int assert type(cfg.collective_timeout) is int
def test_from_kv_pairs_str_field_not_coerced(self): def test_from_kv_pairs_str_field_not_coerced(self):
cfg = _DumperConfig.from_kv_pairs(["server_port=0"]) cfg = DumperConfig.from_kv_pairs(["server_port=0"])
assert cfg.server_port == "0" assert cfg.server_port == "0"
assert type(cfg.server_port) is str assert type(cfg.server_port) is str
def test_from_kv_pairs_str_field_one_stays_str(self): def test_from_kv_pairs_str_field_one_stays_str(self):
cfg = _DumperConfig.from_kv_pairs(["server_port=1"]) cfg = DumperConfig.from_kv_pairs(["server_port=1"])
assert cfg.server_port == "1" assert cfg.server_port == "1"
assert type(cfg.server_port) is str assert type(cfg.server_port) is str
def test_from_kv_pairs_optional_str_field(self): def test_from_kv_pairs_optional_str_field(self):
cfg = _DumperConfig.from_kv_pairs(["filter=layer_id=[0-3]"]) cfg = DumperConfig.from_kv_pairs(["filter=layer_id=[0-3]"])
assert cfg.filter == "layer_id=[0-3]" assert cfg.filter == "layer_id=[0-3]"
def test_from_kv_pairs_optional_str_exp_name(self): def test_from_kv_pairs_optional_str_exp_name(self):
cfg = _DumperConfig.from_kv_pairs(["exp_name=my_experiment"]) cfg = DumperConfig.from_kv_pairs(["exp_name=my_experiment"])
assert cfg.exp_name == "my_experiment" assert cfg.exp_name == "my_experiment"
def test_from_kv_pairs_multiple_fields(self): def test_from_kv_pairs_multiple_fields(self):
cfg = _DumperConfig.from_kv_pairs( cfg = DumperConfig.from_kv_pairs(
[ [
"enable=true", "enable=true",
"dir=/my/dir", "dir=/my/dir",
@@ -180,31 +215,31 @@ class TestKvPairsParsing:
def test_from_kv_pairs_missing_equals_raises(self): def test_from_kv_pairs_missing_equals_raises(self):
with pytest.raises(ValueError, match="missing '='"): with pytest.raises(ValueError, match="missing '='"):
_DumperConfig.from_kv_pairs(["enable"]) DumperConfig.from_kv_pairs(["enable"])
def test_from_kv_pairs_unknown_key_raises(self): def test_from_kv_pairs_unknown_key_raises(self):
with pytest.raises(ValueError, match="Unknown config key"): with pytest.raises(ValueError, match="Unknown config key"):
_DumperConfig.from_kv_pairs(["nonexistent=true"]) DumperConfig.from_kv_pairs(["nonexistent=true"])
def test_kv_pairs_to_dict_returns_only_explicit(self): def test_kv_pairs_to_dict_returns_only_explicit(self):
d = _DumperConfig._kv_pairs_to_dict(["enable=true", "dir=/x"]) d = DumperConfig._kv_pairs_to_dict(["enable=true", "dir=/x"])
assert d == {"enable": True, "dir": "/x"} assert d == {"enable": True, "dir": "/x"}
assert "filter" not in d assert "filter" not in d
assert "collective_timeout" not in d assert "collective_timeout" not in d
def test_kv_pairs_to_dict_none_returns_empty(self): def test_kv_pairs_to_dict_none_returns_empty(self):
assert _DumperConfig._kv_pairs_to_dict(None) == {} assert DumperConfig._kv_pairs_to_dict(None) == {}
def test_kv_pairs_to_dict_empty_returns_empty(self): def test_kv_pairs_to_dict_empty_returns_empty(self):
assert _DumperConfig._kv_pairs_to_dict([]) == {} assert DumperConfig._kv_pairs_to_dict([]) == {}
def test_from_kv_pairs_value_with_equals_in_value(self): def test_from_kv_pairs_value_with_equals_in_value(self):
cfg = _DumperConfig.from_kv_pairs(["filter=name=foo"]) cfg = DumperConfig.from_kv_pairs(["filter=name=foo"])
assert cfg.filter == "name=foo" assert cfg.filter == "name=foo"
def test_from_kv_pairs_type_validation_still_works(self): def test_from_kv_pairs_type_validation_still_works(self):
with pytest.raises(TypeError, match="collective_timeout.*expected int"): with pytest.raises(TypeError, match="collective_timeout.*expected int"):
_DumperConfig.from_kv_pairs(["collective_timeout=not_a_number"]) DumperConfig.from_kv_pairs(["collective_timeout=not_a_number"])
class TestDumperPureFunctions: class TestDumperPureFunctions:
@@ -228,6 +263,21 @@ class TestDumperPureFunctions:
assert result["x"] == 10 assert result["x"] == 10
assert "method" not in result assert "method" not in result
def test_deepcopy_or_clone_tensor(self):
original = torch.randn(3, 3)
cloned = _deepcopy_or_clone(original)
assert torch.equal(cloned, original)
original.fill_(999.0)
assert not torch.equal(cloned, original)
def test_deepcopy_or_clone_non_tensor(self):
original = {"a": [1, 2, 3]}
cloned = _deepcopy_or_clone(original)
assert cloned == original
assert cloned is not original
original["a"].append(4)
assert len(cloned["a"]) == 3
def test_get_tensor_info(self): def test_get_tensor_info(self):
info = get_tensor_info(torch.randn(10, 10)) info = get_tensor_info(torch.randn(10, 10))
for key in ["shape=", "dtype=", "min=", "max=", "mean="]: for key in ["shape=", "dtype=", "min=", "max=", "mean="]:
@@ -339,7 +389,7 @@ class TestDumperDistributed:
@staticmethod @staticmethod
def _test_collective_timeout_func(rank): def _test_collective_timeout_func(rank):
dumper = _Dumper( dumper = _Dumper(
config=_DumperConfig( config=DumperConfig(
enable=True, enable=True,
collective_timeout=3, collective_timeout=3,
enable_http_server=False, enable_http_server=False,
@@ -422,6 +472,13 @@ class TestDumperFileWriteControl:
assert len(_get_filenames(tmpdir)) == 0 assert len(_get_filenames(tmpdir)) == 0
class TestDumpEnableFlags:
def test_all_enables_false_no_output(self, tmp_path):
d = _make_test_dumper(tmp_path, enable_value=False, enable_grad=False)
d.dump("should_skip", torch.randn(3, 3))
assert len(_get_filenames(tmp_path)) == 0
class TestOutputControl: class TestOutputControl:
def test_file_enabled_by_default(self, tmp_path): def test_file_enabled_by_default(self, tmp_path):
d = _make_test_dumper(tmp_path) d = _make_test_dumper(tmp_path)
@@ -493,6 +550,13 @@ class TestOutputControl:
tensor.fill_(999.0) tensor.fill_(999.0)
assert torch.equal(captured["clone_check"]["value"], torch.zeros(3, 3)) assert torch.equal(captured["clone_check"]["value"], torch.zeros(3, 3))
def test_capture_output_nested_raises(self, tmp_path):
d = _make_test_dumper(tmp_path)
with d.capture_output():
with pytest.raises(AssertionError):
with d.capture_output():
pass
def test_capture_output_respects_filter(self, tmp_path): def test_capture_output_respects_filter(self, tmp_path):
d = _make_test_dumper(tmp_path, filter="name=keep") d = _make_test_dumper(tmp_path, filter="name=keep")
@@ -548,7 +612,7 @@ def _make_test_dumper(tmp_path, **overrides) -> _Dumper:
enable_http_server=False, enable_http_server=False,
) )
defaults.update(overrides) defaults.update(overrides)
config = _DumperConfig(**defaults) config = DumperConfig(**defaults)
return _Dumper(config=config) return _Dumper(config=config)
@@ -685,12 +749,12 @@ class TestDumpGrad:
def test_dump_grad_captures_step(self, tmp_path): def test_dump_grad_captures_step(self, tmp_path):
d = _make_test_dumper(tmp_path, enable_grad=True) d = _make_test_dumper(tmp_path, enable_grad=True)
d._step = 42 d._state.step = 42
x = torch.randn(3, 3, requires_grad=True) x = torch.randn(3, 3, requires_grad=True)
y = (x * 2).sum() y = (x * 2).sum()
d.dump("id_test", x) d.dump("id_test", x)
d._step = 999 d._state.step = 999
y.backward() y.backward()
grad_file = _find_dump_file(tmp_path, name="grad__id_test") grad_file = _find_dump_file(tmp_path, name="grad__id_test")
@@ -860,6 +924,34 @@ class TestDumpModel:
filenames = _get_filenames(tmp_path) filenames = _get_filenames(tmp_path)
assert all("grad" not in f for f in filenames) 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)
model = torch.nn.Linear(4, 2, bias=False)
d.dump_model(model, name_prefix="p")
path = _find_dump_file(tmp_path, name="p__weight")
loaded = _load_dump(path)
assert isinstance(loaded["value"], torch.nn.Parameter)
assert torch.equal(loaded["value"], model.weight)
def test_unpicklable_parameter_falls_back_to_data(self, tmp_path):
class BadParam(torch.nn.Parameter):
def __reduce_ex__(self, protocol):
raise RuntimeError("not pickleable")
d = _make_test_dumper(tmp_path, enable_model_grad=False)
model = torch.nn.Linear(4, 2, bias=False)
model.weight = BadParam(model.weight.data)
d.dump_model(model, name_prefix="p")
path = _find_dump_file(tmp_path, name="p__weight")
loaded = _load_dump(path)
assert isinstance(loaded["value"], torch.Tensor)
assert not isinstance(loaded["value"], torch.nn.Parameter)
assert torch.equal(loaded["value"], model.weight.data)
def test_disable_model_value(self, tmp_path): 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_value=False)
model = torch.nn.Linear(4, 2, bias=False) model = torch.nn.Linear(4, 2, bias=False)
@@ -934,9 +1026,9 @@ class TestReset:
d.reset() d.reset()
assert d._dump_index == 0 assert d._state.dump_index == 0
assert d._step == 0 assert d._state.step == 0
assert d._global_ctx == {} assert d._state.global_ctx == {}
def test_dump_works_after_reset(self, tmp_path): def test_dump_works_after_reset(self, tmp_path):
d = _make_test_dumper(tmp_path) d = _make_test_dumper(tmp_path)
@@ -950,6 +1042,123 @@ class TestReset:
post_file = _find_dump_file(tmp_path, name="post") post_file = _find_dump_file(tmp_path, name="post")
assert "dump_index=1" in post_file.name assert "dump_index=1" in post_file.name
def test_cleanup_previous_re_triggers_after_reset(self, tmp_path):
"""Miles pattern: reset() + configure(cleanup_previous=True) should re-clean."""
exp_alpha = "exp_alpha"
exp_beta = "exp_beta"
(tmp_path / exp_alpha).mkdir()
(tmp_path / exp_alpha / "stale.pt").touch()
(tmp_path / exp_beta).mkdir()
(tmp_path / exp_beta / "stale.pt").touch()
d = _make_test_dumper(tmp_path, exp_name=exp_alpha, cleanup_previous=True)
d.dump("phase1", torch.randn(2, 2))
d.reset()
d.configure(exp_name=exp_beta, cleanup_previous=True)
d.dump("phase2", torch.randn(2, 2))
assert not (tmp_path / exp_alpha / "stale.pt").exists()
assert not (tmp_path / exp_beta / "stale.pt").exists()
filenames = _get_filenames(tmp_path)
_assert_files(filenames, exist=["phase1", "phase2"])
def test_no_cleanup_when_config_false(self, tmp_path):
"""cleanup_previous=False: handled stays False but no cleanup runs."""
old_dir = tmp_path / "dump_old"
old_dir.mkdir()
(old_dir / "dummy.pt").touch()
d = _make_test_dumper(tmp_path, cleanup_previous=False)
d.dump("tensor", torch.randn(2, 2))
assert old_dir.exists()
assert d._state.cleanup_previous_handled is False
def test_multi_phase_switch(self, tmp_path):
"""Simulate Miles multi-phase: configure → dump → reset → configure new phase → dump."""
d = _make_test_dumper(tmp_path, cleanup_previous=True)
d.configure(exp_name="fwd_only")
d.dump("weight", torch.randn(2, 2))
d.step()
d.configure(enable=False)
d.reset()
d.configure(exp_name="fwd_bwd", enable=True, cleanup_previous=True)
d.dump("weight", torch.randn(2, 2))
d.step()
fwd_only_files = list(Path(tmp_path).glob("fwd_only/*.pt"))
fwd_bwd_files = list(Path(tmp_path).glob("fwd_bwd/*.pt"))
assert len(fwd_only_files) > 0
assert len(fwd_bwd_files) > 0
assert d._state.step == 1
assert d._state.dump_index == 1
def _dumper_worker(rank, http_port: int, stop_event):
"""Minimal distributed dumper worker: configure, step (triggers ZMQ setup), then wait."""
dumper.configure(enable=False, server_port=str(http_port))
dumper.step()
stop_event.wait()
def _wait_for_dumper_http(url: str, timeout: float = 30) -> None:
deadline = time.time() + timeout
while time.time() < deadline:
try:
requests.post(f"{url}/dumper/configure", json={}, timeout=2)
return
except requests.ConnectionError:
time.sleep(0.5)
raise TimeoutError(f"Dumper HTTP server not reachable at {url}")
class TestZmqPortIsolation:
"""Multiple independent dumper instances (each with 2 ranks) must not conflict on ZMQ ports."""
NUM_INSTANCES = 3
def test_concurrent_instances_no_port_conflict(self):
ports = [
find_available_port(40000 + i * 1000) for i in range(self.NUM_INSTANCES)
]
stop_events = []
threads = []
ctx = multiprocessing.get_context("spawn")
for port in ports:
stop_event = ctx.Event()
stop_events.append(stop_event)
thread = threading.Thread(
target=run_distributed_test,
args=(_dumper_worker,),
kwargs={"http_port": port, "stop_event": stop_event},
)
thread.start()
threads.append(thread)
try:
for port in ports:
_wait_for_dumper_http(f"http://127.0.0.1:{port}")
for i, port in enumerate(ports):
resp = requests.post(
f"http://127.0.0.1:{port}/dumper/get_state", json={}
)
resp.raise_for_status()
states = resp.json()
assert (
len(states) == 2
), f"Instance {i} (port {port}): expected 2 ranks, got {len(states)}"
finally:
for event in stop_events:
event.set()
for thread in threads:
thread.join(timeout=10)
def _dumper_worker(rank, http_port: int, stop_event): def _dumper_worker(rank, http_port: int, stop_event):
"""Minimal distributed dumper worker: configure, step (triggers ZMQ setup), then wait.""" """Minimal distributed dumper worker: configure, step (triggers ZMQ setup), then wait."""
@@ -1124,6 +1333,13 @@ class TestDumperHttp:
) )
assert resp.status_code == 400 assert resp.status_code == 400
def test_error_unknown_method(self, dumper_http_url: str):
resp = requests.post(
f"{dumper_http_url}/dumper/nonexistent",
json={},
)
assert resp.status_code == 400
def test_error_wrong_type(self, dumper_http_url: str): def test_error_wrong_type(self, dumper_http_url: str):
resp = requests.post( resp = requests.post(
f"{dumper_http_url}/dumper/configure", f"{dumper_http_url}/dumper/configure",
@@ -1132,6 +1348,18 @@ class TestDumperHttp:
assert resp.status_code == 400 assert resp.status_code == 400
class TestRegisterForwardHookOrReplaceFn:
def test_unknown_mode_raises(self):
module = torch.nn.Linear(4, 4)
with pytest.raises(ValueError, match="Unknown mode"):
_register_forward_hook_or_replace_fn(
module,
pre_hook=lambda _mod, _input: None,
hook=lambda _mod, _input, _output: None,
mode="bad",
)
class _NonIntrusiveTestBase: class _NonIntrusiveTestBase:
_PREFIX = "non_intrusive__" _PREFIX = "non_intrusive__"