Handle recompute and verify closeness in dumper (#19564)

This commit is contained in:
fzyzcjy
2026-02-28 18:07:44 +08:00
committed by GitHub
parent 63a4778542
commit 40facdb28c
11 changed files with 557 additions and 13 deletions

View File

@@ -16,6 +16,7 @@ from sglang.srt.debug_utils.dumper import (
DumperConfig,
_collective_with_timeout,
_deepcopy_or_clone,
_detect_recompute_status,
_Dumper,
_format_tags,
_get_default_exp_name,
@@ -23,6 +24,7 @@ from sglang.srt.debug_utils.dumper import (
_materialize_value,
_MegatronPlugin,
_obj_to_dict,
_RecomputeStatus,
_register_forward_hook_or_replace_fn,
_SGLangPlugin,
_torch_save,
@@ -2400,5 +2402,113 @@ class TestCtxDecorator:
d.ctx()
class TestRecomputeStatus:
def test_disabled_by_default(self, tmp_path: Path) -> None:
d = _make_test_dumper(tmp_path)
tensor = torch.randn(3, 3)
d.dump("test_tensor", tensor)
filenames = _get_filenames(tmp_path)
_assert_files(filenames, exist=["recompute_status=disabled"])
def test_recompute_status_in_embedded_meta(self, tmp_path: Path) -> None:
d = _make_test_dumper(tmp_path)
tensor = torch.randn(3, 3)
d.dump("test_tensor", tensor)
path = _find_dump_file(tmp_path, rank=0, name="test_tensor")
raw = _load_dump(path)
assert raw["meta"]["recompute_status"] == "disabled"
def test_recompute_status_recompute(self, tmp_path: Path, monkeypatch) -> None:
import sglang.srt.debug_utils.dumper as dumper_mod
monkeypatch.setattr(
dumper_mod, "_detect_recompute_status", lambda: _RecomputeStatus.RECOMPUTE
)
d = _make_test_dumper(tmp_path)
tensor = torch.randn(3, 3)
d.dump("test_tensor", tensor)
filenames = _get_filenames(tmp_path)
_assert_files(filenames, exist=["recompute_status=recompute"])
path = _find_dump_file(tmp_path, rank=0, name="test_tensor")
raw = _load_dump(path)
assert raw["meta"]["recompute_status"] == "recompute"
assert raw["meta"]["recompute_pseudo_rank"] == 1
assert raw["meta"]["recompute_pseudo_size"] == 2
def test_recompute_status_original(self, tmp_path: Path, monkeypatch) -> None:
import sglang.srt.debug_utils.dumper as dumper_mod
monkeypatch.setattr(
dumper_mod,
"_detect_recompute_status",
lambda: _RecomputeStatus.ORIGINAL,
)
d = _make_test_dumper(tmp_path)
tensor = torch.randn(3, 3)
d.dump("test_tensor", tensor)
filenames = _get_filenames(tmp_path)
_assert_files(filenames, exist=["recompute_status=original"])
path = _find_dump_file(tmp_path, rank=0, name="test_tensor")
raw = _load_dump(path)
assert raw["meta"]["recompute_status"] == "original"
assert raw["meta"]["recompute_pseudo_rank"] == 0
assert raw["meta"]["recompute_pseudo_size"] == 2
def test_disabled_no_recompute_pseudo_fields(self, tmp_path: Path) -> None:
d = _make_test_dumper(tmp_path)
tensor = torch.randn(3, 3)
d.dump("test_tensor", tensor)
path = _find_dump_file(tmp_path, rank=0, name="test_tensor")
raw = _load_dump(path)
assert "recompute_pseudo_rank" not in raw["meta"]
assert "recompute_pseudo_size" not in raw["meta"]
def test_grad_hook_has_no_recompute_status(self, tmp_path: Path) -> None:
d = _make_test_dumper(tmp_path, enable_grad=True)
x = torch.randn(3, 3, requires_grad=True)
y = (x * 2).sum()
d.dump("test_tensor", x)
y.backward()
grad_files = [f for f in _get_filenames(tmp_path) if "grad__test_tensor" in f]
assert len(grad_files) == 1
assert "recompute_status" not in grad_files[0]
def test_non_intrusive_hooks_have_recompute_status(self, tmp_path: Path) -> None:
class Simple(torch.nn.Module):
def __init__(self):
super().__init__()
self.linear = torch.nn.Linear(4, 4)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.linear(x)
model = Simple()
d = _make_test_dumper(tmp_path, non_intrusive_mode="all")
d.register_non_intrusive_dumper(model)
with d.capture_output() as captured:
model(torch.randn(2, 4))
for key, data in captured.items():
assert (
"recompute_status" in data["meta"]
), f"missing recompute_status in {key}"
assert data["meta"]["recompute_status"] == "disabled"
def test_detect_recompute_status_default(self) -> None:
assert _detect_recompute_status() == _RecomputeStatus.DISABLED
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))