diff --git a/dpdata/system.py b/dpdata/system.py index 36a01111..2ff210fd 100644 --- a/dpdata/system.py +++ b/dpdata/system.py @@ -1389,16 +1389,17 @@ def __init__(self, *systems, type_map=None): def from_fmt_obj( self, fmtobj: Format, directory, labeled: bool = True, **kwargs: Any ): - if not isinstance(fmtobj, dpdata.plugins.deepmd.DeePMDMixedFormat): - for dd in fmtobj.from_multi_systems(directory, **kwargs): - if labeled: - system = LabeledSystem().from_fmt_obj(fmtobj, dd, **kwargs) - else: - system = System().from_fmt_obj(fmtobj, dd, **kwargs) - system.sort_atom_names() - self.append(system) - return self - else: + try: + if not isinstance(fmtobj, dpdata.plugins.deepmd.DeePMDMixedFormat): + for dd in fmtobj.from_multi_systems(directory, **kwargs): + if labeled: + system = LabeledSystem().from_fmt_obj(fmtobj, dd, **kwargs) + else: + system = System().from_fmt_obj(fmtobj, dd, **kwargs) + system.sort_atom_names() + self.append(system) + return self + system_list = [] for dd in fmtobj.from_multi_systems(directory, **kwargs): if labeled: @@ -1411,6 +1412,16 @@ def from_fmt_obj( system_list.append(System(data=data_item, **kwargs)) self.append(*system_list) return self + except DataError as exc: + if ( + labeled + and isinstance(fmtobj, dpdata.plugins.deepmd.DeePMDMixedFormat) + and str(exc) == "energies not found in data" + ): + raise DataError( + f"{exc}. For coordinate-only mixed datasets, pass labeled=False." + ) from exc + raise def to_fmt_obj(self, fmtobj: Format, directory, *args: Any, **kwargs: Any): if not isinstance(fmtobj, dpdata.plugins.deepmd.DeePMDMixedFormat): @@ -1481,9 +1492,32 @@ def __add__(self, others): raise RuntimeError("Unspported data structure") @classmethod - def from_file(cls, file_name, fmt: str, **kwargs: Any): + def from_file( + cls, file_name, fmt: str, *, labeled: bool = True, **kwargs: Any + ) -> MultiSystems: + """Load multiple systems from a file or directory. + + Parameters + ---------- + file_name + Source accepted by the selected format backend. + fmt : str + Format identifier, such as ``"deepmd/npy/mixed"``. + labeled : bool, default=True + Load :class:`LabeledSystem` objects when true. Set this to false for + coordinate-only data that does not contain energies or forces. + **kwargs + Additional arguments forwarded to the format backend. + + Returns + ------- + MultiSystems + Systems reconstructed from the source. + """ multi_systems = cls() - multi_systems.load_systems_from_file(file_name=file_name, fmt=fmt, **kwargs) + multi_systems.load_systems_from_file( + file_name=file_name, fmt=fmt, labeled=labeled, **kwargs + ) return multi_systems @classmethod @@ -1506,10 +1540,22 @@ def from_dir( ) return multi_systems - def load_systems_from_file(self, file_name=None, fmt: str | None = None, **kwargs): + def load_systems_from_file( + self, + file_name=None, + fmt: str | None = None, + *, + labeled: bool = True, + **kwargs, + ): + """Load systems into this collection. + + ``labeled=False`` selects regular :class:`System` objects, which is + required for DeepMD datasets that omit label arrays. + """ assert fmt is not None fmt = fmt.lower() - return self.from_fmt_obj(load_format(fmt), file_name, **kwargs) + return self.from_fmt_obj(load_format(fmt), file_name, labeled=labeled, **kwargs) def get_nframes(self) -> int: """Returns number of frames in all systems.""" diff --git a/tests/test_deepmd_mixed.py b/tests/test_deepmd_mixed.py index d5b0dec6..d83eaee3 100644 --- a/tests/test_deepmd_mixed.py +++ b/tests/test_deepmd_mixed.py @@ -2,8 +2,10 @@ import os import shutil +import tempfile import unittest from glob import glob +from inspect import Parameter, signature import numpy as np from comp_sys import ( @@ -17,8 +19,68 @@ from dpdata.data_type import ( Axis, + DataError, DataType, ) +from dpdata.plugins.deepmd import DeePMDMixedFormat + + +class TestMixedMultiSystemsUnlabeled(unittest.TestCase): + """Regression coverage for loading coordinate-only mixed datasets.""" + + def test_from_file_unlabeled_round_trip(self): + labeled_parameter = signature(dpdata.MultiSystems.from_file).parameters[ + "labeled" + ] + self.assertIs(labeled_parameter.kind, Parameter.KEYWORD_ONLY) + self.assertIs(labeled_parameter.default, True) + + system = dpdata.System("poscars/POSCAR.h2o.md", fmt="vasp/poscar") + + with tempfile.TemporaryDirectory() as tmpdir: + mixed_dir = os.path.join(tmpdir, "mixed") + dpdata.MultiSystems(system).to("deepmd/npy/mixed", mixed_dir) + + # The default remains labeled for backward compatibility, but the + # error now points coordinate-only users to the public flag. + with self.assertRaisesRegex(DataError, "pass labeled=False"): + dpdata.MultiSystems.from_file( + mixed_dir, + fmt="deepmd/npy/mixed", + ) + + with self.assertRaisesRegex(DataError, "energies not found in data\\. For"): + dpdata.MultiSystems().from_fmt_obj( + DeePMDMixedFormat(), mixed_dir, labeled=True + ) + + # The explicit, discoverable flag selects ordinary System objects. + systems = dpdata.MultiSystems.from_file( + mixed_dir, fmt="deepmd/npy/mixed", labeled=False + ) + + self.assertEqual(len(systems), 1) + loaded = next(iter(systems.systems.values())) + self.assertIs(type(loaded), dpdata.System) + self.assertNotIn("energies", loaded.data) + np.testing.assert_array_equal( + loaded.data["atom_types"], system.data["atom_types"] + ) + np.testing.assert_allclose(loaded.data["cells"], system.data["cells"]) + np.testing.assert_allclose(loaded.data["coords"], system.data["coords"]) + + def test_unrelated_mixed_data_error_does_not_suggest_unlabeled_loading(self): + class CorruptMixedFormat(DeePMDMixedFormat): + def from_multi_systems(self, directory, **kwargs): + yield directory + + def from_labeled_system_mix(self, file_name, type_map=None, **kwargs): + raise DataError("force.npy is truncated") + + with self.assertRaisesRegex(DataError, "force.npy is truncated") as caught: + dpdata.MultiSystems().from_fmt_obj(CorruptMixedFormat(), "mixed") + + self.assertNotIn("labeled=False", str(caught.exception)) class TestMixedMultiSystemsDumpLoad(