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
17 changes: 17 additions & 0 deletions dpgen/generator/arginfo.py
Original file line number Diff line number Diff line change
Expand Up @@ -410,6 +410,10 @@ def model_devi_lmp_args() -> list[Argument]:
doc_model_devi_perc_candi_f = "See model_devi_adapt_trust_lo."
doc_model_devi_perc_candi_v = "See model_devi_adapt_trust_lo."
doc_model_devi_f_avg_relative = "Normalized the force model deviations by the RMS force magnitude along the trajectory. This key should not be used with use_relative."
doc_model_devi_recovery = (
"Opt-in recovery for failed exploration tasks. Set enabled=true to "
"validate and salvage complete prefix frames."
)
doc_model_devi_clean_traj = "If type of model_devi_clean_traj is bool type then it denote whether to clean traj folders in MD since they are too large. If it is Int type, then the most recent n iterations of traj folders will be retained, others will be removed."
doc_model_devi_merge_traj = "If model_devi_merge_traj is set as True, only all.lammpstrj will be generated, instead of lots of small traj files."
doc_model_devi_nopbc = "Assume open boundary condition in MD simulations."
Expand Down Expand Up @@ -498,6 +502,19 @@ def model_devi_lmp_args() -> list[Argument]:
optional=True,
doc=doc_model_devi_f_avg_relative,
),
Argument(
"model_devi_recovery",
dict,
optional=True,
default={
"enabled": False,
"max_failed_tasks": 0,
"max_failed_ratio": 0.0,
"salvage_prefix": False,
"min_valid_frames": 1,
},
doc=doc_model_devi_recovery,
),
Argument(
"model_devi_clean_traj",
[bool, int],
Expand Down
231 changes: 229 additions & 2 deletions dpgen/generator/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -3260,7 +3260,21 @@ def run_md_model_devi(iter_index, jdata, mdata):
outlog="model_devi.log",
errlog="model_devi.log",
)
submission.run_submission()
recovery_policy = _normalize_model_devi_recovery(jdata)
if recovery_policy["enabled"]:
try:
submission.run_submission(
continue_on_failure=True,
raise_on_failure=False,
include_failed_results=True,
)
except TypeError as error:
raise RuntimeError(
"model_devi_recovery requires a DPDispatcher with "
"continue_on_failure/raise_on_failure/include_failed_results support"
) from error
else:
submission.run_submission()


def run_model_devi(iter_index, jdata, mdata):
Expand All @@ -3277,7 +3291,46 @@ def run_model_devi(iter_index, jdata, mdata):


def post_model_devi(iter_index, jdata, mdata):
pass
"""Validate/summarize failed exploration tasks when recovery is enabled."""
policy = _normalize_model_devi_recovery(jdata)
if not policy["enabled"]:
return
iter_name = make_iter_name(iter_index)
work_path = Path(iter_name) / model_devi_name
tasks = sorted(work_path.glob("task.*"))
if not tasks:
raise RuntimeError("model_devi_recovery found no exploration tasks")
model_devi_merge_traj = jdata.get("model_devi_merge_traj", False)
reports = [
_recovery_task_report(task, model_devi_merge_traj, policy["salvage_prefix"])
for task in tasks
]
failed = sum(item["status"] != "completed" for item in reports)
failed_ratio = failed / len(reports)
if failed > policy["max_failed_tasks"] or failed_ratio > float(
policy["max_failed_ratio"]
):
raise RuntimeError(
"model_devi_recovery failure budget exceeded: "
f"{failed}/{len(reports)} tasks failed ({failed_ratio:.3f})"
)
valid_frames = sum(len(item["valid_steps"]) for item in reports)
if valid_frames < policy["min_valid_frames"]:
raise RuntimeError(
"model_devi_recovery found no sufficient valid exploration frames: "
f"{valid_frames} < {policy['min_valid_frames']}"
)
report = {
"version": 1,
"iteration": iter_index,
"policy": policy,
"tasks": reports,
"failed_tasks": failed,
"failed_ratio": failed_ratio,
"valid_frames": valid_frames,
}
report_path = Path(iter_name) / "exploration_report.json"
report_path.write_text(json.dumps(report, indent=2) + "\n")


def _to_face_dist(box_):
Expand Down Expand Up @@ -3361,6 +3414,173 @@ def check_bad_box(conf_name, criteria, fmt="lammps/dump"):
return is_bad


