diff --git a/python/sglang/srt/debug_utils/dumper.py b/python/sglang/srt/debug_utils/dumper.py index 3edc59f91..43f970ba3 100644 --- a/python/sglang/srt/debug_utils/dumper.py +++ b/python/sglang/srt/debug_utils/dumper.py @@ -102,6 +102,7 @@ class _DumperConfig(_FrozenConfig): cleanup_previous: bool = False collective_timeout: int = 60 server_port: str = "-1" + non_intrusive_mode: str = "core" @classmethod def _env_prefix(cls) -> str: @@ -234,8 +235,11 @@ class _Dumper: def register_non_intrusive_dumper( self, model: "torch.nn.Module", - ) -> "_NonIntrusiveDumper": - return _NonIntrusiveDumper(dumper=self, model=model) + ) -> Optional["_NonIntrusiveDumper"]: + mode = self._config.non_intrusive_mode + if mode == "off": + return None + return _NonIntrusiveDumper(dumper=self, model=model, mode=mode) # ------------------------------- public :: secondary --------------------------------- @@ -475,39 +479,50 @@ class _Dumper: class _NonIntrusiveDumper: - """Registers forward hooks on model modules to non-invasively dump tensor outputs.""" - _NAME_PREFIX = "non_intrusive__" + _CORE_FIELDS: frozenset[str] = frozenset({"input_ids", "positions"}) def __init__( self, dumper: _Dumper, model: "torch.nn.Module", + mode: str, ): self._dumper = dumper + self._mode = mode for module_name, module in model.named_modules(): module.register_forward_hook( - self._make_forward_hook(module_name=module_name) + self._make_forward_hook( + module_name=module_name, + is_root=(module_name == ""), + ) ) - def _make_forward_hook(self, module_name: str): + def _make_forward_hook(self, *, module_name: str, is_root: bool): def _hook(_module, input, output): for i, item in enumerate(input): - self._dump_value(module_name, item, role=f"inputs.{i}") + self._dump_value(module_name, item, role=f"inputs.{i}", is_root=is_root) if output is not None: - self._dump_value(module_name, output, role="output") + self._dump_value(module_name, output, role="output", is_root=False) return _hook - def _dump_value(self, module_name: str, value, role: str) -> None: - for key, tensor in self._convert_value(value).items(): - parts = [p for p in (module_name, role, key) if p] - self._dumper.dump(self._NAME_PREFIX + ".".join(parts), tensor) + def _dump_value(self, module_name: str, value, role: str, *, is_root: bool) -> None: + for key, tensor in self._convert_value( + value, skip_forward_batch=(not is_root) + ).items(): + if key in self._CORE_FIELDS: + self._dumper.dump(key, tensor) + elif self._mode == "all": + parts = [p for p in (module_name, role, key) if p] + self._dumper.dump(self._NAME_PREFIX + ".".join(parts), tensor) @staticmethod - def _convert_value(value) -> dict[str, torch.Tensor]: + def _convert_value( + value, *, skip_forward_batch: bool = False + ) -> dict[str, torch.Tensor]: if isinstance(value, torch.Tensor): return {"": value} @@ -528,6 +543,8 @@ class _NonIntrusiveDumper: if isinstance(value, LogitsProcessorOutput): return {"next_token_logits": value.next_token_logits} if isinstance(value, ForwardBatch): + if skip_forward_batch: + return {} return { "input_ids": value.input_ids, "seq_lens": value.seq_lens, diff --git a/test/registered/debug_utils/test_dumper.py b/test/registered/debug_utils/test_dumper.py index 575d9f0eb..7dfbd3f2d 100644 --- a/test/registered/debug_utils/test_dumper.py +++ b/test/registered/debug_utils/test_dumper.py @@ -953,7 +953,7 @@ class TestDumperHttp: assert resp.status_code == 400 -class TestNonIntrusiveDumper: +class _NonIntrusiveTestBase: _PREFIX = "non_intrusive__" @staticmethod @@ -978,6 +978,14 @@ class TestNonIntrusiveDumper: return OuterModel() + @staticmethod + def _make_dumper(tmp_path, **overrides) -> "_Dumper": + return _make_test_dumper(tmp_path, non_intrusive_mode="all", **overrides) + + +class TestNonIntrusiveDumper(_NonIntrusiveTestBase): + """Tests for mode='all' — hooks on every module, non_intrusive__ prefix.""" + def test_basic_inputs_and_outputs(self, tmp_path): class Inner(torch.nn.Module): def __init__(self): @@ -988,7 +996,7 @@ class TestNonIntrusiveDumper: def forward(self, x): return self.relu(self.linear(x)) - d = _make_test_dumper(tmp_path) + d = self._make_dumper(tmp_path) model = self._wrap_as_outer(Inner) d.register_non_intrusive_dumper(model) @@ -1043,7 +1051,7 @@ class TestNonIntrusiveDumper: x = layer(x) return x - d = _make_test_dumper(tmp_path) + d = self._make_dumper(tmp_path) model = self._wrap_as_outer(Inner) d.register_non_intrusive_dumper(model) @@ -1084,7 +1092,7 @@ class TestNonIntrusiveDumper: a, b = self.split(x) return self.linear(a + b) - d = _make_test_dumper(tmp_path) + d = self._make_dumper(tmp_path) model = self._wrap_as_outer(Inner) d.register_non_intrusive_dumper(model) @@ -1111,7 +1119,7 @@ class TestNonIntrusiveDumper: def forward(self, x): return self.wrap(x)[0] - d = _make_test_dumper(tmp_path) + d = self._make_dumper(tmp_path) model = self._wrap_as_outer(Inner) d.register_non_intrusive_dumper(model) @@ -1136,7 +1144,7 @@ class TestNonIntrusiveDumper: mask = torch.ones_like(x) return self.mul(x, mask) - d = _make_test_dumper(tmp_path) + d = self._make_dumper(tmp_path) model = self._wrap_as_outer(Inner) d.register_non_intrusive_dumper(model) @@ -1161,7 +1169,7 @@ class TestNonIntrusiveDumper: self.sink(x) return x - d = _make_test_dumper(tmp_path) + d = self._make_dumper(tmp_path) model = self._wrap_as_outer(Inner) d.register_non_intrusive_dumper(model) @@ -1188,7 +1196,7 @@ class TestNonIntrusiveDumper: self.const(x) return x - d = _make_test_dumper(tmp_path) + d = self._make_dumper(tmp_path) model = self._wrap_as_outer(Inner) d.register_non_intrusive_dumper(model) @@ -1202,7 +1210,7 @@ class TestNonIntrusiveDumper: ) def test_root_module_name_no_malformed_dots(self, tmp_path): - d = _make_test_dumper(tmp_path) + d = self._make_dumper(tmp_path) model = torch.nn.Linear(4, 4) d.register_non_intrusive_dumper(model) @@ -1227,7 +1235,7 @@ class TestNonIntrusiveDumper: def forward(self, x): return self.relu(self.linear(x)) - d = _make_test_dumper( + d = self._make_dumper( tmp_path, filter="name=non_intrusive__model.linear.output" ) model = self._wrap_as_outer(Inner) @@ -1250,7 +1258,7 @@ class TestNonIntrusiveDumper: def forward(self, x): return self.linear(x) - d = _make_test_dumper(tmp_path) + d = self._make_dumper(tmp_path) d.configure(enable=False) model = self._wrap_as_outer(Inner) d.register_non_intrusive_dumper(model) @@ -1262,5 +1270,105 @@ class TestNonIntrusiveDumper: assert len(captured) == 0 +def _make_forward_batch(): + from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode + + return ForwardBatch( + forward_mode=ForwardMode.DECODE, + batch_size=2, + input_ids=torch.tensor([10, 20]), + req_pool_indices=torch.zeros(2, dtype=torch.long), + seq_lens=torch.tensor([5, 6]), + out_cache_loc=torch.zeros(2, dtype=torch.long), + seq_lens_sum=11, + positions=torch.tensor([0, 1]), + ) + + +class TestNonIntrusiveDumperConfigMode(_NonIntrusiveTestBase): + + @staticmethod + def _build_model() -> torch.nn.Module: + class SubLayer(torch.nn.Module): + def __init__(self): + super().__init__() + self.linear = torch.nn.Linear(4, 4) + + def forward(self, forward_batch): + return self.linear( + forward_batch.input_ids.float().unsqueeze(-1).expand(-1, 4) + ) + + class Root(torch.nn.Module): + def __init__(self): + super().__init__() + self.layer = SubLayer() + + def forward(self, forward_batch): + return self.layer(forward_batch) + + return Root() + + def _run(self, tmp_path, mode: str) -> tuple: + d = _make_test_dumper(tmp_path, non_intrusive_mode=mode) + model = self._build_model() + d.register_non_intrusive_dumper(model) + forward_batch = _make_forward_batch() + with d.capture_output() as captured: + model(forward_batch) + return captured, forward_batch + + def test_off_mode(self, tmp_path): + captured, _ = self._run(tmp_path, "off") + assert len(captured) == 0 + + def test_core_mode(self, tmp_path): + captured, fb = self._run(tmp_path, "core") + + # core fields dumped with clean names + assert "input_ids" in captured + assert "positions" in captured + assert torch.equal(captured["input_ids"]["value"], fb.input_ids) + assert torch.equal(captured["positions"]["value"], fb.positions) + + # nothing with non_intrusive__ prefix + assert not any(k.startswith("non_intrusive__") for k in captured) + + def test_all_mode(self, tmp_path): + captured, fb = self._run(tmp_path, "all") + + # core fields dumped with clean names + assert "input_ids" in captured + assert "positions" in captured + assert torch.equal(captured["input_ids"]["value"], fb.input_ids) + assert torch.equal(captured["positions"]["value"], fb.positions) + + # non-core ForwardBatch fields dumped with prefix + assert "non_intrusive__inputs.0.seq_lens" in captured + assert torch.equal( + captured["non_intrusive__inputs.0.seq_lens"]["value"], fb.seq_lens + ) + + # core fields NOT duplicated with prefix + assert not any( + k.startswith("non_intrusive__") and k.endswith("input_ids") + for k in captured + ) + assert not any( + k.startswith("non_intrusive__") and k.endswith("positions") + for k in captured + ) + + # ForwardBatch skipped on sub-modules (no duplication) + assert not any( + k.startswith("non_intrusive__layer.inputs.") and "seq_lens" in k + for k in captured + ), f"ForwardBatch skipped on sub-module, got: {list(captured.keys())}" + + # regular tensor outputs on sub-modules still dumped + assert "non_intrusive__layer.linear.output" in captured + assert "non_intrusive__layer.output" in captured + + if __name__ == "__main__": sys.exit(pytest.main([__file__]))