Support dumping rid for correlation across passes in dump comparator (#19372)

This commit is contained in:
fzyzcjy
2026-02-26 09:57:57 +08:00
committed by GitHub
parent 7c9e8e2def
commit 46321ee70e
5 changed files with 69 additions and 45 deletions

View File

@@ -2002,7 +2002,7 @@ class TestDumperE2E:
assert len(dump_files) > 0, f"No dump files in {dump_dir}"
filenames = {f.name for f in dump_files}
for field in ("input_ids", "positions"):
for field in ("input_ids", "positions", "rids"):
assert any(f"name={field}" in f for f in filenames), (
f"Missing {field} dump from non-intrusive hooks, "
f"got: {sorted(filenames)[:10]}"
@@ -2022,6 +2022,19 @@ class TestDumperE2E:
assert "name" in loaded["meta"]
assert "rank" in loaded["meta"]
assert "step" in loaded["meta"]
rids_files = [f for f in dump_files if "name=rids" in f.name]
rids_loaded = torch.load(
rids_files[0], map_location="cpu", weights_only=False
)
rids_value = rids_loaded["value"]
assert isinstance(
rids_value, list
), f"rids should be a list, got {type(rids_value)}"
assert len(rids_value) > 0, "rids should be non-empty"
assert all(
isinstance(r, str) for r in rids_value
), f"each rid should be a str, got {[type(r) for r in rids_value]}"
finally:
kill_process_tree(proc.pid)