def _normalize_model_devi_recovery(jdata):
"""Validate and normalize the opt-in exploration recovery policy."""
raw = jdata.get("model_devi_recovery") or {}
if not isinstance(raw, dict):
raise TypeError("model_devi_recovery must be a mapping")
policy = {
"enabled": bool(raw.get("enabled", False)),
"max_failed_tasks": raw.get("max_failed_tasks", 0),
"max_failed_ratio": raw.get("max_failed_ratio", 0.0),
"salvage_prefix": bool(raw.get("salvage_prefix", False)),
"min_valid_frames": raw.get("min_valid_frames", 1),
}
if (
not isinstance(policy["max_failed_tasks"], int)
or isinstance(policy["max_failed_tasks"], bool)
or policy["max_failed_tasks"] < 0
):
raise ValueError("model_devi_recovery.max_failed_tasks must be nonnegative")
if (
not isinstance(policy["max_failed_ratio"], (int, float))
or not 0 <= float(policy["max_failed_ratio"]) <= 1
):
raise ValueError("model_devi_recovery.max_failed_ratio must be in [0, 1]")
if (
not isinstance(policy["min_valid_frames"], int)
or isinstance(policy["min_valid_frames"], bool)
or policy["min_valid_frames"] < 1
):
raise ValueError("model_devi_recovery.min_valid_frames must be positive")
return policy


def _read_lammps_dump_steps(filename):
"""Return valid timestep/frame metadata from a LAMMPS dump file."""
lines = Path(filename).read_text(errors="replace").splitlines()
steps, invalid = [], 0
expected_atoms, expected_mapping = None, None
cursor = 0
while cursor < len(lines):
if lines[cursor].strip() != "ITEM: TIMESTEP":
cursor += 1
continue
try:
step = int(lines[cursor + 1].strip())
if lines[cursor + 2].strip() != "ITEM: NUMBER OF ATOMS":
raise ValueError("missing atom-count header")
natoms = int(lines[cursor + 3].strip())
if not lines[cursor + 4].startswith("ITEM: BOX BOUNDS"):
raise ValueError("missing box header")
box_start = cursor + 5
box = [
float(x)
for line in lines[box_start : box_start + 3]
for x in line.split()
]
if len(box) < 6 or not np.all(np.isfinite(box)):
raise ValueError("invalid box")
header_index = box_start + 3
if not lines[header_index].startswith("ITEM: ATOMS"):
raise ValueError("missing atom header")
fields = lines[header_index].split()[2:]
required = {"id", "type", "x", "y", "z", "fx", "fy", "fz"}
if not required.issubset(fields):
raise ValueError("atom header lacks id/type/coordinates/forces")
offsets = {name: fields.index(name) for name in required}
rows = lines[header_index + 1 : header_index + 1 + natoms]
if len(rows) != natoms:
raise ValueError("truncated atom frame")
mapping = []
for row in rows:
values = row.split()
if len(values) < len(fields):
raise ValueError("truncated atom row")
atom_id = int(values[offsets["id"]])
atom_type = int(values[offsets["type"]])
numbers = [
float(values[offsets[name]])
for name in ("x", "y", "z", "fx", "fy", "fz")
]
if not np.all(np.isfinite(numbers)):
raise ValueError("non-finite atom data")
mapping.append((atom_id, atom_type))
mapping = tuple(sorted(mapping))
if expected_atoms is None:
expected_atoms, expected_mapping = natoms, mapping
if natoms != expected_atoms or mapping != expected_mapping:
raise ValueError("inconsistent atom count/type mapping")
steps.append(step)
cursor = header_index + 1 + natoms
except (IndexError, TypeError, ValueError, OverflowError):
invalid += 1
cursor += 1
return {"steps": sorted(set(steps)), "invalid_frames": invalid}


def _recovery_task_report(task_path, model_devi_merge_traj, salvage_prefix):
"""Validate one model-deviation task and return an auditable report."""
task_path = Path(task_path)
model_file = task_path / "model_devi.out"
if not model_file.is_file():
return {
"task": task_path.name,
"status": "missing",
"valid_steps": [],
"excluded_steps": [],
"invalid_frames": 0,
"reason": "model_devi.out is missing",
}
try:
model_devi = _read_model_devi_file(str(task_path), False, model_devi_merge_traj)
if model_devi.ndim == 1:
model_devi = model_devi.reshape(1, -1)
deviation_steps = {
int(row[0]) for row in model_devi if np.all(np.isfinite(row))
}
dump_files = (
[task_path / "all.lammpstrj"]
if model_devi_merge_traj
else sorted((task_path / "traj").glob("*.lammpstrj"))
)
dump_steps, invalid_frames = set(), 0
for dump_file in dump_files:
metadata = _read_lammps_dump_steps(dump_file)
dump_steps.update(metadata["steps"])
invalid_frames += metadata["invalid_frames"]
valid_steps = sorted(deviation_steps & dump_steps)
excluded_steps = sorted(deviation_steps - set(valid_steps))
status = "completed" if not invalid_frames and not excluded_steps else "failed"
eligible = valid_steps if (status == "completed" or salvage_prefix) else []
return {
"task": task_path.name,
"status": status,
"valid_steps": eligible,
"excluded_steps": excluded_steps,
"invalid_frames": invalid_frames,
"last_valid_step": eligible[-1] if eligible else None,
"reason": ""
if status == "completed"
else "dump/deviation mismatch or invalid frame",
}
except (OSError, ValueError, IndexError, AssertionError) as error:
return {
"task": task_path.name,
"status": "corrupt",
"valid_steps": [],
"excluded_steps": [],
"invalid_frames": 0,
"reason": str(error),
}


