Support dims annotation and enhance dump loader in dumper (#19276)

This commit is contained in:
fzyzcjy
2026-02-25 09:41:48 +08:00
committed by GitHub
parent 8b1ab4aaf9
commit 02ca107b2c
9 changed files with 437 additions and 107 deletions
@@ -2256,5 +2256,74 @@ class TestNonIntrusiveKwargsModel(_NonIntrusiveTestBase):
assert captured["qkv_format"]["value"] == "thd"
class TestDumperDims:
def test_dims_in_meta_not_filename(self, tmp_path) -> None:
dumper = _make_test_dumper(tmp_path)
tensor = torch.randn(4, 8)
dumper.dump("hidden", tensor, dims="b h(tp)")
dumper.step()
exp_dir = tmp_path / dumper._config.exp_name
pt_files = list(exp_dir.glob("*.pt"))
assert len(pt_files) == 1
assert "dims" not in pt_files[0].stem
data = torch.load(pt_files[0], weights_only=False)
assert "dims" in data["meta"]
assert data["meta"]["dims"] == "b h(tp)"
def test_dims_grad_override(self, tmp_path) -> None:
dumper = _Dumper(
config=DumperConfig(
enable=True,
dir=str(tmp_path),
enable_http_server=False,
enable_grad=True,
)
)
tensor = torch.randn(4, 8, requires_grad=True)
dumper.dump("hidden", tensor, dims="b h(tp)", dims_grad="b h(tp,partial)")
dumper.step()
tensor.backward(torch.ones_like(tensor))
exp_dir = tmp_path / dumper._config.exp_name
pt_files = sorted(exp_dir.glob("*.pt"))
assert len(pt_files) == 2
value_file = [f for f in pt_files if "grad__" not in f.stem][0]
grad_file = [f for f in pt_files if "grad__" in f.stem][0]
value_data = torch.load(value_file, weights_only=False)
assert value_data["meta"]["dims"] == "b h(tp)"
assert value_data["meta"]["dims_grad"] == "b h(tp,partial)"
grad_data = torch.load(grad_file, weights_only=False)
assert grad_data["meta"]["dims"] == "b h(tp,partial)"
def test_dims_grad_inherits(self, tmp_path) -> None:
dumper = _Dumper(
config=DumperConfig(
enable=True,
dir=str(tmp_path),
enable_http_server=False,
enable_grad=True,
)
)
tensor = torch.randn(4, 8, requires_grad=True)
dumper.dump("hidden", tensor, dims="b h(tp)")
dumper.step()
tensor.backward(torch.ones_like(tensor))
exp_dir = tmp_path / dumper._config.exp_name
grad_file = [f for f in exp_dir.glob("*.pt") if "grad__" in f.stem][0]
grad_data = torch.load(grad_file, weights_only=False)
assert grad_data["meta"]["dims"] == "b h(tp)"
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))