diff --git a/deepmd/dpmodel/atomic_model/base_atomic_model.py b/deepmd/dpmodel/atomic_model/base_atomic_model.py index a894eac142..f2f2218443 100644 --- a/deepmd/dpmodel/atomic_model/base_atomic_model.py +++ b/deepmd/dpmodel/atomic_model/base_atomic_model.py @@ -758,7 +758,7 @@ def change_out_bias( delta_bias, out_std = compute_output_stats( sample_merged, self.get_ntypes(), - keys=list(self.atomic_output_def().keys()), + keys=self.bias_keys, stat_file_path=stat_file_path, model_forward=self._get_forward_wrapper_func(), rcond=self.rcond, @@ -771,7 +771,7 @@ def change_out_bias( bias_out, std_out = compute_output_stats( sample_merged, self.get_ntypes(), - keys=list(self.atomic_output_def().keys()), + keys=self.bias_keys, stat_file_path=stat_file_path, rcond=self.rcond, preset_bias=self.preset_out_bias, diff --git a/deepmd/dpmodel/descriptor/descriptor.py b/deepmd/dpmodel/descriptor/descriptor.py index ef605062f2..1dce8fb7d1 100644 --- a/deepmd/dpmodel/descriptor/descriptor.py +++ b/deepmd/dpmodel/descriptor/descriptor.py @@ -110,6 +110,37 @@ def get_stats(self) -> dict[str, StatItem]: """Get the statistics of the descriptor.""" raise NotImplementedError + def set_stat_mean_and_stddev( + self, + mean: Array, + stddev: Array, + ) -> None: + """Update the normalization arrays of the descriptor block. + + Parameters + ---------- + mean + Mean of the environment matrix. + stddev + Standard deviation of the environment matrix. + + Returns + ------- + None + """ + self["davg"] = mean + self["dstd"] = stddev + + def get_stat_mean_and_stddev(self) -> tuple[Array, Array]: + """Return the normalization arrays of the descriptor block. + + Returns + ------- + tuple[Array, Array] + Mean and standard deviation of the environment matrix. + """ + return self["davg"], self["dstd"] + def share_params( self, base_class: Any, shared_level: Any, resume: bool = False ) -> None: diff --git a/deepmd/dpmodel/descriptor/dpa1.py b/deepmd/dpmodel/descriptor/dpa1.py index eb8eeccbf6..dcb499fe4e 100644 --- a/deepmd/dpmodel/descriptor/dpa1.py +++ b/deepmd/dpmodel/descriptor/dpa1.py @@ -1461,15 +1461,7 @@ def compute_input_stats( env_mat_stat = EnvMatStatSe(self, use_graph=True) if path is not None: path = path / env_mat_stat.get_hash() - if path is None or not path.is_dir(): - if callable(merged): - # only get data for once - sampled = merged() - else: - sampled = merged - else: - sampled = [] - env_mat_stat.load_or_compute_stats(sampled, path) + env_mat_stat.load_or_compute_stats(merged, path) self.stats = env_mat_stat.stats mean, stddev = env_mat_stat() xp = array_api_compat.array_namespace(self.stddev) diff --git a/deepmd/dpmodel/descriptor/repflows.py b/deepmd/dpmodel/descriptor/repflows.py index 627184dd05..7c4c7e968a 100644 --- a/deepmd/dpmodel/descriptor/repflows.py +++ b/deepmd/dpmodel/descriptor/repflows.py @@ -494,15 +494,7 @@ def compute_input_stats( env_mat_stat = EnvMatStatSe(self) if path is not None: path = path / env_mat_stat.get_hash() - if path is None or not path.is_dir(): - if callable(merged): - # only get data for once - sampled = merged() - else: - sampled = merged - else: - sampled = [] - env_mat_stat.load_or_compute_stats(sampled, path) + env_mat_stat.load_or_compute_stats(merged, path) self.stats = env_mat_stat.stats mean, stddev = env_mat_stat() xp = array_api_compat.array_namespace(self.stddev) diff --git a/deepmd/dpmodel/descriptor/repformers.py b/deepmd/dpmodel/descriptor/repformers.py index 09d4309bd5..7d7b46d53a 100644 --- a/deepmd/dpmodel/descriptor/repformers.py +++ b/deepmd/dpmodel/descriptor/repformers.py @@ -459,15 +459,7 @@ def compute_input_stats( env_mat_stat = EnvMatStatSe(self) if path is not None: path = path / env_mat_stat.get_hash() - if path is None or not path.is_dir(): - if callable(merged): - # only get data for once - sampled = merged() - else: - sampled = merged - else: - sampled = [] - env_mat_stat.load_or_compute_stats(sampled, path) + env_mat_stat.load_or_compute_stats(merged, path) self.stats = env_mat_stat.stats mean, stddev = env_mat_stat() xp = array_api_compat.array_namespace(self.stddev) diff --git a/deepmd/dpmodel/descriptor/se_e2_a.py b/deepmd/dpmodel/descriptor/se_e2_a.py index 1030aa23ff..521dc3246c 100644 --- a/deepmd/dpmodel/descriptor/se_e2_a.py +++ b/deepmd/dpmodel/descriptor/se_e2_a.py @@ -349,15 +349,7 @@ def compute_input_stats( env_mat_stat = EnvMatStatSe(self) if path is not None: path = path / env_mat_stat.get_hash() - if path is None or not path.is_dir(): - if callable(merged): - # only get data for once - sampled = merged() - else: - sampled = merged - else: - sampled = [] - env_mat_stat.load_or_compute_stats(sampled, path) + env_mat_stat.load_or_compute_stats(merged, path) self.stats = env_mat_stat.stats mean, stddev = env_mat_stat() xp = array_api_compat.array_namespace(self.dstd) diff --git a/deepmd/dpmodel/descriptor/se_r.py b/deepmd/dpmodel/descriptor/se_r.py index 4c4d86b258..76c1680b41 100644 --- a/deepmd/dpmodel/descriptor/se_r.py +++ b/deepmd/dpmodel/descriptor/se_r.py @@ -328,15 +328,7 @@ def compute_input_stats( env_mat_stat = EnvMatStatSe(self) if path is not None: path = path / env_mat_stat.get_hash() - if path is None or not path.is_dir(): - if callable(merged): - # only get data for once - sampled = merged() - else: - sampled = merged - else: - sampled = [] - env_mat_stat.load_or_compute_stats(sampled, path) + env_mat_stat.load_or_compute_stats(merged, path) self.stats = env_mat_stat.stats mean, stddev = env_mat_stat() xp = array_api_compat.array_namespace(self.dstd) diff --git a/deepmd/dpmodel/descriptor/se_t.py b/deepmd/dpmodel/descriptor/se_t.py index 8eff7b81d1..f9395aa86b 100644 --- a/deepmd/dpmodel/descriptor/se_t.py +++ b/deepmd/dpmodel/descriptor/se_t.py @@ -303,15 +303,7 @@ def compute_input_stats( env_mat_stat = EnvMatStatSe(self) if path is not None: path = path / env_mat_stat.get_hash() - if path is None or not path.is_dir(): - if callable(merged): - # only get data for once - sampled = merged() - else: - sampled = merged - else: - sampled = [] - env_mat_stat.load_or_compute_stats(sampled, path) + env_mat_stat.load_or_compute_stats(merged, path) self.stats = env_mat_stat.stats mean, stddev = env_mat_stat() xp = array_api_compat.array_namespace(self.dstd) diff --git a/deepmd/dpmodel/descriptor/se_t_tebd.py b/deepmd/dpmodel/descriptor/se_t_tebd.py index cb174896cb..f8ff5d1955 100644 --- a/deepmd/dpmodel/descriptor/se_t_tebd.py +++ b/deepmd/dpmodel/descriptor/se_t_tebd.py @@ -791,15 +791,7 @@ def compute_input_stats( env_mat_stat = EnvMatStatSe(self) if path is not None: path = path / env_mat_stat.get_hash() - if path is None or not path.is_dir(): - if callable(merged): - # only get data for once - sampled = merged() - else: - sampled = merged - else: - sampled = [] - env_mat_stat.load_or_compute_stats(sampled, path) + env_mat_stat.load_or_compute_stats(merged, path) self.stats = env_mat_stat.stats mean, stddev = env_mat_stat() xp = array_api_compat.array_namespace(self.stddev) diff --git a/deepmd/dpmodel/fitting/general_fitting.py b/deepmd/dpmodel/fitting/general_fitting.py index bc820828ec..4474f9e6db 100644 --- a/deepmd/dpmodel/fitting/general_fitting.py +++ b/deepmd/dpmodel/fitting/general_fitting.py @@ -35,9 +35,6 @@ from deepmd.dpmodel.utils.seed import ( child_seed, ) -from deepmd.dpmodel.utils.stat import ( - _require_stat_file_items, -) from deepmd.env import ( GLOBAL_NP_FLOAT_PRECISION, ) @@ -51,6 +48,9 @@ from deepmd.utils.path import ( DPPath, ) +from deepmd.utils.stat_file import ( + load_required_items, +) from .base_fitting import ( BaseFitting, @@ -264,14 +264,10 @@ def compute_input_stats( return # stat fparam if self.numb_fparam > 0: - _require_stat_file_items(stat_file_path, ["fparam"]) - if ( - stat_file_path is not None - and stat_file_path.is_dir() - and (stat_file_path / "fparam").is_file() - ): - fparam_stats = self._load_param_stats_from_file( - stat_file_path, "fparam", self.numb_fparam + cached = load_required_items(stat_file_path, ["fparam"]) + if cached is not None: + fparam_stats = self._load_param_stats( + cached["fparam"], "fparam", self.numb_fparam ) else: sampled = merged() if callable(merged) else merged @@ -323,14 +319,10 @@ def compute_input_stats( ) # stat aparam if self.numb_aparam > 0: - _require_stat_file_items(stat_file_path, ["aparam"]) - if ( - stat_file_path is not None - and stat_file_path.is_dir() - and (stat_file_path / "aparam").is_file() - ): - aparam_stats = self._load_param_stats_from_file( - stat_file_path, "aparam", self.numb_aparam + cached = load_required_items(stat_file_path, ["aparam"]) + if cached is not None: + aparam_stats = self._load_param_stats( + cached["aparam"], "aparam", self.numb_aparam ) else: sampled = merged() if callable(merged) else merged @@ -402,14 +394,15 @@ def _save_param_stats_to_file( fp.save_numpy(arr) @staticmethod - def _load_param_stats_from_file( - stat_file_path: DPPath, + def _load_param_stats( + arr: np.ndarray, name: str, numb: int, ) -> list[StatItem]: - fp = stat_file_path / name - arr = fp.load_numpy() - assert arr.shape == (numb, 3) + if arr.shape != (numb, 3): + raise ValueError( + f"Invalid {name} statistics shape {arr.shape}; expected ({numb}, 3)." + ) return [ StatItem(number=arr[ii][0], sum=arr[ii][1], squared_sum=arr[ii][2]) for ii in range(numb) diff --git a/deepmd/dpmodel/train/data.py b/deepmd/dpmodel/train/data.py index e26ed01d48..3e91dcfc53 100644 --- a/deepmd/dpmodel/train/data.py +++ b/deepmd/dpmodel/train/data.py @@ -14,6 +14,11 @@ Any, ) +from deepmd.utils.stat_file import ( + StatFileMode, + StatFileSpec, +) + from .trainer import ( DEFAULT_TASK_KEY, ) @@ -36,6 +41,23 @@ class TrainingTaskConfig: validation_data_params: Mapping[str, Any] | None stat_file: str | None valid_numb_batch: int + stat_file_mode: StatFileMode = "update" + + @property + def stat_file_spec(self) -> StatFileSpec: + """Return the unopened statistics-cache configuration. + + Returns + ------- + StatFileSpec + Validated cache path and access mode for this task. + + Raises + ------ + ValueError + If the cache configuration is invalid. + """ + return StatFileSpec(self.stat_file, self.stat_file_mode) def iter_training_task_configs( @@ -53,6 +75,7 @@ def iter_training_task_configs( validation_data_params=validation_data_params, stat_file=training_params.get("stat_file"), valid_numb_batch=_valid_numb_batch(validation_data_params), + stat_file_mode=training_params.get("stat_file_mode", "update"), ) return @@ -67,6 +90,7 @@ def iter_training_task_configs( validation_data_params=validation_data_params, stat_file=task_data_params.get("stat_file"), valid_numb_batch=_valid_numb_batch(validation_data_params), + stat_file_mode=task_data_params.get("stat_file_mode", "update"), ) diff --git a/deepmd/dpmodel/utils/env_mat_stat.py b/deepmd/dpmodel/utils/env_mat_stat.py index b6e8108bd1..8288ede2e7 100644 --- a/deepmd/dpmodel/utils/env_mat_stat.py +++ b/deepmd/dpmodel/utils/env_mat_stat.py @@ -89,28 +89,18 @@ def merge_env_stat( # Update base_obj stats for chaining base_obj.stats = merged_stats - # Update buffers in-place: davg/dstd (simple) or mean/stddev (blocks) - # mean/stddev are numpy arrays; convert to match the buffer's backend - if hasattr(base_obj, "davg"): - xp = array_api_compat.array_namespace(base_obj.dstd) - device = array_api_compat.device(base_obj.dstd) - if not getattr(base_obj, "set_davg_zero", False): - base_obj.davg[...] = xp.asarray( - mean, dtype=base_obj.davg.dtype, device=device - ) - base_obj.dstd[...] = xp.asarray( - stddev, dtype=base_obj.dstd.dtype, device=device - ) - elif hasattr(base_obj, "mean"): - xp = array_api_compat.array_namespace(base_obj.stddev) - device = array_api_compat.device(base_obj.stddev) - if not getattr(base_obj, "set_davg_zero", False): - base_obj.mean[...] = xp.asarray( - mean, dtype=base_obj.mean.dtype, device=device - ) - base_obj.stddev[...] = xp.asarray( - stddev, dtype=base_obj.stddev.dtype, device=device - ) + current_mean, current_stddev = base_obj.get_stat_mean_and_stddev() + xp = array_api_compat.array_namespace(current_stddev) + device = array_api_compat.device(current_stddev) + merged_mean = current_mean + if not getattr(base_obj, "set_davg_zero", False): + merged_mean = xp.asarray(mean, dtype=current_mean.dtype, device=device) + merged_stddev = xp.asarray( + stddev, + dtype=current_stddev.dtype, + device=device, + ) + base_obj.set_stat_mean_and_stddev(merged_mean, merged_stddev) class EnvMatStat(BaseEnvMatStat): diff --git a/deepmd/dpmodel/utils/stat.py b/deepmd/dpmodel/utils/stat.py index 536eb1e21f..de01a10792 100644 --- a/deepmd/dpmodel/utils/stat.py +++ b/deepmd/dpmodel/utils/stat.py @@ -25,39 +25,15 @@ from deepmd.utils.path import ( DPPath, ) +from deepmd.utils.stat_file import ( + load_paired_items, + load_required_items, + replace_paired_items, +) log = logging.getLogger(__name__) -def _require_stat_file_items( - stat_file_path: DPPath | None, - items: list[str], -) -> None: - """Require named statistics items when a cache is opened read-only. - - Parameters - ---------- - stat_file_path : DPPath | None - Statistics cache path. - items : list[str] - Relative item names required by the current statistics consumer. - - Raises - ------ - FileNotFoundError - If a read-only cache does not contain one or more required items. - """ - if stat_file_path is None or getattr(stat_file_path, "mode", None) != "r": - return - missing = [item for item in items if not (stat_file_path / item).is_file()] - if missing: - missing_items = ", ".join(repr(item) for item in missing) - raise FileNotFoundError( - f"Read-only statistics cache {stat_file_path} is missing " - f"required item(s): {missing_items}." - ) - - def collect_observed_types(sampled: list[dict], type_map: list[str]) -> list[str]: """Collect observed element types from sampled training data. @@ -93,15 +69,12 @@ def _restore_observed_type_from_file( stat_file_path: DPPath | None, ) -> list[str] | None: """Try to load observed_type from stat file.""" - if stat_file_path is None: + items = load_required_items(stat_file_path, ["observed_type"]) + if items is None: return None - _require_stat_file_items(stat_file_path, ["observed_type"]) - fp = stat_file_path / "observed_type" - if fp.is_file(): - arr = fp.load_numpy() - # Decode bytes back to str if stored as bytes (for h5py compatibility) - return [x.decode() if isinstance(x, bytes) else x for x in arr.tolist()] - return None + arr = items["observed_type"] + # HDF5 string datasets may return bytes instead of str. + return [x.decode() if isinstance(x, bytes) else x for x in arr.tolist()] def _save_observed_type_to_file( @@ -117,50 +90,34 @@ def _save_observed_type_to_file( def _restore_from_file( - stat_file_path: DPPath, + stat_file_path: DPPath | None, keys: list[str], ) -> tuple[dict | None, dict | None]: """Restore bias and std from stat file.""" - if stat_file_path is None: - return None, None - _require_stat_file_items( - stat_file_path, - [item for key in keys for item in (f"bias_atom_{key}", f"std_atom_{key}")], - ) - stat_files = [stat_file_path / f"bias_atom_{kk}" for kk in keys] - if all(not (ii.is_file()) for ii in stat_files): - return None, None - stat_files = [stat_file_path / f"std_atom_{kk}" for kk in keys] - if all(not (ii.is_file()) for ii in stat_files): + pairs = [(f"bias_atom_{key}", f"std_atom_{key}") for key in keys] + items = load_paired_items(stat_file_path, pairs) + if items is None: return None, None - - ret_bias = {} - ret_std = {} - for kk in keys: - fp = stat_file_path / f"bias_atom_{kk}" - if fp.is_file(): - ret_bias[kk] = fp.load_numpy() - for kk in keys: - fp = stat_file_path / f"std_atom_{kk}" - if fp.is_file(): - ret_std[kk] = fp.load_numpy() + cached_keys = [key for key in keys if f"bias_atom_{key}" in items] + ret_bias = {key: items[f"bias_atom_{key}"] for key in cached_keys} + ret_std = {key: items[f"std_atom_{key}"] for key in cached_keys} return ret_bias, ret_std def _save_to_file( stat_file_path: DPPath, + requested_keys: list[str], bias_out: dict, std_out: dict, ) -> None: """Save bias and std to stat file.""" assert stat_file_path is not None - stat_file_path.mkdir(exist_ok=True, parents=True) - for kk, vv in bias_out.items(): - fp = stat_file_path / f"bias_atom_{kk}" - fp.save_numpy(vv) - for kk, vv in std_out.items(): - fp = stat_file_path / f"std_atom_{kk}" - fp.save_numpy(vv) + pairs = [(f"bias_atom_{key}", f"std_atom_{key}") for key in requested_keys] + items = { + **{f"bias_atom_{key}": value for key, value in bias_out.items()}, + **{f"std_atom_{key}": value for key, value in std_out.items()}, + } + replace_paired_items(stat_file_path, pairs, items) def _post_process_stat( @@ -311,6 +268,7 @@ def compute_output_stats( # normalize keys to list keys = [keys] if isinstance(keys, str) else keys assert isinstance(keys, list) + requested_keys = list(keys) # try to restore the bias from stat file bias_atom_e, std_atom_e = _restore_from_file(stat_file_path, keys) @@ -425,7 +383,12 @@ def compute_output_stats( raise RuntimeError("Fail to compute stat.") if stat_file_path is not None: - _save_to_file(stat_file_path, bias_atom_e, std_atom_e) + _save_to_file( + stat_file_path, + requested_keys, + bias_atom_e, + std_atom_e, + ) return bias_atom_e, std_atom_e diff --git a/deepmd/jax/common.py b/deepmd/jax/common.py index aacf375a74..70f307ce72 100644 --- a/deepmd/jax/common.py +++ b/deepmd/jax/common.py @@ -28,6 +28,7 @@ ) from deepmd.jax.env import ( flax_version, + jax, jnp, nnx, ) @@ -210,6 +211,16 @@ def dpmodel_setattr(obj: nnx.Module, name: str, value: Any) -> tuple[bool, Any]: if name in getattr(obj, "_jax_skip_auto_convert_attrs", ()): return False, value + current = vars(obj).get(name) + if isinstance(current, nnx.Variable) and isinstance(value, (np.ndarray, jax.Array)): + if isinstance(value, np.ndarray): + value = to_jax_array(value) + if Version(flax_version) >= _FLAX_0_12: + current.set_value(value) + else: + current.value = value + return True, current + if ( isinstance(value, list) and name in getattr(obj, "_jax_data_list_attrs", ()) diff --git a/deepmd/jax/train/trainer.py b/deepmd/jax/train/trainer.py index 41c0628e69..ce2fb8d3f2 100644 --- a/deepmd/jax/train/trainer.py +++ b/deepmd/jax/train/trainer.py @@ -9,6 +9,7 @@ import shutil import time from collections.abc import ( + Callable, Mapping, ) from copy import ( @@ -104,6 +105,12 @@ from deepmd.utils.model_stat import ( make_stat_input, ) +from deepmd.utils.stat_file import ( + StatFileSpec, + open_stat_file, + run_stat_on_chief, + stat_file_specs_by_task, +) log = logging.getLogger(__name__) @@ -154,6 +161,21 @@ def __init__( else [DEFAULT_TASK_KEY] ) self.model_params_by_task = self._model_params_by_task(self.model_def_script) + stat_file_specs = {} + for model_key in self.model_keys: + task_training = ( + self.training_param["data_dict"][model_key] + if self.multi_task + else self.training_param + ) + stat_file_specs[model_key] = StatFileSpec( + task_training.get("stat_file"), + task_training.get("stat_file_mode", "update"), + ) + self.stat_file_specs = stat_file_specs_by_task( + stat_file_specs, + self.model_keys, + ) if init_model is not None or restart is not None: checkpoint_path = init_model if init_model is not None else restart @@ -468,6 +490,8 @@ def _setup_training( self.model_params_by_task[model_key].get("data_stat_nbatch", 10), ) + synchronize_model_state = False + merge_shared_statistics = False if self.init_model is None and self.restart is None: for model_key in self.model_keys: finetune_has_new_type = ( @@ -477,17 +501,23 @@ def _setup_training( and self.finetune_links[model_key].get_has_new_type() ) if self.finetune_model is None or finetune_has_new_type: - self.models[model_key].atomic_model.compute_or_load_stat( - self._sample_funcs[model_key] + self._run_on_chief( + lambda _model_key=model_key: self._initialize_stat(_model_key), + operation=f"statistics initialization for task {model_key!r}", ) + synchronize_model_state = True + merge_shared_statistics = True if self.finetune_model is not None: - self._apply_finetune() + self._run_on_chief( + self._apply_finetune, + operation="fine-tuning initialization", + ) + synchronize_model_state = True - self._share_model_params( - resume=self.init_model is not None - or self.restart is not None - or self.finetune_model is not None + self._synchronize_initial_model_state( + state_changed=synchronize_model_state, + merge_shared_statistics=merge_shared_statistics, ) for model_key in self.model_keys: @@ -574,6 +604,83 @@ def _apply_finetune(self) -> None: bias_adjust_mode=bias_mode, ) + def _initialize_stat(self, model_key: str) -> None: + """Initialize one model from its scoped statistics cache.""" + with open_stat_file( + self.stat_file_specs[model_key], + ) as stat_file_path: + self.models[model_key].atomic_model.compute_or_load_stat( + self._sample_funcs[model_key], + stat_file_path=stat_file_path, + ) + + def _run_on_chief( + self, + action: Callable[[], None], + *, + operation: str, + ) -> None: + """Run one statistics action on rank 0 and propagate failure status.""" + synchronize_failure: Callable[[bool], bool] | None = None + if self.rank_context.world_size > 1: + from jax.experimental import ( + multihost_utils, + ) + + def broadcast_failure(failed: bool) -> bool: + return bool( + np.asarray( + multihost_utils.broadcast_one_to_all( + np.asarray(failed, dtype=np.bool_), + is_source=self.rank_context.is_chief, + ) + ).item() + ) + + synchronize_failure = broadcast_failure + + run_stat_on_chief( + action, + is_chief=self.rank_context.is_chief, + synchronize_failure=synchronize_failure, + operation=operation, + ) + + def _broadcast_model_states(self) -> None: + """Broadcast complete model states from rank 0 to every JAX process.""" + if self.rank_context.world_size <= 1: + return + from jax.experimental import ( + multihost_utils, + ) + + for model_key in self.model_keys: + _, state = nnx.split(self.models[model_key]) + state = multihost_utils.broadcast_one_to_all( + state.to_pure_dict(), + is_source=self.rank_context.is_chief, + ) + nnx.update(self.models[model_key], state) + + def _synchronize_initial_model_state( + self, + *, + state_changed: bool, + merge_shared_statistics: bool, + ) -> None: + """Merge, broadcast, and bind initial multi-task model state.""" + has_shared_parameters = self.multi_task and bool(self.shared_links) + if has_shared_parameters and merge_shared_statistics: + self._run_on_chief( + lambda: self._share_model_params(resume=False), + operation="shared statistics merge", + ) + state_changed = True + if state_changed: + self._broadcast_model_states() + if has_shared_parameters: + self._share_model_params(resume=True) + def _share_model_params(self, *, resume: bool = False) -> None: """Apply multi-task shared_dict links to JAX model branches.""" if not self.multi_task or not self.shared_links: @@ -808,19 +915,7 @@ def _change_bias_after_training(self) -> None: self.model_keys, bias_adjust_mode="change-by-statistic", ) - if self.rank_context.world_size <= 1: - return - from jax.experimental import ( - multihost_utils, - ) - - for model_key in self.model_keys: - _, state = nnx.split(self.models[model_key]) - state = multihost_utils.broadcast_one_to_all( - state.to_pure_dict(), - is_source=self.rank_context.is_chief, - ) - nnx.update(self.models[model_key], state) + self._broadcast_model_states() def run_full_validation( self, diff --git a/deepmd/pt/entrypoints/main.py b/deepmd/pt/entrypoints/main.py index 98a354cba9..9e99a7f52f 100644 --- a/deepmd/pt/entrypoints/main.py +++ b/deepmd/pt/entrypoints/main.py @@ -11,7 +11,6 @@ Any, ) -import h5py import torch import torch.distributed as dist import torch.version @@ -92,68 +91,14 @@ get_data, process_systems, ) -from deepmd.utils.path import ( - DPPath, +from deepmd.utils.stat_file import ( + StatFileSpec, ) from deepmd.utils.summary import SummaryPrinter as BaseSummaryPrinter log = logging.getLogger(__name__) -def _prepare_stat_file_path( - stat_file: str | None, - stat_file_mode: str = "update", -) -> DPPath | None: - """Prepare a statistics cache with the requested access mode. - - Parameters - ---------- - stat_file - Path to an HDF5 statistics file or a directory-based statistics cache. - stat_file_mode - ``"update"`` creates the cache when needed and permits missing - statistics to be written. ``"read"`` requires an existing cache and - prevents all writes. - - Returns - ------- - DPPath or None - The prepared statistics path, or ``None`` when no cache is configured. - - Raises - ------ - FileNotFoundError - If ``stat_file_mode`` is ``"read"`` and the cache does not exist. - ValueError - If the access mode is invalid or read mode has no cache path. - """ - if stat_file_mode not in {"read", "update"}: - raise ValueError( - "`stat_file_mode` must be either 'read' or 'update', " - f"but received {stat_file_mode!r}." - ) - if stat_file is None: - if stat_file_mode == "read": - raise ValueError("`stat_file_mode='read'` requires `stat_file`.") - return None - - path = Path(stat_file) - if stat_file_mode == "read": - if not path.exists(): - raise FileNotFoundError( - f"Statistics cache {stat_file!r} does not exist in read mode." - ) - return DPPath(stat_file, "r") - - if not path.exists(): - if stat_file.endswith((".h5", ".hdf5")): - with h5py.File(stat_file, "w"): - pass - else: - path.mkdir() - return DPPath(stat_file, "a") - - def _update_changed_model_tensors( target_state_dict: dict[str, Any], source_state_dict: dict[str, Any], @@ -204,7 +149,9 @@ def prepare_trainer_input_single( rank: int = 0, seed: int | None = None, ) -> tuple[ - DpLoaderSet | LmdbDataset, DpLoaderSet | LmdbDataset | None, DPPath | None + DpLoaderSet | LmdbDataset, + DpLoaderSet | LmdbDataset | None, + StatFileSpec, ]: # get data modifier modifier = None @@ -219,14 +166,10 @@ def prepare_trainer_input_single( ) training_systems = training_dataset_params["systems"] - # stat files - if rank != 0: - stat_file_path_single = None - else: - stat_file_path_single = _prepare_stat_file_path( - data_dict_single.get("stat_file"), - data_dict_single.get("stat_file_mode", "update"), - ) + stat_file_spec = StatFileSpec( + data_dict_single.get("stat_file"), + data_dict_single.get("stat_file_mode", "update"), + ) rank_seed = [rank, seed % (2**32)] if seed is not None else None @@ -283,7 +226,7 @@ def _make_dp_loader_set( return ( train_data_single, validation_data_single, - stat_file_path_single, + stat_file_spec, ) rank = dist.get_rank() if dist.is_available() and dist.is_initialized() else 0 @@ -292,7 +235,7 @@ def _make_dp_loader_set( ( train_data, validation_data, - stat_file_path, + stat_file_spec, ) = prepare_trainer_input_single( config["model"], config["training"], @@ -300,12 +243,12 @@ def _make_dp_loader_set( seed=data_seed, ) else: - train_data, validation_data, stat_file_path = {}, {}, {} + train_data, validation_data, stat_file_spec = {}, {}, {} for model_key in config["model"]["model_dict"]: ( train_data[model_key], validation_data[model_key], - stat_file_path[model_key], + stat_file_spec[model_key], ) = prepare_trainer_input_single( config["model"]["model_dict"][model_key], config["training"]["data_dict"][model_key], @@ -316,7 +259,7 @@ def _make_dp_loader_set( trainer = training.Trainer( config, train_data, - stat_file_path=stat_file_path, + stat_file_spec=stat_file_spec, validation_data=validation_data, init_model=init_model, restart_model=restart_model, diff --git a/deepmd/pt/model/atomic_model/sezm_atomic_model.py b/deepmd/pt/model/atomic_model/sezm_atomic_model.py index f3e0914e47..e7014b6963 100644 --- a/deepmd/pt/model/atomic_model/sezm_atomic_model.py +++ b/deepmd/pt/model/atomic_model/sezm_atomic_model.py @@ -15,9 +15,6 @@ import numpy as np import torch -from deepmd.dpmodel.utils.stat import ( - _require_stat_file_items, -) from deepmd.pt.model.atomic_model.dp_atomic_model import ( DPAtomicModel, ) @@ -41,6 +38,9 @@ from deepmd.pt.utils.utils import ( to_torch_tensor, ) +from deepmd.utils.stat_file import ( + load_required_items, +) from deepmd.utils.version import ( check_version_compatibility, ) @@ -188,12 +188,9 @@ def _compute_or_load_dens_force_stat( ValueError If force labels are unavailable for SeZM `dens` statistics. """ - force_stat_path = ( - None if stat_file_path is None else stat_file_path / "rmsd_dforce" - ) - _require_stat_file_items(stat_file_path, ["rmsd_dforce"]) - if force_stat_path is not None and force_stat_path.is_file(): - force_rmsd = float(np.asarray(force_stat_path.load_numpy()).reshape(-1)[0]) + cached = load_required_items(stat_file_path, ["rmsd_dforce"]) + if cached is not None: + force_rmsd = float(np.asarray(cached["rmsd_dforce"]).reshape(-1)[0]) else: sampled = sampled_func() if callable(sampled_func) else sampled_func force_square_sum = 0.0 @@ -244,7 +241,8 @@ def _compute_or_load_dens_force_stat( "the global direct-force RMSD can be computed." ) force_rmsd = math.sqrt(force_square_sum / force_atom_count) - if force_stat_path is not None: + if stat_file_path is not None: + force_stat_path = stat_file_path / "rmsd_dforce" force_stat_path.save_numpy(np.asarray([force_rmsd], dtype=np.float64)) if force_rmsd <= 0.0: diff --git a/deepmd/pt/model/task/fitting.py b/deepmd/pt/model/task/fitting.py index 71b947fedb..d4918149c0 100644 --- a/deepmd/pt/model/task/fitting.py +++ b/deepmd/pt/model/task/fitting.py @@ -20,9 +20,6 @@ from deepmd.dpmodel.utils.seed import ( child_seed, ) -from deepmd.dpmodel.utils.stat import ( - _require_stat_file_items, -) from deepmd.pt.model.network.mlp import ( FittingNet, NetworkCollection, @@ -54,6 +51,9 @@ from deepmd.utils.path import ( DPPath, ) +from deepmd.utils.stat_file import ( + load_required_items, +) dtype = env.GLOBAL_PT_FLOAT_PRECISION device = env.DEVICE @@ -212,13 +212,7 @@ def restore_fparam_from_file(self, stat_file_path: DPPath) -> None: """ fp = stat_file_path / "fparam" arr = fp.load_numpy() - assert arr.shape == (self.numb_fparam, 3) - _fparam_stat = [] - for ii in range(self.numb_fparam): - _fparam_stat.append( - StatItem(number=arr[ii][0], sum=arr[ii][1], squared_sum=arr[ii][2]) - ) - self.stats["fparam"] = _fparam_stat + self._restore_param_stats("fparam", arr, self.numb_fparam) log.info(f"Load fparam stats from {fp}.") def restore_aparam_from_file(self, stat_file_path: DPPath) -> None: @@ -231,15 +225,24 @@ def restore_aparam_from_file(self, stat_file_path: DPPath) -> None: """ fp = stat_file_path / "aparam" arr = fp.load_numpy() - assert arr.shape == (self.numb_aparam, 3) - _aparam_stat = [] - for ii in range(self.numb_aparam): - _aparam_stat.append( - StatItem(number=arr[ii][0], sum=arr[ii][1], squared_sum=arr[ii][2]) - ) - self.stats["aparam"] = _aparam_stat + self._restore_param_stats("aparam", arr, self.numb_aparam) log.info(f"Load aparam stats from {fp}.") + def _restore_param_stats( + self, + name: str, + arr: np.ndarray, + dimension: int, + ) -> None: + if arr.shape != (dimension, 3): + raise ValueError( + f"Invalid {name} statistics shape {arr.shape}; " + f"expected ({dimension}, 3)." + ) + self.stats[name] = [ + StatItem(number=row[0], sum=row[1], squared_sum=row[2]) for row in arr + ] + def compute_input_stats( self, merged: Callable[[], list[dict]] | list[dict], @@ -272,13 +275,9 @@ def compute_input_stats( # stat fparam if self.numb_fparam > 0: - _require_stat_file_items(stat_file_path, ["fparam"]) - if ( - stat_file_path is not None - and stat_file_path.is_dir() - and (stat_file_path / "fparam").is_file() - ): - self.restore_fparam_from_file(stat_file_path) + cached = load_required_items(stat_file_path, ["fparam"]) + if cached is not None: + self._restore_param_stats("fparam", cached["fparam"], self.numb_fparam) else: sampled = merged() if callable(merged) else merged self.stats["fparam"] = [] @@ -311,13 +310,9 @@ def compute_input_stats( # stat aparam if self.numb_aparam > 0: - _require_stat_file_items(stat_file_path, ["aparam"]) - if ( - stat_file_path is not None - and stat_file_path.is_dir() - and (stat_file_path / "aparam").is_file() - ): - self.restore_aparam_from_file(stat_file_path) + cached = load_required_items(stat_file_path, ["aparam"]) + if cached is not None: + self._restore_param_stats("aparam", cached["aparam"], self.numb_aparam) else: sampled = merged() if callable(merged) else merged self.stats["aparam"] = [] diff --git a/deepmd/pt/train/training.py b/deepmd/pt/train/training.py index 8dd68c8cdb..ec92301b9d 100644 --- a/deepmd/pt/train/training.py +++ b/deepmd/pt/train/training.py @@ -8,6 +8,7 @@ Callable, Generator, Iterable, + Mapping, ) from contextlib import ( nullcontext, @@ -151,8 +152,11 @@ DataLoader, ) -from deepmd.utils.path import ( - DPH5Path, +from deepmd.utils.stat_file import ( + StatFileSpec, + open_stat_file, + run_stat_on_chief, + stat_file_specs_by_task, ) log = logging.getLogger(__name__) @@ -163,7 +167,7 @@ def __init__( self, config: dict[str, Any], training_data: DpLoaderSet, - stat_file_path: str | None = None, + stat_file_spec: StatFileSpec | Mapping[str, StatFileSpec] | None = None, validation_data: DpLoaderSet | None = None, init_model: str | None = None, restart_model: str | None = None, @@ -187,6 +191,7 @@ def __init__( else: resume_model = None resuming = resume_model is not None + has_initial_state = resuming or init_frz_model is not None self.restart_training = restart_model is not None model_params = config["model"] training_params = config["training"] @@ -202,7 +207,9 @@ def __init__( infer_env_defaults["DP_AMP_INFER"] = "1" self.multi_task = "model_dict" in model_params self.finetune_links = finetune_links - self.finetune_update_stat = False + finetune_updates_statistics = finetune_links is not None and any( + rule.get_has_new_type() for rule in finetune_links.values() + ) self.model_keys = ( list(model_params["model_dict"]) if self.multi_task else ["Default"] ) @@ -211,6 +218,10 @@ def __init__( self.world_size = dist.get_world_size() if self.is_distributed else 1 self.num_model = len(self.model_keys) self.model_prob = None + self.stat_file_specs = stat_file_specs_by_task( + stat_file_spec, + self.model_keys, + ) # Iteration config self.num_steps = training_params.get("numb_steps") @@ -413,7 +424,7 @@ def single_model_stat( _model: Any, _data_stat_nbatch: int, _training_data: DpLoaderSet, - _stat_file_path: str | None, + _stat_file_spec: StatFileSpec, finetune_has_new_type: bool = False, preset_observed_type: list[str] | None = None, ) -> Callable[[], Any]: @@ -426,14 +437,20 @@ def get_sample() -> Any: ) return sampled - if (not resuming or finetune_has_new_type) and self.rank == 0: - _model.compute_or_load_stat( - sampled_func=get_sample, - stat_file_path=_stat_file_path, - preset_observed_type=preset_observed_type, + if not has_initial_state or finetune_has_new_type: + + def initialize_statistics() -> None: + with open_stat_file(_stat_file_spec) as stat_file_path: + _model.compute_or_load_stat( + sampled_func=get_sample, + stat_file_path=stat_file_path, + preset_observed_type=preset_observed_type, + ) + + self._run_stat_on_chief( + initialize_statistics, + operation="statistics initialization", ) - if isinstance(_stat_file_path, DPH5Path): - _stat_file_path.root.close() return get_sample def get_lr(lr_params: dict[str, Any]) -> BaseLR: @@ -523,7 +540,7 @@ def get_lr(lr_params: dict[str, Any]) -> BaseLR: self.model, model_params.get("data_stat_nbatch", 10), training_data, - stat_file_path, + self.stat_file_specs["Default"], finetune_has_new_type=self.finetune_links["Default"].get_has_new_type() if self.finetune_links is not None else False, @@ -603,7 +620,7 @@ def get_lr(lr_params: dict[str, Any]) -> BaseLR: self.model[model_key], model_params["model_dict"][model_key].get("data_stat_nbatch", 10), training_data[model_key], - stat_file_path[model_key], + self.stat_file_specs[model_key], finetune_has_new_type=self.finetune_links[ model_key ].get_has_new_type() @@ -845,7 +862,6 @@ def get_lr(lr_params: dict[str, Any]) -> BaseLR: ): model_with_new_type_stat = None if finetune_rule_single.get_has_new_type(): - self.finetune_update_stat = True model_with_new_type_stat = self.wrapper.model[model_key] pretrained_model_wrapper.model[ _model_key_from @@ -939,44 +955,45 @@ def collect_single_finetune_params( state_dict["_extra_state"] = self.wrapper.state_dict()["_extra_state"] self.wrapper.load_state_dict(state_dict) - # change bias for fine-tuning - if finetune_model is not None: + if finetune_model is not None: + for model_key in self.model_keys: + finetune_rule = self.finetune_links[model_key] + if self.multi_task and finetune_rule.get_resuming(): + if self.rank == 0: + log.info("Model branch %s will resume training.", model_key) + continue + if self.multi_task and self.rank == 0: + log.info( + "Model branch %s will be fine-tuned. " + "This may take a long time...", + model_key, + ) - def single_model_finetune( - _model: Any, - _finetune_rule_single: Any, - _sample_func: Callable, - ) -> Any: - _model = model_change_out_bias( - _model, - _sample_func, + def update_finetune_bias( + _model_key: str = model_key, + _finetune_rule: Any = finetune_rule, + ) -> None: + model = ( + self.model[_model_key] if self.multi_task else self.model + ) + model = model_change_out_bias( + model, + self.get_sample_func[_model_key] + if self.multi_task + else self.get_sample_func, _bias_adjust_mode="change-by-statistic" - if not _finetune_rule_single.get_random_fitting() + if not _finetune_rule.get_random_fitting() else "set-by-statistic", ) - return _model + if self.multi_task: + self.model[_model_key] = model + else: + self.model = model - if not self.multi_task: - finetune_rule_single = self.finetune_links["Default"] - self.model = single_model_finetune( - self.model, finetune_rule_single, self.get_sample_func - ) - else: - for model_key in self.model_keys: - finetune_rule_single = self.finetune_links[model_key] - if not finetune_rule_single.get_resuming(): - log.info( - f"Model branch {model_key} will be fine-tuned. This may take a long time..." - ) - self.model[model_key] = single_model_finetune( - self.model[model_key], - finetune_rule_single, - self.get_sample_func[model_key], - ) - else: - log.info( - f"Model branch {model_key} will resume training." - ) + self._run_stat_on_chief( + update_finetune_bias, + operation=f"fine-tuning statistics for task {model_key!r}", + ) if init_frz_model is not None: frz_model = torch.jit.load(init_frz_model, map_location=DEVICE) @@ -998,11 +1015,25 @@ def single_model_finetune( assert np.allclose(_data_stat_protect, _data_stat_protect[0]), ( "Model key 'data_stat_protect' must be the same in each branch when multitask!" ) + share_kwargs = { + "model_key_prob_map": dict( + zip(self.model_keys, self.model_prob, strict=True) + ), + "data_stat_protect": _data_stat_protect[0], + } + if not has_initial_state or finetune_updates_statistics: + self._run_stat_on_chief( + lambda: self.wrapper.share_params( + shared_links, + resume=False, + **share_kwargs, + ), + operation="shared statistics merge", + ) self.wrapper.share_params( shared_links, - resume=(resuming and not self.finetune_update_stat) or self.rank != 0, - model_key_prob_map=dict(zip(self.model_keys, self.model_prob)), - data_stat_protect=_data_stat_protect[0], + resume=True, + **share_kwargs, ) # LoRA injection (single-task only; argcheck rejects multi-task). @@ -1169,6 +1200,30 @@ def single_model_finetune( if self.rank == 0: self._log_parameter_count() + def _run_stat_on_chief( + self, + action: Callable[[], None], + *, + operation: str, + ) -> None: + """Run a statistics action on rank 0 and synchronize its outcome.""" + synchronize_failure: Callable[[bool], bool] | None = None + if self.is_distributed: + + def broadcast_failure(failed: bool) -> bool: + holder = [failed if self.rank == 0 else False] + dist.broadcast_object_list(holder, src=0, device=DEVICE) + return bool(holder[0]) + + synchronize_failure = broadcast_failure + + run_stat_on_chief( + action, + is_chief=self.rank == 0, + synchronize_failure=synchronize_failure, + operation=operation, + ) + def _broadcast_value_from_rank0(self, value: Any) -> Any: """Return rank 0's copy of ``value`` on every rank. diff --git a/deepmd/pt/utils/stat.py b/deepmd/pt/utils/stat.py index b268ca2ba7..e0ef13de4c 100644 --- a/deepmd/pt/utils/stat.py +++ b/deepmd/pt/utils/stat.py @@ -32,12 +32,15 @@ from deepmd.utils.path import ( DPPath, ) +from deepmd.utils.stat_file import ( + load_paired_items, + replace_paired_items, +) log = logging.getLogger(__name__) # Re-export from dpmodel (backend-agnostic implementations) from deepmd.dpmodel.utils.stat import ( - _require_stat_file_items, _restore_observed_type_from_file, _save_observed_type_to_file, collect_observed_types, @@ -109,50 +112,32 @@ def make_stat_input( def _restore_from_file( - stat_file_path: DPPath, - keys: list[str] = ["energy"], -) -> dict | None: - if stat_file_path is None: - return None, None - _require_stat_file_items( - stat_file_path, - [item for key in keys for item in (f"bias_atom_{key}", f"std_atom_{key}")], - ) - stat_files = [stat_file_path / f"bias_atom_{kk}" for kk in keys] - if all(not (ii.is_file()) for ii in stat_files): - return None, None - stat_files = [stat_file_path / f"std_atom_{kk}" for kk in keys] - if all(not (ii.is_file()) for ii in stat_files): + stat_file_path: DPPath | None, + keys: list[str], +) -> tuple[dict | None, dict | None]: + pairs = [(f"bias_atom_{key}", f"std_atom_{key}") for key in keys] + items = load_paired_items(stat_file_path, pairs) + if items is None: return None, None - - ret_bias = {} - ret_std = {} - for kk in keys: - fp = stat_file_path / f"bias_atom_{kk}" - # only read the key that exists - if fp.is_file(): - ret_bias[kk] = fp.load_numpy() - for kk in keys: - fp = stat_file_path / f"std_atom_{kk}" - # only read the key that exists - if fp.is_file(): - ret_std[kk] = fp.load_numpy() + cached_keys = [key for key in keys if f"bias_atom_{key}" in items] + ret_bias = {key: items[f"bias_atom_{key}"] for key in cached_keys} + ret_std = {key: items[f"std_atom_{key}"] for key in cached_keys} return ret_bias, ret_std def _save_to_file( stat_file_path: DPPath, + requested_keys: list[str], bias_out: dict, std_out: dict, ) -> None: assert stat_file_path is not None - stat_file_path.mkdir(exist_ok=True, parents=True) - for kk, vv in bias_out.items(): - fp = stat_file_path / f"bias_atom_{kk}" - fp.save_numpy(vv) - for kk, vv in std_out.items(): - fp = stat_file_path / f"std_atom_{kk}" - fp.save_numpy(vv) + pairs = [(f"bias_atom_{key}", f"std_atom_{key}") for key in requested_keys] + items = { + **{f"bias_atom_{key}": value for key, value in bias_out.items()}, + **{f"std_atom_{key}": value for key, value in std_out.items()}, + } + replace_paired_items(stat_file_path, pairs, items) def _post_process_stat( @@ -312,6 +297,10 @@ def compute_output_stats( intensive : bool, optional Whether the fitting target is intensive. """ + keys = [keys] if isinstance(keys, str) else keys + assert isinstance(keys, list) + requested_keys = list(keys) + # try to restore the bias from stat file bias_atom_e, std_atom_e = _restore_from_file(stat_file_path, keys) @@ -325,14 +314,11 @@ def compute_output_stats( model_pred = None # remove the keys that are not in the sample - keys = [keys] if isinstance(keys, str) else keys - assert isinstance(keys, list) new_keys = [ ii for ii in keys if (ii in sampled[0].keys()) or ("atom_" + ii in sampled[0].keys()) ] - del keys keys = new_keys # split system based on label atomic_sampled_idx = defaultdict(list) @@ -428,7 +414,12 @@ def compute_output_stats( raise RuntimeError("Fail to compute stat.") if stat_file_path is not None: - _save_to_file(stat_file_path, bias_atom_e, std_atom_e) + _save_to_file( + stat_file_path, + requested_keys, + bias_atom_e, + std_atom_e, + ) bias_atom_e = {kk: to_torch_tensor(vv) for kk, vv in bias_atom_e.items()} std_atom_e = {kk: to_torch_tensor(vv) for kk, vv in std_atom_e.items()} diff --git a/deepmd/pt_expt/entrypoints/main.py b/deepmd/pt_expt/entrypoints/main.py index 6b394b3229..06fd41af0c 100644 --- a/deepmd/pt_expt/entrypoints/main.py +++ b/deepmd/pt_expt/entrypoints/main.py @@ -14,8 +14,6 @@ Any, ) -import h5py - from deepmd.dpmodel.train import ( AbstractTrainEntrypoint, TrainEntrypointOptions, @@ -38,8 +36,8 @@ get_data, process_systems, ) -from deepmd.utils.path import ( - DPPath, +from deepmd.utils.stat_file import ( + StatFileSpec, ) from deepmd.utils.summary import SummaryPrinter as BaseSummaryPrinter @@ -164,21 +162,6 @@ def _build_data_system( ) -def _ensure_stat_file_path(stat_file_path: str | None) -> DPPath | None: - """Create a stat-file target and return a DPPath wrapper.""" - if stat_file_path is None: - return None - path = Path(stat_file_path) - if not path.exists(): - if stat_file_path.endswith((".h5", ".hdf5")): - path.parent.mkdir(parents=True, exist_ok=True) - with h5py.File(path, "w"): - pass - else: - path.mkdir(parents=True, exist_ok=True) - return DPPath(stat_file_path, "a") - - def get_trainer( config: dict[str, Any], init_model: str | None = None, @@ -195,7 +178,7 @@ def get_trainer( def factory( task_config: TrainingTaskConfig, - ) -> tuple[DeepmdDataSystem | LmdbDataSystem, Any | None, DPPath | None]: + ) -> tuple[DeepmdDataSystem | LmdbDataSystem, Any | None, StatFileSpec]: type_map = list(task_config.model_params["type_map"]) train_data = _build_data_system( dict(task_config.training_data_params), type_map, seed=data_seed @@ -208,27 +191,27 @@ def factory( return ( train_data, validation_data, - _ensure_stat_file_path(task_config.stat_file), + task_config.stat_file_spec, ) - train_data_map, validation_data_map, stat_file_path_map = make_task_maps( + train_data_map, validation_data_map, stat_file_spec_map = make_task_maps( config, factory ) print_data_summaries(train_data_map, validation_data_map) if multi_task: train_data = train_data_map validation_data = validation_data_map - stat_file_path = stat_file_path_map + stat_file_spec = stat_file_spec_map else: task_key = next(iter(train_data_map)) train_data = train_data_map[task_key] validation_data = validation_data_map[task_key] - stat_file_path = stat_file_path_map[task_key] + stat_file_spec = stat_file_spec_map[task_key] trainer = training.Trainer( config, train_data, - stat_file_path=stat_file_path, + stat_file_spec=stat_file_spec, validation_data=validation_data, init_model=init_model, restart_model=restart_model, diff --git a/deepmd/pt_expt/train/training.py b/deepmd/pt_expt/train/training.py index 697bf46dfb..1129696ed1 100644 --- a/deepmd/pt_expt/train/training.py +++ b/deepmd/pt_expt/train/training.py @@ -10,6 +10,10 @@ import logging import os import time +from collections.abc import ( + Callable, + Mapping, +) from copy import ( deepcopy, ) @@ -87,8 +91,11 @@ from deepmd.utils.finetune import ( warn_configuration_mismatch_during_finetune, ) -from deepmd.utils.path import ( - DPPath, +from deepmd.utils.stat_file import ( + StatFileSpec, + open_stat_file, + run_stat_on_chief, + stat_file_specs_by_task, ) log = logging.getLogger(__name__) @@ -1316,8 +1323,8 @@ class Trainer(AbstractTrainer): Full training configuration. training_data : DeepmdDataSystem or dict Training data. Dict of ``{model_key: DeepmdDataSystem}`` for multi-task. - stat_file_path : DPPath or dict or None - Path for saving / loading statistics. + stat_file_spec : StatFileSpec or dict or None + Unopened statistics-cache configuration. validation_data : DeepmdDataSystem or dict or None Validation data. init_model : str or None @@ -1332,7 +1339,7 @@ def __init__( self, config: dict[str, Any], training_data: DeepmdDataSystem | dict, - stat_file_path: DPPath | dict | None = None, + stat_file_spec: StatFileSpec | Mapping[str, StatFileSpec] | None = None, validation_data: DeepmdDataSystem | dict | None = None, init_model: str | None = None, restart_model: str | None = None, @@ -1378,10 +1385,9 @@ def __init__( multi_task=self.multi_task, model_keys=self.model_keys, ) - self.stat_file_path_by_task = _as_task_map( - stat_file_path, - multi_task=self.multi_task, - model_keys=self.model_keys, + self.stat_file_specs = stat_file_specs_by_task( + stat_file_spec, + self.model_keys, ) # Distributed training detection @@ -1476,7 +1482,6 @@ def __init__( for model_key in self.model_keys: _nbatch = self.model_params_by_task[model_key].get("data_stat_nbatch", 10) _data = self.training_data_by_task[model_key] - _stat_path = self.stat_file_path_by_task[model_key] @functools.lru_cache def _make_sample( @@ -1494,14 +1499,26 @@ def _make_sample( ) if _finetune_has_new_type: self._finetune_update_stat = True - if (not resuming or _finetune_has_new_type) and self.rank == 0: - self.models[model_key].compute_or_load_stat( - sampled_func=_make_sample, - stat_file_path=_stat_path, + if not resuming or _finetune_has_new_type: + + def initialize_statistics( + _model_key: str = model_key, + _sample_func: Callable[[], list[dict[str, np.ndarray]]] = ( + _make_sample + ), + ) -> None: + with open_stat_file( + self.stat_file_specs[_model_key] + ) as stat_file_path: + self.models[_model_key].compute_or_load_stat( + sampled_func=_sample_func, + stat_file_path=stat_file_path, + ) + + self._run_stat_on_chief( + initialize_statistics, + operation=f"statistics initialization for task {model_key!r}", ) - if self.is_distributed: - for model_key in self.model_keys: - self._broadcast_model_stat(self.models[model_key]) # Model probability (multi-task) -------------------------------------- if self.multi_task: @@ -1531,6 +1548,7 @@ def _make_sample( # Shared params (multi-task) ------------------------------------------ self._shared_links = shared_links + synchronize_model_state = not resuming or self._finetune_update_stat if shared_links is not None: _data_stat_protect = np.array( [ @@ -1542,13 +1560,31 @@ def _make_sample( raise ValueError( "Model key 'data_stat_protect' must be the same in each branch when multitask!" ) - self.wrapper.share_params( - shared_links, - resume=(resuming and not self._finetune_update_stat) or self.rank != 0, - model_key_prob_map=dict( + share_kwargs = { + "model_key_prob_map": dict( zip(self.model_keys, self.model_prob, strict=True) ), - data_stat_protect=_data_stat_protect[0], + "data_stat_protect": _data_stat_protect[0], + } + if synchronize_model_state: + self._run_stat_on_chief( + lambda: self.wrapper.share_params( + shared_links, + resume=False, + **share_kwargs, + ), + operation="shared statistics merge", + ) + + if synchronize_model_state and self.is_distributed: + for model_key in self.model_keys: + self._broadcast_model_stat(self.models[model_key]) + + if shared_links is not None: + self.wrapper.share_params( + shared_links, + resume=True, + **share_kwargs, ) # DDP wrapping -------------------------------------------------------- @@ -1755,12 +1791,21 @@ def _make_sample( if not finetune_rule.get_random_fitting() else "set-by-statistic" ) - if self.rank == 0: - self.models[model_key] = model_change_out_bias( - self.models[model_key], - self._sample_funcs[model_key], - _bias_adjust_mode=bias_mode, + + def update_finetune_bias( + _model_key: str = model_key, + _bias_mode: str = bias_mode, + ) -> None: + self.models[_model_key] = model_change_out_bias( + self.models[_model_key], + self._sample_funcs[_model_key], + _bias_adjust_mode=_bias_mode, ) + + self._run_stat_on_chief( + update_finetune_bias, + operation=f"fine-tuning statistics for task {model_key!r}", + ) if self.is_distributed: self._broadcast_model_stat(self.models[model_key]) self.model = ( @@ -2091,6 +2136,30 @@ def _unwrapped(self) -> "ModelWrapper": return self.wrapper.module return self.wrapper + def _run_stat_on_chief( + self, + action: Callable[[], None], + *, + operation: str, + ) -> None: + """Run a statistics action on rank 0 and synchronize its outcome.""" + synchronize_failure: Callable[[bool], bool] | None = None + if self.is_distributed: + + def broadcast_failure(failed: bool) -> bool: + holder = [failed if self.rank == 0 else False] + dist.broadcast_object_list(holder, src=0, device=DEVICE) + return bool(holder[0]) + + synchronize_failure = broadcast_failure + + run_stat_on_chief( + action, + is_chief=self.rank == 0, + synchronize_failure=synchronize_failure, + operation=operation, + ) + @staticmethod def _broadcast_model_stat(model: torch.nn.Module) -> None: """Broadcast model parameters and buffers from rank 0 to all ranks.""" diff --git a/deepmd/tf2/entrypoints/train.py b/deepmd/tf2/entrypoints/train.py index 873c7df89a..a026bcc03f 100644 --- a/deepmd/tf2/entrypoints/train.py +++ b/deepmd/tf2/entrypoints/train.py @@ -7,15 +7,11 @@ import logging import time -from pathlib import ( - Path, -) from typing import ( + TYPE_CHECKING, Any, ) -import h5py - from deepmd.dpmodel.model.base_model import ( BaseModel, ) @@ -40,11 +36,13 @@ from deepmd.utils.data_system import ( get_data, ) -from deepmd.utils.path import ( - DPPath, -) from deepmd.utils.summary import SummaryPrinter as BaseSummaryPrinter +if TYPE_CHECKING: + from deepmd.utils.stat_file import ( + StatFileSpec, + ) + __all__ = ["train", "update_sel"] log = logging.getLogger(__name__) @@ -157,7 +155,7 @@ def run_training( def factory( task_config: TrainingTaskConfig, - ) -> tuple[Any, Any | None, DPPath | None]: + ) -> tuple[Any, Any | None, StatFileSpec]: type_map = list(task_config.model_params.get("type_map", [])) ipt_type_map = type_map if type_map else None train_data = get_data( @@ -174,9 +172,9 @@ def factory( train_data.type_map, None, ) - return train_data, valid_data, _make_stat_file_path(task_config.stat_file) + return train_data, valid_data, task_config.stat_file_spec - train_data_map, valid_data_map, stat_file_map = make_task_maps( + train_data_map, valid_data_map, stat_file_spec_map = make_task_maps( config, factory, ) @@ -185,7 +183,7 @@ def factory( trainer = DPTrainer( config, train_data_map, - stat_file_path=stat_file_map, + stat_file_spec=stat_file_spec_map, validation_data=valid_data_map, init_model=options.init_model, restart_model=options.restart, @@ -263,17 +261,3 @@ def update_sel( jdata_cpy["model"] = updated_model return jdata_cpy, task_min_nbor_dist return jdata_cpy, min_nbor_dist - - -def _make_stat_file_path(stat_file_raw: str | None) -> DPPath | None: - if stat_file_raw is None: - return None - stat_file_target = Path(stat_file_raw) - stat_file_target.parent.mkdir(parents=True, exist_ok=True) - if not stat_file_target.exists(): - if stat_file_raw.endswith((".h5", ".hdf5")): - with h5py.File(stat_file_raw, "w"): - pass - else: - stat_file_target.mkdir(parents=True, exist_ok=True) - return DPPath(stat_file_raw, "a") diff --git a/deepmd/tf2/train/trainer.py b/deepmd/tf2/train/trainer.py index f631c36620..6fde31e8f8 100644 --- a/deepmd/tf2/train/trainer.py +++ b/deepmd/tf2/train/trainer.py @@ -86,14 +86,16 @@ from deepmd.utils.model_stat import ( make_stat_input, ) +from deepmd.utils.stat_file import ( + StatFileSpec, + open_stat_file, + stat_file_specs_by_task, +) if TYPE_CHECKING: from deepmd.utils.data_system import ( DeepmdDataSystem, ) - from deepmd.utils.path import ( - DPPath, - ) log = logging.getLogger(__name__) @@ -211,7 +213,7 @@ def __init__( self, config: dict[str, Any], training_data: DeepmdDataSystem | Mapping[str, DeepmdDataSystem], - stat_file_path: DPPath | Mapping[str, DPPath | None] | None = None, + stat_file_spec: StatFileSpec | Mapping[str, StatFileSpec] | None = None, validation_data: DeepmdDataSystem | Mapping[str, DeepmdDataSystem | None] | None = None, @@ -267,10 +269,9 @@ def __init__( multi_task=self.multi_task, model_keys=self.model_keys, ) - self.stat_file_path_by_task = _as_task_map( - stat_file_path, - multi_task=self.multi_task, - model_keys=self.model_keys, + self.stat_file_specs = stat_file_specs_by_task( + stat_file_spec, + self.model_keys, ) self.num_steps = int(training_params["numb_steps"]) @@ -361,10 +362,11 @@ def sample( "data stating for task %s... (this step may take long time)", model_key, ) - self.models[model_key].compute_or_load_stat( - self._sample_funcs[model_key], - stat_file_path=self.stat_file_path_by_task[model_key], - ) + with open_stat_file(self.stat_file_specs[model_key]) as stat_file_path: + self.models[model_key].compute_or_load_stat( + self._sample_funcs[model_key], + stat_file_path=stat_file_path, + ) if self.finetune_model is not None: self._apply_finetune() diff --git a/deepmd/utils/argcheck.py b/deepmd/utils/argcheck.py index 2b0d5b9132..fc8c11b635 100644 --- a/deepmd/utils/argcheck.py +++ b/deepmd/utils/argcheck.py @@ -5305,7 +5305,8 @@ def training_args( "otherwise, a directory containing NumPy binary files are used." ) doc_stat_file_mode = ( - doc_only_pt_supported + "The access mode for `stat_file`. " + "Supported by the PyTorch, JAX, TensorFlow 2, and experimental PyTorch " + "backends. The access mode for `stat_file`. " "`update` creates the cache when needed and writes any missing statistics; " "this is the behavior used when the option is omitted. " "`read` requires a complete existing cache and opens it read-only, allowing " diff --git a/deepmd/utils/stat_file.py b/deepmd/utils/stat_file.py new file mode 100644 index 0000000000..98e02f72e7 --- /dev/null +++ b/deepmd/utils/stat_file.py @@ -0,0 +1,629 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""Scoped access to persistent training-statistics caches.""" + +from __future__ import ( + annotations, +) + +from collections.abc import ( + Mapping, +) +from contextlib import ( + contextmanager, +) +from dataclasses import ( + dataclass, +) +from pathlib import ( + Path, +) +from typing import ( + TYPE_CHECKING, + Any, + Literal, +) + +import h5py +import numpy as np +from wcmatch.glob import ( + globfilter, +) + +from deepmd.utils.path import ( + DPH5Path, + DPOSPath, + DPPath, +) + +if TYPE_CHECKING: + from collections.abc import ( + Callable, + Iterator, + Sequence, + ) + +StatFileMode = Literal["read", "update"] + +_HDF5_SUFFIXES = {".h5", ".hdf5"} +_PAIR_TRANSACTION_MARKER = "__deepmd_output_stat_transaction__" + + +@dataclass(frozen=True) +class StatFileSpec: + """Describe a statistics cache without opening it. + + Parameters + ---------- + path + Cache path. ``None`` disables persistent statistics storage. + mode + ``read`` opens an existing cache read-only. ``update`` creates the + cache when necessary and permits statistics writers to update it. + + Raises + ------ + ValueError + If the mode is invalid, the path is empty, or read mode has no path. + """ + + path: str | None + mode: StatFileMode = "update" + + def __post_init__(self) -> None: + if self.mode not in {"read", "update"}: + raise ValueError( + "`stat_file_mode` must be either 'read' or 'update', " + f"but received {self.mode!r}." + ) + if self.path is not None and not self.path.strip(): + raise ValueError("`stat_file` must not be empty.") + if self.path is None and self.mode == "read": + raise ValueError("`stat_file_mode='read'` requires `stat_file`.") + + +def stat_file_specs_by_task( + spec: StatFileSpec | Mapping[str, StatFileSpec] | None, + task_names: Sequence[str], +) -> dict[str, StatFileSpec]: + """Normalize statistics-cache configuration by task. + + Parameters + ---------- + spec + One single-task specification, a mapping for multi-task training, or + ``None`` when persistent caching is disabled. + task_names + Ordered task names owned by the trainer. + + Returns + ------- + dict[str, StatFileSpec] + One validated specification per task. + + Raises + ------ + TypeError + If a single specification is supplied for multiple tasks. + KeyError + If a task is absent from a supplied mapping. + ValueError + If multiple tasks reference the same physical cache path. + """ + if spec is None: + specs = {task: StatFileSpec(None) for task in task_names} + elif isinstance(spec, Mapping): + specs = {task: spec[task] for task in task_names} + else: + if len(task_names) != 1: + raise TypeError( + "Multi-task training requires one statistics-cache " + "configuration per task." + ) + specs = {task_names[0]: spec} + + tasks_by_path: dict[Path, list[str]] = {} + for task, task_spec in specs.items(): + if task_spec.path is None: + continue + path = Path(task_spec.path).expanduser().resolve(strict=False) + tasks_by_path.setdefault(path, []).append(task) + duplicates = { + path: tasks for path, tasks in tasks_by_path.items() if len(tasks) > 1 + } + if duplicates: + details = "; ".join( + f"{str(path)!r}: {', '.join(repr(task) for task in tasks)}" + for path, tasks in duplicates.items() + ) + raise ValueError( + "Each training task must use a distinct statistics-cache path; " + f"duplicate path(s): {details}." + ) + return specs + + +@contextmanager +def open_stat_file( + spec: StatFileSpec, +) -> Iterator[DPPath | None]: + """Open one statistics cache for a bounded initialization scope. + + Existing HDF5 caches in update mode remain read-only until the first + write. A cache hit therefore neither acquires a writer lock nor changes + the file's HDF5 write-status metadata. + + Parameters + ---------- + spec + Statistics-cache configuration. + + Yields + ------ + DPPath or None + Scoped cache root, or ``None`` when persistent storage is disabled. + + Raises + ------ + FileNotFoundError + If read mode targets a cache that does not exist. + ValueError + If the target has an unsupported type. + """ + if spec.path is None: + yield None + return + + target = Path(spec.path).expanduser().resolve(strict=False) + if not target.exists() and spec.mode == "read": + raise FileNotFoundError( + f"Statistics cache {str(target)!r} does not exist in read mode." + ) + + if target.is_dir() or ( + not target.exists() and target.suffix.lower() not in _HDF5_SUFFIXES + ): + if not target.exists(): + target.mkdir(parents=True, exist_ok=True) + root: DPPath = DPOSPath(target, mode="r" if spec.mode == "read" else "a") + yield root + return + + if target.exists() and not target.is_file(): + raise ValueError( + f"Statistics cache {str(target)!r} is neither a file nor a directory." + ) + if not target.exists(): + target.parent.mkdir(parents=True, exist_ok=True) + + owner = _H5StatFile(target, spec.mode) + try: + yield _H5StatPath(owner, "/") + finally: + owner.close() + + +def load_required_items( + path: DPPath | None, + names: Sequence[str], +) -> dict[str, np.ndarray] | None: + """Load a complete group of statistics datasets. + + Parameters + ---------- + path + Statistics-cache root used by the current consumer. + names + Dataset names that form one indivisible statistics group. + + Returns + ------- + dict[str, numpy.ndarray] or None + Loaded arrays when every dataset exists. ``None`` indicates that the + group must be recomputed in update mode or that caching is disabled. + + Raises + ------ + FileNotFoundError + If a read-only cache is missing one or more required datasets. + """ + if path is None: + return None + missing = [name for name in names if not (path / name).is_file()] + if missing: + if getattr(path, "mode", None) == "r": + missing_items = ", ".join(repr(name) for name in missing) + raise FileNotFoundError( + f"Read-only statistics cache {path} is missing required " + f"item(s): {missing_items}." + ) + return None + return {name: (path / name).load_numpy() for name in names} + + +def load_paired_items( + path: DPPath | None, + pairs: Sequence[tuple[str, str]], +) -> dict[str, np.ndarray] | None: + """Load complete pairs from a legacy statistics cache. + + Read mode requires every requested pair. Update mode preserves the legacy + convention that both datasets may be absent for an output unavailable in + the sampled data. An output represented by either dataset is valid only + when its pair is also present. Statistics writers store every first item + before any second item, so an interrupted write leaves at least one + incomplete pair. + + Parameters + ---------- + path + Statistics-cache root used by the current consumer. + pairs + Dataset-name pairs that represent independently optional outputs. + + Returns + ------- + dict[str, numpy.ndarray] or None + Arrays for every represented pair. ``None`` indicates that the cache + contains no represented pair or must be recomputed in update mode. + + Raises + ------ + FileNotFoundError + If a read-only cache is missing any requested dataset. + """ + if path is None: + return None + + marker = path / _PAIR_TRANSACTION_MARKER + if marker.is_file() or marker.is_dir(): + if getattr(path, "mode", None) == "r": + raise FileNotFoundError( + f"Read-only statistics cache {path} contains an incomplete " + "output-statistics transaction." + ) + return None + + present = {name: (path / name).is_file() for pair in pairs for name in pair} + if getattr(path, "mode", None) == "r": + missing = [name for pair in pairs for name in pair if not present[name]] + if missing: + missing_items = ", ".join(repr(name) for name in missing) + raise FileNotFoundError( + f"Read-only statistics cache {path} is missing required " + f"item(s): {missing_items}." + ) + return {name: (path / name).load_numpy() for pair in pairs for name in pair} + + represented = [pair for pair in pairs if any(present[name] for name in pair)] + missing = [name for pair in represented for name in pair if not present[name]] + if not represented or missing: + return None + + return {name: (path / name).load_numpy() for pair in represented for name in pair} + + +def replace_paired_items( + path: DPPath, + pairs: Sequence[tuple[str, str]], + items: Mapping[str, np.ndarray], +) -> None: + """Replace optional pairs as one recoverable statistics transaction. + + All requested datasets are removed before the newly computed complete + pairs are written. A transient marker makes an interrupted replacement + distinguishable from a valid legacy cache. The marker is removed after + every dataset has been flushed, so successful caches retain the legacy + layout. + + Parameters + ---------- + path + Writable statistics-cache group that owns the pairs. + pairs + Dataset-name pairs requested by the statistics consumer. + items + Newly computed datasets. Each requested pair must be either complete + or entirely absent. + + Returns + ------- + None + + Raises + ------ + ValueError + If the path is read-only, names are duplicated or reserved, an item + does not belong to a requested pair, or exactly one item of a pair is + supplied. + TypeError + If the path implementation is unsupported. + """ + pair_list = list(pairs) + requested_names = [name for pair in pair_list for name in pair] + if len(requested_names) != len(set(requested_names)): + raise ValueError("Statistics pair names must be unique.") + if _PAIR_TRANSACTION_MARKER in requested_names: + raise ValueError( + f"{_PAIR_TRANSACTION_MARKER!r} is reserved for cache transactions." + ) + + unknown_names = set(items).difference(requested_names) + if unknown_names: + names = ", ".join(repr(name) for name in sorted(unknown_names)) + raise ValueError( + f"Statistics items do not belong to a requested pair: {names}." + ) + + complete_pairs: list[tuple[str, str]] = [] + for first, second in pair_list: + first_present = first in items + second_present = second in items + if first_present != second_present: + raise ValueError( + f"Statistics pair ({first!r}, {second!r}) must be supplied together." + ) + if first_present: + complete_pairs.append((first, second)) + + if getattr(path, "mode", None) == "r": + raise ValueError("Cannot write to a read-only statistics cache.") + + ordered_names = [pair[0] for pair in complete_pairs] + [ + pair[1] for pair in complete_pairs + ] + path.mkdir(parents=True, exist_ok=True) + if isinstance(path, _H5StatPath): + file = path._owner.file(write=True) + _replace_h5_items( + file, + path._connect_path, + requested_names, + [(name, items[name]) for name in ordered_names], + ) + return + if isinstance(path, DPH5Path): + try: + _replace_h5_items( + path.root, + path._connect_path, + requested_names, + [(name, items[name]) for name in ordered_names], + ) + finally: + DPH5Path._file_keys.cache_clear() + path._new_keys.clear() + return + if isinstance(path, DPOSPath): + _replace_os_items( + path, + requested_names, + [(name, items[name]) for name in ordered_names], + ) + return + raise TypeError(f"Unsupported statistics-cache path type: {type(path).__name__}.") + + +def _replace_h5_items( + file: h5py.File, + connect_path: Callable[[str], str], + requested_names: Sequence[str], + ordered_items: Sequence[tuple[str, np.ndarray]], +) -> None: + """Replace HDF5 datasets while retaining an interruption marker.""" + marker_name = connect_path(_PAIR_TRANSACTION_MARKER) + if marker_name not in file: + file.create_dataset(marker_name, data=np.array([1], dtype=np.uint8)) + file.flush() + + for name in requested_names: + item_name = connect_path(name) + if item_name in file: + del file[item_name] + for name, value in ordered_items: + file.create_dataset(connect_path(name), data=value) + file.flush() + + del file[marker_name] + file.flush() + + +def _replace_os_items( + path: DPOSPath, + requested_names: Sequence[str], + ordered_items: Sequence[tuple[str, np.ndarray]], +) -> None: + """Replace directory datasets while retaining an interruption marker.""" + marker = path / _PAIR_TRANSACTION_MARKER + assert isinstance(marker, DPOSPath) + if marker.is_dir(): + raise ValueError(f"Statistics transaction marker {marker} is a directory.") + if not marker.is_file(): + marker.save_numpy(np.array([1], dtype=np.uint8)) + + for name in requested_names: + item = path / name + assert isinstance(item, DPOSPath) + item.path.unlink(missing_ok=True) + for name, value in ordered_items: + (path / name).save_numpy(value) + + marker.path.unlink() + + +def run_stat_on_chief( + action: Callable[[], None], + *, + is_chief: bool, + synchronize_failure: Callable[[bool], bool] | None, + operation: str, +) -> None: + """Execute a statistics action on rank 0 and synchronize its outcome. + + Parameters + ---------- + action + Statistics operation executed only by the chief process. + is_chief + Whether the current process is the chief. + synchronize_failure + Backend callback that broadcasts the chief failure flag and returns + the synchronized value. ``None`` selects single-process execution. + operation + Human-readable operation name included in peer-rank errors. + + Returns + ------- + None + + Raises + ------ + Exception + Re-raises the original exception on the chief process. + RuntimeError + If the chief reports failure to a peer process. + """ + completed = False + try: + if is_chief: + action() + completed = True + finally: + local_failure = not completed + failed = ( + synchronize_failure(local_failure) + if synchronize_failure is not None + else local_failure + ) + if failed and not local_failure: + raise RuntimeError(f"Rank 0 failed during {operation}; see rank-0 logs.") + + +class _H5StatFile: + """Own one HDF5 handle for a single statistics initialization scope.""" + + def __init__(self, path: Path, mode: StatFileMode) -> None: + self.path = path + self.mode = mode + self._writable = not path.exists() + self._file = h5py.File(path, "a" if self._writable else "r") + + def file(self, *, write: bool = False) -> h5py.File: + """Return the live handle, promoting it for a requested write.""" + if not self._file.id.valid: + raise RuntimeError("The statistics HDF5 file is closed.") + if write: + self._promote_to_writer() + return self._file + + def flush(self) -> None: + """Flush pending HDF5 writes.""" + self.file().flush() + + def close(self) -> None: + """Close the owned handle if it remains open.""" + if self._file.id.valid: + self._file.close() + + def _promote_to_writer(self) -> None: + if self.mode == "read": + raise ValueError("Cannot write to a read-only statistics cache.") + if self._writable: + return + self._file.close() + self._file = h5py.File(self.path, "r+") + self._writable = True + + +class _H5StatPath(DPPath): + """Provide a non-owning DPPath view over a scoped HDF5 handle.""" + + def __init__(self, owner: _H5StatFile, name: str) -> None: + self._owner = owner + self._name = name + self.mode = "r" if owner.mode == "read" else "a" + self.root_path = str(owner.path) + + def __getnewargs__(self) -> tuple[str, str]: + raise TypeError("Scoped statistics paths cannot be serialized.") + + def load_numpy(self) -> np.ndarray: + return self._file[self._name][:] + + def load_txt(self, dtype: np.dtype | None = None, **kwargs: Any) -> np.ndarray: + array = self.load_numpy() + return array.astype(dtype) if dtype is not None else array + + def save_numpy(self, arr: np.ndarray) -> None: + file = self._owner.file(write=True) + if self._name in file: + del file[self._name] + file.create_dataset(self._name, data=arr) + self._owner.flush() + + def glob(self, pattern: str) -> list[DPPath]: + file = self._file + if self._name == "/": + group = file + elif self._name not in file or not isinstance(file[self._name], h5py.Group): + return [] + else: + group = file[self._name] + + keys: list[str] = [] + group.visit(lambda key: keys.append(self._connect_path(key))) + return [ + type(self)(self._owner, key) + for key in globfilter(keys, self._connect_path(pattern)) + ] + + def rglob(self, pattern: str) -> list[DPPath]: + return self.glob("**/" + pattern) + + def is_file(self) -> bool: + return self._name in self._file and isinstance( + self._file[self._name], h5py.Dataset + ) + + def is_dir(self) -> bool: + if self._name == "/": + self._owner.file() + return True + return self._name in self._file and isinstance( + self._file[self._name], h5py.Group + ) + + def __truediv__(self, key: str) -> DPPath: + return type(self)(self._owner, self._connect_path(key)) + + def __lt__(self, other: DPPath) -> bool: + return str(self) < str(other) + + def __str__(self) -> str: + return f"{self.root_path}#{self._name}" + + @property + def name(self) -> str: + return self._name.rsplit("/", 1)[-1] + + def mkdir(self, parents: bool = False, exist_ok: bool = False) -> None: + if self._owner.mode == "read": + raise ValueError("Cannot write to a read-only statistics cache.") + read_file = self._owner.file() + if self._name in read_file: + if not isinstance(read_file[self._name], h5py.Group) or not exist_ok: + raise FileExistsError(f"Statistics path {self} already exists.") + return + + file = self._owner.file(write=True) + if parents: + file.require_group(self._name) + else: + file.create_group(self._name) + self._owner.flush() + + @property + def _file(self) -> h5py.File: + return self._owner.file() + + def _connect_path(self, key: str) -> str: + return f"{self._name.rstrip('/')}/{key.lstrip('/')}" diff --git a/source/tests/common/dpmodel/test_train_data.py b/source/tests/common/dpmodel/test_train_data.py index ee9880a491..0cae6d888a 100644 --- a/source/tests/common/dpmodel/test_train_data.py +++ b/source/tests/common/dpmodel/test_train_data.py @@ -3,6 +3,10 @@ from deepmd.dpmodel.train.data import ( _print_summary, + iter_training_task_configs, +) +from deepmd.utils.stat_file import ( + StatFileSpec, ) @@ -28,3 +32,18 @@ def print_summary(self, name: str, prob: list[float] | None) -> None: with pytest.raises(TypeError, match="internal summary failure"): _print_summary(BrokenSummary(), "training", [1.0]) + + +def test_training_task_config_preserves_stat_file_mode() -> None: + config = { + "model": {}, + "training": { + "training_data": {}, + "stat_file": "stat.hdf5", + "stat_file_mode": "read", + }, + } + + task = next(iter_training_task_configs(config)) + + assert task.stat_file_spec == StatFileSpec("stat.hdf5", "read") diff --git a/source/tests/common/stat_file.py b/source/tests/common/stat_file.py new file mode 100644 index 0000000000..c7145a47bd --- /dev/null +++ b/source/tests/common/stat_file.py @@ -0,0 +1,172 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +"""Shared fixtures for statistics-cache round-trip tests.""" + +from collections.abc import ( + Callable, +) +from pathlib import ( + Path, +) +from typing import ( + Any, +) + +import h5py +import numpy as np + +from deepmd.dpmodel.common import ( + to_numpy_array, +) +from deepmd.utils.stat_file import ( + StatFileSpec, + open_stat_file, +) + + +def energy_model_params() -> dict[str, Any]: + """Return a minimal energy-model configuration. + + Returns + ------- + dict[str, Any] + Model parameters shared by backend round-trip tests. + """ + return { + "type_map": ["O", "H"], + "descriptor": { + "type": "se_e2_a", + "sel": [4, 4], + "rcut": 3.0, + "rcut_smth": 2.5, + "neuron": [4, 8], + "axis_neuron": 4, + "precision": "float64", + }, + "fitting_net": { + "type": "ener", + "neuron": [8], + "numb_fparam": 1, + "numb_aparam": 1, + "precision": "float64", + }, + } + + +def energy_stat_sample() -> list[dict[str, Any]]: + """Return NumPy statistics input for the minimal energy model. + + Returns + ------- + list[dict[str, Any]] + One sampled system containing energy labels for two atom types. + """ + return [ + { + "coord": np.array( + [ + [[0.0, 0.0, 0.0], [1.0, 0.0, 0.0]], + [[0.0, 0.0, 0.0], [1.5, 0.0, 0.0]], + ], + dtype=np.float64, + ), + "atype": np.array([[0, 0], [1, 1]], dtype=np.int64), + "box": None, + "natoms": np.array( + [[2, 2, 2, 0], [2, 2, 0, 2]], + dtype=np.int64, + ), + "energy": np.array([[2.0], [4.0]], dtype=np.float64), + "find_energy": np.float32(1.0), + "fparam": np.array([[1.0], [3.0]], dtype=np.float64), + "find_fparam": np.float32(1.0), + "aparam": np.array( + [ + [[1.0], [2.0]], + [[4.0], [7.0]], + ], + dtype=np.float64, + ), + "find_aparam": np.float32(1.0), + } + ] + + +def _model_stat_values(model: Any) -> dict[str, np.ndarray]: + """Collect numerical statistics applied to a backend model.""" + descriptor_mean, descriptor_stddev = ( + model.get_descriptor().get_stat_mean_and_stddev() + ) + fitting = model.get_fitting_net() + values = { + "descriptor_mean": descriptor_mean, + "descriptor_stddev": descriptor_stddev, + "fparam_avg": fitting.fparam_avg, + "fparam_inv_std": fitting.fparam_inv_std, + "aparam_avg": fitting.aparam_avg, + "aparam_inv_std": fitting.aparam_inv_std, + "out_bias": model.atomic_model.out_bias, + "out_std": model.atomic_model.out_std, + } + result = {} + for name, value in values.items(): + array = to_numpy_array(value) + if array is None: + raise AssertionError(f"Model statistic {name!r} was not initialized.") + result[name] = np.array(array, copy=True) + return result + + +def assert_energy_stat_cache_round_trip( + model_factory: Callable[[], Any], + stat_file: Path, + *, + sample_factory: Callable[[], list[dict[str, Any]]] = energy_stat_sample, +) -> None: + """Verify that an energy model reads the cache it writes without sampling. + + Parameters + ---------- + model_factory + Factory returning a new backend model for each cache pass. + stat_file + HDF5 statistics-cache path. + sample_factory + Factory returning statistics input in the backend's array format. + + Returns + ------- + None + + Raises + ------ + AssertionError + If the generated cache has an invalid key set or read mode samples data. + """ + update_model = model_factory() + with open_stat_file(StatFileSpec(str(stat_file))) as stat_path: + assert stat_path is not None + update_model.compute_or_load_stat(sample_factory, stat_path) + expected_values = _model_stat_values(update_model) + + with h5py.File(stat_file, "r") as file: + assert "tasks" not in file + type_map_cache = file["O H"] + assert "bias_atom_energy" in type_map_cache + assert "std_atom_energy" in type_map_cache + assert "bias_atom_mask" not in type_map_cache + assert "std_atom_mask" not in type_map_cache + original = stat_file.read_bytes() + + def unexpected_sample() -> list[dict[str, Any]]: + raise AssertionError("A complete read-only cache must not sample data.") + + for mode in ("update", "read"): + read_model = model_factory() + with open_stat_file(StatFileSpec(str(stat_file), mode)) as stat_path: + assert stat_path is not None + read_model.compute_or_load_stat(unexpected_sample, stat_path) + actual_values = _model_stat_values(read_model) + assert actual_values.keys() == expected_values.keys() + for name, expected in expected_values.items(): + np.testing.assert_allclose(actual_values[name], expected) + assert stat_file.read_bytes() == original diff --git a/source/tests/common/test_stat_file.py b/source/tests/common/test_stat_file.py new file mode 100644 index 0000000000..a4bdfa0c88 --- /dev/null +++ b/source/tests/common/test_stat_file.py @@ -0,0 +1,470 @@ +# SPDX-License-Identifier: LGPL-3.0-or-later +from pathlib import ( + Path, +) +from unittest.mock import ( + Mock, +) + +import h5py +import numpy as np +import pytest + +from deepmd.dpmodel.model.model import ( + get_model, +) +from deepmd.dpmodel.utils.stat import ( + compute_output_stats, +) +from deepmd.utils.path import ( + DPH5Path, + DPPath, +) +from deepmd.utils.stat_file import ( + StatFileSpec, + load_paired_items, + load_required_items, + open_stat_file, + replace_paired_items, + run_stat_on_chief, + stat_file_specs_by_task, +) + +from .stat_file import ( + assert_energy_stat_cache_round_trip, + energy_model_params, + energy_stat_sample, +) + + +@pytest.mark.parametrize("mode", ["read", "invalid"]) +def test_stat_file_spec_rejects_invalid_disabled_mode(mode: str) -> None: + with pytest.raises(ValueError): + StatFileSpec(None, mode) # type: ignore[arg-type] + + +def test_stat_file_spec_rejects_empty_path() -> None: + with pytest.raises(ValueError, match="must not be empty"): + StatFileSpec(" ") + + +def test_disabled_stat_file_yields_none() -> None: + with open_stat_file(StatFileSpec(None)) as path: + assert path is None + + +def test_hdf5_handle_is_scoped(tmp_path: Path) -> None: + target = tmp_path / "nested" / "stat.hdf5" + with open_stat_file(StatFileSpec(str(target))) as path: + assert path is not None + value = path / "value" + value.save_numpy(np.array([1.0])) + assert value.load_numpy().tolist() == [1.0] + + with pytest.raises(RuntimeError, match="closed"): + value.load_numpy() + with h5py.File(target, "r+") as file: + assert file["value"][:].tolist() == [1.0] + + +def test_hdf5_handle_closes_after_exception(tmp_path: Path) -> None: + target = tmp_path / "stat.hdf5" + path = None + with pytest.raises(RuntimeError, match="statistics failed"): + with open_stat_file(StatFileSpec(str(target))) as path: + assert path is not None + (path / "value").save_numpy(np.array([1.0])) + raise RuntimeError("statistics failed") + + assert path is not None + with pytest.raises(RuntimeError, match="closed"): + (path / "value").load_numpy() + + +def test_complete_update_cache_remains_byte_identical(tmp_path: Path) -> None: + target = tmp_path / "stat.hdf5" + with h5py.File(target, "w") as file: + file.create_dataset("value", data=[1.0]) + original = target.read_bytes() + + with open_stat_file(StatFileSpec(str(target))) as path: + assert path is not None + assert (path / "value").load_numpy().tolist() == [1.0] + + assert target.read_bytes() == original + + +def test_existing_update_cache_promotes_on_first_write(tmp_path: Path) -> None: + target = tmp_path / "stat.hdf5" + with h5py.File(target, "w") as file: + file.create_dataset("existing", data=[1.0]) + + with open_stat_file(StatFileSpec(str(target))) as path: + assert path is not None + (path / "added").save_numpy(np.array([2.0])) + + with h5py.File(target, "r") as file: + assert file["existing"][:].tolist() == [1.0] + assert file["added"][:].tolist() == [2.0] + + +def test_dpmodel_energy_cache_round_trip_uses_fitting_outputs_only( + tmp_path: Path, +) -> None: + assert_energy_stat_cache_round_trip( + lambda: get_model(energy_model_params()), + tmp_path / "stat.hdf5", + ) + + +def test_read_mode_rejects_writes(tmp_path: Path) -> None: + target = tmp_path / "stat.hdf5" + with h5py.File(target, "w"): + pass + + with open_stat_file(StatFileSpec(str(target), "read")) as path: + assert path is not None + with pytest.raises(ValueError, match="read-only"): + (path / "value").save_numpy(np.array([1.0])) + with pytest.raises(ValueError, match="read-only"): + path.mkdir(exist_ok=True) + + +def test_directory_cache_uses_same_scope_interface(tmp_path: Path) -> None: + target = tmp_path / "stat" + with open_stat_file(StatFileSpec(str(target))) as path: + assert path is not None + path.mkdir(parents=True, exist_ok=True) + (path / "value").save_numpy(np.array([1.0])) + + with open_stat_file(StatFileSpec(str(target), "read")) as path: + assert path is not None + assert (path / "value").load_numpy().tolist() == [1.0] + + +def test_multi_task_cache_paths_must_be_distinct(tmp_path: Path) -> None: + target = str(tmp_path / "stat.hdf5") + specs = { + "task/one": StatFileSpec(target), + "task_two": StatFileSpec(str(tmp_path / "." / "stat.hdf5")), + } + + with pytest.raises(ValueError, match="distinct statistics-cache path"): + stat_file_specs_by_task(specs, ["task/one", "task_two"]) + + +def test_stat_file_specs_are_normalized_by_task() -> None: + disabled = stat_file_specs_by_task(None, ["one", "two"]) + assert disabled == {"one": StatFileSpec(None), "two": StatFileSpec(None)} + + configured = { + "one": StatFileSpec("one.hdf5"), + "two": StatFileSpec("two.hdf5", "read"), + } + assert stat_file_specs_by_task(configured, ["one", "two"]) == configured + with pytest.raises(TypeError, match="Multi-task"): + stat_file_specs_by_task(StatFileSpec("shared.hdf5"), ["one", "two"]) + + +def test_read_mode_reports_all_missing_required_items(tmp_path: Path) -> None: + target = tmp_path / "stat.hdf5" + with h5py.File(target, "w") as file: + file.create_dataset("present", data=[1.0]) + + with open_stat_file(StatFileSpec(str(target), "read")) as path: + assert path is not None + with pytest.raises(FileNotFoundError) as error: + load_required_items(path, ["missing_one", "present", "missing_two"]) + + message = str(error.value) + assert "'missing_one'" in message + assert "'missing_two'" in message + + +def test_update_mode_recomputes_partial_output_group(tmp_path: Path) -> None: + target = tmp_path / "stat.hdf5" + sampled = energy_stat_sample() + sampled[0]["property"] = np.array([[1.0], [3.0]], dtype=np.float64) + sampled[0]["find_property"] = np.float32(1.0) + + with open_stat_file(StatFileSpec(str(target))) as path: + assert path is not None + (path / "bias_atom_energy").save_numpy(np.full((2, 1), 100.0)) + (path / "bias_atom_property").save_numpy(np.full((2, 1), 100.0)) + (path / "std_atom_energy").save_numpy(np.full((2, 1), 100.0)) + sampler = Mock(return_value=sampled) + + bias, std = compute_output_stats( + sampler, + ntypes=2, + keys=["energy", "property"], + stat_file_path=path, + ) + + sampler.assert_called_once_with() + assert set(bias) == {"energy", "property"} + assert set(std) == {"energy", "property"} + + with h5py.File(target, "r") as file: + assert set(file) == { + "bias_atom_energy", + "bias_atom_property", + "std_atom_energy", + "std_atom_property", + } + assert not np.all(file["bias_atom_energy"][:] == 100.0) + + +def test_update_mode_replaces_orphaned_output_pair(tmp_path: Path) -> None: + target = tmp_path / "stat.hdf5" + with h5py.File(target, "w") as file: + file.create_dataset("bias_atom_property", data=np.zeros((2, 1))) + + sampler = Mock(return_value=energy_stat_sample()) + with open_stat_file(StatFileSpec(str(target))) as path: + assert path is not None + bias, std = compute_output_stats( + sampler, + ntypes=2, + keys=["energy", "property"], + stat_file_path=path, + ) + + sampler.assert_called_once_with() + assert set(bias) == {"energy"} + assert set(std) == {"energy"} + with h5py.File(target, "r") as file: + assert set(file) == {"bias_atom_energy", "std_atom_energy"} + original = target.read_bytes() + + sampler.reset_mock(side_effect=True) + sampler.side_effect = AssertionError( + "A complete update cache must not sample data." + ) + with open_stat_file(StatFileSpec(str(target))) as path: + assert path is not None + bias, std = compute_output_stats( + sampler, + ntypes=2, + keys=["energy", "property"], + stat_file_path=path, + ) + + sampler.assert_not_called() + assert set(bias) == {"energy"} + assert set(std) == {"energy"} + assert target.read_bytes() == original + + +def test_interrupted_pair_replacement_requires_recovery(tmp_path: Path) -> None: + target = tmp_path / "stat.hdf5" + pairs = [("bias_atom_energy", "std_atom_energy")] + invalid_items = { + "bias_atom_energy": np.zeros((2, 1)), + "std_atom_energy": np.array([object()], dtype=object), + } + + with open_stat_file(StatFileSpec(str(target))) as path: + assert path is not None + with pytest.raises(TypeError): + replace_paired_items(path, pairs, invalid_items) + assert load_paired_items(path, pairs) is None + + with open_stat_file(StatFileSpec(str(target), "read")) as path: + assert path is not None + with pytest.raises(FileNotFoundError, match=r"incomplete.*transaction"): + load_paired_items(path, pairs) + + expected = { + "bias_atom_energy": np.zeros((2, 1)), + "std_atom_energy": np.ones((2, 1)), + } + with open_stat_file(StatFileSpec(str(target))) as path: + assert path is not None + assert load_paired_items(path, pairs) is None + replace_paired_items(path, pairs, expected) + + with h5py.File(target, "r") as file: + assert set(file) == set(expected) + for name, value in expected.items(): + np.testing.assert_array_equal(file[name][:], value) + + +def test_read_mode_rejects_entirely_missing_output_pair(tmp_path: Path) -> None: + target = tmp_path / "stat.hdf5" + with h5py.File(target, "w") as file: + file.create_dataset("bias_atom_energy", data=np.zeros((2, 1))) + file.create_dataset("std_atom_energy", data=np.ones((2, 1))) + + sampler = Mock() + with open_stat_file(StatFileSpec(str(target), "read")) as path: + assert path is not None + with pytest.raises(FileNotFoundError, match="bias_atom_property"): + compute_output_stats( + sampler, + ntypes=2, + keys=["energy", "property"], + stat_file_path=path, + ) + + sampler.assert_not_called() + + +def test_update_mode_preserves_absent_legacy_output_pair(tmp_path: Path) -> None: + target = tmp_path / "stat.hdf5" + with h5py.File(target, "w") as file: + file.create_dataset("bias_atom_energy", data=np.zeros((2, 1))) + file.create_dataset("std_atom_energy", data=np.ones((2, 1))) + original = target.read_bytes() + + sampler = Mock() + with open_stat_file(StatFileSpec(str(target))) as path: + assert path is not None + bias, std = compute_output_stats( + sampler, + ntypes=2, + keys=["energy", "property"], + stat_file_path=path, + ) + + sampler.assert_not_called() + assert set(bias) == {"energy"} + assert set(std) == {"energy"} + assert target.read_bytes() == original + + +def test_update_mode_recomputes_partial_descriptor_group(tmp_path: Path) -> None: + target = tmp_path / "stat.hdf5" + with open_stat_file(StatFileSpec(str(target))) as path: + assert path is not None + get_model(energy_model_params()).compute_or_load_stat(energy_stat_sample, path) + + with h5py.File(target, "r+") as file: + type_map_group = file["O H"] + descriptor_group = next( + item for item in type_map_group.values() if isinstance(item, h5py.Group) + ) + del descriptor_group["r_0"] + + read_sampler = Mock() + with open_stat_file(StatFileSpec(str(target), "read")) as path: + assert path is not None + with pytest.raises(FileNotFoundError, match="'r_0'"): + get_model(energy_model_params()).compute_or_load_stat(read_sampler, path) + read_sampler.assert_not_called() + + sampler = Mock(return_value=energy_stat_sample()) + with open_stat_file(StatFileSpec(str(target))) as path: + assert path is not None + get_model(energy_model_params()).compute_or_load_stat(sampler, path) + + sampler.assert_called_once_with() + with h5py.File(target, "r") as file: + type_map_group = file["O H"] + descriptor_group = next( + item for item in type_map_group.values() if isinstance(item, h5py.Group) + ) + assert "r_0" in descriptor_group + + +def test_scoped_cache_matches_legacy_hdf5_layout(tmp_path: Path) -> None: + legacy_file = tmp_path / "legacy.hdf5" + scoped_file = tmp_path / "scoped.hdf5" + with h5py.File(legacy_file, "w"): + pass + + legacy_path = DPPath(str(legacy_file), "a") + assert isinstance(legacy_path, DPH5Path) + legacy_handle = legacy_path.root + try: + get_model(energy_model_params()).compute_or_load_stat( + energy_stat_sample, + legacy_path, + ) + finally: + legacy_handle.close() + DPH5Path._load_h5py.cache_clear() + DPH5Path._file_keys.cache_clear() + DPH5Path._DPH5Path__file_new_keys.pop(legacy_handle, None) + + with open_stat_file(StatFileSpec(str(scoped_file))) as scoped_path: + assert scoped_path is not None + get_model(energy_model_params()).compute_or_load_stat( + energy_stat_sample, + scoped_path, + ) + + def read_layout( + path: Path, + ) -> tuple[set[str], dict[str, tuple[np.dtype, tuple[int, ...], np.ndarray]]]: + groups: set[str] = set() + datasets: dict[str, tuple[np.dtype, tuple[int, ...], np.ndarray]] = {} + with h5py.File(path, "r") as file: + + def collect(name: str, item: h5py.Group | h5py.Dataset) -> None: + if isinstance(item, h5py.Group): + groups.add(name) + else: + datasets[name] = (item.dtype, item.shape, item[:]) + + file.visititems(collect) + return groups, datasets + + legacy_groups, legacy_datasets = read_layout(legacy_file) + scoped_groups, scoped_datasets = read_layout(scoped_file) + assert scoped_groups == legacy_groups + assert scoped_datasets.keys() == legacy_datasets.keys() + for name, (legacy_dtype, legacy_shape, legacy_value) in legacy_datasets.items(): + scoped_dtype, scoped_shape, scoped_value = scoped_datasets[name] + assert scoped_dtype == legacy_dtype + assert scoped_shape == legacy_shape + np.testing.assert_array_equal(scoped_value, legacy_value) + + original = legacy_file.read_bytes() + + def unexpected_sample() -> list[dict[str, object]]: + raise AssertionError("A complete legacy cache must not sample data.") + + for mode in ("update", "read"): + with open_stat_file(StatFileSpec(str(legacy_file), mode)) as stat_path: + assert stat_path is not None + get_model(energy_model_params()).compute_or_load_stat( + unexpected_sample, + stat_path, + ) + assert legacy_file.read_bytes() == original + + +def test_chief_failure_is_synchronized_before_reraising() -> None: + synchronized: list[bool] = [] + + def fail() -> None: + raise ValueError("invalid statistics") + + def synchronize(failed: bool) -> bool: + synchronized.append(failed) + return failed + + with pytest.raises(ValueError, match="invalid statistics"): + run_stat_on_chief( + fail, + is_chief=True, + synchronize_failure=synchronize, + operation="statistics initialization", + ) + + assert synchronized == [True] + + +def test_peer_raises_when_chief_reports_failure() -> None: + action = Mock() + + with pytest.raises(RuntimeError, match="Rank 0 failed during statistics"): + run_stat_on_chief( + action, + is_chief=False, + synchronize_failure=lambda _: True, + operation="statistics initialization", + ) + + action.assert_not_called() diff --git a/source/tests/jax/test_model_factory.py b/source/tests/jax/test_model_factory.py index 75ffc519a1..b44e7064af 100644 --- a/source/tests/jax/test_model_factory.py +++ b/source/tests/jax/test_model_factory.py @@ -10,7 +10,16 @@ """ import unittest +from pathlib import ( + Path, +) + +import numpy as np +from deepmd.jax.env import ( + jnp, + nnx, +) from deepmd.jax.model.ener_model import ( EnergyModel, ) @@ -18,6 +27,12 @@ get_model, ) +from ..common.stat_file import ( + assert_energy_stat_cache_round_trip, + energy_model_params, + energy_stat_sample, +) + def _base_config() -> dict: return { @@ -62,5 +77,56 @@ def test_explicit_fitting_type_preserved(self) -> None: self.assertIsInstance(model, EnergyModel) +def test_jax_array_assignment_preserves_variable_for_shape_change() -> None: + """Backend array assignment updates an existing NNX variable container.""" + descriptor = get_model(energy_model_params()).get_descriptor() + variable = descriptor.davg + new_shape = (3, *variable.shape[1:]) + + descriptor.davg = jnp.zeros(new_shape, dtype=variable.dtype) + + assert descriptor.davg is variable + assert descriptor.davg.shape == new_shape + + +def test_jax_energy_cache_round_trip_uses_fitting_outputs_only( + tmp_path: Path, +) -> None: + models = [] + + def model_factory(): + model = get_model(energy_model_params()) + models.append(model) + return model + + def sample_factory(): + return [ + { + key: jnp.asarray(value) if isinstance(value, np.ndarray) else value + for key, value in sample.items() + } + for sample in energy_stat_sample() + ] + + assert_energy_stat_cache_round_trip( + model_factory, + tmp_path / "stat.hdf5", + sample_factory=sample_factory, + ) + for model in models: + descriptor = model.get_descriptor() + fitting = model.get_fitting_net() + for value in ( + *descriptor.get_stat_mean_and_stddev(), + fitting.fparam_avg, + fitting.fparam_inv_std, + fitting.aparam_avg, + fitting.aparam_inv_std, + model.atomic_model.out_bias, + model.atomic_model.out_std, + ): + assert isinstance(value, nnx.Variable) + + if __name__ == "__main__": unittest.main() diff --git a/source/tests/jax/test_training.py b/source/tests/jax/test_training.py index b720a71789..c175de8da4 100644 --- a/source/tests/jax/test_training.py +++ b/source/tests/jax/test_training.py @@ -11,6 +11,9 @@ import tempfile import textwrap import unittest +from collections.abc import ( + Callable, +) from copy import ( deepcopy, ) @@ -21,11 +24,13 @@ SimpleNamespace, ) from unittest.mock import ( + Mock, patch, ) import numpy as np import optax +import pytest from deepmd.dpmodel.output_def import ( OutputVariableCategory, @@ -425,6 +430,17 @@ def __init__(self, stats: dict[str, StatItem]) -> None: self.davg = np.asarray([0.0], dtype=np.float64) self.dstd = np.asarray([1.0], dtype=np.float64) + def set_stat_mean_and_stddev( + self, + mean: np.ndarray, + stddev: np.ndarray, + ) -> None: + self.davg = mean + self.dstd = stddev + + def get_stat_mean_and_stddev(self) -> tuple[np.ndarray, np.ndarray]: + return self.davg, self.dstd + def test_jax_shared_descriptor_stats_merge_weighted_values() -> None: """Shared descriptor merge recomputes weighted avg/std for nested stats.""" @@ -448,6 +464,44 @@ def test_jax_shared_descriptor_stats_merge_weighted_values() -> None: assert base.se_atten.stats["env"].number == 4 +def test_jax_shared_descriptor_stats_preserve_real_nnx_state() -> None: + """Shared statistics merge keeps real JAX descriptor arrays registered.""" + trainer = DPTrainer( + _minimal_jax_multitask_config(_shared_jax_model_config()), + ) + sampled = [ + { + "coord": jnp.asarray( + [ + [ + [0.0, 0.0, 0.0], + [0.8, 0.0, 0.0], + [0.0, 0.8, 0.0], + ] + ] + ), + "atype": jnp.asarray([[0, 1, 2]], dtype=jnp.int32), + "box": jnp.asarray([np.eye(3) * 8.0]), + } + ] + descriptors = [ + trainer.models[model_key].get_descriptor() for model_key in ("task_a", "task_b") + ] + for descriptor in descriptors: + descriptor.compute_input_stats(sampled) + assert isinstance(descriptor.davg, nnx.Variable) + assert isinstance(descriptor.dstd, nnx.Variable) + + base_count = descriptors[0].stats["r_0"].number + link_count = descriptors[1].stats["r_0"].number + trainer._share_model_params(resume=False) + + shared_descriptor = trainer.models["task_a"].get_descriptor() + assert isinstance(shared_descriptor.davg, nnx.Variable) + assert isinstance(shared_descriptor.dstd, nnx.Variable) + assert shared_descriptor.stats["r_0"].number == base_count + link_count + + class _FittingWithStats: def __init__(self, param_stats: dict[str, list[StatItem]]) -> None: self.numb_fparam = len(param_stats.get("fparam", [])) @@ -679,6 +733,49 @@ def test_jax_change_bias_after_training_uses_broadcast_on_peer_rank() -> None: ) +def test_jax_statistics_failure_reaches_peer_rank() -> None: + """Peer ranks fail before model-state synchronization when rank 0 fails.""" + trainer = _bias_sync_trainer(rank=1) + action = Mock() + + with ( + patch( + "jax.experimental.multihost_utils.broadcast_one_to_all", + return_value=np.asarray(True), + ), + pytest.raises(RuntimeError, match="Rank 0 failed during statistics"), + ): + trainer._run_on_chief(action, operation="statistics initialization") + + action.assert_not_called() + + +def test_jax_shared_statistics_merge_precedes_broadcast_and_binding() -> None: + """The chief broadcasts merged statistics before shared parameters bind.""" + trainer = DPTrainer.__new__(DPTrainer) + trainer.multi_task = True + trainer.shared_links = {"shared_fitting": object()} + events: list[str] = [] + + def run_on_chief(action: Callable[[], None], *, operation: str) -> None: + assert operation == "shared statistics merge" + action() + + def share_model_params(*, resume: bool = False) -> None: + events.append("bind" if resume else "merge") + + trainer._run_on_chief = run_on_chief + trainer._share_model_params = share_model_params + trainer._broadcast_model_states = lambda: events.append("broadcast") + + trainer._synchronize_initial_model_state( + state_changed=True, + merge_shared_statistics=True, + ) + + assert events == ["merge", "broadcast", "bind"] + + class TestJAXTraining(unittest.TestCase): """Regression tests for complete JAX training runs.""" diff --git a/source/tests/pt/test_stat_file_mode.py b/source/tests/pt/test_stat_file_mode.py index 58f0e09c7f..706ce1381b 100644 --- a/source/tests/pt/test_stat_file_mode.py +++ b/source/tests/pt/test_stat_file_mode.py @@ -6,6 +6,10 @@ from typing import ( Any, ) +from unittest.mock import ( + Mock, + patch, +) import h5py import numpy as np @@ -16,11 +20,14 @@ ) from deepmd.pt.entrypoints.main import ( - _prepare_stat_file_path, + get_trainer, ) from deepmd.pt.model.model import ( get_model, ) +from deepmd.pt.train.training import ( + Trainer, +) from deepmd.pt.utils.env import ( DEVICE, ) @@ -30,20 +37,12 @@ from deepmd.utils.argcheck import ( normalize, ) -from deepmd.utils.path import ( - DPH5Path, - DPPath, +from deepmd.utils.stat_file import ( + StatFileSpec, + open_stat_file, ) -def _close_stat_path(stat_path: DPPath) -> None: - """Close a test HDF5 path and reset its process-local caches.""" - assert isinstance(stat_path, DPH5Path) - stat_path.root.close() - DPH5Path._load_h5py.cache_clear() - DPH5Path._file_keys.cache_clear() - - def _load_dpa4_example() -> dict[str, Any]: """Load the DPA4 example used by configuration validation tests.""" example_path = ( @@ -57,6 +56,18 @@ def _load_dpa4_example() -> dict[str, Any]: return json.load(stream) +def _load_pt_training_example() -> dict[str, Any]: + """Load a minimal PT training configuration with absolute data paths.""" + test_root = Path(__file__).resolve().parent + with (test_root / "water" / "se_atten.json").open(encoding="utf-8") as stream: + config = json.load(stream) + systems = [str(test_root / "water" / "data" / "single")] + config["training"]["training_data"]["systems"] = systems + config["training"]["validation_data"]["systems"] = systems + config["training"]["numb_steps"] = 0 + return config + + def _energy_model_params() -> dict[str, Any]: """Build a minimal PT energy-model configuration.""" return { @@ -113,14 +124,11 @@ def _energy_stat_sample() -> list[dict[str, Any]]: def test_default_stat_file_mode_remains_writable(tmp_path: Path) -> None: stat_file = tmp_path / "stat.hdf5" - stat_path = _prepare_stat_file_path(str(stat_file)) - assert isinstance(stat_path, DPH5Path) - try: + with open_stat_file(StatFileSpec(str(stat_file))) as stat_path: + assert stat_path is not None assert stat_path.mode == "a" (stat_path / "value").save_numpy(np.array([1.0])) assert (stat_path / "value").load_numpy().tolist() == [1.0] - finally: - _close_stat_path(stat_path) def test_read_stat_file_mode_reads_existing_cache(tmp_path: Path) -> None: @@ -128,25 +136,41 @@ def test_read_stat_file_mode_reads_existing_cache(tmp_path: Path) -> None: with h5py.File(stat_file, "w") as file: file.create_dataset("value", data=[1.0]) - stat_path = _prepare_stat_file_path(str(stat_file), "read") - assert isinstance(stat_path, DPH5Path) - try: + with open_stat_file(StatFileSpec(str(stat_file), "read")) as stat_path: + assert stat_path is not None assert stat_path.mode == "r" - assert stat_path.root.mode == "r" assert (stat_path / "value").load_numpy().tolist() == [1.0] - finally: - _close_stat_path(stat_path) def test_read_stat_file_mode_requires_existing_cache(tmp_path: Path) -> None: stat_file = tmp_path / "missing.hdf5" with pytest.raises(FileNotFoundError, match="does not exist in read mode"): - _prepare_stat_file_path(str(stat_file), "read") + with open_stat_file(StatFileSpec(str(stat_file), "read")): + pass def test_read_stat_file_mode_requires_cache_path() -> None: with pytest.raises(ValueError, match="requires `stat_file`"): - _prepare_stat_file_path(None, "read") + StatFileSpec(None, "read") + + +def test_frozen_initialization_does_not_open_read_cache(tmp_path: Path) -> None: + config = _load_pt_training_example() + missing_cache = tmp_path / "missing.hdf5" + config["training"]["stat_file"] = str(missing_cache) + config["training"]["stat_file_mode"] = "read" + frozen_model = Mock() + frozen_model.state_dict.return_value = {} + + with patch( + "deepmd.pt.train.training.torch.jit.load", + return_value=frozen_model, + ) as load_frozen: + trainer = get_trainer(config, init_frz_model="model.pth") + + load_frozen.assert_called_once() + assert trainer.model is not None + assert not missing_cache.exists() def test_read_stat_file_mode_rejects_incomplete_statistics_cache( @@ -156,8 +180,8 @@ def test_read_stat_file_mode_rejects_incomplete_statistics_cache( with h5py.File(stat_file, "w") as file: file.create_dataset("bias_atom_energy", data=np.zeros((1, 1))) - stat_path = _prepare_stat_file_path(str(stat_file), "read") - try: + with open_stat_file(StatFileSpec(str(stat_file), "read")) as stat_path: + assert stat_path is not None with pytest.raises(FileNotFoundError, match="std_atom_energy"): compute_output_stats( lambda: pytest.fail("read-only statistics must not sample data"), @@ -165,8 +189,6 @@ def test_read_stat_file_mode_rejects_incomplete_statistics_cache( keys=["energy"], stat_file_path=stat_path, ) - finally: - _close_stat_path(stat_path) def test_read_stat_file_mode_loads_complete_cache_from_two_readers( @@ -177,12 +199,10 @@ def test_read_stat_file_mode_loads_complete_cache_from_two_readers( file.create_dataset("bias_atom_energy", data=np.zeros((1, 1))) file.create_dataset("std_atom_energy", data=np.ones((1, 1))) - reader_one = _prepare_stat_file_path(str(stat_file), "read") - DPH5Path._load_h5py.cache_clear() - reader_two = _prepare_stat_file_path(str(stat_file), "read") - assert isinstance(reader_one, DPH5Path) - assert isinstance(reader_two, DPH5Path) - try: + spec = StatFileSpec(str(stat_file), "read") + with open_stat_file(spec) as reader_one, open_stat_file(spec) as reader_two: + assert reader_one is not None + assert reader_two is not None for reader in (reader_one, reader_two): bias, std = compute_output_stats( lambda: pytest.fail("complete read-only cache must not sample data"), @@ -192,38 +212,154 @@ def test_read_stat_file_mode_loads_complete_cache_from_two_readers( ) assert bias["energy"].shape == (1, 1) assert std["energy"].shape == (1, 1) - finally: - reader_one.root.close() - reader_two.root.close() - DPH5Path._load_h5py.cache_clear() - DPH5Path._file_keys.cache_clear() + + +def test_distributed_statistics_failure_reaches_peer_rank() -> None: + trainer = Trainer.__new__(Trainer) + trainer.is_distributed = True + trainer.rank = 1 + action = Mock() + + def report_chief_failure(holder: list[bool], **_: Any) -> None: + holder[0] = True + + with ( + patch( + "deepmd.pt.train.training.dist.broadcast_object_list", + side_effect=report_chief_failure, + ), + pytest.raises(RuntimeError, match="Rank 0 failed during statistics"), + ): + trainer._run_stat_on_chief(action, operation="statistics initialization") + + action.assert_not_called() + + +def test_update_mode_recomputes_partial_multi_output_cache(tmp_path: Path) -> None: + stat_file = tmp_path / "stat.hdf5" + sampled = _energy_stat_sample() + sampled[0]["property"] = torch.tensor( + [[1.0], [3.0]], + dtype=torch.float64, + device=DEVICE, + ) + sampled[0]["find_property"] = np.float32(1.0) + sample_count = 0 + + def sample() -> list[dict[str, Any]]: + nonlocal sample_count + sample_count += 1 + return sampled + + with open_stat_file(StatFileSpec(str(stat_file))) as stat_path: + assert stat_path is not None + (stat_path / "bias_atom_energy").save_numpy(np.full((2, 1), 100.0)) + (stat_path / "bias_atom_property").save_numpy(np.full((2, 1), 100.0)) + (stat_path / "std_atom_energy").save_numpy(np.full((2, 1), 100.0)) + bias, std = compute_output_stats( + sample, + ntypes=2, + keys=["energy", "property"], + stat_file_path=stat_path, + ) + + assert sample_count == 1 + assert set(bias) == {"energy", "property"} + assert set(std) == {"energy", "property"} + + with h5py.File(stat_file, "r") as file: + assert set(file) == { + "bias_atom_energy", + "bias_atom_property", + "std_atom_energy", + "std_atom_property", + } + assert not np.all(file["bias_atom_energy"][:] == 100.0) + + +def test_update_mode_replaces_orphaned_output_pair(tmp_path: Path) -> None: + stat_file = tmp_path / "stat.hdf5" + with h5py.File(stat_file, "w") as file: + file.create_dataset("bias_atom_property", data=np.zeros((2, 1))) + + sampler = Mock(return_value=_energy_stat_sample()) + with open_stat_file(StatFileSpec(str(stat_file))) as stat_path: + assert stat_path is not None + bias, std = compute_output_stats( + sampler, + ntypes=2, + keys=["energy", "property"], + stat_file_path=stat_path, + ) + + sampler.assert_called_once_with() + assert set(bias) == {"energy"} + assert set(std) == {"energy"} + with h5py.File(stat_file, "r") as file: + assert set(file) == {"bias_atom_energy", "std_atom_energy"} + original = stat_file.read_bytes() + + sampler.reset_mock(side_effect=True) + sampler.side_effect = AssertionError( + "A complete update cache must not sample data." + ) + with open_stat_file(StatFileSpec(str(stat_file))) as stat_path: + assert stat_path is not None + bias, std = compute_output_stats( + sampler, + ntypes=2, + keys=["energy", "property"], + stat_file_path=stat_path, + ) + + sampler.assert_not_called() + assert set(bias) == {"energy"} + assert set(std) == {"energy"} + assert stat_file.read_bytes() == original def test_energy_model_reloads_update_cache_in_read_mode(tmp_path: Path) -> None: sampled = _energy_stat_sample() stat_file = tmp_path / "stat.hdf5" - update_path = _prepare_stat_file_path(str(stat_file), "update") - assert isinstance(update_path, DPH5Path) - try: + with open_stat_file(StatFileSpec(str(stat_file))) as update_path: + assert update_path is not None update_model = get_model(_energy_model_params()).to(DEVICE) update_model.compute_or_load_stat(lambda: sampled, update_path) stat_root = update_path / "O H" assert (stat_root / "bias_atom_energy").is_file() assert not (stat_root / "bias_atom_mask").is_file() - finally: - _close_stat_path(update_path) + original = stat_file.read_bytes() - read_path = _prepare_stat_file_path(str(stat_file), "read") - assert isinstance(read_path, DPH5Path) - try: + with open_stat_file(StatFileSpec(str(stat_file), "read")) as read_path: + assert read_path is not None read_model = get_model(_energy_model_params()).to(DEVICE) read_model.compute_or_load_stat( lambda: pytest.fail("complete read-only cache must not sample data"), read_path, ) - finally: - _close_stat_path(read_path) + assert stat_file.read_bytes() == original + + +def test_energy_model_reuses_update_cache_without_modifying_hdf5( + tmp_path: Path, +) -> None: + stat_file = tmp_path / "stat.hdf5" + with open_stat_file(StatFileSpec(str(stat_file))) as stat_path: + assert stat_path is not None + model = get_model(_energy_model_params()).to(DEVICE) + model.compute_or_load_stat(_energy_stat_sample, stat_path) + original = stat_file.read_bytes() + + with open_stat_file(StatFileSpec(str(stat_file))) as stat_path: + assert stat_path is not None + model = get_model(_energy_model_params()).to(DEVICE) + model.compute_or_load_stat( + lambda: pytest.fail("complete update cache must not sample data"), + stat_path, + ) + + assert stat_file.read_bytes() == original def test_read_mode_rejects_missing_descriptor_stats_before_sampling( @@ -233,17 +369,14 @@ def test_read_mode_rejects_missing_descriptor_stats_before_sampling( with h5py.File(stat_file, "w") as file: file.create_group("O H") - read_path = _prepare_stat_file_path(str(stat_file), "read") - assert isinstance(read_path, DPH5Path) - try: + with open_stat_file(StatFileSpec(str(stat_file), "read")) as read_path: + assert read_path is not None model = get_model(_energy_model_params()).to(DEVICE) with pytest.raises(FileNotFoundError, match="environment statistics"): model.compute_or_load_stat( lambda: pytest.fail("read-only cache miss must not sample data"), read_path, ) - finally: - _close_stat_path(read_path) with h5py.File(stat_file, "r") as file: assert list(file.keys()) == ["O H"] @@ -254,13 +387,10 @@ def test_read_mode_rejects_partial_descriptor_stats_before_sampling( tmp_path: Path, ) -> None: stat_file = tmp_path / "stat.hdf5" - update_path = _prepare_stat_file_path(str(stat_file), "update") - assert isinstance(update_path, DPH5Path) - try: + with open_stat_file(StatFileSpec(str(stat_file))) as update_path: + assert update_path is not None model = get_model(_energy_model_params()).to(DEVICE) model.compute_or_load_stat(_energy_stat_sample, update_path) - finally: - _close_stat_path(update_path) with h5py.File(stat_file, "r+") as file: type_map_group = file["O H"] @@ -270,17 +400,14 @@ def test_read_mode_rejects_partial_descriptor_stats_before_sampling( assert len(descriptor_groups) == 1 del descriptor_groups[0]["r_0"] - read_path = _prepare_stat_file_path(str(stat_file), "read") - assert isinstance(read_path, DPH5Path) - try: + with open_stat_file(StatFileSpec(str(stat_file), "read")) as read_path: + assert read_path is not None model = get_model(_energy_model_params()).to(DEVICE) with pytest.raises(FileNotFoundError, match="'r_0'"): model.compute_or_load_stat( lambda: pytest.fail("partial read-only cache must not sample data"), read_path, ) - finally: - _close_stat_path(read_path) def test_stat_file_mode_configuration_validation() -> None: diff --git a/source/tests/pt_expt/test_entrypoint.py b/source/tests/pt_expt/test_entrypoint.py index aba7fcb339..9068df58d3 100644 --- a/source/tests/pt_expt/test_entrypoint.py +++ b/source/tests/pt_expt/test_entrypoint.py @@ -7,7 +7,6 @@ from deepmd.pt_expt.entrypoints.main import ( PTExptTrainEntrypoint, _ensure_pt_expt_model_suffix, - _ensure_stat_file_path, train, ) @@ -222,12 +221,3 @@ def state_dict(self) -> dict[str, object]: assert latest.is_symlink() assert latest.resolve() == ckpt_path assert latest.readlink().as_posix() == "model-1.pt" - - -def test_pt_expt_stat_file_path_creates_hdf5_parent(tmp_path) -> None: - stat_file = tmp_path / "stats" / "model_stat.hdf5" - - stat_path = _ensure_stat_file_path(str(stat_file)) - - assert stat_file.exists() - assert stat_path is not None diff --git a/source/tests/pt_expt/test_training.py b/source/tests/pt_expt/test_training.py index 193fa81911..a0dee795e9 100644 --- a/source/tests/pt_expt/test_training.py +++ b/source/tests/pt_expt/test_training.py @@ -14,7 +14,11 @@ import shutil import tempfile import unittest +from pathlib import ( + Path, +) from unittest.mock import ( + Mock, patch, ) @@ -37,6 +41,11 @@ update_deepmd_input, ) +from ..common.stat_file import ( + assert_energy_stat_cache_round_trip, + energy_model_params, +) + EXAMPLE_DIR = os.path.join( os.path.dirname(__file__), "..", @@ -257,6 +266,40 @@ def _make_config(data_dir: str, numb_steps: int = 5) -> dict: return config +def test_pt_expt_energy_cache_round_trip_uses_fitting_outputs_only( + tmp_path: Path, +) -> None: + assert_energy_stat_cache_round_trip( + lambda: get_model(energy_model_params()), + tmp_path / "stat.hdf5", + ) + + +def test_pt_expt_distributed_statistics_failure_reaches_peer_rank() -> None: + from deepmd.pt_expt.train.training import ( + Trainer, + ) + + trainer = Trainer.__new__(Trainer) + trainer.is_distributed = True + trainer.rank = 1 + action = Mock() + + def report_chief_failure(holder: list[bool], **_: object) -> None: + holder[0] = True + + with ( + patch( + "deepmd.pt_expt.train.training.dist.broadcast_object_list", + side_effect=report_chief_failure, + ), + pytest.raises(RuntimeError, match="Rank 0 failed during statistics"), + ): + trainer._run_stat_on_chief(action, operation="statistics initialization") + + action.assert_not_called() + + class TestTraining(unittest.TestCase): """Basic smoke test for the pt_expt training loop.""" diff --git a/source/tests/tf2/test_training.py b/source/tests/tf2/test_training.py index e282e983db..0c0a15b2b2 100644 --- a/source/tests/tf2/test_training.py +++ b/source/tests/tf2/test_training.py @@ -6,6 +6,9 @@ from contextlib import ( nullcontext, ) +from pathlib import ( + Path, +) from types import ( SimpleNamespace, ) @@ -49,12 +52,19 @@ from deepmd.tf2.model.base_model import ( forward_common_atomic, ) +from deepmd.tf2.model.model import ( + get_model, +) from deepmd.tf2.train.trainer import ( Trainer, ) from deepmd.tf2.utils.jit import ( default_jit_compile, ) +from source.tests.common.stat_file import ( + assert_energy_stat_cache_round_trip, + energy_model_params, +) pytestmark = [ pytest.mark.filterwarnings( @@ -169,6 +179,15 @@ def _make_minimal_trainer() -> tuple[Trainer, _LinearModel]: return trainer, model +def test_tf2_energy_cache_round_trip_uses_fitting_outputs_only( + tmp_path: Path, +) -> None: + assert_energy_stat_cache_round_trip( + lambda: get_model(energy_model_params()), + tmp_path / "stat.hdf5", + ) + + def test_forward_common_atomic_reuses_taped_atomic_forward() -> None: model = _FakeEnergyModel() coord = tf.constant(