diff --git a/deepmd/dpmodel/utils/lmdb_data.py b/deepmd/dpmodel/utils/lmdb_data.py index 0396640a4a..759ddf3bc9 100644 --- a/deepmd/dpmodel/utils/lmdb_data.py +++ b/deepmd/dpmodel/utils/lmdb_data.py @@ -393,7 +393,9 @@ def _raw_frame_availability( explicit_true |= bit else: bit = key_bits.get(name) - if bit is not None: + if bit is not None and ( + name != "box" or not np.allclose(_decode_value(value), 0.0) + ): present |= bit return present & (~explicit_known | explicit_true) @@ -737,6 +739,11 @@ def _frame_source_available(frame: dict[str, Any], key: str) -> bool: source_present = _is_encoded_array(value) or isinstance( value, (np.ndarray, np.generic, int, float, bool) ) + if source_present and key == "box": + # A zero cell is DeePMD's non-periodic sentinel, not an available + # periodic box. Treating it as present would mix PBC and non-PBC + # frames under one scalar find_box flag. + source_present = not np.allclose(_decode_value(value), 0.0) find_key = f"find_{key}" if find_key in frame: return source_present and bool( @@ -1970,8 +1977,17 @@ def __init__( self, lmdb_path: str, type_map: list[str], - batch_size: int | str = "auto", + batch_size: int | str | Sequence[int | str] = "auto", ) -> None: + if isinstance(batch_size, (Sequence, np.ndarray)) and not isinstance( + batch_size, str + ): + if len(batch_size) != 1: + raise ValueError( + "One LMDB path is one training dataset and therefore " + "requires exactly one batch_size value." + ) + batch_size = batch_size[0] self.lmdb_path = str(Path(lmdb_path).resolve()) self._type_map = type_map # Read before opening the frame-serving environment, which disables @@ -2235,14 +2251,8 @@ def get_batch_size_for_nloc(self, nloc: int) -> int: def __len__(self) -> int: return self.nframes - def __getitem__(self, index: int) -> dict[str, Any]: - """Read frame from LMDB, decode, remap keys, return dict of numpy arrays. - - ``index`` is a dataset-level index in ``[0, len(self))``. Under - ``filter:N`` the LMDB key space may have gaps (dropped frames), so - we translate through ``self._retained_keys`` before hitting LMDB. - """ - self._data_requirements_frozen = True + def _read_frame(self, index: int) -> dict[str, Any]: + """Decode one frame without changing requirement-registration state.""" if index < 0 or index >= self.nframes: raise IndexError(f"dataset index {index} out of range [0, {self.nframes})") original_key = int(self._retained_keys[index]) @@ -2259,6 +2269,25 @@ def __getitem__(self, index: int) -> dict[str, Any]: copy_arrays=True, ) + def peek_frame(self, index: int) -> dict[str, Any]: + """Inspect one frame without freezing later requirement registration. + + Structural probes such as periodic-boundary detection happen before a + model supplies its label requirements. They may inspect a frame, but + must not turn that inspection into the first training read. + """ + return self._read_frame(index) + + def __getitem__(self, index: int) -> dict[str, Any]: + """Read and decode one frame, freezing the registered data contract. + + ``index`` is a dataset-level index in ``[0, len(self))``. Under + ``filter:N`` the LMDB key space may have gaps, which :meth:`_read_frame` + translates through ``self._retained_keys``. + """ + self._data_requirements_frozen = True + return self._read_frame(index) + def original_keys(self, indices: Sequence[int]) -> list[int]: """Translate dataset indices to original integer LMDB keys.""" keys: list[int] = [] @@ -2854,17 +2883,17 @@ def compute_block_targets( Each element is ``(system_indices_in_block, target_frame_count)``. Returns empty list if no expansion is needed (all targets == actual). """ - from deepmd.utils.data_system import ( - prob_sys_size_ext, - ) - - # Parse block definitions from the auto_prob string - # Format: "prob_sys_size;stt:end:weight;stt:end:weight;..." - block_str = auto_prob_style.split(";")[1:] + # ``prob_uniform`` is one equal-weight block per original system. The + # extended ``prob_sys_size`` form names arbitrary ranges explicitly. blocks: list[tuple[int, int, float]] = [] - for part in block_str: - stt, end, weight = part.split(":") - blocks.append((int(stt), int(end), float(weight))) + if auto_prob_style == "prob_uniform": + blocks = [(system_id, system_id + 1, 1.0) for system_id in range(nsystems)] + elif auto_prob_style.startswith("prob_sys_size"): + for part in auto_prob_style.split(";")[1:]: + stt, end, weight = part.split(":") + blocks.append((int(stt), int(end), float(weight))) + else: + raise RuntimeError(f"Unknown auto prob style: {auto_prob_style}") # A bare ``prob_sys_size`` names no blocks: it asks for a probability # proportional to system size, which is what sampling the merged frames @@ -2905,13 +2934,29 @@ def compute_block_targets( f"0 frames, likely after filter:N): {dropped}. Remaining block " "weights will be renormalised to sum to 1.0." ) - auto_prob_style = "prob_sys_size;" + ";".join( - f"{stt}:{end}:{weight}" for stt, end, weight in nonempty - ) blocks = nonempty - # Compute per-system probabilities using the standard function - sys_probs = prob_sys_size_ext(auto_prob_style, nsystems, system_nframes) + # Compute the same per-system probabilities as prob_sys_size_ext locally. + # Keeping this framework-agnostic LMDB module independent of data_system + # avoids an import cycle when the legacy adapter imports the LMDB reader. + block_weights = np.asarray([weight for _, _, weight in blocks], dtype=float) + if not np.all(np.isfinite(block_weights)): + raise ValueError("block weights must be finite") + if np.any(block_weights < 0): + raise ValueError("the weight of a block should be no less than 0") + total_block_weight = np.sum(block_weights) + if total_block_weight <= 0: + raise ValueError("the sum of block weights should be greater than 0") + block_probs = block_weights / total_block_weight + sys_probs = np.zeros(nsystems, dtype=np.float64) + for block_idx, (stt, end, _weight) in enumerate(blocks): + block_frames = np.asarray(system_nframes[stt:end], dtype=float) + total_block_frames = np.sum(block_frames) + if total_block_frames <= 0: + raise ValueError( + f"block {stt}:{end} must contain at least one retained frame" + ) + sys_probs[stt:end] = block_frames / total_block_frames * block_probs[block_idx] # Group systems by block, compute block-level frames and prob block_info: list[tuple[list[int], int, float]] = [] # (sys_ids, frames, prob) @@ -3673,22 +3718,29 @@ def make_neighbor_stat_data( ) reader = LmdbDataReader(lmdb_path, type_map=type_map) - nframes = len(reader) - rng = np.random.RandomState(42) - if nframes > max_frames: - indices = np.sort(rng.choice(nframes, max_frames, replace=False)) - else: - indices = np.arange(nframes, dtype=np.int64) - - # Read sampled frames, group by nloc - nloc_frames: dict[int, list[tuple[np.ndarray, np.ndarray, np.ndarray | None]]] = {} - for idx in indices: - frame = reader[int(idx)] - atype = frame["atype"] - nloc = len(atype) - nloc_frames.setdefault(nloc, []).append( - (frame["coord"], atype, frame.get("box")) - ) + try: + nframes = len(reader) + rng = np.random.RandomState(42) + if nframes > max_frames: + indices = np.sort(rng.choice(nframes, max_frames, replace=False)) + else: + indices = np.arange(nframes, dtype=np.int64) + + # The copied arrays remain valid after the reader is closed, so this + # helper does not leave an LMDB transaction or mmap alive in callers. + nloc_frames: dict[ + int, list[tuple[np.ndarray, np.ndarray, np.ndarray | None]] + ] = {} + for idx in indices: + frame = reader[int(idx)] + atype = frame["atype"] + nloc = len(atype) + nloc_frames.setdefault(nloc, []).append( + (frame["coord"], atype, frame.get("box")) + ) + ntypes = len(type_map) if type_map else reader._ntypes + finally: + reader.close() # Build per-nloc data_system proxies data_systems = [] @@ -3710,7 +3762,6 @@ def make_neighbor_stat_data( data_systems.append(proxy) system_dirs.append(label) - ntypes = len(type_map) if type_map else reader._ntypes return SimpleNamespace( system_dirs=system_dirs, data_systems=data_systems, @@ -3888,19 +3939,17 @@ def _read_frames(self, frame_indices: Sequence[int]) -> list[dict[str, Any]]: ) return frames - def __del__(self) -> None: - """Release the LMDB environment ref-count on garbage collection. - - The count is released only once, and only if construction got as far - as taking it: an instance that failed earlier holds no reference, and - releasing one it never took would close the environment underneath - whichever reader does hold it. - """ + def close(self) -> None: + """Release the LMDB environment ref-count idempotently.""" if getattr(self, "_env", None) is None: return self._env = None _close_lmdb(self.lmdb_path) + def __del__(self) -> None: + """Release the LMDB environment ref-count on garbage collection.""" + self.close() + @property def nloc_groups(self) -> dict[int, np.ndarray]: """Nloc → the LMDB frame indices retained for that atom count.""" @@ -4263,14 +4312,33 @@ def __init__( lmdb_test_data: "LmdbTestData", nloc: int, frame_indices: Sequence[int] | None = None, + *, + pbc: bool | None = None, + stat_groups: dict[str, Sequence[int]] | None = None, ) -> None: self._inner = lmdb_test_data self._nloc = nloc self._frame_indices = frame_indices + self._pbc = pbc + self._stat_groups = stat_groups or {} + self.dirs = list(self._stat_groups) def __getattr__(self, name: str) -> Any: return getattr(self._inner, name) + @property + def pbc(self) -> bool: + """Whether every frame represented by this view is periodic.""" + return self._inner.pbc if self._pbc is None else self._pbc + + def get_natoms(self) -> int: + """Return the fixed atom count of this stack-compatible view.""" + return self._nloc + + def _load_set(self, set_name: str) -> dict[str, Any]: + """Load one bounded neighbor-stat chunk from this view.""" + return self._inner.get_test_by_indices(self._stat_groups[str(set_name)]) + def get_test(self) -> dict[str, Any]: if self._frame_indices is not None: return self._inner.get_test_by_indices(self._frame_indices) diff --git a/deepmd/entrypoints/test.py b/deepmd/entrypoints/test.py index d9972b967e..c08751cbf4 100644 --- a/deepmd/entrypoints/test.py +++ b/deepmd/entrypoints/test.py @@ -111,7 +111,12 @@ def test( else: systems = [str((root / Path(ss)).resolve()) for ss in systems] patterns = data_params.get("rglob_patterns", None) - all_sys = process_systems(systems, patterns=patterns) + all_sys = process_systems( + systems, + patterns=patterns, + fmt=data_params.get("format"), + out_fmt=data_params.get("out_format", data_params.get("output_format")), + ) elif valid_json is not None: jdata = j_loader(valid_json) jdata = update_deepmd_input(jdata) @@ -125,7 +130,12 @@ def test( else: systems = [str((root / Path(ss)).resolve()) for ss in systems] patterns = data_params.get("rglob_patterns", None) - all_sys = process_systems(systems, patterns=patterns) + all_sys = process_systems( + systems, + patterns=patterns, + fmt=data_params.get("format"), + out_fmt=data_params.get("out_format", data_params.get("output_format")), + ) elif datafile is not None: with open(datafile) as datalist: all_sys = datalist.read().splitlines() diff --git a/deepmd/jax/entrypoints/train.py b/deepmd/jax/entrypoints/train.py index 61d415a340..1cad2af5bc 100644 --- a/deepmd/jax/entrypoints/train.py +++ b/deepmd/jax/entrypoints/train.py @@ -36,6 +36,7 @@ ) from deepmd.utils import random as dp_random from deepmd.utils.data_system import ( + close_data_systems, get_data, ) from deepmd.utils.summary import SummaryPrinter as BaseSummaryPrinter @@ -201,11 +202,14 @@ def factory( train_data_map, valid_data_map, _ = make_task_maps(config, factory) print_data_summaries(train_data_map, valid_data_map) - start_time = time.time() - model.train(train_data_map, valid_data_map) - end_time = time.time() - log.info("finished training") - log.info(f"wall time: {(end_time - start_time):.3f} s") + try: + start_time = time.time() + model.train(train_data_map, valid_data_map) + end_time = time.time() + log.info("finished training") + log.info(f"wall time: {(end_time - start_time):.3f} s") + finally: + close_data_systems(train_data_map, valid_data_map) def train( @@ -296,9 +300,12 @@ def update_sel( type_map, None, # not used ) - updated_model, task_min_nbor_dist = BaseModel.update_sel( - train_data, type_map, dict(task_config.model_params) - ) + try: + updated_model, task_min_nbor_dist = BaseModel.update_sel( + train_data, type_map, dict(task_config.model_params) + ) + finally: + close_data_systems(train_data) if multi_task: jdata_cpy["model"]["model_dict"][task_config.key] = updated_model min_nbor_dist[task_config.key] = task_min_nbor_dist diff --git a/deepmd/pd/entrypoints/main.py b/deepmd/pd/entrypoints/main.py index acd49c589c..60698416a7 100644 --- a/deepmd/pd/entrypoints/main.py +++ b/deepmd/pd/entrypoints/main.py @@ -69,6 +69,8 @@ from deepmd.utils.data_system import ( get_data, process_systems, + validate_backend_data_config, + validate_lmdb_systems, ) from deepmd.utils.path import ( DPPath, @@ -108,11 +110,39 @@ def prepare_trainer_input_single( validation_dataset_params["systems"] if validation_dataset_params else None ) training_systems = training_dataset_params["systems"] + validate_backend_data_config( + training_dataset_params, + backend_name="Paddle", + lmdb_supported=False, + ) trn_patterns = training_dataset_params.get("rglob_patterns", None) - training_systems = process_systems(training_systems, patterns=trn_patterns) + training_systems = process_systems( + training_systems, + patterns=trn_patterns, + fmt=training_dataset_params.get("format", None), + out_fmt=training_dataset_params.get( + "out_format", training_dataset_params.get("output_format", None) + ), + ) + validate_lmdb_systems(training_systems, backend_name="Paddle", supported=False) if validation_systems is not None: + validate_backend_data_config( + validation_dataset_params, + backend_name="Paddle", + lmdb_supported=False, + ) val_patterns = validation_dataset_params.get("rglob_patterns", None) - validation_systems = process_systems(validation_systems, val_patterns) + validation_systems = process_systems( + validation_systems, + val_patterns, + fmt=validation_dataset_params.get("format", None), + out_fmt=validation_dataset_params.get( + "out_format", validation_dataset_params.get("output_format", None) + ), + ) + validate_lmdb_systems( + validation_systems, backend_name="Paddle", supported=False + ) # stat files stat_file_path_single = data_dict_single.get("stat_file") diff --git a/deepmd/pt/entrypoints/main.py b/deepmd/pt/entrypoints/main.py index 9e99a7f52f..08a3a183c3 100644 --- a/deepmd/pt/entrypoints/main.py +++ b/deepmd/pt/entrypoints/main.py @@ -88,8 +88,11 @@ update_deepmd_input, ) from deepmd.utils.data_system import ( + conversion_will_write_lmdb, get_data, process_systems, + validate_lmdb_sampling_options, + validate_lmdb_systems, ) from deepmd.utils.stat_file import ( StatFileSpec, @@ -176,10 +179,28 @@ def prepare_trainer_input_single( def _make_dp_loader_set( systems: str | list[str], dataset_params: dict[str, Any], - ) -> DpLoaderSet: - """Create a DpLoaderSet from systems with pattern expansion.""" + ) -> DpLoaderSet | LmdbDataset: + """Create a dataset from systems with pattern expansion/conversion.""" + if conversion_will_write_lmdb(dataset_params): + validate_lmdb_sampling_options(dataset_params) patterns = dataset_params.get("rglob_patterns") - systems = process_systems(systems, patterns=patterns) + systems = process_systems( + systems, + patterns=patterns, + fmt=dataset_params.get("format"), + out_fmt=dataset_params.get( + "out_format", dataset_params.get("output_format") + ), + ) + lmdb_path = validate_lmdb_systems(systems, backend_name="PyTorch") + if lmdb_path is not None: + validate_lmdb_sampling_options(dataset_params) + return LmdbDataset( + lmdb_path, + model_params_single["type_map"], + dataset_params["batch_size"], + auto_prob_style=dataset_params.get("auto_prob"), + ) return DpLoaderSet( systems, dataset_params["batch_size"], @@ -189,7 +210,12 @@ def _make_dp_loader_set( ) # LMDB path: single string → LmdbDataset - if isinstance(training_systems, str) and is_lmdb(training_systems): + if ( + training_dataset_params.get("format", None) is None + and isinstance(training_systems, str) + and is_lmdb(training_systems) + ): + validate_lmdb_sampling_options(training_dataset_params) auto_prob = training_dataset_params.get("auto_prob", None) train_data_single = LmdbDataset( training_systems, @@ -199,9 +225,11 @@ def _make_dp_loader_set( ) if ( validation_systems is not None + and validation_dataset_params.get("format", None) is None and isinstance(validation_systems, str) and is_lmdb(validation_systems) ): + validate_lmdb_sampling_options(validation_dataset_params) validation_data_single = LmdbDataset( validation_systems, model_params_single["type_map"], @@ -390,23 +418,44 @@ def train( "Calculate neighbor statistics... (add --skip-neighbor-stat to skip this step)" ) - if not multi_task: - type_map = config["model"].get("type_map") - training_systems = config["training"]["training_data"].get("systems") - if ( - training_systems is not None + def _get_neighbor_stat_data_from_params( + dataset_params: dict[str, Any], + type_map: list[str] | None, + ) -> Any: + training_systems = dataset_params.get("systems") + direct_lmdb = ( + dataset_params.get("format") is None + and training_systems is not None and isinstance(training_systems, str) and is_lmdb(training_systems) - ): + ) + if direct_lmdb or conversion_will_write_lmdb(dataset_params): + validate_lmdb_sampling_options(dataset_params) + if direct_lmdb: + systems = [training_systems] + else: + systems = process_systems( + training_systems, + patterns=dataset_params.get("rglob_patterns"), + fmt=dataset_params.get("format"), + out_fmt=dataset_params.get( + "out_format", dataset_params.get("output_format") + ), + ) + lmdb_path = validate_lmdb_systems(systems, backend_name="PyTorch") + if lmdb_path is not None: from deepmd.dpmodel.utils.lmdb_data import ( make_neighbor_stat_data, ) - train_data = make_neighbor_stat_data(training_systems, type_map) - else: - train_data = get_data( - config["training"]["training_data"], 0, type_map, None - ) + return make_neighbor_stat_data(lmdb_path, type_map) + return get_data(dataset_params, 0, type_map, None) + + if not multi_task: + type_map = config["model"].get("type_map") + train_data = _get_neighbor_stat_data_from_params( + config["training"]["training_data"], type_map + ) config["model"], min_nbor_dist = BaseModel.update_sel( train_data, type_map, config["model"] ) @@ -414,26 +463,10 @@ def train( min_nbor_dist = {} for model_item in config["model"]["model_dict"]: type_map = config["model"]["model_dict"][model_item].get("type_map") - training_systems = config["training"]["data_dict"][model_item][ - "training_data" - ].get("systems") - if ( - training_systems is not None - and isinstance(training_systems, str) - and is_lmdb(training_systems) - ): - from deepmd.dpmodel.utils.lmdb_data import ( - make_neighbor_stat_data, - ) - - train_data = make_neighbor_stat_data(training_systems, type_map) - else: - train_data = get_data( - config["training"]["data_dict"][model_item]["training_data"], - 0, - type_map, - None, - ) + train_data = _get_neighbor_stat_data_from_params( + config["training"]["data_dict"][model_item]["training_data"], + type_map, + ) config["model"]["model_dict"][model_item], min_nbor_dist[model_item] = ( BaseModel.update_sel( train_data, type_map, config["model"]["model_dict"][model_item] diff --git a/deepmd/pt_expt/entrypoints/main.py b/deepmd/pt_expt/entrypoints/main.py index 5a5f8d4e06..1381750e0d 100644 --- a/deepmd/pt_expt/entrypoints/main.py +++ b/deepmd/pt_expt/entrypoints/main.py @@ -33,8 +33,11 @@ ) from deepmd.utils.data_system import ( DeepmdDataSystem, + conversion_will_write_lmdb, get_data, process_systems, + validate_lmdb_sampling_options, + validate_lmdb_systems, ) from deepmd.utils.stat_file import ( StatFileSpec, @@ -116,13 +119,35 @@ def _get_neighbor_stat_data( ``make_neighbor_stat_data``; falls back to the legacy ``get_data`` for npy/HDF5 directories. """ - lmdb_path = _detect_lmdb_path(dataset_params.get("systems")) + lmdb_path = ( + None + if dataset_params.get("format") is not None + else _detect_lmdb_path(dataset_params.get("systems")) + ) if lmdb_path is not None: + validate_lmdb_sampling_options(dataset_params) from deepmd.dpmodel.utils.lmdb_data import ( make_neighbor_stat_data, ) return make_neighbor_stat_data(lmdb_path, type_map) + if conversion_will_write_lmdb(dataset_params): + validate_lmdb_sampling_options(dataset_params) + systems = process_systems( + dataset_params["systems"], + patterns=dataset_params.get("rglob_patterns"), + fmt=dataset_params.get("format"), + out_fmt=dataset_params.get("out_format", dataset_params.get("output_format")), + ) + converted_lmdb_path = validate_lmdb_systems( + systems, backend_name="PyTorch exportable" + ) + if converted_lmdb_path is not None: + from deepmd.dpmodel.utils.lmdb_data import ( + make_neighbor_stat_data, + ) + + return make_neighbor_stat_data(converted_lmdb_path, type_map) return get_data(dataset_params, 0, type_map, None) @@ -142,8 +167,13 @@ def _build_data_system( systems. """ systems_raw = dataset_params["systems"] - lmdb_path = _detect_lmdb_path(systems_raw) + lmdb_path = ( + None + if dataset_params.get("format") is not None + else _detect_lmdb_path(systems_raw) + ) if lmdb_path is not None: + validate_lmdb_sampling_options(dataset_params) return LmdbDataSystem( lmdb_path=lmdb_path, type_map=type_map, @@ -153,10 +183,28 @@ def _build_data_system( rank=rank, world_size=world_size, ) + if conversion_will_write_lmdb(dataset_params): + validate_lmdb_sampling_options(dataset_params) systems = process_systems( systems_raw, patterns=dataset_params.get("rglob_patterns"), + fmt=dataset_params.get("format"), + out_fmt=dataset_params.get("out_format", dataset_params.get("output_format")), ) + converted_lmdb_path = validate_lmdb_systems( + systems, backend_name="PyTorch exportable" + ) + if converted_lmdb_path is not None: + validate_lmdb_sampling_options(dataset_params) + return LmdbDataSystem( + lmdb_path=converted_lmdb_path, + type_map=type_map, + batch_size=dataset_params["batch_size"], + auto_prob_style=dataset_params.get("auto_prob"), + seed=seed, + rank=rank, + world_size=world_size, + ) return DeepmdDataSystem( systems=systems, batch_size=dataset_params["batch_size"], diff --git a/deepmd/tf/entrypoints/train.py b/deepmd/tf/entrypoints/train.py index d7b40ebfee..2dc1feba9f 100755 --- a/deepmd/tf/entrypoints/train.py +++ b/deepmd/tf/entrypoints/train.py @@ -52,6 +52,7 @@ replace_model_params_with_pretrained_model, ) from deepmd.utils.data_system import ( + close_data_systems, get_data, ) from deepmd.utils.path import ( @@ -308,26 +309,31 @@ def _do_work( if ( origin_type_map is not None and not origin_type_map ): # get the type_map from data if not provided - origin_type_map = get_data( - jdata["training"]["training_data"], rcut, None, modifier - ).get_type_map() - model.build( - train_data, - stop_batch, - origin_type_map=origin_type_map, - stat_file_path=stat_file_path, - ) + origin_data = get_data(jdata["training"]["training_data"], rcut, None, modifier) + try: + origin_type_map = origin_data.get_type_map() + finally: + close_data_systems(origin_data) + try: + model.build( + train_data, + stop_batch, + origin_type_map=origin_type_map, + stat_file_path=stat_file_path, + ) - if not is_compress: - # train the model with the provided systems in a cyclic way - start_time = time.time() - model.train(train_data, valid_data) - end_time = time.time() - log.info("finished training") - log.info(f"wall time: {(end_time - start_time):.3f} s") - else: - model.save_compressed() - log.info("finished compressing") + if not is_compress: + # train the model with the provided systems in a cyclic way + start_time = time.time() + model.train(train_data, valid_data) + end_time = time.time() + log.info("finished training") + log.info(f"wall time: {(end_time - start_time):.3f} s") + else: + model.save_compressed() + log.info("finished compressing") + finally: + close_data_systems(train_data, valid_data) def get_modifier(modi_data: dict | None = None) -> BaseModifier | None: @@ -355,9 +361,12 @@ def update_sel(jdata: dict) -> dict: type_map, None, # not used ) - jdata_cpy["model"], min_nbor_dist = Model.update_sel( - train_data, type_map, jdata["model"] - ) + try: + jdata_cpy["model"], min_nbor_dist = Model.update_sel( + train_data, type_map, jdata["model"] + ) + finally: + close_data_systems(train_data) if min_nbor_dist is not None: tf.constant( diff --git a/deepmd/tf2/entrypoints/train.py b/deepmd/tf2/entrypoints/train.py index a026bcc03f..941bca8984 100644 --- a/deepmd/tf2/entrypoints/train.py +++ b/deepmd/tf2/entrypoints/train.py @@ -34,6 +34,7 @@ ) from deepmd.utils import random as dp_random from deepmd.utils.data_system import ( + close_data_systems, get_data, ) from deepmd.utils.summary import SummaryPrinter as BaseSummaryPrinter @@ -192,11 +193,14 @@ def factory( shared_links=self.shared_links, min_nbor_dist=neighbor_stat, ) - start_time = time.time() - trainer.run() - end_time = time.time() - log.info("finished training") - log.info("wall time: %.3f s", end_time - start_time) + try: + start_time = time.time() + trainer.run() + end_time = time.time() + log.info("finished training") + log.info("wall time: %.3f s", end_time - start_time) + finally: + close_data_systems(train_data_map, valid_data_map) def train( @@ -249,11 +253,14 @@ def update_sel( type_map, None, ) - updated_model, task_min_nbor_dist = BaseModel.update_sel( - train_data, - type_map, - dict(task_config.model_params), - ) + try: + updated_model, task_min_nbor_dist = BaseModel.update_sel( + train_data, + type_map, + dict(task_config.model_params), + ) + finally: + close_data_systems(train_data) min_nbor_dist[task_config.key] = task_min_nbor_dist if multi_task: jdata_cpy["model"]["model_dict"][task_config.key] = updated_model diff --git a/deepmd/utils/argcheck.py b/deepmd/utils/argcheck.py index f0132348f1..8b9ac03a80 100644 --- a/deepmd/utils/argcheck.py +++ b/deepmd/utils/argcheck.py @@ -5366,6 +5366,19 @@ def training_data_args() -> list[ doc_patterns = ( "The customized patterns used in `rglob` to collect all training systems. " ) + doc_format = ( + "The input data format passed to dpdata for automatic conversion. " + "If this key is not set, `systems` must already point to DeePMD data. " + "If this key is set to a non-DeePMD format, each selected input path is " + "loaded by dpdata and converted before training. Use dpdata format names " + "such as `extxyz`, `ase/structure`, `ase/traj`, or `auto`." + ) + doc_out_format = ( + "The output data format passed to dpdata for automatic conversion. " + "When `format` requests conversion from a non-DeePMD format, this key " + "defaults to `deepmd/lmdb`. Use a DeePMD format supported by dpdata, " + "such as `deepmd/lmdb`, `deepmd/hdf5`, or `deepmd/npy`." + ) doc_batch_size = f'This key can be \n\n\ - list: the length of which is the same as the {link_sys}. The batch size of each system is given by the elements of the list.\n\n\ - int: all {link_sys} use the same batch size.\n\n\ @@ -5405,6 +5418,20 @@ def training_data_args() -> list[ doc=supported_backends("tf", "pt", "jax", "pd", "pt_expt", "tf2") + doc_patterns, ), + Argument( + "format", + [str, None], + optional=True, + doc=doc_format, + ), + Argument( + "out_format", + [str, None], + optional=True, + default="deepmd/lmdb", + doc=doc_out_format, + alias=["output_format"], + ), Argument( "batch_size", [list[int], int, str], @@ -5464,6 +5491,19 @@ def validation_data_args() -> list[ doc_patterns = ( "The customized patterns used in `rglob` to collect all validation systems. " ) + doc_format = ( + "The input data format passed to dpdata for automatic conversion. " + "If this key is not set, `systems` must already point to DeePMD data. " + "If this key is set to a non-DeePMD format, each selected input path is " + "loaded by dpdata and converted before validation. Use dpdata format names " + "such as `extxyz`, `ase/structure`, `ase/traj`, or `auto`." + ) + doc_out_format = ( + "The output data format passed to dpdata for automatic conversion. " + "When `format` requests conversion from a non-DeePMD format, this key " + "defaults to `deepmd/lmdb`. Use a DeePMD format supported by dpdata, " + "such as `deepmd/lmdb`, `deepmd/hdf5`, or `deepmd/npy`." + ) doc_batch_size = f'This key can be \n\n\ - list: the length of which is the same as the {link_sys}. The batch size of each system is given by the elements of the list.\n\n\ - int: all {link_sys} use the same batch size.\n\n\ @@ -5491,6 +5531,20 @@ def validation_data_args() -> list[ doc=supported_backends("tf", "pt", "jax", "pd", "pt_expt", "tf2") + doc_patterns, ), + Argument( + "format", + [str, None], + optional=True, + doc=doc_format, + ), + Argument( + "out_format", + [str, None], + optional=True, + default="deepmd/lmdb", + doc=doc_out_format, + alias=["output_format"], + ), Argument( "batch_size", [list[int], int, str], diff --git a/deepmd/utils/data_system.py b/deepmd/utils/data_system.py index cea270bd37..90ec9b31da 100644 --- a/deepmd/utils/data_system.py +++ b/deepmd/utils/data_system.py @@ -1,12 +1,24 @@ # SPDX-License-Identifier: LGPL-3.0-or-later import collections +import hashlib +import importlib.metadata +import json import logging +import os +import shutil +import socket +import threading +import time import warnings from functools import ( cached_property, ) +from pathlib import ( + Path, +) from typing import ( Any, + Self, ) import numpy as np @@ -30,6 +42,51 @@ log = logging.getLogger(__name__) +_DPDATA_CACHE_DIR = ".deepmd_dpdata_cache" +_DPDATA_DEFAULT_OUT_FORMAT = "deepmd/lmdb" +_DPDATA_CONVERSION_SCHEMA_VERSION = "2" +_DPDATA_CONVERSION_CACHE: dict[tuple[str, str, str, str, str], list[str]] = {} +_DPDATA_SOURCE_MTIME_CACHE: dict[tuple[str, str], tuple[float, float]] = {} +# Neighbor-stat, trainer construction, and multi-task routing commonly query +# the same directory within one startup. Reuse that O(file-count) scan while +# still letting a long-lived process notice later source rewrites. +_DPDATA_SOURCE_MTIME_CACHE_TTL = 60.0 +_CONVERSION_LOCK_HEARTBEAT_SECONDS = 5.0 +_CONVERSION_LOCK_STALE_SECONDS = 30.0 + + +def validate_lmdb_systems( + systems: list[str], + *, + backend_name: str, + supported: bool = True, +) -> str | None: + """Validate expanded systems and return the sole resolved LMDB path. + + LMDB stores multiple logical systems inside one database, so mixing an + LMDB path with other expanded paths is ambiguous and unsupported. + """ + # Import after data_system has initialized. Importing the dpmodel package + # at module load time re-enters this module through descriptor utilities. + from deepmd.dpmodel.utils.lmdb_data import ( + is_lmdb, + ) + + lmdb_paths = [path for path in systems if is_lmdb(path)] + if not lmdb_paths: + return None + if not supported: + raise NotImplementedError( + f"{backend_name} backend does not support LMDB data yet. " + "Choose out_format='deepmd/hdf5' for automatic conversion." + ) + if len(systems) != 1: + raise ValueError( + f"{backend_name} backend requires an LMDB dataset to resolve to " + "exactly one path; LMDB paths cannot be mixed with other systems." + ) + return lmdb_paths[0] + class DeepmdDataSystem: """Class for manipulating many data systems. @@ -701,6 +758,418 @@ def _check_type_map_consistency( return ret +class LmdbDataSystem: + """A DeepmdDataSystem-compatible adapter for LMDB datasets. + + The adapter returns raw DeePMD-style numpy batches (``type``, + ``natoms_vec``, ``default_mesh``) so it can be consumed by the legacy + TensorFlow/JAX training paths. Consumers that need the dpmodel canonical + format can still call ``normalize_batch`` on its output. + """ + + def __init__( + self, + lmdb_path: str, + type_map: list[str], + batch_size: int | str | list[int | str] = "auto", + auto_prob_style: str | None = None, + seed: int | None = None, + ) -> None: + # Keep the framework-agnostic LMDB implementation lazy so importing a + # legacy backend cannot create a data_system <-> dpmodel import cycle. + from deepmd.dpmodel.utils.lmdb_data import ( + LmdbBatchSampler, + LmdbDataReader, + LmdbTestData, + compute_block_targets, + ) + + if not type_map: + raise ValueError( + "LMDB datasets require a non-empty model/type_map because " + "LMDB stores atom type indices and the training data adapter " + "must map them to element names." + ) + + self.lmdb_path = lmdb_path + self._type_map = list(type_map) + self._closed = False + self._data_dict = { + "box": DataRequirementItem( + "box", + 9, + atomic=False, + must=False, + default=np.zeros(9, dtype=GLOBAL_NP_FLOAT_PRECISION), + ).dict, + "coord": { + "ndof": 3, + "atomic": True, + "must": True, + "high_prec": False, + "type_sel": None, + "repeat": 1, + "default": 0.0, + "dtype": None, + "output_natoms_for_type_sel": False, + }, + "numb_copy": { + "ndof": 1, + "atomic": False, + "must": False, + "high_prec": False, + "type_sel": None, + "repeat": 1, + "default": 1, + "dtype": int, + "output_natoms_for_type_sel": False, + }, + } + + self._reader = LmdbDataReader(lmdb_path, type_map, batch_size) + # Box availability is part of stack compatibility. Register it before + # any grouping so periodic and non-periodic frames never share the + # scalar ``find_box`` flag of one legacy batch. + box_requirement = DataRequirementItem( + "box", + 9, + atomic=False, + must=False, + default=np.zeros(9, dtype=GLOBAL_NP_FLOAT_PRECISION), + ) + self._reader.add_data_requirement([box_requirement]) + self._test_data = LmdbTestData( + lmdb_path, + type_map=type_map, + shuffle_test=False, + ) + self._test_data.add_data_requirement([box_requirement]) + # LMDB is defined as mixed-type by its reader contract; determining + # this must not scan every frame during data-system initialization. + self.mixed_type = self._reader.mixed_type + self.nsystems = 1 + self.natoms = [ + int(self._reader.frame_nlocs.max()) if len(self._reader.frame_nlocs) else 0 + ] + self.batch_size = [self._reader.batch_size] + self.sys_probs = [1.0] + + block_targets = None + if auto_prob_style is not None and self._reader.frame_system_ids is not None: + block_targets = compute_block_targets( + auto_prob_style, + self._reader.nsystems, + self._reader.system_nframes, + ) + self._sampler = LmdbBatchSampler( + self._reader, + shuffle=True, + seed=seed, + block_targets=block_targets, + ) + self.nbatches = [self._sampler.total_batches] + self._iter = iter(self._sampler) + self._refresh_groups() + + def _refresh_groups(self) -> None: + """Refresh bounded statistics chunks and full-validation views.""" + from deepmd.dpmodel.utils.lmdb_data import ( + LmdbTestDataNlocView, + collect_lmdb_sampling_groups, + ) + + groups = collect_lmdb_sampling_groups(self._reader) + self._stat_groups = groups + self._stat_offsets = [0] * len(groups) + + # Neighbor statistics are a bounded sample, matching the dedicated + # LMDB path. Chunks additionally cap decoded atoms so one large-nloc + # group cannot create a large transient Python/NumPy allocation. + selected = np.zeros(len(self._reader), dtype=bool) + max_frames = min(len(self._reader), 2000) + if max_frames: + rng = np.random.RandomState(42) + chosen = ( + rng.choice(len(self._reader), max_frames, replace=False) + if max_frames < len(self._reader) + else np.arange(len(self._reader), dtype=np.int64) + ) + selected[np.asarray(chosen, dtype=np.int64)] = True + + self._nloc_set_indices: dict[str, np.ndarray] = {} + data_systems = [] + system_dirs: list[str] = [] + any_pbc = False + for group_idx, (nloc, indices) in enumerate(groups): + group_label = f"{self.lmdb_path}#group={group_idx}:nloc={nloc}" + original_indices = np.asarray( + self._reader.original_keys(indices), dtype=np.int64 + ) + frame = self._reader.peek_frame(int(indices[0])) + group_pbc = bool(float(frame.get("find_box", 0.0)) > 0.5) + any_pbc = any_pbc or group_pbc + + stat_groups: dict[str, np.ndarray] = {} + sampled_indices = np.asarray(indices)[selected[np.asarray(indices)]] + chunk_size = max(1, min(128, 20000 // max(int(nloc), 1))) + for chunk_idx, start in enumerate( + range(0, len(sampled_indices), chunk_size) + ): + chunk = sampled_indices[start : start + chunk_size] + set_name = f"{group_label}:chunk={chunk_idx}" + self._nloc_set_indices[set_name] = np.asarray(chunk, dtype=np.int64) + stat_groups[set_name] = np.asarray( + self._reader.original_keys(chunk), dtype=np.int64 + ) + + data_systems.append( + LmdbTestDataNlocView( + self._test_data, + int(nloc), + original_indices, + pbc=group_pbc, + stat_groups=stat_groups, + ) + ) + system_dirs.append(group_label) + + # These views do not point back to this adapter, avoiding the + # ``data_systems=[self]`` reference cycle while satisfying both the + # neighbor-stat and JAX/TF2 full-validation contracts. + self.data_systems = data_systems + self.system_dirs = system_dirs + self.dirs = list(self._nloc_set_indices) + self.pbc = any_pbc + + def _detect_pbc(self) -> bool: + """Return True when LMDB frames contain a non-zero simulation box.""" + if len(self._reader) == 0: + return False + frame = self._reader.peek_frame(0) + return bool(float(frame.get("find_box", 0.0)) > 0.5) + + def add_data_requirements( + self, data_requirements: list[DataRequirementItem] + ) -> None: + """Add label/auxiliary data requirements.""" + self._reader.add_data_requirement(data_requirements) + self._test_data.add_data_requirement(data_requirements) + for item in data_requirements: + self._data_dict[item.key] = item.dict + self._refresh_groups() + self.nbatches = [self._sampler.total_batches] + self._iter = iter(self._sampler) + + def add_data_requirement(self, data_requirement: list[DataRequirementItem]) -> None: + """Alias used by DataLoader-style backends.""" + self.add_data_requirements(data_requirement) + + def add( + self, + key: str, + ndof: int, + atomic: bool = False, + must: bool = False, + high_prec: bool = False, + type_sel: list[int] | None = None, + repeat: int = 1, + default: float = 0.0, + dtype: np.dtype | None = None, + output_natoms_for_type_sel: bool = False, + ) -> None: + item = DataRequirementItem( + key, + ndof, + atomic=atomic, + must=must, + high_prec=high_prec, + type_sel=type_sel, + repeat=repeat, + default=default, + dtype=dtype, + output_natoms_for_type_sel=output_natoms_for_type_sel, + ) + self.add_data_requirements([item]) + + def get_data_dict(self, ii: int = 0) -> dict[str, dict[str, Any]]: + del ii + return self._data_dict + + def _load_set(self, set_name: str) -> dict[str, Any]: + """Load one bounded same-nloc chunk for legacy neighbor statistics.""" + indices = self._nloc_set_indices[str(set_name)] + return self._legacy_batch(self._reader.decode_batch(indices, ragged=False)) + + def _next_indices(self) -> list[int]: + try: + return next(self._iter) + except StopIteration: + self._iter = iter(self._sampler) + return next(self._iter) + + def _legacy_batch(self, batch: dict[str, Any]) -> dict[str, Any]: + """Translate a canonical LMDB batch to the legacy data-system shape.""" + coord = np.asarray(batch["coord"]) + nframes = coord.shape[0] + out: dict[str, Any] = {} + structural_keys = {"coord", "box"} + for key, value in batch.items(): + if key in { + "atype", + "natoms", + "real_natoms_vec", + "fid", + "sid", + "n_node", + }: + continue + if key.startswith("find_") and key[5:] not in self._data_dict: + continue + if ( + not key.startswith("find_") + and key not in structural_keys + and key not in self._data_dict + ): + continue + if value is None: + out[key] = None + continue + array = np.asarray(value) + data_info = self._data_dict.get(key) + if key == "coord" or ( + data_info is not None and data_info["atomic"] and array.ndim >= 3 + ): + array = array.reshape(nframes, -1) + out[key] = array + + atype = np.asarray(batch["atype"], dtype=np.int32) + real_natoms_vec = np.asarray( + batch.get("real_natoms_vec", batch["natoms"]), dtype=np.int32 + ) + if real_natoms_vec.ndim == 1: + real_natoms_vec = np.tile(real_natoms_vec, (nframes, 1)) + pad_nloc = int(atype.shape[1]) + natoms_vec = np.concatenate( + ( + np.array([pad_nloc, pad_nloc], dtype=np.int32), + real_natoms_vec[:, 2:].max(axis=0).astype(np.int32), + ) + ) + + out["type"] = atype + out["natoms_vec"] = natoms_vec + out["real_natoms_vec"] = real_natoms_vec + if "box" not in out or out["box"] is None: + out["box"] = np.zeros((nframes, 9), dtype=GLOBAL_NP_FLOAT_PRECISION) + out["find_box"] = np.float32(0.0) + elif "find_box" not in out: + out["find_box"] = np.float32(0.0 if np.allclose(out["box"], 0.0) else 1.0) + out.setdefault("find_coord", np.float32(1.0)) + if "numb_copy" not in out: + out["numb_copy"] = np.ones((nframes, 1), dtype=np.int64) + out["find_numb_copy"] = np.float32(0.0) + out["default_mesh"] = np.asarray( + make_default_mesh(bool(float(out["find_box"]) > 0.5), self.mixed_type), + dtype=np.int32, + ) + return out + + def _stack_frames(self, frames: list[dict[str, Any]]) -> dict[str, Any]: + """Collate already-decoded frames with the reader's flag semantics.""" + if not frames: + raise ValueError("Cannot stack an empty LMDB frame batch.") + from deepmd.dpmodel.utils.lmdb_data import ( + collate_lmdb_frames, + resolve_per_atom_keys, + ) + + per_atom_keys = resolve_per_atom_keys(frames[0], self._reader.decode_config) + return self._legacy_batch(collate_lmdb_frames(frames, per_atom_keys)) + + def get_batch(self, sys_idx: int | None = None) -> dict[str, Any]: + del sys_idx + indices = self._next_indices() + return self._legacy_batch(self._reader.decode_batch(indices, ragged=False)) + + def get_stat_batch(self, sys_idx: int) -> dict[str, Any]: + """Return one bounded batch from a homogeneous statistical group.""" + if not 0 <= sys_idx < len(self._stat_groups): + raise IndexError(f"Statistical system index {sys_idx} is out of range") + nloc, indices = self._stat_groups[sys_idx] + batch_size = self._get_stat_batch_size(nloc) + start = self._stat_offsets[sys_idx] + if start >= len(indices): + start = 0 + stop = min(start + batch_size, len(indices)) + self._stat_offsets[sys_idx] = stop + return self._legacy_batch( + self._reader.decode_batch(indices[start:stop], ragged=False) + ) + + def get_stat_nsystems(self) -> int: + """Return the number of stack-compatible statistical groups.""" + return len(self._stat_groups) + + def _get_stat_batch_size(self, nloc: int) -> int: + """Cap model-stat decoding by both frames and decoded atom rows.""" + configured = self._reader.get_batch_size_for_nloc(nloc) + return max(1, min(configured, 128, 20000 // max(int(nloc), 1))) + + def get_stat_numb_batches(self, sys_idx: int) -> int: + """Return the finite batch count of one statistical group.""" + if not 0 <= sys_idx < len(self._stat_groups): + raise IndexError(f"Statistical system index {sys_idx} is out of range") + nloc, indices = self._stat_groups[sys_idx] + batch_size = self._get_stat_batch_size(nloc) + return (len(indices) + batch_size - 1) // batch_size + + def get_nsystems(self) -> int: + return self.nsystems + + def get_natoms(self) -> int: + return self.natoms[0] + + def get_ntypes(self) -> int: + return len(self._type_map) + + def get_type_map(self) -> list[str]: + return self._type_map + + @property + def type_map(self) -> list[str]: + """Model-side atom names exposed by the legacy data-system API.""" + return self._type_map + + def get_batch_size(self) -> list[int]: + return self.batch_size + + def print_summary(self, name: str, prob: Any | None = None) -> None: + del prob + self._reader.print_summary(name, self.sys_probs) + + def close(self) -> None: + """Release LMDB readers idempotently.""" + if getattr(self, "_closed", True): + return + self.data_systems = [] + test_data = getattr(self, "_test_data", None) + if test_data is not None: + test_data.close() + reader = getattr(self, "_reader", None) + if reader is not None: + reader.close() + self._closed = True + + def __enter__(self) -> Self: + return self + + def __exit__(self, *args: object) -> None: + self.close() + + def __del__(self) -> None: + self.close() + + def _format_name_length(name: str, width: int) -> str: if len(name) <= width: return "{: >{}}".format(name, width) @@ -844,14 +1313,483 @@ def prob_sys_size_ext(keywords: str, nsystems: int, nbatch: int) -> list[float]: return sys_probs +def _is_deepmd_data_format(fmt: str) -> bool: + return fmt in { + "deepmd", + "deepmd/raw", + "deepmd/npy", + "deepmd/comp", + "deepmd/npy/mixed", + "deepmd/hdf5", + "deepmd/lmdb", + "lmdb", + } + + +def _is_dpdata_lmdb_format(fmt: str) -> bool: + """Return whether *fmt* names dpdata's DeePMD-compatible LMDB format.""" + return fmt in {"deepmd/lmdb", "lmdb"} + + +def _canonical_dpdata_out_format(out_fmt: str | None) -> str: + """Return the canonical dpdata output format used by conversion caches.""" + if out_fmt is None: + return _DPDATA_DEFAULT_OUT_FORMAT + out_fmt = out_fmt.lower() + return _DPDATA_DEFAULT_OUT_FORMAT if out_fmt == "lmdb" else out_fmt + + +def conversion_will_write_lmdb(data_config: dict[str, Any]) -> bool: + """Whether a non-DeePMD input config will be converted to LMDB.""" + data_format = data_config.get("format") + if data_format is None or _is_deepmd_data_format(data_format.lower()): + return False + out_format = data_config.get("out_format", data_config.get("output_format")) + return _is_dpdata_lmdb_format(_canonical_dpdata_out_format(out_format)) + + +def validate_backend_data_config( + data_config: dict[str, Any], + *, + backend_name: str, + lmdb_supported: bool, +) -> None: + """Reject unsupported converted output before dpdata performs any I/O.""" + if not lmdb_supported and conversion_will_write_lmdb(data_config): + raise NotImplementedError( + f"{backend_name} backend does not support LMDB data yet. " + "Choose out_format='deepmd/hdf5' for automatic conversion." + ) + + +def validate_lmdb_sampling_options(data_config: dict[str, Any]) -> None: + """Reject sampling options that cannot be represented by one LMDB route.""" + if data_config.get("sys_probs") is not None: + raise ValueError( + "LMDB data does not support explicit sys_probs yet. Use auto_prob " + "('prob_sys_size', 'prob_uniform', or block weights) so sampling " + "can be derived from LMDB frame_system_ids." + ) + + +def close_data_systems(*values: Any) -> None: + """Close nested data-system mappings/sequences, ignoring shared objects.""" + seen: set[int] = set() + + def close_one(value: Any) -> None: + if value is None or id(value) in seen: + return + seen.add(id(value)) + if isinstance(value, dict): + for child in value.values(): + close_one(child) + return + if isinstance(value, (list, tuple)): + for child in value: + close_one(child) + return + close = getattr(value, "close", None) + if callable(close): + try: + close() + except Exception: + log.warning( + "Failed to close data system %r during cleanup", + value, + exc_info=True, + ) + + for value in values: + close_one(value) + + +def _looks_like_extxyz(path: Path) -> bool: + if not path.is_file(): + return False + try: + with path.open() as fp: + fp.readline() + comment = fp.readline() + except OSError: + return False + return "Properties=" in comment or "Lattice=" in comment + + +def _normalize_dpdata_format(fmt: str, source: Path) -> str: + fmt = fmt.lower() + if fmt == "ase": + return "ase/structure" + if fmt != "auto": + return fmt + suffix = source.suffix.lower().lstrip(".") + if suffix == "traj": + return "ase/traj" + if suffix == "extxyz" or (suffix == "xyz" and _looks_like_extxyz(source)): + return "extxyz" + return suffix or fmt + + +def _iter_conversion_inputs(path: str, patterns: list[str] | None) -> list[str]: + if patterns is None: + return [path] + root = Path(path) + if not root.is_dir(): + return [path] + matches = [] + for pattern in patterns: + matches.extend(str(match) for match in root.rglob(pattern)) + return sorted(set(matches)) + + +def _conversion_cache_path(source: Path, fmt: str, out_fmt: str) -> Path: + source_resolved = source.resolve(strict=False) + try: + dpdata_version = importlib.metadata.version("dpdata") + except importlib.metadata.PackageNotFoundError: + dpdata_version = "unknown" + digest = hashlib.sha1( + ( + f"{source_resolved}|{fmt}|{out_fmt}|" + f"schema={_DPDATA_CONVERSION_SCHEMA_VERSION}|dpdata={dpdata_version}" + ).encode() + ).hexdigest()[:16] + stem = source_resolved.stem or source_resolved.name or "dataset" + safe_out_fmt = out_fmt.replace("/", "-") + suffix = ".lmdb" if _is_dpdata_lmdb_format(out_fmt) else "" + return Path.cwd() / _DPDATA_CACHE_DIR / f"{stem}-{safe_out_fmt}-{digest}{suffix}" + + +def _source_mtime(source: Path, cache_file: Path, *, force: bool = False) -> float: + if source.is_file(): + return source.stat().st_mtime + if not source.is_dir(): + return 0.0 + cache_key = ( + str(source.resolve(strict=False)), + str(cache_file.parent.resolve(strict=False)), + ) + now = time.monotonic() + cached = _DPDATA_SOURCE_MTIME_CACHE.get(cache_key) + if ( + not force + and cached is not None + and now - cached[0] < _DPDATA_SOURCE_MTIME_CACHE_TTL + ): + return cached[1] + cache_dir = cache_file.parent.resolve(strict=False) + latest = source.stat().st_mtime + for item in source.rglob("*"): + try: + item_resolved = item.resolve(strict=False) + if item_resolved == cache_file or cache_dir in item_resolved.parents: + continue + latest = max(latest, item.stat().st_mtime) + except OSError: + continue + _DPDATA_SOURCE_MTIME_CACHE[cache_key] = (now, latest) + return latest + + +def _is_conversion_current( + source: Path, output: Path, *, force_source_scan: bool = False +) -> bool: + if not output.exists(): + return False + return output.stat().st_mtime >= _source_mtime( + source, output, force=force_source_scan + ) + + +def _process_start_time(pid: int) -> str | None: + """Return Linux's stable process start token, if available.""" + try: + stat_text = Path(f"/proc/{pid}/stat").read_text() + except OSError: + return None + fields_after_name = stat_text.rsplit(")", 1)[1].split() + return fields_after_name[19] if len(fields_after_name) > 19 else None + + +def _same_lock_file(lock_path: Path, expected: os.stat_result) -> bool: + """Whether *lock_path* still names the inode originally acquired/read.""" + try: + current = lock_path.stat(follow_symlinks=False) + except FileNotFoundError: + return False + return (current.st_dev, current.st_ino) == (expected.st_dev, expected.st_ino) + + +class _ConversionLock: + """Owned conversion lock with a heartbeat for cross-host stale recovery.""" + + def __init__(self, lock_path: Path, lock_fd: int) -> None: + self.path = lock_path + self._stat = os.fstat(lock_fd) + payload = { + "hostname": socket.gethostname(), + "pid": os.getpid(), + "process_start": _process_start_time(os.getpid()), + "created": time.time(), + } + with os.fdopen(lock_fd, "w") as fp: + json.dump(payload, fp) + fp.flush() + os.fsync(fp.fileno()) + self._stop = threading.Event() + self._thread = threading.Thread( + target=self._heartbeat, + name="deepmd-dpdata-conversion-lock", + daemon=True, + ) + try: + self._thread.start() + except Exception: + if _same_lock_file(self.path, self._stat): + self.path.unlink(missing_ok=True) + raise + + def _heartbeat(self) -> None: + while not self._stop.wait(_CONVERSION_LOCK_HEARTBEAT_SECONDS): + if not _same_lock_file(self.path, self._stat): + return + try: + os.utime(self.path, None, follow_symlinks=False) + except FileNotFoundError: + return + + def release(self) -> None: + """Stop heartbeating and remove only the lock inode we own.""" + self._stop.set() + self._thread.join() + if _same_lock_file(self.path, self._stat): + try: + self.path.unlink() + except FileNotFoundError: + pass + + +def _lock_owner_is_alive(payload: dict[str, Any]) -> bool | None: + """Return owner liveness locally, or None for another host/invalid data.""" + if payload.get("hostname") != socket.gethostname(): + return None + try: + pid = int(payload["pid"]) + except (KeyError, TypeError, ValueError): + return None + expected_start = payload.get("process_start") + current_start = _process_start_time(pid) + if expected_start is not None and current_start is not None: + return str(expected_start) == current_start + try: + os.kill(pid, 0) + except ProcessLookupError: + return False + except PermissionError: + return True + return True + + +def _recover_stale_conversion_lock(lock_path: Path) -> bool: + """Remove a dead-owner or expired-lease lock without touching a successor.""" + try: + lock_stat = lock_path.stat(follow_symlinks=False) + except FileNotFoundError: + return True + except OSError: + return False + try: + payload = json.loads(lock_path.read_text()) + except (OSError, json.JSONDecodeError): + payload = {} + + owner_alive = _lock_owner_is_alive(payload) + lease_expired = time.time() - lock_stat.st_mtime > _CONVERSION_LOCK_STALE_SECONDS + stale = owner_alive is False or (owner_alive is None and lease_expired) + if not stale or not _same_lock_file(lock_path, lock_stat): + return False + log.warning("Recovering stale dpdata conversion lock %s", lock_path) + try: + lock_path.unlink() + except FileNotFoundError: + pass + return True + + +def _wait_for_conversion(source: Path, output: Path, lock_path: Path) -> bool: + """Wait without rescanning the source tree while a valid writer owns it.""" + while lock_path.exists(): + if _recover_stale_conversion_lock(lock_path): + continue + time.sleep(1.0) + # Freshness is checked once after publication, not once per waiter-second. + return _is_conversion_current(source, output, force_source_scan=True) + + +def _remove_path(path: Path) -> None: + # Check links before directories: Path.is_dir follows a directory symlink, + # while cache cleanup must never recurse into a target outside the cache. + if path.is_symlink(): + path.unlink() + elif path.is_dir(): + shutil.rmtree(path) + elif path.exists(): + path.unlink() + + +def _publish_conversion_output(tmp_output: Path, output: Path) -> None: + """Publish a non-LMDB conversion without discarding a valid old cache.""" + if not output.exists() or output.is_symlink() or not output.is_dir(): + os.replace(tmp_output, output) + backup = output.with_name(f".{output.name}.backup") + _remove_path(backup) + return + + # POSIX cannot replace a non-empty directory directly. Preserve the old + # cache under a sibling name and restore it if the second rename fails. + backup = output.with_name(f".{output.name}.backup") + _remove_path(backup) + os.replace(output, backup) + try: + os.replace(tmp_output, output) + except Exception: + os.replace(backup, output) + raise + else: + _remove_path(backup) + + +def _write_dpdata_conversion( + source: Path, fmt: str, out_fmt: str, output: Path +) -> None: + """Load *source* with dpdata and publish it in a DeePMD format. + + dpdata 1.1.0 makes ``deepmd/lmdb`` writes transactional: it stages and + validates the complete database before atomically publishing it. Use that + writer directly so DeePMD-kit does not duplicate or weaken its overwrite + guarantees. Other dpdata formats do not share that contract, so they keep + the cache-level temporary output used for failure isolation. + """ + try: + import dpdata + except ImportError as exc: + raise ImportError( + "dpdata is required when training_data.format or " + "validation_data.format is specified. Install dpdata to enable " + "automatic dataset conversion." + ) from exc + + multi_systems = dpdata.MultiSystems() + try: + multi_systems.load_systems_from_file(str(source), fmt=fmt) + except (NotImplementedError, ValueError) as labeled_error: + # dpdata 1.1 exposes an explicit unlabeled path. This matters for + # structure-only EXTXYZ/ASE inputs, which are valid descriptor data + # even though a supervised loss may later require labels. + unlabeled_systems = dpdata.MultiSystems() + try: + unlabeled_systems.load_systems_from_file( + str(source), fmt=fmt, labeled=False + ) + except (NotImplementedError, TypeError, ValueError): + try: + labeled_system = dpdata.LabeledSystem(str(source), fmt=fmt) + except Exception: + raise labeled_error from None + multi_systems = dpdata.MultiSystems(labeled_system) + else: + log.info("Loaded unlabeled dpdata input %s using format %s", source, fmt) + multi_systems = unlabeled_systems + if len(multi_systems) == 0: + raise RuntimeError(f"No frames were loaded by dpdata from {source}") + + if _is_dpdata_lmdb_format(out_fmt): + multi_systems.to(out_fmt, str(output), overwrite=True) + return + + tmp_output = output.with_name(f".{output.name}.{os.getpid()}.tmp") + _remove_path(tmp_output) + try: + multi_systems.to(out_fmt, str(tmp_output)) + _publish_conversion_output(tmp_output, output) + except Exception: + _remove_path(tmp_output) + raise + + +def _convert_system_by_dpdata( + source_path: str, fmt: str, out_fmt: str | None +) -> list[str]: + source = Path(source_path) + fmt = _normalize_dpdata_format(fmt, source) + out_fmt = _canonical_dpdata_out_format(out_fmt) + output = _conversion_cache_path(source, fmt, out_fmt) + cache_key = ( + str(Path.cwd().resolve(strict=False)), + str(source.resolve(strict=False)), + fmt, + out_fmt, + str(output), + ) + cached_systems = _DPDATA_CONVERSION_CACHE.get(cache_key) + if cached_systems is not None: + if _is_conversion_current(source, output): + return cached_systems + # A long-lived training/validation process may observe source files + # rewritten in place. Drop the fast-path entry so the normal locked + # conversion flow refreshes the on-disk result before it is reused. + del _DPDATA_CONVERSION_CACHE[cache_key] + + output.parent.mkdir(parents=True, exist_ok=True) + lock_path = output.with_suffix(output.suffix + ".lock") + if not _is_conversion_current(source, output): + while True: + try: + lock_fd = os.open(lock_path, os.O_CREAT | os.O_EXCL | os.O_WRONLY) + except FileExistsError: + if _wait_for_conversion(source, output, lock_path): + break + continue + else: + conversion_lock = _ConversionLock(lock_path, lock_fd) + try: + if not _is_conversion_current( + source, output, force_source_scan=True + ): + log.info( + "Converting %s from dpdata format %s to %s at %s", + source, + fmt, + out_fmt, + output, + ) + _write_dpdata_conversion(source, fmt, out_fmt, output) + finally: + conversion_lock.release() + break + + if _is_dpdata_lmdb_format(out_fmt): + converted_systems = [str(output)] + else: + converted_systems = expand_sys_str(str(output)) + if not converted_systems: + raise RuntimeError(f"No DeePMD systems were found in converted file {output}") + _DPDATA_CONVERSION_CACHE[cache_key] = converted_systems + return converted_systems + + def process_systems( - systems: str | list[str], patterns: list[str] | None = None + systems: str | list[str], + patterns: list[str] | None = None, + fmt: str | None = None, + out_fmt: str | None = None, ) -> list[str]: """Process the user-input systems. If it is a single directory, search for all the systems in the directory. If it is a list, each item in the list is treated as a directory to search. If it is a single LMDB path, return it directly without expansion. + If fmt is specified and is not a DeePMD data format, each input path is + converted by dpdata and the converted systems are returned. Check if the systems are valid. Parameters @@ -860,20 +1798,23 @@ def process_systems( The user-input systems patterns : list of str, optional The patterns to match the systems, by default None + fmt : str, optional + The dpdata input format. If None, no conversion is performed. + out_fmt : str, optional + The dpdata output format. If None, ``deepmd/lmdb`` is used when fmt + triggers conversion. Returns ------- result_systems: list of str The valid systems """ + # See validate_lmdb_systems: this must remain a local import because + # deepmd.dpmodel initializes descriptors that depend on data_system. from deepmd.dpmodel.utils.lmdb_data import ( is_lmdb, ) - # LMDB path: return directly without expansion - if isinstance(systems, str) and is_lmdb(systems): - return [systems] - # Normalize input to a list of paths to search if isinstance(systems, str): search_paths = [systems] @@ -885,15 +1826,40 @@ def process_systems( f"Invalid systems type: {type(systems)}. Must be str or list[str]." ) + if fmt is not None: + fmt = fmt.lower() + if _is_deepmd_data_format(fmt): + fmt = None + + conversion_inputs: list[str] = [] + if fmt is not None: + for path in search_paths: + conversion_inputs.extend(_iter_conversion_inputs(path, patterns)) + if ( + _is_dpdata_lmdb_format(_canonical_dpdata_out_format(out_fmt)) + and len(conversion_inputs) != 1 + ): + raise ValueError( + "Automatic LMDB conversion requires exactly one resolved input " + "path. Merge multiple inputs with dpdata first or choose " + "out_format='deepmd/hdf5'." + ) + # Iterate over the search_paths list and apply expansion logic to each path result_systems = [] - for path in search_paths: - if patterns is None: - expanded_paths = expand_sys_str(path) - else: - expanded_paths = rglob_sys_str(path, patterns) - - result_systems.extend(expanded_paths) + if fmt is not None: + for input_path in conversion_inputs: + result_systems.extend(_convert_system_by_dpdata(input_path, fmt, out_fmt)) + else: + for path in search_paths: + if is_lmdb(path): + result_systems.append(path) + elif patterns is None: + expanded_paths = expand_sys_str(path) + result_systems.extend(expanded_paths) + else: + expanded_paths = rglob_sys_str(path, patterns) + result_systems.extend(expanded_paths) return result_systems @@ -904,7 +1870,7 @@ def get_data( type_map: list[str] | None, modifier: Any | None, multi_task_mode: bool = False, -) -> DeepmdDataSystem: +) -> DeepmdDataSystem | LmdbDataSystem: """Get the data system. Parameters @@ -927,13 +1893,35 @@ def get_data( """ systems = jdata["systems"] rglob_patterns = jdata.get("rglob_patterns") - systems = process_systems(systems, patterns=rglob_patterns) + data_format = jdata.get("format") + out_format = jdata.get("out_format", jdata.get("output_format")) + if conversion_will_write_lmdb(jdata): + validate_lmdb_sampling_options(jdata) + systems = process_systems( + systems, patterns=rglob_patterns, fmt=data_format, out_fmt=out_format + ) batch_size = jdata["batch_size"] sys_probs = jdata.get("sys_probs") auto_prob = jdata.get("auto_prob", "prob_sys_size") optional_type_map = not multi_task_mode + lmdb_path = validate_lmdb_systems(systems, backend_name="legacy data loader") + if lmdb_path is not None: + validate_lmdb_sampling_options(jdata) + if type_map is None: + raise ValueError( + "LMDB training data requires model/type_map to be set. " + "Set model/type_map or choose training_data.out_format=" + "'deepmd/hdf5' for automatic conversion." + ) + return LmdbDataSystem( + lmdb_path=lmdb_path, + type_map=type_map, + batch_size=batch_size, + auto_prob_style=auto_prob, + ) + data = DeepmdDataSystem( systems=systems, batch_size=batch_size, diff --git a/pyproject.toml b/pyproject.toml index 3b64d3ec1b..59e83f9217 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -56,6 +56,7 @@ dependencies = [ 'array-api-compat', 'lmdb', 'msgpack', + 'dpdata>=1.1.0', ] requires-python = ">=3.10" keywords = ["deepmd"] @@ -79,7 +80,6 @@ repository = "https://github.com/deepmodeling/deepmd-kit" # which can be read by the build backend. [tool.deepmd_build_backend.optional-dependencies] test = [ - "dpdata>=0.2.7", # ASE issue: https://gitlab.com/ase/ase/-/merge_requests/2843 # fixed in 3.23.0 "ase>=3.23.0", diff --git a/source/tests/common/dpmodel/test_lmdb_data.py b/source/tests/common/dpmodel/test_lmdb_data.py index 16dd1da1d1..bc17b876d4 100644 --- a/source/tests/common/dpmodel/test_lmdb_data.py +++ b/source/tests/common/dpmodel/test_lmdb_data.py @@ -2021,6 +2021,13 @@ def test_compute_block_targets_unequal(self): self.assertEqual(result[0], ([0], 500)) self.assertEqual(result[1], ([1], 500)) + def test_compute_block_targets_prob_uniform(self): + """Bare prob_uniform equalizes unequal original system sizes.""" + result = compute_block_targets( + "prob_uniform", nsystems=2, system_nframes=[100, 500] + ) + self.assertEqual(result, [([0], 500), ([1], 500)]) + def test_compute_block_targets_multi_sys_block(self): result = compute_block_targets( "prob_sys_size;0:2:0.5;2:3:0.5", @@ -2055,6 +2062,22 @@ def test_compute_block_targets_logs_dropped_block(self): self.assertTrue(any("empty blocks" in msg for msg in cm.output)) self.assertEqual(result, []) + def test_compute_block_targets_rejects_invalid_weights(self): + """Weights must define a finite, nonnegative probability mass.""" + invalid_styles = ( + ("prob_sys_size;0:1:-0.1;1:2:1.1", "no less than 0"), + ("prob_sys_size;0:1:0;1:2:0", "greater than 0"), + ("prob_sys_size;0:1:nan;1:2:1", "finite"), + ) + for style, message in invalid_styles: + with self.subTest(style=style): + with self.assertRaisesRegex(ValueError, message): + compute_block_targets( + style, + nsystems=2, + system_nframes=[100, 100], + ) + def test_expand_indices_basic(self): frame_system_ids = [0] * 5 + [1] * 5 block_targets = [([0], 25), ([1], 25)] diff --git a/source/tests/common/test_data_system_conversion.py b/source/tests/common/test_data_system_conversion.py new file mode 100644 index 0000000000..658a97519f --- /dev/null +++ b/source/tests/common/test_data_system_conversion.py @@ -0,0 +1,705 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +import json +import os +import sys +import tempfile +import types +import unittest +from pathlib import ( + Path, +) +from typing import ( + ClassVar, +) +from unittest.mock import ( + patch, +) + +import h5py +import lmdb +import msgpack +import numpy as np + +from deepmd.dpmodel.utils.lmdb_data import ( + is_lmdb, +) +from deepmd.entrypoints.test import test as run_model_test +from deepmd.utils import ( + data_system, +) +from deepmd.utils.data import ( + DataRequirementItem, +) +from deepmd.utils.data_system import ( + LmdbDataSystem, + get_data, + process_systems, + validate_backend_data_config, + validate_lmdb_systems, +) + + +def _write_minimal_deepmd_hdf5(file_name: str) -> None: + with h5py.File(file_name, "w") as fp: + system = fp.create_group("H") + system.create_dataset("type.raw", data=np.array([0], dtype=np.int32)) + string_dtype = h5py.string_dtype(encoding="utf-8") + system.create_dataset("type_map.raw", data=np.array(["H"], dtype=string_dtype)) + set_dir = system.create_group("set.000") + set_dir.create_dataset("coord.npy", data=np.zeros((1, 3), dtype=np.float32)) + set_dir.create_dataset( + "box.npy", data=np.eye(3, dtype=np.float32).reshape(1, 9) + ) + + +def _encode_array(arr: np.ndarray) -> dict: + return { + "type": str(arr.dtype), + "shape": list(arr.shape), + "data": arr.tobytes(), + } + + +def _write_minimal_lmdb(path: str) -> None: + env = lmdb.open(path, map_size=10 * 1024 * 1024) + frame = { + "atom_names": ["H"], + "atom_numbs": [1], + "atom_types": _encode_array(np.array([0], dtype=np.int64)), + "cells": _encode_array(np.eye(3, dtype=np.float64) * 8.0), + "coords": _encode_array(np.zeros((1, 3), dtype=np.float64)), + "energies": _encode_array(np.array([0.0], dtype=np.float64)), + "forces": _encode_array(np.zeros((1, 3), dtype=np.float64)), + } + metadata = { + "nframes": 1, + "frame_idx_fmt": "012d", + "type_map": ["H"], + "system_info": { + "formula": "H", + "natoms": [1], + "nframes": 1, + }, + } + with env.begin(write=True) as txn: + txn.put(b"__metadata__", msgpack.packb(metadata, use_bin_type=True)) + txn.put(b"000000000000", msgpack.packb(frame, use_bin_type=True)) + env.close() + + +def _write_repeated_lmdb(path: str, nframes: int) -> None: + """Write a compact same-nloc LMDB suitable for bounded-I/O assertions.""" + env = lmdb.open(path, map_size=64 * 1024 * 1024) + frame = { + "atom_names": ["H"], + "atom_numbs": [1], + "atom_types": _encode_array(np.array([0], dtype=np.int64)), + "cells": _encode_array(np.eye(3, dtype=np.float64) * 8.0), + "coords": _encode_array(np.zeros((1, 3), dtype=np.float64)), + "energies": _encode_array(np.array([0.0], dtype=np.float64)), + "forces": _encode_array(np.zeros((1, 3), dtype=np.float64)), + } + packed_frame = msgpack.packb(frame, use_bin_type=True) + metadata = { + "nframes": nframes, + "frame_idx_fmt": "012d", + "frame_nlocs": [1] * nframes, + "type_map": ["H"], + "system_info": { + "formula": "H", + "natoms": [1], + "nframes": nframes, + }, + } + with env.begin(write=True) as txn: + txn.put(b"__metadata__", msgpack.packb(metadata, use_bin_type=True)) + for index in range(nframes): + txn.put(f"{index:012d}".encode(), packed_frame) + env.close() + + +def _write_mixed_nloc_lmdb(path: str) -> None: + """Write two frames that exercise legacy mix:N padding.""" + env = lmdb.open(path, map_size=10 * 1024 * 1024) + frames = [] + for nloc in (1, 2): + frames.append( + { + "atom_names": ["H"], + "atom_numbs": [nloc], + "atom_types": _encode_array(np.zeros(nloc, dtype=np.int64)), + "cells": _encode_array(np.eye(3, dtype=np.float64) * 8.0), + "coords": _encode_array(np.zeros((nloc, 3), dtype=np.float64)), + "energies": _encode_array(np.array([0.0], dtype=np.float64)), + "forces": _encode_array(np.zeros((nloc, 3), dtype=np.float64)), + } + ) + metadata = { + "nframes": 2, + "frame_idx_fmt": "012d", + "frame_nlocs": [1, 2], + "type_map": ["H"], + "system_info": { + "formula": "H", + "natoms": [1], + "nframes": 2, + }, + } + with env.begin(write=True) as txn: + txn.put(b"__metadata__", msgpack.packb(metadata, use_bin_type=True)) + for index, frame in enumerate(frames): + txn.put(f"{index:012d}".encode(), msgpack.packb(frame, use_bin_type=True)) + env.close() + + +def _write_mixed_pbc_lmdb(path: str) -> None: + """Write equal-nloc periodic and non-periodic frames.""" + env = lmdb.open(path, map_size=10 * 1024 * 1024) + frames = [] + for cell in (np.eye(3, dtype=np.float64) * 8.0, np.zeros((3, 3))): + frames.append( + { + "atom_names": ["H"], + "atom_numbs": [1], + "atom_types": _encode_array(np.array([0], dtype=np.int64)), + "cells": _encode_array(cell), + "coords": _encode_array(np.zeros((1, 3), dtype=np.float64)), + "energies": _encode_array(np.array([0.0], dtype=np.float64)), + "forces": _encode_array(np.zeros((1, 3), dtype=np.float64)), + } + ) + metadata = { + "nframes": 2, + "frame_idx_fmt": "012d", + "frame_nlocs": [1, 1], + "type_map": ["H"], + "system_info": {"formula": "H", "natoms": [1], "nframes": 2}, + } + with env.begin(write=True) as txn: + txn.put(b"__metadata__", msgpack.packb(metadata, use_bin_type=True)) + for index, frame in enumerate(frames): + txn.put(f"{index:012d}".encode(), msgpack.packb(frame, use_bin_type=True)) + env.close() + + +class _FakeMultiSystems: + write_count = 0 + load_calls: ClassVar[list[tuple[str, str]]] = [] + to_calls: ClassVar[list[tuple[str, str, dict]]] = [] + + def __init__(self, *systems) -> None: + self.systems = list(systems) + self.loaded = False + + def load_systems_from_file(self, file_name: str, fmt: str): + self.load_calls.append((file_name, fmt)) + self.loaded = True + return self + + def __len__(self) -> int: + return 1 if self.loaded or self.systems else 0 + + def to(self, fmt: str, file_name: str, **kwargs) -> None: + type(self).write_count += 1 + self.to_calls.append((fmt, file_name, kwargs)) + if fmt == "deepmd/hdf5": + _write_minimal_deepmd_hdf5(file_name) + elif fmt == "deepmd/lmdb": + _write_minimal_lmdb(file_name) + else: + raise AssertionError(fmt) + + +class _FakeLabeledSystem: + def __init__(self, file_name: str, fmt: str) -> None: + self.file_name = file_name + self.fmt = fmt + + +class TestDpdataFormatConversion(unittest.TestCase): + def setUp(self) -> None: + self.tmpdir = tempfile.TemporaryDirectory() + self.root = Path(self.tmpdir.name) + self.old_cwd = Path.cwd() + os.chdir(self.root) + self.source = self.root / "data.extxyz" + self.source.write_text("1\nProperties=species:S:1:pos:R:3\nH 0 0 0\n") + _FakeMultiSystems.write_count = 0 + _FakeMultiSystems.load_calls = [] + _FakeMultiSystems.to_calls = [] + data_system._DPDATA_CONVERSION_CACHE.clear() + data_system._DPDATA_SOURCE_MTIME_CACHE.clear() + self.fake_dpdata = types.SimpleNamespace( + MultiSystems=_FakeMultiSystems, + LabeledSystem=_FakeLabeledSystem, + ) + + def tearDown(self) -> None: + os.chdir(self.old_cwd) + self.tmpdir.cleanup() + data_system._DPDATA_CONVERSION_CACHE.clear() + data_system._DPDATA_SOURCE_MTIME_CACHE.clear() + + def test_process_systems_defaults_to_deepmd_lmdb_and_reuses_cache(self) -> None: + with patch.dict(sys.modules, {"dpdata": self.fake_dpdata}): + systems = process_systems(str(self.source), fmt="extxyz") + systems_again = process_systems(str(self.source), fmt="extxyz") + + self.assertEqual(systems, systems_again) + self.assertEqual(_FakeMultiSystems.write_count, 1) + self.assertEqual(_FakeMultiSystems.load_calls, [(str(self.source), "extxyz")]) + self.assertEqual(len(systems), 1) + self.assertTrue(systems[0].endswith(".lmdb")) + self.assertTrue(is_lmdb(systems[0])) + self.assertTrue(Path(systems[0]).is_relative_to(self.root)) + self.assertEqual(Path(systems[0]).parent, self.root / ".deepmd_dpdata_cache") + self.assertEqual( + _FakeMultiSystems.to_calls, + [("deepmd/lmdb", systems[0], {"overwrite": True})], + ) + + def test_process_systems_cache_is_scoped_to_cwd(self) -> None: + other_cwd = self.root / "run2" + other_cwd.mkdir() + + with patch.dict(sys.modules, {"dpdata": self.fake_dpdata}): + systems = process_systems(str(self.source), fmt="extxyz") + os.chdir(other_cwd) + systems_other = process_systems(str(self.source), fmt="extxyz") + + self.assertNotEqual(systems, systems_other) + self.assertEqual(_FakeMultiSystems.write_count, 2) + self.assertEqual(Path(systems[0]).parent, self.root / ".deepmd_dpdata_cache") + self.assertEqual( + Path(systems_other[0]).parent, + other_cwd / ".deepmd_dpdata_cache", + ) + + def test_process_systems_revalidates_in_memory_cache(self) -> None: + """A source rewrite in one process must refresh its converted output.""" + with patch.dict(sys.modules, {"dpdata": self.fake_dpdata}): + systems = process_systems(str(self.source), fmt="extxyz") + output_mtime = Path(systems[0]).stat().st_mtime + refreshed_mtime = output_mtime + 1.0 + os.utime(self.source, (refreshed_mtime, refreshed_mtime)) + systems_again = process_systems(str(self.source), fmt="extxyz") + + self.assertEqual(systems, systems_again) + self.assertEqual(_FakeMultiSystems.write_count, 2) + self.assertTrue( + all( + call == ("deepmd/lmdb", systems[0], {"overwrite": True}) + for call in _FakeMultiSystems.to_calls + ) + ) + + def test_lmdb_alias_uses_canonical_dpdata_writer(self) -> None: + """The legacy alias must share the canonical dpdata cache entry.""" + with patch.dict(sys.modules, {"dpdata": self.fake_dpdata}): + systems = process_systems(str(self.source), fmt="extxyz", out_fmt="lmdb") + + self.assertEqual( + _FakeMultiSystems.to_calls, + [("deepmd/lmdb", systems[0], {"overwrite": True})], + ) + + def test_real_dpdata_lmdb_writer_is_deepmd_compatible(self) -> None: + """Verify dpdata 1.1 writes the schema consumed by DeePMD's LMDB reader.""" + extxyz = ( + "1\n" + "Properties=species:S:1:pos:R:3:forces:R:3 energy=0.0 " + 'Lattice="8 0 0 0 8 0 0 0 8"\n' + "H 0 0 0 0 0 0\n" + ) + self.source.write_text(extxyz) + + systems = process_systems(str(self.source), fmt="extxyz") + # A changed source exercises dpdata's own transactional overwrite path; + # DeePMD-kit must not remove or rename the LMDB directory around it. + self.source.write_text(extxyz.replace("energy=0.0", "energy=1.0")) + systems_again = process_systems(str(self.source), fmt="extxyz") + + self.assertEqual(len(systems), 1) + self.assertEqual(systems_again, systems) + self.assertTrue(is_lmdb(systems[0])) + data = LmdbDataSystem(systems[0], ["H"], batch_size=1) + batch = data.get_batch() + self.assertEqual(batch["coord"].shape, (1, 3)) + self.assertEqual(batch["type"].shape, (1, 1)) + + def test_cache_cleanup_unlinks_directory_symlink(self) -> None: + """Cleanup must not recurse through a symlink outside the cache.""" + target = self.root / "outside" + target.mkdir() + sentinel = target / "keep.txt" + sentinel.write_text("keep") + link = self.root / ".deepmd_dpdata_cache" / "stale.tmp" + link.parent.mkdir() + link.symlink_to(target, target_is_directory=True) + + data_system._remove_path(link) + + self.assertFalse(link.exists()) + self.assertTrue(target.is_dir()) + self.assertEqual(sentinel.read_text(), "keep") + + def test_process_systems_converts_to_explicit_hdf5(self) -> None: + with patch.dict(sys.modules, {"dpdata": self.fake_dpdata}): + systems = process_systems( + str(self.source), fmt="extxyz", out_fmt="deepmd/hdf5" + ) + + self.assertEqual(_FakeMultiSystems.write_count, 1) + self.assertEqual(_FakeMultiSystems.load_calls, [(str(self.source), "extxyz")]) + self.assertEqual(len(systems), 1) + self.assertTrue(systems[0].endswith("#/H")) + + def test_get_data_uses_format_conversion(self) -> None: + with patch.dict(sys.modules, {"dpdata": self.fake_dpdata}): + data = get_data( + { + "systems": str(self.source), + "format": "auto", + "batch_size": 1, + }, + 0.0, + ["H"], + None, + ) + + self.assertEqual(data.get_nsystems(), 1) + self.assertIsInstance(data, LmdbDataSystem) + self.assertTrue(data.mixed_type) + self.assertEqual(_FakeMultiSystems.load_calls, [(str(self.source), "extxyz")]) + batch = data.get_batch() + self.assertIn("type", batch) + self.assertIn("natoms_vec", batch) + self.assertEqual(batch["coord"].shape, (1, 3)) + self.assertIsNot(data.data_systems[0], data) + self.assertEqual(data.type_map, ["H"]) + self.assertEqual(data.data_systems[0].get_test()["coord"].shape, (1, 3)) + stat_set = data._load_set(data.dirs[0]) + self.assertEqual(stat_set["coord"].shape, (1, 3)) + self.assertEqual(stat_set["type"].shape, (1, 1)) + data.close() + data.close() + + def test_multiple_lmdb_paths_are_rejected(self) -> None: + lmdb_a = self.root / "a.lmdb" + lmdb_b = self.root / "b.lmdb" + _write_minimal_lmdb(str(lmdb_a)) + _write_minimal_lmdb(str(lmdb_b)) + + with self.assertRaisesRegex(ValueError, "exactly one path"): + get_data( + { + "systems": [str(lmdb_a), str(lmdb_b)], + "batch_size": 1, + }, + 0.0, + ["H"], + None, + ) + + def test_backend_without_lmdb_support_rejects_any_resolved_path(self) -> None: + lmdb_path = self.root / "unsupported.lmdb" + _write_minimal_lmdb(str(lmdb_path)) + + with self.assertRaisesRegex(NotImplementedError, "Paddle backend"): + validate_lmdb_systems( + [str(lmdb_path)], backend_name="Paddle", supported=False + ) + + def test_lmdb_stack_frames_rejects_empty_batch(self) -> None: + lmdb_path = self.root / "empty-batch.lmdb" + _write_minimal_lmdb(str(lmdb_path)) + data = LmdbDataSystem(str(lmdb_path), ["H"], batch_size=1) + + with self.assertRaisesRegex(ValueError, "empty LMDB frame batch"): + data._stack_frames([]) + + def test_requirements_can_be_registered_after_adapter_construction(self) -> None: + """PBC probing must not freeze the model's label contract.""" + lmdb_path = self.root / "requirements.lmdb" + _write_minimal_lmdb(str(lmdb_path)) + data = LmdbDataSystem(str(lmdb_path), ["H"], batch_size=[1]) + + data.add_data_requirements( + [DataRequirementItem("energy", 1, atomic=False, must=True)] + ) + batch = data.get_batch() + + self.assertEqual(batch["energy"].shape, (1, 1)) + self.assertEqual(float(batch["find_energy"]), 1.0) + + def test_legacy_collation_ands_find_flags(self) -> None: + lmdb_path = self.root / "find-flags.lmdb" + _write_minimal_lmdb(str(lmdb_path)) + data = LmdbDataSystem(str(lmdb_path), ["H"], batch_size=2) + data.add_data_requirements( + [DataRequirementItem("energy", 1, atomic=False, must=False)] + ) + first = data._reader.peek_frame(0) + second = { + key: value.copy() if isinstance(value, np.ndarray) else value + for key, value in first.items() + } + second["find_energy"] = np.float32(0.0) + + batch = data._stack_frames([first, second]) + + self.assertEqual(float(batch["find_energy"]), 0.0) + + def test_neighbor_stat_reads_are_sampled_and_chunked(self) -> None: + """A large same-nloc group must never be decoded as one Python list.""" + lmdb_path = self.root / "bounded.lmdb" + _write_repeated_lmdb(str(lmdb_path), 2101) + data = LmdbDataSystem(str(lmdb_path), ["H"], batch_size=1000) + + total_sampled = sum(len(indices) for indices in data._nloc_set_indices.values()) + self.assertEqual(total_sampled, 2000) + self.assertLessEqual( + max(len(indices) for indices in data._nloc_set_indices.values()), 128 + ) + first = data._load_set(data.dirs[0]) + self.assertLessEqual(first["coord"].shape[0], 128) + self.assertEqual(data.get_stat_nsystems(), 1) + self.assertEqual(data.get_stat_numb_batches(0), 17) + self.assertEqual(data.get_stat_batch(0)["coord"].shape, (128, 3)) + + def test_mix_batch_pads_different_atom_counts(self) -> None: + lmdb_path = self.root / "mixed.lmdb" + _write_mixed_nloc_lmdb(str(lmdb_path)) + data = LmdbDataSystem(str(lmdb_path), ["H"], batch_size="mix:4", seed=0) + + batch = data.get_batch() + + self.assertEqual(batch["coord"].shape, (2, 6)) + self.assertEqual(batch["type"].shape, (2, 2)) + self.assertIn(-1, batch["type"]) + + def test_periodic_and_nonperiodic_frames_are_separate_views(self) -> None: + lmdb_path = self.root / "mixed-pbc.lmdb" + _write_mixed_pbc_lmdb(str(lmdb_path)) + data = LmdbDataSystem(str(lmdb_path), ["H"], batch_size=2, seed=0) + + self.assertEqual(len(data.data_systems), 2) + self.assertEqual({view.pbc for view in data.data_systems}, {False, True}) + batches = [data.get_batch(), data.get_batch()] + self.assertEqual({float(batch["find_box"]) for batch in batches}, {0.0, 1.0}) + + def test_multiple_conversion_inputs_fail_before_dpdata_io(self) -> None: + second = self.root / "second.extxyz" + second.write_text(self.source.read_text()) + + with patch.dict(sys.modules, {"dpdata": self.fake_dpdata}): + with self.assertRaisesRegex(ValueError, "exactly one resolved input"): + process_systems( + [str(self.source), str(second)], + fmt="extxyz", + ) + + self.assertEqual(_FakeMultiSystems.load_calls, []) + self.assertEqual(_FakeMultiSystems.write_count, 0) + + def test_lmdb_sys_probs_fail_before_conversion(self) -> None: + with self.assertRaisesRegex(ValueError, "does not support explicit sys_probs"): + get_data( + { + "systems": str(self.source), + "format": "extxyz", + "batch_size": 1, + "sys_probs": [1.0], + }, + 0.0, + ["H"], + None, + ) + + self.assertEqual(_FakeMultiSystems.write_count, 0) + + def test_paddle_rejects_default_lmdb_before_conversion(self) -> None: + with self.assertRaisesRegex(NotImplementedError, "Paddle backend"): + validate_backend_data_config( + {"systems": str(self.source), "format": "extxyz"}, + backend_name="Paddle", + lmdb_supported=False, + ) + + self.assertEqual(_FakeMultiSystems.write_count, 0) + + def test_waiter_checks_source_only_after_lock_release(self) -> None: + output = self.root / "wait.lmdb" + output.mkdir() + lock_path = output.with_suffix(".lmdb.lock") + lock_path.write_text( + json.dumps( + { + "hostname": data_system.socket.gethostname(), + "pid": os.getpid(), + "process_start": data_system._process_start_time(os.getpid()), + } + ) + ) + waits = 0 + + def release_after_long_valid_conversion(_seconds: float) -> None: + nonlocal waits + waits += 1 + if waits == 301: + lock_path.unlink() + + with ( + patch.object( + data_system.time, "sleep", release_after_long_valid_conversion + ), + patch.object( + data_system, "_is_conversion_current", return_value=True + ) as is_current, + ): + self.assertTrue( + data_system._wait_for_conversion(self.source, output, lock_path) + ) + + self.assertEqual(waits, 301) + is_current.assert_called_once_with(self.source, output, force_source_scan=True) + + def test_dead_conversion_owner_is_recovered(self) -> None: + lock_path = self.root / "dead.lock" + lock_path.write_text( + json.dumps( + { + "hostname": data_system.socket.gethostname(), + "pid": 2147483647, + "process_start": "missing", + } + ) + ) + + self.assertTrue(data_system._recover_stale_conversion_lock(lock_path)) + self.assertFalse(lock_path.exists()) + + def test_directory_publication_rolls_back_on_failure(self) -> None: + output = self.root / "published" + output.mkdir() + (output / "old.txt").write_text("old") + staged = self.root / "staged" + staged.mkdir() + (staged / "new.txt").write_text("new") + real_replace = os.replace + calls = 0 + + def fail_publication(source: str | Path, target: str | Path) -> None: + nonlocal calls + calls += 1 + if calls == 2: + raise OSError("publish failed") + real_replace(source, target) + + with patch.object(data_system.os, "replace", fail_publication): + with self.assertRaisesRegex(OSError, "publish failed"): + data_system._publish_conversion_output(staged, output) + + self.assertEqual((output / "old.txt").read_text(), "old") + self.assertFalse((output / "new.txt").exists()) + + def test_unlabeled_input_uses_dpdata_unlabeled_loader(self) -> None: + class UnlabeledMultiSystems(_FakeMultiSystems): + calls: ClassVar[list[bool]] = [] + + def load_systems_from_file( + self, file_name: str, fmt: str, *, labeled: bool = True + ): + type(self).calls.append(labeled) + if labeled: + raise ValueError("missing labels") + self.loaded = True + return self + + UnlabeledMultiSystems.calls = [] + fake_dpdata = types.SimpleNamespace( + MultiSystems=UnlabeledMultiSystems, + LabeledSystem=_FakeLabeledSystem, + ) + + with patch.dict(sys.modules, {"dpdata": fake_dpdata}): + systems = process_systems(str(self.source), fmt="extxyz") + + self.assertEqual(UnlabeledMultiSystems.calls, [True, False]) + self.assertTrue(is_lmdb(systems[0])) + + def test_cache_path_changes_with_dpdata_version(self) -> None: + with patch.object( + data_system.importlib.metadata, "version", side_effect=["1.1.0", "1.2.0"] + ): + first = data_system._conversion_cache_path( + self.source, "extxyz", "deepmd/lmdb" + ) + second = data_system._conversion_cache_path( + self.source, "extxyz", "deepmd/lmdb" + ) + + self.assertNotEqual(first, second) + + def test_directory_freshness_scan_is_reused_within_one_routing_pass(self) -> None: + source_dir = self.root / "source-dir" + source_dir.mkdir() + (source_dir / "frame.xyz").write_text("frame") + output = self.root / ".deepmd_dpdata_cache" / "cached.lmdb" + output.parent.mkdir() + + first = data_system._source_mtime(source_dir, output) + with patch.object( + Path, "rglob", side_effect=AssertionError("unexpected rescan") + ): + second = data_system._source_mtime(source_dir, output) + + self.assertEqual(first, second) + + def test_dp_test_forwards_conversion_format_from_training_config(self) -> None: + config_path = self.root / "input.json" + config_path.write_text("{}") + config = { + "training": { + "training_data": { + "systems": "data.extxyz", + "format": "extxyz", + "out_format": "deepmd/lmdb", + "rglob_patterns": ["*.extxyz"], + } + } + } + with ( + patch("deepmd.entrypoints.test.j_loader", return_value=config), + patch( + "deepmd.entrypoints.test.update_deepmd_input", + side_effect=lambda value: value, + ), + patch( + "deepmd.entrypoints.test.process_systems", + side_effect=RuntimeError("stop after routing"), + ) as process, + ): + with self.assertRaisesRegex(RuntimeError, "stop after routing"): + run_model_test( + model="unused.pb", + system=None, + datafile=None, + train_json=str(config_path), + numb_test=1, + rand_seed=None, + shuffle_test=False, + detail_file="detail", + atomic=False, + ) + + process.assert_called_once_with( + str(self.source.resolve()), + patterns=["*.extxyz"], + fmt="extxyz", + out_fmt="deepmd/lmdb", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/source/tests/pt_expt/test_lmdb_training.py b/source/tests/pt_expt/test_lmdb_training.py index fc7a64e7ee..0e350daa5a 100644 --- a/source/tests/pt_expt/test_lmdb_training.py +++ b/source/tests/pt_expt/test_lmdb_training.py @@ -30,6 +30,8 @@ collate_lmdb_frames, ) from deepmd.pt_expt.entrypoints.main import ( + _build_data_system, + _get_neighbor_stat_data, get_trainer, ) from deepmd.pt_expt.loss import ( @@ -69,6 +71,70 @@ def _encode_array(arr: np.ndarray) -> dict: } +class TestConvertedLmdbValidation(unittest.TestCase): + """Reject format conversion that resolves to multiple LMDB databases.""" + + def test_neighbor_stat_and_training_data_reject_multiple_lmdb(self) -> None: + params = { + "systems": "input.extxyz", + "format": "extxyz", + "batch_size": 1, + } + converted = ["first.lmdb", "second.lmdb"] + with ( + patch( + "deepmd.pt_expt.entrypoints.main.process_systems", + return_value=converted, + ), + # is_lmdb is imported inside the validating function, so patch + # it at the source module rather than at an import site. + patch( + "deepmd.dpmodel.utils.lmdb_data.is_lmdb", + return_value=True, + ), + ): + with self.assertRaisesRegex(ValueError, "exactly one path"): + _get_neighbor_stat_data(params, ["O", "H"]) + with self.assertRaisesRegex(ValueError, "exactly one path"): + _build_data_system(params, ["O", "H"]) + + def test_converted_lmdb_forwards_distributed_sharding(self) -> None: + """Converted and direct LMDB routes must pass identical DDP metadata.""" + params = { + "systems": "input.extxyz", + "format": "extxyz", + "batch_size": 2, + } + with ( + patch( + "deepmd.pt_expt.entrypoints.main.process_systems", + return_value=["converted.lmdb"], + ), + patch( + "deepmd.pt_expt.entrypoints.main.validate_lmdb_systems", + return_value="converted.lmdb", + ), + patch("deepmd.pt_expt.entrypoints.main.LmdbDataSystem") as lmdb_data_system, + ): + _build_data_system( + params, + ["O", "H"], + seed=7, + rank=2, + world_size=4, + ) + + lmdb_data_system.assert_called_once_with( + lmdb_path="converted.lmdb", + type_map=["O", "H"], + batch_size=2, + auto_prob_style=None, + seed=7, + rank=2, + world_size=4, + ) + + def _make_frame(natoms: int, seed: int, *, include_spin: bool = False) -> dict: """Synthetic LMDB frame matching the on-disk schema used by LmdbDataReader.""" rng = np.random.RandomState(seed)