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)) diff --git a/source/tests/pt/model/test_sezm_export.py b/source/tests/pt/model/test_sezm_export.py index 04a920d83f..e8ddd799da 100644 --- a/source/tests/pt/model/test_sezm_export.py +++ b/source/tests/pt/model/test_sezm_export.py @@ -1206,6 +1206,63 @@ 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()