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
16 changes: 14 additions & 2 deletions deepmd/dpmodel/atomic_model/base_atomic_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,9 @@
map_atom_exclude_types,
map_pair_exclude_types,
)
from deepmd.utils.out_stat import (
get_redu_stat_scanner,
)
from deepmd.utils.path import (
DPPath,
)
Expand Down Expand Up @@ -106,15 +109,21 @@ def _collect_and_set_observed_type(
_restore_observed_type_from_file,
_save_observed_type_to_file,
collect_observed_types,
observed_types_from_counts,
)

if preset_observed_type is not None:
self._observed_type = preset_observed_type
else:
observed = _restore_observed_type_from_file(stat_file_path)
if observed is None:
sampled = sampled_func()
observed = collect_observed_types(sampled, self.type_map)
scanner = get_redu_stat_scanner(sampled_func)
if scanner is not None:
observed = observed_types_from_counts(
scanner.natoms_total(len(self.type_map)), self.type_map
)
else:
observed = collect_observed_types(sampled_func(), self.type_map)
_save_observed_type_to_file(stat_file_path, observed)
self._observed_type = observed

Expand Down Expand Up @@ -767,6 +776,9 @@ def wrapped_sampler() -> list[dict]:
sample["find_fparam"] = np.bool_(True)
return sampled

# the full-data scanner, when the trainer attached one, is part of the
# sampler contract and must survive wrapping
wrapped_sampler.redu_stat_scanner = get_redu_stat_scanner(sampled_func)
return wrapped_sampler

def change_out_bias(
Expand Down
26 changes: 26 additions & 0 deletions deepmd/dpmodel/utils/stat.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,32 @@ def collect_observed_types(sampled: list[dict], type_map: list[str]) -> list[str
return sort_element_type(observed_types)


def observed_types_from_counts(
natoms_total: np.ndarray, type_map: list[str]
) -> list[str]:
"""Collect observed element types from per-type atom counts.

Parameters
----------
natoms_total : np.ndarray
Total occurrences of each type, shape ``[ntypes]``.
type_map : list[str]
Mapping from type index to element symbol.

Returns
-------
list[str]
Sorted list of observed element symbols.
"""
from deepmd.utils.econf_embd import (
sort_element_type,
)

return sort_element_type(
[type_map[i] for i in np.flatnonzero(np.asarray(natoms_total) > 0)]
)


def _restore_observed_type_from_file(
stat_file_path: DPPath | None,
) -> list[str] | None:
Expand Down
7 changes: 7 additions & 0 deletions deepmd/pd/model/atomic_model/dp_atomic_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,9 @@
from deepmd.pd.model.task.base_fitting import (
BaseFitting,
)
from deepmd.utils.out_stat import (
get_redu_stat_scanner,
)
from deepmd.utils.path import (
DPPath,
)
Expand Down Expand Up @@ -409,6 +412,10 @@ def wrapped_sampler() -> list[dict]:
sample["atom_exclude_types"] = list(atom_exclude_types)
return sampled

# the full-data scanner, when the trainer attached one, is part of the
# sampler contract and must survive wrapping
wrapped_sampler.redu_stat_scanner = get_redu_stat_scanner(sampled_func)

self.descriptor.compute_input_stats(wrapped_sampler, stat_file_path)
self.compute_fitting_input_stat(wrapped_sampler, stat_file_path)
if compute_or_load_out_stat:
Expand Down
18 changes: 18 additions & 0 deletions deepmd/pd/train/training.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,7 @@
)
from deepmd.pd.utils.stat import (
make_stat_input,
scan_redu_stats,
)
from deepmd.pd.utils.utils import (
nvprof_context,
Expand All @@ -84,6 +85,9 @@
from deepmd.utils.finetune import (
warn_configuration_mismatch_during_finetune,
)
from deepmd.utils.out_stat import (
ReduStatScanner,
)
from deepmd.utils.path import (
DPH5Path,
)
Expand Down Expand Up @@ -233,6 +237,7 @@ def single_model_stat(
_validation_data: Any | None,
_stat_file_path: str | Path | None,
_data_requirement: list[DataRequirementItem],
_data_stat_full: bool = False,
finetune_has_new_type: bool = False,
) -> Any:
_data_requirement += get_additional_data_requirement(_model)
Expand All @@ -249,6 +254,15 @@ def get_sample() -> dict[str, Any]:
)
return sampled

if _data_stat_full:
# sampling a few batches per system can miss rare elements
# entirely; scan every frame for the output statistics instead
get_sample.redu_stat_scanner = ReduStatScanner(
lambda ntypes, keys, intensive: scan_redu_stats(
_training_data.dataloaders, ntypes, keys, intensive=intensive
)
)

if (not resuming or finetune_has_new_type) and self.rank == 0:
_model.compute_or_load_stat(
sampled_func=get_sample,
Expand Down Expand Up @@ -310,6 +324,7 @@ def get_lr(lr_params: dict[str, Any]) -> BaseLR:
validation_data,
stat_file_path,
self.loss.label_requirement,
_data_stat_full=model_params.get("data_stat_full", False),
finetune_has_new_type=self.finetune_links["Default"].get_has_new_type()
if self.finetune_links is not None
else False,
Expand Down Expand Up @@ -349,6 +364,9 @@ def get_lr(lr_params: dict[str, Any]) -> BaseLR:
validation_data[model_key],
stat_file_path[model_key],
self.loss[model_key].label_requirement,
_data_stat_full=model_params["model_dict"][model_key].get(
"data_stat_full", False
),
finetune_has_new_type=self.finetune_links[
model_key
].get_has_new_type()
Expand Down
Loading
Loading