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
+23 -18
View File
@@ -63,7 +63,6 @@ class _Dumper:
): ):
# Config # Config
self._enable = enable self._enable = enable
# TODO (1) support filtering kv instead of name only (2) allow HTTP req change it
self._filter = filter self._filter = filter
self._base_dir = base_dir self._base_dir = base_dir
self._enable_output_file = enable_output_file self._enable_output_file = enable_output_file
@@ -216,8 +215,11 @@ class _Dumper:
if not (self._enable and (self._override_enable is not False)): if not (self._enable and (self._override_enable is not False)):
return 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 return
if not (enable_value or enable_curr_grad or enable_future_grad): if not (enable_value or enable_curr_grad or enable_future_grad):
return return
@@ -229,9 +231,8 @@ class _Dumper:
if enable_value: if enable_value:
self._dump_single( self._dump_single(
tag=value_tag, tag=value_tag,
name=name, tags=tags,
value=value, value=value,
extra_kwargs=extra_kwargs,
save=save, save=save,
) )
@@ -242,9 +243,8 @@ class _Dumper:
): ):
self._dump_single( self._dump_single(
tag=grad_tag, tag=grad_tag,
name=f"grad__{name}", tags={**tags, "name": f"grad__{name}"},
value=g, value=g,
extra_kwargs=extra_kwargs,
save=save, save=save,
) )
@@ -252,12 +252,17 @@ class _Dumper:
self._register_dump_grad_hook( self._register_dump_grad_hook(
name=name, name=name,
tensor=value, tensor=value,
extra_kwargs=extra_kwargs,
save=save, save=save,
**extra_kwargs,
) )
def _register_dump_grad_hook( def _register_dump_grad_hook(
self, *, name: str, tensor, save: bool, **kwargs self,
*,
name: str,
tensor,
extra_kwargs: dict,
save: bool,
) -> None: ) -> None:
if not isinstance(tensor, torch.Tensor): if not isinstance(tensor, torch.Tensor):
return return
@@ -265,14 +270,13 @@ class _Dumper:
return return
captured_forward_pass_id = self._forward_pass_id 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: def grad_hook(grad: torch.Tensor) -> None:
self._dump_single( self._dump_single(
tag="Dumper.Grad", tag="Dumper.Grad",
name=f"grad__{name}", tags=captured_tags,
value=grad, value=grad,
extra_kwargs=captured_extra,
save=save, save=save,
forward_pass_id=captured_forward_pass_id, forward_pass_id=captured_forward_pass_id,
) )
@@ -283,9 +287,8 @@ class _Dumper:
self, self,
*, *,
tag: str, tag: str,
name: str, tags: dict,
value, value,
extra_kwargs: dict,
save: bool, save: bool,
forward_pass_id: Optional[int] = None, forward_pass_id: Optional[int] = None,
) -> None: ) -> None:
@@ -300,12 +303,10 @@ class _Dumper:
else self._forward_pass_id else self._forward_pass_id
), ),
rank=rank, rank=rank,
name=name,
dump_index=self._dump_index, dump_index=self._dump_index,
**extra_kwargs, **tags,
**self._global_ctx,
) )
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 path = self._base_dir / f"sglang_dump_{self._partial_name}" / full_filename
if self._enable_output_console: if self._enable_output_console:
@@ -328,7 +329,7 @@ class _Dumper:
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[name] = output_data self._captured_output_data[tags["name"]] = output_data
else: else:
if self._pending_cleanup: if self._pending_cleanup:
self._pending_cleanup = False self._pending_cleanup = False
@@ -441,6 +442,10 @@ def _materialize_value(value):
return 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): def _deepcopy_or_clone(x):
if isinstance(x, torch.Tensor): if isinstance(x, torch.Tensor):
return x.clone() return x.clone()
+52 -2
View File
@@ -15,6 +15,7 @@ from sglang.srt.debug_utils.dumper import (
_collect_sglang_parallel_info, _collect_sglang_parallel_info,
_collective_with_timeout, _collective_with_timeout,
_Dumper, _Dumper,
_format_tags,
_materialize_value, _materialize_value,
_obj_to_dict, _obj_to_dict,
_torch_save, _torch_save,
@@ -248,7 +249,7 @@ class TestDumperFileWriteControl:
allow_sglang=True, allow_sglang=True,
SGLANG_DUMPER_ENABLE="1", SGLANG_DUMPER_ENABLE="1",
SGLANG_DUMPER_DIR=str(tmp_path), 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)) 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)) assert torch.equal(captured["clone_check"]["value"], torch.zeros(3, 3))
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="^keep") d = _make_test_dumper(tmp_path, filter="name=keep")
with d.capture_output() as captured: with d.capture_output() as captured:
d.dump("keep_this", torch.randn(3, 3)) 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: class TestDumpModel:
def test_grad_basic(self, tmp_path): def test_grad_basic(self, tmp_path):
d = _make_test_dumper(tmp_path, enable_model_value=False) d = _make_test_dumper(tmp_path, enable_model_value=False)