diff --git a/python/sglang/srt/debug_utils/dumper.py b/python/sglang/srt/debug_utils/dumper.py index 3be898034..4e8ccd431 100644 --- a/python/sglang/srt/debug_utils/dumper.py +++ b/python/sglang/srt/debug_utils/dumper.py @@ -4,6 +4,7 @@ import re import socket import threading import time +from copy import deepcopy from functools import cached_property from http.server import BaseHTTPRequestHandler, HTTPServer from pathlib import Path @@ -46,6 +47,10 @@ class _Dumper: base_dir: Path, filter: Optional[str] = None, enable_write_file: bool = True, + enable_value: bool = True, + enable_grad: bool = False, + enable_model_value: bool = True, + enable_model_grad: bool = True, partial_name: Optional[str] = None, enable_http_server: bool = True, ): @@ -55,6 +60,10 @@ class _Dumper: self._filter = filter self._base_dir = base_dir self._enable_write_file = enable_write_file + self._enable_value = enable_value + self._enable_grad = enable_grad + self._enable_model_value = enable_model_value + self._enable_model_grad = enable_model_grad # States self._partial_name = partial_name @@ -71,6 +80,12 @@ class _Dumper: base_dir=Path(_get_str_env_var("SGLANG_DUMPER_DIR", "/tmp")), filter=_get_str_env_var("SGLANG_DUMPER_FILTER"), enable_write_file=get_bool_env_var("SGLANG_DUMPER_WRITE_FILE", "1"), + enable_value=get_bool_env_var("SGLANG_DUMPER_ENABLE_VALUE", "1"), + enable_grad=get_bool_env_var("SGLANG_DUMPER_ENABLE_GRAD", "0"), + enable_model_value=get_bool_env_var( + "SGLANG_DUMPER_ENABLE_MODEL_VALUE", "1" + ), + enable_model_grad=get_bool_env_var("SGLANG_DUMPER_ENABLE_MODEL_GRAD", "1"), partial_name=_get_str_env_var("SGLANG_DUMPER_PARTIAL_NAME"), enable_http_server=get_bool_env_var( "SGLANG_ENABLE_DUMPER_HTTP_SERVER", "1" @@ -126,41 +141,156 @@ class _Dumper: for name, value in data.items(): self.dump(f"{name_prefix}_{name}", value, save=save, **kwargs) - def dump(self, name, value, save: bool = True, **kwargs): + def dump(self, name: str, value, save: bool = True, **kwargs) -> None: + self._dump_inner( + name=name, + value=value, + extra_kwargs=kwargs, + save=save, + enable_value=self._enable_value, + enable_curr_grad=False, + enable_future_grad=self._enable_grad, + value_tag="Dumper.Value", + grad_tag="Dumper.Grad", + ) + + def dump_model( + self, + model: "torch.nn.Module", + name_prefix: str = "param", + save: bool = True, + **kwargs, + ) -> None: + for param_name, param in model.named_parameters(): + self._dump_inner( + name=f"{name_prefix}__{param_name}", + value=param, + extra_kwargs=kwargs, + save=save, + enable_value=self._enable_model_value, + enable_curr_grad=self._enable_model_grad, + enable_future_grad=False, + value_tag="Dumper.ParamValue", + grad_tag="Dumper.ParamGrad", + ) + + def _dump_inner( + self, + *, + name: str, + value, + extra_kwargs: dict, + save: bool, + enable_value: bool, + enable_curr_grad: bool, + enable_future_grad: bool, + value_tag: str, + grad_tag: str, + ) -> None: self._ensure_http_server() if not (self._enable and (self._override_enable is not False)): return if (f := self._filter) is not None and re.search(f, name) is None: return + if not (enable_value or enable_curr_grad or enable_future_grad): + return if self._forward_pass_id < 1: print("Dump without on_forward_pass_start()") + + value = _materialize_value(value) + + if enable_value: + self._dump_single( + tag=value_tag, + name=name, + value=value, + extra_kwargs=extra_kwargs, + save=save, + ) + + if ( + enable_curr_grad + and isinstance(value, torch.Tensor) + and (g := value.grad) is not None + ): + self._dump_single( + tag=grad_tag, + name=f"grad__{name}", + value=g, + extra_kwargs=extra_kwargs, + save=save, + ) + + if enable_future_grad: + self._register_dump_grad_hook( + name=name, + tensor=value, + save=save, + **extra_kwargs, + ) + + def _register_dump_grad_hook( + self, *, name: str, tensor, save: bool, **kwargs + ) -> None: + if not isinstance(tensor, torch.Tensor): + return + if not tensor.requires_grad: + return + + captured_forward_pass_id = self._forward_pass_id + captured_extra = deepcopy(dict(**kwargs)) + + def grad_hook(grad: torch.Tensor) -> None: + self._dump_single( + tag="Dumper.Grad", + name=f"grad__{name}", + value=grad, + extra_kwargs=captured_extra, + save=save, + forward_pass_id=captured_forward_pass_id, + ) + + tensor.register_hook(grad_hook) + + def _dump_single( + self, + *, + tag: str, + name: str, + value, + extra_kwargs: dict, + save: bool, + forward_pass_id: Optional[int] = None, + ) -> None: self._ensure_partial_name() self._dump_index += 1 rank = _get_rank() full_kwargs = dict( - forward_pass_id=self._forward_pass_id, + forward_pass_id=( + forward_pass_id + if forward_pass_id is not None + else self._forward_pass_id + ), rank=rank, name=name, dump_index=self._dump_index, - **kwargs, + **extra_kwargs, **self._global_ctx, ) full_filename = "___".join(f"{k}={v}" for k, v in full_kwargs.items()) + ".pt" path = self._base_dir / f"sglang_dump_{self._partial_name}" / full_filename - sample_value = get_truncated_value(value) - print( - f"[Dumper] [{rank}, {time.time()}] {path} " + f"[{tag}] [{rank}, {time.time()}] {path} " f"type={type(value)} " f"shape={value.shape if isinstance(value, torch.Tensor) else None} " f"dtype={value.dtype if isinstance(value, torch.Tensor) else None} " f"device={value.device if isinstance(value, torch.Tensor) else None} " f"id={id(value)} " - f"sample_value={sample_value}" + f"sample_value={get_truncated_value(value)}" ) if self._enable_write_file and save: @@ -230,7 +360,13 @@ def _obj_to_dict(obj): return ret -# -------------------------------------- static metadata ------------------------------------------ +def _materialize_value(value): + if callable(value): + value = value() + return value + + +# -------------------------------------- static meta ------------------------------------------ def _compute_static_meta(): diff --git a/test/registered/debug_utils/test_dumper.py b/test/registered/debug_utils/test_dumper.py index 22cb8c25d..4409859db 100644 --- a/test/registered/debug_utils/test_dumper.py +++ b/test/registered/debug_utils/test_dumper.py @@ -11,6 +11,7 @@ from sglang.srt.debug_utils.dumper import ( _collect_megatron_parallel_info, _collect_sglang_parallel_info, _Dumper, + _materialize_value, _obj_to_dict, _torch_save, get_tensor_info, @@ -315,6 +316,28 @@ def _find_dump_file(tmpdir, *, rank: int = 0, name: str) -> Path: return matches[0] +class TestMaterializeValue: + def test_materialize_value_callable(self): + tensor = torch.randn(3, 3) + result = _materialize_value(lambda: tensor) + assert torch.equal(result, tensor) + + def test_materialize_value_passthrough(self): + tensor = torch.randn(3, 3) + result = _materialize_value(tensor) + assert result is tensor + + def test_dump_with_callable_value(self, tmp_path): + d = _make_test_dumper(tmp_path) + tensor = torch.randn(4, 4) + d.dump("lazy_tensor", lambda: tensor) + + _assert_files(_get_filenames(tmp_path), exist=["name=lazy_tensor"]) + + path = _find_dump_file(tmp_path, rank=0, name="lazy_tensor") + assert torch.equal(_load_dump(path)["value"], tensor) + + class TestSaveValue: def test_dump_output_format(self, tmp_path): dumper = _make_test_dumper(tmp_path) @@ -364,5 +387,176 @@ class TestStaticMetadata: assert "world_size" in meta +class TestDumpGrad: + def test_dump_grad_basic(self, tmp_path): + d = _make_test_dumper(tmp_path, enable_grad=True) + x = torch.randn(3, 3, requires_grad=True) + y = (x * 2).sum() + + d.dump("test_tensor", x) + y.backward() + + filenames = _get_filenames(tmp_path) + assert any("name=test_tensor" in f and "grad__" not in f for f in filenames) + _assert_files(filenames, exist=["grad__test_tensor"]) + + def test_dump_grad_non_tensor_skipped(self, tmp_path): + d = _make_test_dumper(tmp_path, enable_grad=True) + d.dump("not_tensor", 42) + + _assert_files(_get_filenames(tmp_path), not_exist=["grad__"]) + + def test_dump_grad_no_requires_grad_skipped(self, tmp_path): + d = _make_test_dumper(tmp_path, enable_grad=True) + x = torch.randn(3, 3, requires_grad=False) + d.dump("no_grad_tensor", x) + + _assert_files( + _get_filenames(tmp_path), + exist=["name=no_grad_tensor"], + not_exist=["grad__"], + ) + + def test_dump_grad_captures_forward_pass_id(self, tmp_path): + d = _make_test_dumper(tmp_path, enable_grad=True) + d._forward_pass_id = 42 + x = torch.randn(3, 3, requires_grad=True) + y = (x * 2).sum() + + d.dump("id_test", x) + d._forward_pass_id = 999 + y.backward() + + grad_file = _find_dump_file(tmp_path, name="grad__id_test") + assert "forward_pass_id=42" in grad_file.name + + def test_dump_grad_file_content(self, tmp_path): + d = _make_test_dumper(tmp_path, enable_grad=True) + x = torch.tensor([[1.0, 2.0], [3.0, 4.0]], requires_grad=True) + y = (x * 3).sum() + + d.dump("content_check", x) + y.backward() + + grad_path = _find_dump_file(tmp_path, name="grad__content_check") + expected_grad = torch.full((2, 2), 3.0) + assert torch.equal(_load_dump(grad_path)["value"], expected_grad) + + def test_disable_value(self, tmp_path): + d = _make_test_dumper(tmp_path, enable_value=False, enable_grad=True) + x = torch.randn(3, 3, requires_grad=True) + y = (x * 2).sum() + + d.dump("fwd_disabled", x) + y.backward() + + filenames = _get_filenames(tmp_path) + assert not any( + "name=fwd_disabled" in f and "grad__" not in f for f in filenames + ) + _assert_files(filenames, exist=["grad__fwd_disabled"]) + + def test_disable_grad(self, tmp_path): + d = _make_test_dumper(tmp_path, enable_grad=False) + x = torch.randn(3, 3, requires_grad=True) + y = (x * 2).sum() + + d.dump("grad_disabled", x) + y.backward() + + _assert_files( + _get_filenames(tmp_path), + exist=["name=grad_disabled"], + not_exist=["grad__"], + ) + + +class TestDumpModel: + def test_grad_basic(self, tmp_path): + d = _make_test_dumper(tmp_path, enable_model_value=False) + model = torch.nn.Linear(4, 2) + x = torch.randn(3, 4) + y = model(x).sum() + y.backward() + + d.dump_model(model, name_prefix="model") + + _assert_files( + _get_filenames(tmp_path), + exist=["grad__model__weight", "grad__model__bias"], + ) + + def test_value_basic(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="model") + + _assert_files( + _get_filenames(tmp_path), + exist=["model__weight"], + ) + + def test_no_grad_skipped(self, tmp_path): + d = _make_test_dumper(tmp_path, enable_model_value=False) + model = torch.nn.Linear(4, 2) + + d.dump_model(model, name_prefix="model") + + filenames = _get_filenames(tmp_path) + assert len(filenames) == 0 + + def test_filter(self, tmp_path): + d = _make_test_dumper(tmp_path, filter="weight") + model = torch.nn.Linear(4, 2) + x = torch.randn(3, 4) + y = model(x).sum() + y.backward() + + d.dump_model(model, name_prefix="model") + + _assert_files( + _get_filenames(tmp_path), + exist=["model__weight", "grad__model__weight"], + not_exist=["model__bias", "grad__model__bias"], + ) + + def test_grad_file_content(self, tmp_path): + d = _make_test_dumper(tmp_path, enable_model_value=False) + model = torch.nn.Linear(4, 2, bias=False) + x = torch.ones(1, 4) + y = model(x).sum() + y.backward() + + d.dump_model(model, name_prefix="p") + + path = _find_dump_file(tmp_path, name="grad__p__weight") + 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) + model = torch.nn.Linear(4, 2) + x = torch.randn(3, 4) + y = model(x).sum() + y.backward() + + d.dump_model(model, name_prefix="model") + + filenames = _get_filenames(tmp_path) + assert all("grad" not in f for f in filenames) + + def test_disable_model_value(self, tmp_path): + d = _make_test_dumper(tmp_path, enable_model_value=False) + model = torch.nn.Linear(4, 2, bias=False) + x = torch.ones(1, 4) + y = model(x).sum() + y.backward() + + d.dump_model(model, name_prefix="model") + + filenames = _get_filenames(tmp_path) + assert all("grad" in f for f in filenames) + + if __name__ == "__main__": sys.exit(pytest.main([__file__]))