Support dumping gradients, parameters, lazy values (#18881)
Co-authored-by: Yueming Yuan <112649537+yueming-yuan@users.noreply.github.com>
This commit is contained in:
@@ -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():
|
||||
|
||||
@@ -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__]))
|
||||
|
||||
Reference in New Issue
Block a user