Skip to content

Commit 55d7e79

Browse files
njzjzpre-commit-ci[bot]njzjz-bot
authored
feat(dpmodel): add backend-independent trainer abstraction (#5603)
## Summary - Add backend-independent training abstractions under `deepmd.dpmodel.train` for task/rank normalization, display scheduling, learning-curve output, checkpoint cadence, lifecycle hooks, and shared train entrypoint orchestration. - Factor common training-data helpers so single-task training is handled as a one-task collection and multi-task data construction/summary/probability handling is shared where possible. - Add a backend-independent finetune rule builder in `deepmd.utils.finetune`, and reduce the PT, PT-exportable, Paddle, and JAX backend finetune modules to backend-specific checkpoint loading plus shared rule generation. - Migrate JAX train entrypoint/trainer onto the shared pipeline and add JAX finetune plus multi-task support on top of the new abstractions. - Migrate `pt_expt` train entrypoint/trainer behavior further onto the shared pipeline, including single-task-as-multi-task normalization, data summaries, checkpoint retention, stat-file parent creation, relative latest checkpoint symlinks, and checkpoint parent creation. - Address PR review comments around task-key validation, learning-curve metric ordering, lifecycle cleanup, `print_summary` fallback behavior, broken `__len__` handling, JAX finetune branch/alias validation, numeric-looking JAX task keys, HDF5 stat paths, and `pt_expt` checkpoint symlinks. - Move the new dpmodel trainer/entrypoint tests from `source/tests/test_dpmodel_*.py` into `source/tests/common/dpmodel/`. Refs #5229, #5230, #5231 ## Tests - `ruff format .` - `ruff check .` - `git diff --check` - `PYTHONPATH=/home/jzzeng/codes/deepmd-kit pytest source/tests/common/dpmodel/test_train_abstract_trainer.py source/tests/common/dpmodel/test_train_entrypoint.py source/tests/common/dpmodel/test_train_data.py source/tests/common/dpmodel/test_training_utils.py source/tests/common/test_finetune_utils.py source/tests/jax/test_training.py source/tests/pt_expt/test_entrypoint.py source/tests/pt_expt/test_multitask.py::TestMultiTaskSeA::test_multitask_finetune source/tests/pt_expt/test_multitask.py::TestMultiTaskSeA::test_multitask_finetune_from_single_task source/tests/pt_expt/test_multitask.py::TestMultiTaskSeA::test_multitask_finetune_no_change_model_params -q` (`53 passed, 2 subtests passed`) - `PYTHONPATH=/home/jzzeng/codes/deepmd-kit timeout 180 srun --gres=gpu:1 dp --jax train input.json --skip-neighbor-stat --finetune pretrain.jax --use-pretrain-script` on a temporary 1-step water finetune smoke; completed on NVIDIA GeForce RTX 5090 and saved `ft-model-1.jax`. - `PYTHONPATH=/home/jzzeng/codes/deepmd-kit timeout 180 srun --gres=gpu:1 dp --pt-expt train input.json --skip-neighbor-stat` on a temporary 2-step water smoke; completed on NVIDIA GeForce RTX 5090, saved `ckpts/pt-model-2.pt`, created `stats/stat.hdf5`, and verified `ckpts/pt-model.pt -> pt-model-2.pt` with old step checkpoint pruned by `max_ckpt_keep=1`. ## Notes - Paddle-specific runtime tests were not run locally because `paddle` is not installed in this environment. - Plain PyTorch backend test collection is blocked in this environment by external `deepmd_gnn`/CUDA initialization, not by the shared finetune rule builder changes. <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Introduced a unified, backend-independent training framework with consistent single-task and multi-task handling, learning-curve output, and structured training/validation steps. * Added a common training entrypoint abstraction that standardizes config normalization, neighbor-stat updates, and lifecycle teardown. * Implemented full-validation with best-checkpoint tracking, top-K selection, and `val.log` reporting (including backend-specific checkpoint suffixes). * **Bug Fixes** * Improved checkpoint save/restore and retention (including “latest” link updates and older checkpoint cleanup). * Improved task-weighting logic to better handle datasets with/without sizing information. * Fixed multi-task neighbor-stat updates and JAX full-validation error propagation across ranks. * **Tests** * Expanded unit and smoke tests for training orchestration, finetuning, validation, and checkpoint reconciliation. * **Documentation** * Updated validation-configuration help text to reflect broader backend support. <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: njzjz-bot (driven by OpenClaw (model: custom-chat-jinzhezeng-group/gpt-5.5))[bot] <48687836+njzjz-bot@users.noreply.github.com>
1 parent 4c94171 commit 55d7e79

28 files changed

Lines changed: 5542 additions & 1495 deletions

deepmd/dpmodel/train/__init__.py

Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,40 @@
1+
# SPDX-License-Identifier: LGPL-3.0-or-later
2+
"""Backend-independent training abstractions."""
3+
4+
from .data import (
5+
TrainingTaskConfig,
6+
iter_training_task_configs,
7+
make_task_maps,
8+
print_data_summaries,
9+
)
10+
from .entrypoint import (
11+
AbstractTrainEntrypoint,
12+
TrainEntrypointOptions,
13+
)
14+
from .trainer import (
15+
DEFAULT_TASK_KEY,
16+
AbstractTrainer,
17+
LearningCurveWriter,
18+
RankContext,
19+
TrainerConfig,
20+
TrainingTask,
21+
TrainingTaskCollection,
22+
TrainStepResult,
23+
)
24+
25+
__all__ = [
26+
"DEFAULT_TASK_KEY",
27+
"AbstractTrainEntrypoint",
28+
"AbstractTrainer",
29+
"LearningCurveWriter",
30+
"RankContext",
31+
"TrainEntrypointOptions",
32+
"TrainStepResult",
33+
"TrainerConfig",
34+
"TrainingTask",
35+
"TrainingTaskCollection",
36+
"TrainingTaskConfig",
37+
"iter_training_task_configs",
38+
"make_task_maps",
39+
"print_data_summaries",
40+
]

deepmd/dpmodel/train/data.py

Lines changed: 137 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,137 @@
1+
# SPDX-License-Identifier: LGPL-3.0-or-later
2+
"""Shared training-data helpers for backend entrypoints."""
3+
4+
from __future__ import (
5+
annotations,
6+
)
7+
8+
import inspect
9+
from dataclasses import (
10+
dataclass,
11+
)
12+
from typing import (
13+
TYPE_CHECKING,
14+
Any,
15+
)
16+
17+
from .trainer import (
18+
DEFAULT_TASK_KEY,
19+
)
20+
21+
if TYPE_CHECKING:
22+
from collections.abc import (
23+
Callable,
24+
Iterator,
25+
Mapping,
26+
)
27+
28+
29+
@dataclass(frozen=True)
30+
class TrainingTaskConfig:
31+
"""Normalized config view for one training task."""
32+
33+
key: str
34+
model_params: Mapping[str, Any]
35+
training_data_params: Mapping[str, Any]
36+
validation_data_params: Mapping[str, Any] | None
37+
stat_file: str | None
38+
valid_numb_batch: int
39+
40+
41+
def iter_training_task_configs(
42+
config: Mapping[str, Any],
43+
) -> Iterator[TrainingTaskConfig]:
44+
"""Yield task configs, treating single-task input as one ``Default`` task."""
45+
model_params = config["model"]
46+
training_params = config["training"]
47+
if "model_dict" not in model_params:
48+
validation_data_params = training_params.get("validation_data")
49+
yield TrainingTaskConfig(
50+
key=DEFAULT_TASK_KEY,
51+
model_params=model_params,
52+
training_data_params=training_params["training_data"],
53+
validation_data_params=validation_data_params,
54+
stat_file=training_params.get("stat_file"),
55+
valid_numb_batch=_valid_numb_batch(validation_data_params),
56+
)
57+
return
58+
59+
data_dict = training_params["data_dict"]
60+
for task_key, task_model_params in model_params["model_dict"].items():
61+
task_data_params = data_dict[task_key]
62+
validation_data_params = task_data_params.get("validation_data")
63+
yield TrainingTaskConfig(
64+
key=task_key,
65+
model_params=task_model_params,
66+
training_data_params=task_data_params["training_data"],
67+
validation_data_params=validation_data_params,
68+
stat_file=task_data_params.get("stat_file"),
69+
valid_numb_batch=_valid_numb_batch(validation_data_params),
70+
)
71+
72+
73+
def make_task_maps(
74+
config: Mapping[str, Any],
75+
factory: Callable[[TrainingTaskConfig], tuple[Any, Any | None, Any | None]],
76+
) -> tuple[dict[str, Any], dict[str, Any | None], dict[str, Any | None]]:
77+
"""Build training, validation, and stat maps from normalized task configs."""
78+
training_data: dict[str, Any] = {}
79+
validation_data: dict[str, Any | None] = {}
80+
stat_data: dict[str, Any | None] = {}
81+
for task_config in iter_training_task_configs(config):
82+
train_item, valid_item, stat_item = factory(task_config)
83+
training_data[task_config.key] = train_item
84+
validation_data[task_config.key] = valid_item
85+
stat_data[task_config.key] = stat_item
86+
return training_data, validation_data, stat_data
87+
88+
89+
def print_data_summaries(
90+
training_data: Mapping[str, Any],
91+
validation_data: Mapping[str, Any | None],
92+
*,
93+
probabilities: Mapping[str, float] | None = None,
94+
) -> None:
95+
"""Print train/validation data summaries for one or more tasks."""
96+
multi_task = len(training_data) > 1
97+
for task_key, data in training_data.items():
98+
name = f"training data({task_key})" if multi_task else "training"
99+
_print_summary(data, name, _task_probability(probabilities, task_key))
100+
valid_data = validation_data.get(task_key)
101+
if valid_data is not None:
102+
name = f"validation data({task_key})" if multi_task else "validation"
103+
_print_summary(valid_data, name, None)
104+
105+
106+
def _valid_numb_batch(validation_data_params: Mapping[str, Any] | None) -> int:
107+
if validation_data_params is None:
108+
return 1
109+
return max(int(validation_data_params.get("numb_btch", 1)), 1)
110+
111+
112+
def _task_probability(
113+
probabilities: Mapping[str, float] | None,
114+
task_key: str,
115+
) -> list[float] | None:
116+
if probabilities is None or task_key not in probabilities:
117+
return None
118+
return [float(probabilities[task_key])]
119+
120+
121+
def _print_summary(data: Any, name: str, prob: list[float] | None) -> None:
122+
printer = data.print_summary
123+
try:
124+
signature = inspect.signature(printer)
125+
except (TypeError, ValueError):
126+
printer(name, prob)
127+
return
128+
try:
129+
signature.bind(name, prob)
130+
except TypeError as exc:
131+
try:
132+
signature.bind(name)
133+
except TypeError:
134+
raise exc from None
135+
printer(name)
136+
else:
137+
printer(name, prob)

deepmd/dpmodel/train/entrypoint.py

Lines changed: 173 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,173 @@
1+
# SPDX-License-Identifier: LGPL-3.0-or-later
2+
"""Backend-independent training entrypoint pipeline."""
3+
4+
from __future__ import (
5+
annotations,
6+
)
7+
8+
import json
9+
import logging
10+
from abc import (
11+
ABC,
12+
abstractmethod,
13+
)
14+
from dataclasses import (
15+
dataclass,
16+
)
17+
from typing import (
18+
Any,
19+
)
20+
21+
from deepmd.common import (
22+
j_loader,
23+
)
24+
from deepmd.utils.argcheck import (
25+
normalize,
26+
)
27+
from deepmd.utils.compat import (
28+
update_deepmd_input,
29+
)
30+
31+
log = logging.getLogger(__name__)
32+
33+
34+
@dataclass
35+
class TrainEntrypointOptions:
36+
"""Common command options for backend train entrypoints."""
37+
38+
input_file: str
39+
output: str = "out.json"
40+
init_model: str | None = None
41+
restart: str | None = None
42+
init_frz_model: str | None = None
43+
finetune: str | None = None
44+
model_branch: str = ""
45+
use_pretrain_script: bool = False
46+
skip_neighbor_stat: bool = False
47+
48+
49+
class AbstractTrainEntrypoint(ABC):
50+
"""Shared pipeline for backend train entrypoints.
51+
52+
Backend subclasses keep ownership of backend-specific feature handling,
53+
neighbor-stat updates, distributed setup, data construction, and trainer
54+
construction. This pipeline only coordinates the common command flow.
55+
"""
56+
57+
def run(self, options: TrainEntrypointOptions) -> None:
58+
"""Run the training entrypoint."""
59+
log.info("Configuration path: %s", options.input_file)
60+
options = self.prepare_options(options)
61+
config = self.load_config(options.input_file)
62+
self.validate_options(config, options)
63+
64+
config = self.preprocess_config(config, options)
65+
multi_task = self.is_multi_task(config)
66+
config = self.update_input(config)
67+
config = self.normalize_config(config, multi_task=multi_task)
68+
69+
neighbor_stat = None
70+
if not options.skip_neighbor_stat:
71+
config, neighbor_stat = self.update_neighbor_stat(
72+
config,
73+
options,
74+
multi_task=multi_task,
75+
)
76+
77+
self.dump_config(config, options.output)
78+
self.print_summary()
79+
80+
try:
81+
self.setup_run(options, config)
82+
self.run_training(config, options, neighbor_stat)
83+
finally:
84+
self.teardown_run(options, config)
85+
86+
def prepare_options(
87+
self,
88+
options: TrainEntrypointOptions,
89+
) -> TrainEntrypointOptions:
90+
"""Normalize command options before reading or preprocessing config."""
91+
return options
92+
93+
def load_config(self, input_file: str) -> dict[str, Any]:
94+
"""Load the JSON/YAML training config."""
95+
return j_loader(input_file)
96+
97+
def validate_options(
98+
self,
99+
config: dict[str, Any],
100+
options: TrainEntrypointOptions,
101+
) -> None:
102+
"""Validate backend feature support before mutating the config."""
103+
return None
104+
105+
def preprocess_config(
106+
self,
107+
config: dict[str, Any],
108+
options: TrainEntrypointOptions,
109+
) -> dict[str, Any]:
110+
"""Apply backend-specific config preprocessing before argcheck."""
111+
return config
112+
113+
def is_multi_task(self, config: dict[str, Any]) -> bool:
114+
"""Return whether the config is in multi-task layout."""
115+
return "model_dict" in config.get("model", {})
116+
117+
def update_input(self, config: dict[str, Any]) -> dict[str, Any]:
118+
"""Apply DeePMD input-version compatibility conversion."""
119+
return update_deepmd_input(config, warning=True, dump="input_v2_compat.json")
120+
121+
def normalize_config(
122+
self,
123+
config: dict[str, Any],
124+
*,
125+
multi_task: bool,
126+
) -> dict[str, Any]:
127+
"""Run DeePMD argcheck normalization."""
128+
return normalize(config, multi_task=multi_task)
129+
130+
def update_neighbor_stat(
131+
self,
132+
config: dict[str, Any],
133+
options: TrainEntrypointOptions,
134+
*,
135+
multi_task: bool,
136+
) -> tuple[dict[str, Any], Any]:
137+
"""Update descriptor selections from neighbor statistics."""
138+
return config, None
139+
140+
def dump_config(self, config: dict[str, Any], output: str) -> None:
141+
"""Dump the normalized config used for training."""
142+
with open(output, "w") as fp:
143+
json.dump(config, fp, indent=4)
144+
145+
def print_summary(self) -> None:
146+
"""Print backend summary information."""
147+
return None
148+
149+
def setup_run(
150+
self,
151+
options: TrainEntrypointOptions,
152+
config: dict[str, Any],
153+
) -> None:
154+
"""Set up backend runtime state before trainer execution."""
155+
return None
156+
157+
def teardown_run(
158+
self,
159+
options: TrainEntrypointOptions,
160+
config: dict[str, Any],
161+
) -> None:
162+
"""Tear down backend runtime state after trainer execution."""
163+
return None
164+
165+
@abstractmethod
166+
def run_training(
167+
self,
168+
config: dict[str, Any],
169+
options: TrainEntrypointOptions,
170+
neighbor_stat: Any,
171+
) -> None:
172+
"""Build backend data/trainer objects and run training."""
173+
raise NotImplementedError

0 commit comments

Comments
 (0)