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
7 changes: 5 additions & 2 deletions examples/water/dpa4/input_multitask.json
Original file line number Diff line number Diff line change
@@ -1,8 +1,6 @@
{
"_comment": "DPA4-Mini multitask example with a shared descriptor and case-conditioned shared fitting network.",
"model": {
"use_compile": false,
"enable_tf32": true,
"shared_dict": {
"type_map": [
"O",
Expand Down Expand Up @@ -52,6 +50,8 @@
},
"model_dict": {
"water_1": {
"use_compile": false,
"enable_tf32": true,
Comment thread
iProzd marked this conversation as resolved.
"type": "dpa4",
"type_map": "type_map",
"descriptor": "descriptor",
Expand All @@ -65,6 +65,8 @@
}
},
"water_2": {
"use_compile": false,
"enable_tf32": true,
"type": "dpa4",
"type_map": "type_map",
"descriptor": "descriptor",
Expand Down Expand Up @@ -116,6 +118,7 @@
"weight_decay": 0.001
},
"training": {
"enable_tf32": true,
"model_prob": {
"water_1": 0.5,
"water_2": 0.5
Expand Down
7 changes: 5 additions & 2 deletions examples/water/dpa4/input_multitask_preset.json
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,6 @@
"_comment": "DPA4-Nano multitask example with a shared descriptor and case-conditioned shared fitting network, both taken from the named preset; only run-specific keys are written out.",
"model": {
"preset": "dpa4-nano-v20260901",
"use_compile": false,
"enable_tf32": true,
"shared_dict": {
"type_map": [
"O",
Expand All @@ -21,6 +19,8 @@
},
"model_dict": {
"water_1": {
"use_compile": false,
"enable_tf32": true,
Comment thread
iProzd marked this conversation as resolved.
"type_map": "type_map",
"descriptor": "descriptor",
"fitting_net": "shared_fit_with_id",
Expand All @@ -33,6 +33,8 @@
}
},
"water_2": {
"use_compile": false,
"enable_tf32": true,
"type_map": "type_map",
"descriptor": "descriptor",
"fitting_net": "shared_fit_with_id",
Expand Down Expand Up @@ -83,6 +85,7 @@
"weight_decay": 0.001
},
"training": {
"enable_tf32": true,
"model_prob": {
"water_1": 0.5,
"water_2": 0.5
Expand Down
55 changes: 55 additions & 0 deletions source/tests/common/test_examples.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,61 @@ def test_arguments(self) -> None:
jdata["model"], _ = preprocess_shared_params(jdata["model"])
normalize(jdata, multi_task=multi_task)

def test_arguments_pt_expt(self) -> None:
"""The same configurations, through the PyTorch-Exportable path.

That backend does not cascade top-level ``model`` options into the
branches, which the pt one does, so a multi-task example can be valid
for pt and rejected here. ``test_arguments`` uses pt's preprocessing and
cannot see it.
"""
from deepmd.pt_expt.utils.multi_task import (
preprocess_shared_params as preprocess_shared_params_pt_expt,
)

for fn in input_files + input_files_multi:
multi_task = fn in input_files_multi
fn = str(fn)
with self.subTest(fn=fn):
jdata = j_loader(fn)
jdata["model"] = expand_model_preset(jdata["model"])
if multi_task:
jdata["model"], _ = preprocess_shared_params_pt_expt(jdata["model"])
normalize(jdata, multi_task=multi_task)

def test_tf32_says_what_pt_expt_will_do(self) -> None:
"""Schema acceptance is not the same as taking effect.

The PyTorch-Exportable backend reads ``training.enable_tf32`` and
ignores the model-level key, which belongs to pt. An example that sets
the branch-level one alone passes every schema check and then trains
with TF32 off while saying it is on.
"""
from deepmd.pt_expt.utils.multi_task import (
preprocess_shared_params as preprocess_shared_params_pt_expt,
)

for fn in input_files_multi:
jdata = j_loader(str(fn))
jdata["model"] = expand_model_preset(jdata["model"])
jdata["model"], _ = preprocess_shared_params_pt_expt(jdata["model"])
jdata = normalize(jdata, multi_task=True)
advertised = {
branch["enable_tf32"]
for branch in jdata["model"]["model_dict"].values()
if "enable_tf32" in branch
}
if not advertised:
continue
with self.subTest(fn=str(fn)):
self.assertEqual(
len(advertised), 1, "branches disagree about enable_tf32"
)
self.assertEqual(
bool(jdata["training"].get("enable_tf32", False)),
advertised.pop(),
)

def test_data_paths_exist(self) -> None:
"""Each example's data ``systems`` must resolve relative to the example's
own directory, so the example is runnable from that directory.
Expand Down
Loading