Skip to content
Open
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
2 changes: 2 additions & 0 deletions dpgen/generator/arginfo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)."
Expand All @@ -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),
Expand Down
32 changes: 19 additions & 13 deletions dpgen/generator/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand All @@ -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)

Expand Down
28 changes: 28 additions & 0 deletions tests/generator/test_parse_cur_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()