Support enabling partial non intrusive dump in dumper (#19069)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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__]))
|
||||
|
||||
Reference in New Issue
Block a user