Support filtering labels in dumper (#19018)
This commit is contained in:
@@ -63,7 +63,6 @@ class _Dumper:
|
||||
):
|
||||
# Config
|
||||
self._enable = enable
|
||||
# TODO (1) support filtering kv instead of name only (2) allow HTTP req change it
|
||||
self._filter = filter
|
||||
self._base_dir = base_dir
|
||||
self._enable_output_file = enable_output_file
|
||||
@@ -216,8 +215,11 @@ class _Dumper:
|
||||
|
||||
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:
|
||||
|
||||
tags = dict(name=name, **extra_kwargs, **self._global_ctx)
|
||||
if (f := self._filter) is not None and re.search(f, _format_tags(tags)) is None:
|
||||
return
|
||||
|
||||
if not (enable_value or enable_curr_grad or enable_future_grad):
|
||||
return
|
||||
|
||||
@@ -229,9 +231,8 @@ class _Dumper:
|
||||
if enable_value:
|
||||
self._dump_single(
|
||||
tag=value_tag,
|
||||
name=name,
|
||||
tags=tags,
|
||||
value=value,
|
||||
extra_kwargs=extra_kwargs,
|
||||
save=save,
|
||||
)
|
||||
|
||||
@@ -242,9 +243,8 @@ class _Dumper:
|
||||
):
|
||||
self._dump_single(
|
||||
tag=grad_tag,
|
||||
name=f"grad__{name}",
|
||||
tags={**tags, "name": f"grad__{name}"},
|
||||
value=g,
|
||||
extra_kwargs=extra_kwargs,
|
||||
save=save,
|
||||
)
|
||||
|
||||
@@ -252,12 +252,17 @@ class _Dumper:
|
||||
self._register_dump_grad_hook(
|
||||
name=name,
|
||||
tensor=value,
|
||||
extra_kwargs=extra_kwargs,
|
||||
save=save,
|
||||
**extra_kwargs,
|
||||
)
|
||||
|
||||
def _register_dump_grad_hook(
|
||||
self, *, name: str, tensor, save: bool, **kwargs
|
||||
self,
|
||||
*,
|
||||
name: str,
|
||||
tensor,
|
||||
extra_kwargs: dict,
|
||||
save: bool,
|
||||
) -> None:
|
||||
if not isinstance(tensor, torch.Tensor):
|
||||
return
|
||||
@@ -265,14 +270,13 @@ class _Dumper:
|
||||
return
|
||||
|
||||
captured_forward_pass_id = self._forward_pass_id
|
||||
captured_extra = deepcopy(dict(**kwargs))
|
||||
captured_tags = dict(name=f"grad__{name}", **deepcopy(extra_kwargs))
|
||||
|
||||
def grad_hook(grad: torch.Tensor) -> None:
|
||||
self._dump_single(
|
||||
tag="Dumper.Grad",
|
||||
name=f"grad__{name}",
|
||||
tags=captured_tags,
|
||||
value=grad,
|
||||
extra_kwargs=captured_extra,
|
||||
save=save,
|
||||
forward_pass_id=captured_forward_pass_id,
|
||||
)
|
||||
@@ -283,9 +287,8 @@ class _Dumper:
|
||||
self,
|
||||
*,
|
||||
tag: str,
|
||||
name: str,
|
||||
tags: dict,
|
||||
value,
|
||||
extra_kwargs: dict,
|
||||
save: bool,
|
||||
forward_pass_id: Optional[int] = None,
|
||||
) -> None:
|
||||
@@ -300,12 +303,10 @@ class _Dumper:
|
||||
else self._forward_pass_id
|
||||
),
|
||||
rank=rank,
|
||||
name=name,
|
||||
dump_index=self._dump_index,
|
||||
**extra_kwargs,
|
||||
**self._global_ctx,
|
||||
**tags,
|
||||
)
|
||||
full_filename = "___".join(f"{k}={v}" for k, v in full_kwargs.items()) + ".pt"
|
||||
full_filename = _format_tags(full_kwargs) + ".pt"
|
||||
path = self._base_dir / f"sglang_dump_{self._partial_name}" / full_filename
|
||||
|
||||
if self._enable_output_console:
|
||||
@@ -328,7 +329,7 @@ class _Dumper:
|
||||
|
||||
if capturing:
|
||||
output_data["value"] = _deepcopy_or_clone(output_data["value"])
|
||||
self._captured_output_data[name] = output_data
|
||||
self._captured_output_data[tags["name"]] = output_data
|
||||
else:
|
||||
if self._pending_cleanup:
|
||||
self._pending_cleanup = False
|
||||
@@ -441,6 +442,10 @@ def _materialize_value(value):
|
||||
return value
|
||||
|
||||
|
||||
def _format_tags(kwargs: dict) -> str:
|
||||
return "___".join(f"{k}={v}" for k, v in kwargs.items())
|
||||
|
||||
|
||||
def _deepcopy_or_clone(x):
|
||||
if isinstance(x, torch.Tensor):
|
||||
return x.clone()
|
||||
|
||||
@@ -15,6 +15,7 @@ from sglang.srt.debug_utils.dumper import (
|
||||
_collect_sglang_parallel_info,
|
||||
_collective_with_timeout,
|
||||
_Dumper,
|
||||
_format_tags,
|
||||
_materialize_value,
|
||||
_obj_to_dict,
|
||||
_torch_save,
|
||||
@@ -248,7 +249,7 @@ class TestDumperFileWriteControl:
|
||||
allow_sglang=True,
|
||||
SGLANG_DUMPER_ENABLE="1",
|
||||
SGLANG_DUMPER_DIR=str(tmp_path),
|
||||
SGLANG_DUMPER_FILTER="^keep",
|
||||
SGLANG_DUMPER_FILTER="name=keep",
|
||||
):
|
||||
run_distributed_test(self._test_filter_func, tmpdir=str(tmp_path))
|
||||
|
||||
@@ -360,7 +361,7 @@ class TestOutputControl:
|
||||
assert torch.equal(captured["clone_check"]["value"], torch.zeros(3, 3))
|
||||
|
||||
def test_capture_output_respects_filter(self, tmp_path):
|
||||
d = _make_test_dumper(tmp_path, filter="^keep")
|
||||
d = _make_test_dumper(tmp_path, filter="name=keep")
|
||||
|
||||
with d.capture_output() as captured:
|
||||
d.dump("keep_this", torch.randn(3, 3))
|
||||
@@ -603,6 +604,55 @@ class TestDumpGrad:
|
||||
)
|
||||
|
||||
|
||||
class TestKvFilter:
|
||||
def test_format_tags(self):
|
||||
assert _format_tags({"a": 1, "b": "hello"}) == "a=1___b=hello"
|
||||
assert _format_tags({}) == ""
|
||||
|
||||
def test_filter_matches_extra_kwargs(self, tmp_path):
|
||||
d = _make_test_dumper(tmp_path, filter="layer_id=0")
|
||||
d.dump("tensor_a", torch.randn(3), layer_id=0)
|
||||
d.dump("tensor_b", torch.randn(3), layer_id=1)
|
||||
|
||||
filenames = _get_filenames(tmp_path)
|
||||
_assert_files(filenames, exist=["tensor_a"], not_exist=["tensor_b"])
|
||||
|
||||
def test_filter_matches_global_ctx(self, tmp_path):
|
||||
d = _make_test_dumper(tmp_path, filter="ctx_arg=200")
|
||||
d.set_ctx(ctx_arg=200)
|
||||
d.dump("tensor_a", torch.randn(3))
|
||||
d.set_ctx(ctx_arg=None)
|
||||
d.dump("tensor_b", torch.randn(3))
|
||||
|
||||
filenames = _get_filenames(tmp_path)
|
||||
_assert_files(filenames, exist=["tensor_a"], not_exist=["tensor_b"])
|
||||
|
||||
def test_filter_matches_name(self, tmp_path):
|
||||
d = _make_test_dumper(tmp_path, filter="name=keep")
|
||||
d.dump("keep_this", torch.randn(3))
|
||||
d.dump("skip_this", torch.randn(3))
|
||||
|
||||
filenames = _get_filenames(tmp_path)
|
||||
_assert_files(filenames, exist=["keep_this"], not_exist=["skip_this"])
|
||||
|
||||
def test_filter_regex(self, tmp_path):
|
||||
d = _make_test_dumper(tmp_path, filter=r"layer_id=[0-2]")
|
||||
d.dump("t0", torch.randn(3), layer_id=0)
|
||||
d.dump("t1", torch.randn(3), layer_id=1)
|
||||
d.dump("t5", torch.randn(3), layer_id=5)
|
||||
|
||||
filenames = _get_filenames(tmp_path)
|
||||
_assert_files(filenames, exist=["name=t0", "name=t1"], not_exist=["name=t5"])
|
||||
|
||||
def test_no_filter_dumps_all(self, tmp_path):
|
||||
d = _make_test_dumper(tmp_path)
|
||||
d.dump("a", torch.randn(3))
|
||||
d.dump("b", torch.randn(3))
|
||||
|
||||
filenames = _get_filenames(tmp_path)
|
||||
_assert_files(filenames, exist=["name=a", "name=b"])
|
||||
|
||||
|
||||
class TestDumpModel:
|
||||
def test_grad_basic(self, tmp_path):
|
||||
d = _make_test_dumper(tmp_path, enable_model_value=False)
|
||||
|
||||
Reference in New Issue
Block a user