Skip to content

Commit 14183d7

Browse files
committed
fix: honor job-local LAMMPS time settings
Fixes #1904 Coding-Agent: Codex Codex-Version: codex-cli 0.149.0 Model: gpt-5.6-sol Reasoning-Effort: xhigh
1 parent d5ce577 commit 14183d7

3 files changed

Lines changed: 49 additions & 13 deletions

File tree

dpgen/generator/arginfo.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -292,6 +292,7 @@ def model_devi_jobs_args() -> list[Argument]:
292292
doc_nsteps = "Running steps of MD. It is not optional when not using a template."
293293
doc_nbeads = "Number of beads in PIMD. If not given, classical MD will be performed. Only supported for LAMMPS version >= 20230615."
294294
doc_ensemble = "Determining which ensemble used in MD, options include “npt” and “nvt”. It is not optional when not using a template."
295+
doc_dt = "Timestep for this MD job. Overrides the workflow-wide model_devi_dt."
295296
doc_neidelay = "delay building until this many steps since last build."
296297
doc_taut = "Coupling time of thermostat (ps)."
297298
doc_taup = "Coupling time of barostat (ps)."
@@ -311,6 +312,7 @@ def model_devi_jobs_args() -> list[Argument]:
311312
Argument("nsteps", int, optional=True, doc=doc_nsteps),
312313
Argument("nbeads", int, optional=True, doc=doc_nbeads),
313314
Argument("ensemble", str, optional=True, doc=doc_ensemble),
315+
Argument("dt", float, optional=True, doc=doc_dt),
314316
Argument("neidelay", int, optional=True, doc=doc_neidelay),
315317
Argument("taut", float, optional=True, doc=doc_taut),
316318
Argument("taup", float, optional=True, doc=doc_taup),

dpgen/generator/run.py

Lines changed: 19 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -995,6 +995,16 @@ def parse_cur_job(cur_job):
995995
return ensemble, nsteps, trj_freq, temps, press, pka_e, dt, nbeads
996996

997997

998+
def _get_lammps_job_settings(cur_job, jdata):
999+
"""Resolve per-job LAMMPS settings over workflow-wide defaults."""
1000+
return (
1001+
cur_job.get("dt", jdata["model_devi_dt"]),
1002+
cur_job.get("neidelay", jdata.get("model_devi_neidelay")),
1003+
cur_job.get("taut", jdata.get("model_devi_taut", 0.1)),
1004+
cur_job.get("taup", jdata.get("model_devi_taup", 0.5)),
1005+
)
1006+
1007+
9981008
def expand_matrix_values(target_list, cur_idx=0):
9991009
nvar = len(target_list)
10001010
if cur_idx == nvar:
@@ -1667,7 +1677,9 @@ def _make_model_devi_native(iter_index, jdata, mdata, conf_systems):
16671677
if iter_index >= len(model_devi_jobs):
16681678
return False
16691679
cur_job = model_devi_jobs[iter_index]
1670-
ensemble, nsteps, trj_freq, temps, press, pka_e, dt, nbeads = parse_cur_job(cur_job)
1680+
ensemble, nsteps, trj_freq, temps, press, pka_e, _dt, nbeads = parse_cur_job(
1681+
cur_job
1682+
)
16711683
model_devi_f_avg_relative = jdata.get("model_devi_f_avg_relative", False)
16721684
model_devi_merge_traj = jdata.get("model_devi_merge_traj", False)
16731685
if (nbeads is not None) and model_devi_f_avg_relative:
@@ -1682,23 +1694,17 @@ def _make_model_devi_native(iter_index, jdata, mdata, conf_systems):
16821694
raise RuntimeError(
16831695
"trj_freq should be a factor of nsteps for pimd. Please check your input."
16841696
)
1685-
if dt is not None:
1686-
model_devi_dt = dt
16871697
sys_idx = expand_idx(cur_job["sys_idx"])
16881698
if len(sys_idx) != len(list(set(sys_idx))):
16891699
raise RuntimeError("system index should be uniq")
16901700

16911701
use_ele_temp = jdata.get("use_ele_temp", 0)
1692-
model_devi_dt = jdata["model_devi_dt"]
1693-
model_devi_neidelay = None
1694-
if "model_devi_neidelay" in jdata:
1695-
model_devi_neidelay = jdata["model_devi_neidelay"]
1696-
model_devi_taut = 0.1
1697-
if "model_devi_taut" in jdata:
1698-
model_devi_taut = jdata["model_devi_taut"]
1699-
model_devi_taup = 0.5
1700-
if "model_devi_taup" in jdata:
1701-
model_devi_taup = jdata["model_devi_taup"]
1702+
(
1703+
model_devi_dt,
1704+
model_devi_neidelay,
1705+
model_devi_taut,
1706+
model_devi_taup,
1707+
) = _get_lammps_job_settings(cur_job, jdata)
17021708
mass_map = jdata["mass_map"]
17031709
nopbc = jdata.get("model_devi_nopbc", False)
17041710

tests/generator/test_parse_cur_job.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,9 @@
44

55
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
66
__package__ = "generator"
7+
from dpgen.generator.arginfo import model_devi_jobs_args
8+
from dpgen.generator.run import _get_lammps_job_settings
9+
710
from .context import (
811
parse_cur_job,
912
setUpModule, # noqa: F401
@@ -63,6 +66,31 @@ def test_pka(self):
6366
for ii, jj in zip(res, [ens, ns, tf, ts, [-1], pka, dt]):
6467
self.assertEqual(ii, jj)
6568

69+
def test_job_local_lammps_settings_override_global_defaults(self):
70+
cur_job = {"dt": 0.001, "neidelay": 5, "taut": 0.2, "taup": 1.0}
71+
jdata = {
72+
"model_devi_dt": 0.002,
73+
"model_devi_neidelay": 10,
74+
"model_devi_taut": 0.1,
75+
"model_devi_taup": 0.5,
76+
}
77+
self.assertEqual(_get_lammps_job_settings(cur_job, jdata), (0.001, 5, 0.2, 1.0))
78+
79+
def test_job_schema_accepts_local_timestep(self):
80+
arginfo = model_devi_jobs_args()
81+
jobs = [
82+
{
83+
"sys_idx": [0],
84+
"ensemble": "nvt",
85+
"temps": [300.0],
86+
"nsteps": 100,
87+
"trj_freq": 10,
88+
"dt": 0.001,
89+
}
90+
]
91+
normalized = arginfo.normalize_value(jobs, trim_pattern="_*")
92+
arginfo.check_value(normalized, strict=True)
93+
6694

6795
if __name__ == "__main__":
6896
unittest.main()

0 commit comments

Comments
 (0)