diff --git a/dpgen/generator/arginfo.py b/dpgen/generator/arginfo.py index 0f8abb865..a6987ac15 100644 --- a/dpgen/generator/arginfo.py +++ b/dpgen/generator/arginfo.py @@ -292,6 +292,7 @@ def model_devi_jobs_args() -> list[Argument]: doc_nsteps = "Running steps of MD. It is not optional when not using a template." doc_nbeads = "Number of beads in PIMD. If not given, classical MD will be performed. Only supported for LAMMPS version >= 20230615." doc_ensemble = "Determining which ensemble used in MD, options include “npt” and “nvt”. It is not optional when not using a template." + doc_dt = "Timestep for this MD job. Overrides the workflow-wide model_devi_dt." doc_neidelay = "delay building until this many steps since last build." doc_taut = "Coupling time of thermostat (ps)." doc_taup = "Coupling time of barostat (ps)." @@ -311,6 +312,7 @@ def model_devi_jobs_args() -> list[Argument]: Argument("nsteps", int, optional=True, doc=doc_nsteps), Argument("nbeads", int, optional=True, doc=doc_nbeads), Argument("ensemble", str, optional=True, doc=doc_ensemble), + Argument("dt", float, optional=True, doc=doc_dt), Argument("neidelay", int, optional=True, doc=doc_neidelay), Argument("taut", float, optional=True, doc=doc_taut), Argument("taup", float, optional=True, doc=doc_taup), diff --git a/dpgen/generator/run.py b/dpgen/generator/run.py index 0ce513a96..c37cc02cc 100644 --- a/dpgen/generator/run.py +++ b/dpgen/generator/run.py @@ -995,6 +995,16 @@ def parse_cur_job(cur_job): return ensemble, nsteps, trj_freq, temps, press, pka_e, dt, nbeads +def _get_lammps_job_settings(cur_job, jdata): + """Resolve per-job LAMMPS settings over workflow-wide defaults.""" + return ( + cur_job.get("dt", jdata["model_devi_dt"]), + cur_job.get("neidelay", jdata.get("model_devi_neidelay")), + cur_job.get("taut", jdata.get("model_devi_taut", 0.1)), + cur_job.get("taup", jdata.get("model_devi_taup", 0.5)), + ) + + def expand_matrix_values(target_list, cur_idx=0): nvar = len(target_list) if cur_idx == nvar: @@ -1667,7 +1677,9 @@ def _make_model_devi_native(iter_index, jdata, mdata, conf_systems): if iter_index >= len(model_devi_jobs): return False cur_job = model_devi_jobs[iter_index] - ensemble, nsteps, trj_freq, temps, press, pka_e, dt, nbeads = parse_cur_job(cur_job) + ensemble, nsteps, trj_freq, temps, press, pka_e, _dt, nbeads = parse_cur_job( + cur_job + ) model_devi_f_avg_relative = jdata.get("model_devi_f_avg_relative", False) model_devi_merge_traj = jdata.get("model_devi_merge_traj", False) 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): raise RuntimeError( "trj_freq should be a factor of nsteps for pimd. Please check your input." ) - if dt is not None: - model_devi_dt = dt sys_idx = expand_idx(cur_job["sys_idx"]) if len(sys_idx) != len(list(set(sys_idx))): raise RuntimeError("system index should be uniq") use_ele_temp = jdata.get("use_ele_temp", 0) - model_devi_dt = jdata["model_devi_dt"] - model_devi_neidelay = None - if "model_devi_neidelay" in jdata: - model_devi_neidelay = jdata["model_devi_neidelay"] - model_devi_taut = 0.1 - if "model_devi_taut" in jdata: - model_devi_taut = jdata["model_devi_taut"] - model_devi_taup = 0.5 - if "model_devi_taup" in jdata: - model_devi_taup = jdata["model_devi_taup"] + ( + model_devi_dt, + model_devi_neidelay, + model_devi_taut, + model_devi_taup, + ) = _get_lammps_job_settings(cur_job, jdata) mass_map = jdata["mass_map"] nopbc = jdata.get("model_devi_nopbc", False) diff --git a/tests/generator/test_parse_cur_job.py b/tests/generator/test_parse_cur_job.py index fb0a5a911..853a1f973 100644 --- a/tests/generator/test_parse_cur_job.py +++ b/tests/generator/test_parse_cur_job.py @@ -4,6 +4,9 @@ sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))) __package__ = "generator" +from dpgen.generator.arginfo import model_devi_jobs_args +from dpgen.generator.run import _get_lammps_job_settings + from .context import ( parse_cur_job, setUpModule, # noqa: F401 @@ -63,6 +66,31 @@ def test_pka(self): for ii, jj in zip(res, [ens, ns, tf, ts, [-1], pka, dt]): self.assertEqual(ii, jj) + def test_job_local_lammps_settings_override_global_defaults(self): + cur_job = {"dt": 0.001, "neidelay": 5, "taut": 0.2, "taup": 1.0} + jdata = { + "model_devi_dt": 0.002, + "model_devi_neidelay": 10, + "model_devi_taut": 0.1, + "model_devi_taup": 0.5, + } + self.assertEqual(_get_lammps_job_settings(cur_job, jdata), (0.001, 5, 0.2, 1.0)) + + def test_job_schema_accepts_local_timestep(self): + arginfo = model_devi_jobs_args() + jobs = [ + { + "sys_idx": [0], + "ensemble": "nvt", + "temps": [300.0], + "nsteps": 100, + "trj_freq": 10, + "dt": 0.001, + } + ] + normalized = arginfo.normalize_value(jobs, trim_pattern="_*") + arginfo.check_value(normalized, strict=True) + if __name__ == "__main__": unittest.main()