def _recovery_valid_steps(task_path):
"""Read valid timestep allow-list produced by post_model_devi."""
task = Path(task_path)
report_path = task.parent.parent / "exploration_report.json"
if not report_path.is_file():
return None
try:
report = json.loads(report_path.read_text())
except (OSError, json.JSONDecodeError):
return None
for item in report.get("tasks", []):
if item.get("task") == task.name:
return set(item.get("valid_steps", []))
return None


def _read_model_devi_file(
task_path: str,
model_devi_f_avg_relative: bool = False,
Expand Down Expand Up @@ -3446,6 +3666,13 @@ def _read_model_devi_file(
_, reverse_indices = np.unique(model_devi[::-1, 0], return_index=True)
last_indices = model_devi.shape[0] - 1 - reverse_indices
model_devi = model_devi[np.sort(last_indices)]
valid_steps = _recovery_valid_steps(task_path)
if valid_steps is not None:
if model_devi.ndim == 1:
model_devi = model_devi.reshape(1, -1)
model_devi = model_devi[
np.isin(model_devi[:, 0].astype(int), list(valid_steps))
]
if model_devi_f_avg_relative:
if model_devi_merge_traj is True:
all_traj = os.path.join(task_path, "all.lammpstrj")
Expand Down
76 changes: 76 additions & 0 deletions tests/generator/test_model_devi_recovery.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
import json
import os
import sys
import tempfile
import unittest
from pathlib import Path

import numpy as np

sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../..")))

from dpgen.generator.run import (
_normalize_model_devi_recovery,
_read_model_devi_file,
_recovery_task_report,
)

DUMP = """ITEM: TIMESTEP
0
ITEM: NUMBER OF ATOMS
2
ITEM: BOX BOUNDS pp pp pp
0 5
0 5
0 5
ITEM: ATOMS id type x y z fx fy fz
1 1 0 0 0 0 0 0
2 1 1 0 0 0 0 0
"""


class TestModelDeviRecovery(unittest.TestCase):
def test_policy_defaults_and_validation(self):
policy = _normalize_model_devi_recovery(
{"model_devi_recovery": {"enabled": True}}
)
self.assertTrue(policy["enabled"])
self.assertEqual(policy["max_failed_tasks"], 0)
with self.assertRaises(ValueError):
_normalize_model_devi_recovery(
{"model_devi_recovery": {"max_failed_ratio": 2}}
)

def test_task_report_salvages_matching_prefix(self):
with tempfile.TemporaryDirectory() as root:
task = Path(root) / "task.000.000000"
(task / "traj").mkdir(parents=True)
(task / "traj" / "0.lammpstrj").write_text(DUMP)
np.savetxt(
task / "model_devi.out",
np.asarray([[0, 0.1, 0, 0, 0.2, 0, 0], [10, 0.1, 0, 0, 0.2, 0, 0]]),
)
report = _recovery_task_report(task, False, True)
self.assertEqual(report["status"], "failed")
self.assertEqual(report["valid_steps"], [0])
self.assertEqual(report["excluded_steps"], [10])

def test_selection_reads_report_allow_list(self):
with tempfile.TemporaryDirectory() as root:
iteration = Path(root) / "iter.000000"
task = iteration / "01.model_devi" / "task.000.000000"
task.mkdir(parents=True)
np.savetxt(
task / "model_devi.out",
np.asarray([[0, 0.1, 0, 0, 0.2, 0, 0], [10, 0.1, 0, 0, 0.2, 0, 0]]),
)
(iteration / "exploration_report.json").write_text(
json.dumps({"tasks": [{"task": task.name, "valid_steps": [0]}]})
)
result = _read_model_devi_file(str(task))
self.assertEqual(result.shape, (1, 7))
self.assertEqual(int(result[0, 0]), 0)


if __name__ == "__main__":
unittest.main()
Loading