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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions deepmd/dpmodel/atomic_model/base_atomic_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand Down
31 changes: 31 additions & 0 deletions deepmd/dpmodel/descriptor/descriptor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
10 changes: 1 addition & 9 deletions deepmd/dpmodel/descriptor/dpa1.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
10 changes: 1 addition & 9 deletions deepmd/dpmodel/descriptor/repflows.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
10 changes: 1 addition & 9 deletions deepmd/dpmodel/descriptor/repformers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
10 changes: 1 addition & 9 deletions deepmd/dpmodel/descriptor/se_e2_a.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
10 changes: 1 addition & 9 deletions deepmd/dpmodel/descriptor/se_r.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
10 changes: 1 addition & 9 deletions deepmd/dpmodel/descriptor/se_t.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
10 changes: 1 addition & 9 deletions deepmd/dpmodel/descriptor/se_t_tebd.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
41 changes: 17 additions & 24 deletions deepmd/dpmodel/fitting/general_fitting.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand All @@ -51,6 +48,9 @@
from deepmd.utils.path import (
DPPath,
)
from deepmd.utils.stat_file import (
load_required_items,
)

from .base_fitting import (
BaseFitting,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
24 changes: 24 additions & 0 deletions deepmd/dpmodel/train/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,11 @@
Any,
)

from deepmd.utils.stat_file import (
StatFileMode,
StatFileSpec,
)

from .trainer import (
DEFAULT_TASK_KEY,
)
Expand All @@ -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(
Expand All @@ -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

Expand All @@ -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"),
)


Expand Down
34 changes: 12 additions & 22 deletions deepmd/dpmodel/utils/env_mat_stat.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
Loading
Loading