Support enabling partial non intrusive dump in dumper (#19069)

This commit is contained in:
fzyzcjy
2026-02-22 16:07:45 +08:00
committed by GitHub
parent 0384c459a7
commit 8bc0751376
2 changed files with 149 additions and 24 deletions
+119 -11
View File
@@ -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__]))