Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
74 changes: 60 additions & 14 deletions dpdata/system.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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):
Expand Down Expand Up @@ -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
Expand All @@ -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."""
Expand Down
62 changes: 62 additions & 0 deletions tests/test_deepmd_mixed.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand All @@ -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
Comment thread
njzjz-bot marked this conversation as resolved.
)

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(
Expand Down
Loading