Support filtering labels in dumper (#19018)

This commit is contained in:
fzyzcjy
2026-02-20 12:27:12 +08:00
committed by GitHub
parent 261bca3c58
commit df995aab56
2 changed files with 75 additions and 20 deletions

View File

@@ -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()

View File

@@ -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)