From 170dbd6ebb536f5d9c258b4ae2252608f0b7450b Mon Sep 17 00:00:00 2001 From: Fillianore Date: Tue, 15 Sep 2026 12:09:11 +0800 Subject: [PATCH 1/3] fix(pt): fall back to single-rank .pt2 when with-comm export fails `dp --pt freeze` aborts entirely when the optional parallel with-comm artifact export raises, leaving behind an unloadable partial .pt2 that misses `model/extra/metadata.json`. On torch 2.12.1 the export reliably fails: `aot_export_module`'s `detect_fake_mode` sees graph inputs from two different FakeTensorModes (one created by the earlier main-graph `make_fx` trace, one by `torch.export`'s `make_fake_inputs`) and raises AssertionError. The with-comm artifact is optional by design: the .pt2 format and the loader already handle its absence. Catch the failure, log a warning and mark has_comm_artifact=false so freeze still produces a valid single-rank archive. --- deepmd/pt/entrypoints/freeze_pt2.py | 26 ++++++++++++++++++++------ 1 file changed, 20 insertions(+), 6 deletions(-) diff --git a/deepmd/pt/entrypoints/freeze_pt2.py b/deepmd/pt/entrypoints/freeze_pt2.py index 62693d547d..9bd87b95ff 100644 --- a/deepmd/pt/entrypoints/freeze_pt2.py +++ b/deepmd/pt/entrypoints/freeze_pt2.py @@ -1085,18 +1085,32 @@ def _freeze_sezm_to_pt2( "Compiling the parallel with-comm artifact (second AOTInductor " "compilation)..." ) - with_comm_bytes = _export_with_comm_artifact( - model, - target_device=target_device, - compile_options=compile_options, - ) + try: + with_comm_bytes = _export_with_comm_artifact( + model, + target_device=target_device, + compile_options=compile_options, + ) + except Exception as e: + # The with-comm artifact is optional: the .pt2 format and the + # loader already handle its absence (``has_comm_artifact=false``, + # single-rank inference). A failure here must not abort the + # whole freeze, which would leave behind an unloadable partial + # archive without ``model/extra/metadata.json``. + with_comm_bytes = None + log.warning( + "Parallel with-comm artifact export failed (%s); the frozen " + ".pt2 will support single-rank inference only " + "(has_comm_artifact=false).", + e, + ) metadata = _collect_metadata( model, output_keys=output_keys, is_spin=is_spin, do_atomic_virial=atomic_virial, - has_comm_artifact=with_comm, + has_comm_artifact=with_comm and with_comm_bytes is not None, ) with zipfile.ZipFile(out_path_str, "a") as zf: zf.writestr("model/extra/metadata.json", json.dumps(metadata)) From 00241489d131757f0b92a4985346762f5f126709 Mon Sep 17 00:00:00 2001 From: Fillianore Date: Tue, 15 Sep 2026 13:02:11 +0800 Subject: [PATCH 2/3] test(pt): cover the with-comm export failure fallback in SeZM freeze Force _export_with_comm_artifact to raise through freeze_sezm_to_pt2 and verify the archive still completes: metadata.json is present, has_comm_artifact is false, forward_lower_with_comm.pt2 is absent, a single-rank fallback WARNING is logged, and the frozen archive loads and runs via aoti_load_package with finite outputs. --- source/tests/pt/model/test_sezm_export.py | 58 +++++++++++++++++++++++ 1 file changed, 58 insertions(+) diff --git a/source/tests/pt/model/test_sezm_export.py b/source/tests/pt/model/test_sezm_export.py index 04a920d83f..7d3e5401a4 100644 --- a/source/tests/pt/model/test_sezm_export.py +++ b/source/tests/pt/model/test_sezm_export.py @@ -1206,6 +1206,64 @@ def fake_compile(_exported: torch.export.ExportedProgram, package_path: str): self.assertTrue(metadata["has_message_passing"]) self.assertIn("model/extra/forward_lower_with_comm.pt2", names) + @unittest.skipIf(_SKIP_OFF_COMPILE_TORCH, _SKIP_OFF_COMPILE_TORCH_REASON) + def test_with_comm_export_failure_falls_back_to_single_rank(self) -> None: + """A failing with-comm export degrades to a single-rank archive. + + The with-comm artifact is optional: when its export raises (e.g. the + FakeTensorMode mismatch between the main-graph ``make_fx`` trace and + ``torch.export``'s ``make_fake_inputs`` on torch 2.12.1), freeze must + still emit a complete, loadable ``.pt2`` instead of aborting and + leaving behind a partial archive without ``metadata.json``. + """ + + with tempfile.TemporaryDirectory() as tmp: + tmp_path = Path(tmp) + params = _tiny_sezm_model_params() + ckpt_path = _write_tiny_sezm_checkpoint(tmp_path, params) + out = tmp_path / "single_rank.pt2" + + with ( + mock.patch( + "deepmd.pt.entrypoints.freeze_pt2._export_with_comm_artifact", + side_effect=RuntimeError("simulated with-comm export failure"), + ), + self.assertLogs( + "deepmd.pt.entrypoints.freeze_pt2", level="WARNING" + ) as logs, + ): + freeze_sezm_to_pt2(str(ckpt_path), str(out), device=_CPU) + + self.assertTrue( + any("single-rank" in message for message in logs.output), + msg=f"expected a single-rank fallback warning, got {logs.output}", + ) + self.assertTrue(zipfile.is_zipfile(str(out))) + with zipfile.ZipFile(str(out), "r") as zf: + names = zf.namelist() + self.assertIn("model/extra/metadata.json", names) + self.assertNotIn("model/extra/forward_lower_with_comm.pt2", names) + metadata = json.loads( + zf.read("model/extra/metadata.json").decode("utf-8") + ) + self.assertFalse(metadata["has_comm_artifact"]) + + # The fallback archive still loads and runs in single-rank mode. + from torch._inductor import ( + aoti_load_package, + ) + + loader = aoti_load_package(str(out)) + probe = _build_tiny_sezm_model() + outs = loader(*_make_sample(probe, nloc=5, start=2)) + if hasattr(outs, "items"): + out_map = dict(outs.items()) + else: + out_map = dict(zip(metadata["output_keys"], outs, strict=True)) + for key in ("energy_redu", "energy_derv_r"): + self.assertIn(key, out_map) + self.assertTrue(torch.isfinite(out_map[key]).all().item()) + if __name__ == "__main__": unittest.main() From 9c99867f27e89543ee2dab7b4289588d738467fd Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Tue, 15 Sep 2026 05:03:09 +0000 Subject: [PATCH 3/3] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- source/tests/pt/model/test_sezm_export.py | 1 - 1 file changed, 1 deletion(-) diff --git a/source/tests/pt/model/test_sezm_export.py b/source/tests/pt/model/test_sezm_export.py index 7d3e5401a4..e8ddd799da 100644 --- a/source/tests/pt/model/test_sezm_export.py +++ b/source/tests/pt/model/test_sezm_export.py @@ -1216,7 +1216,6 @@ def test_with_comm_export_failure_falls_back_to_single_rank(self) -> None: still emit a complete, loadable ``.pt2`` instead of aborting and leaving behind a partial archive without ``metadata.json``. """ - with tempfile.TemporaryDirectory() as tmp: tmp_path = Path(tmp) params = _tiny_sezm_model_params()