diff --git a/dpgen/generator/arginfo.py b/dpgen/generator/arginfo.py index 943af3a33..e14b475b8 100644 --- a/dpgen/generator/arginfo.py +++ b/dpgen/generator/arginfo.py @@ -410,6 +410,14 @@ 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 task outputs. To salvage prefix frames from failed tasks, also " + "set salvage_prefix=true and choose nonzero max_failed_tasks and " + "max_failed_ratio budgets. Recovery requires a DPDispatcher build with " + "the failed-result policy API (continue_on_failure, raise_on_failure, " + "and include_failed_results)." + ) 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." @@ -498,6 +506,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], diff --git a/dpgen/generator/run.py b/dpgen/generator/run.py index af6ffe6e6..d3d717157 100644 --- a/dpgen/generator/run.py +++ b/dpgen/generator/run.py @@ -3260,7 +3260,28 @@ 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_result = 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 build with " + "continue_on_failure/raise_on_failure/include_failed_results " + "support; install the failed-result policy release from " + "dpdispatcher#687 or a later compatible release" + ) from error + task_states = _recovery_task_states(submission_result) + if task_states: + Path(work_path, "recovery_task_states.json").write_text( + json.dumps({"tasks": task_states}, indent=2) + "\n" + ) + else: + submission.run_submission() def run_model_devi(iter_index, jdata, mdata): @@ -3277,7 +3298,54 @@ 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) + task_states = _recovery_task_states_from_file(work_path) + expected_last_step = _recovery_expected_last_step(work_path) + reports = [ + _recovery_task_report( + task, + model_devi_merge_traj, + policy["salvage_prefix"], + task_state=task_states.get(task.name), + expected_last_step=expected_last_step, + ) + 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_): @@ -3361,10 +3429,320 @@ 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") + if ( + bool(raw.get("enabled", False)) + and jdata.get("model_devi_engine", "lammps") != "lammps" + ): + raise ValueError( + "model_devi_recovery currently supports model_devi_engine='lammps' only" + ) + 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 _recovery_task_states(submission_result): + """Extract task terminal states from a serialized DPDispatcher result.""" + if not isinstance(submission_result, dict): + return {} + states = {} + for serialized_job in submission_result.get("belonging_jobs", []): + if not isinstance(serialized_job, dict): + continue + for job_data in serialized_job.values(): + if not isinstance(job_data, dict): + continue + tasks = job_data.get("job_task_list", []) + task_states = job_data.get("task_states") or [] + job_state = job_data.get("job_state") + for index, task in enumerate(tasks): + if not isinstance(task, dict): + continue + task_path = task.get("task_work_path") + if not task_path: + continue + state = task_states[index] if index < len(task_states) else job_state + states[Path(task_path).name] = state + return states + + +def _recovery_task_states_from_file(work_path): + """Load persisted DPDispatcher task states for a model-deviation run.""" + state_path = Path(work_path) / "recovery_task_states.json" + if not state_path.is_file(): + return {} + try: + payload = json.loads(state_path.read_text()) + except (OSError, json.JSONDecodeError): + return {} + states = payload.get("tasks", {}) if isinstance(payload, dict) else {} + return states if isinstance(states, dict) else {} + + +def _recovery_task_failed(task_state): + """Return whether a dispatcher task state represents terminal failure.""" + if isinstance(task_state, str): + state = task_state.rsplit(".", 1)[-1].lower() + return state in {"failed", "terminated", "unknown", "4", "7", "100"} + try: + return int(task_state) in {4, 7, 100} + except (TypeError, ValueError): + return False + + +def _recovery_expected_last_step(work_path): + """Read the requested model-deviation duration from cur_job.json.""" + cur_job_path = Path(work_path) / "cur_job.json" + if not cur_job_path.is_file(): + return None + try: + cur_job = json.loads(cur_job_path.read_text()) + except (OSError, json.JSONDecodeError): + return None + nsteps = cur_job.get("nsteps") if isinstance(cur_job, dict) else None + if isinstance(nsteps, bool) or not isinstance(nsteps, (int, float)): + return None + nsteps = int(nsteps) + if nsteps < 0: + return None + trj_freq = next( + ( + cur_job.get(name) + for name in ("t_freq", "trj_freq", "traj_freq") + if cur_job.get(name) is not None + ), + None, + ) + if isinstance(trj_freq, bool) or not isinstance(trj_freq, (int, float)): + return nsteps + trj_freq = int(trj_freq) + return nsteps if trj_freq <= 0 else (nsteps // trj_freq) * trj_freq + + +def _read_lammps_dump_steps(filename, expected_atoms=None, expected_mapping=None): + """Return valid timestep/frame metadata from a LAMMPS dump file.""" + lines = Path(filename).read_text(errors="replace").splitlines() + steps, invalid = [], 0 + 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 + bounds = [ + [float(x) for x in line.split()] + for line in lines[box_start : box_start + 3] + ] + box = [value for bound in bounds for value in bound] + if ( + len(bounds) != 3 + or any(len(bound) < 2 or bound[1] <= bound[0] for bound in bounds) + or 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, + "expected_atoms": expected_atoms, + "expected_mapping": expected_mapping, + } + + +def _recovery_task_report( + task_path, + model_devi_merge_traj, + salvage_prefix, + task_state=None, + expected_last_step=None, +): + """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, + "task_state": task_state, + "expected_last_step": expected_last_step, + "reason": "model_devi.out is missing", + } + try: + model_devi = _read_model_devi_file( + str(task_path), False, model_devi_merge_traj, filter_recovery_steps=False + ) + 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 + expected_atoms, expected_mapping = None, None + for dump_file in dump_files: + metadata = _read_lammps_dump_steps( + dump_file, + expected_atoms=expected_atoms, + expected_mapping=expected_mapping, + ) + dump_steps.update(metadata["steps"]) + invalid_frames += metadata["invalid_frames"] + expected_atoms = metadata["expected_atoms"] + expected_mapping = metadata["expected_mapping"] + valid_steps = sorted(deviation_steps & dump_steps) + excluded_steps = sorted(deviation_steps - set(valid_steps)) + terminal_failed = _recovery_task_failed(task_state) + truncated = expected_last_step is not None and ( + not dump_steps or max(dump_steps) < expected_last_step + ) + if salvage_prefix and ( + invalid_frames or excluded_steps or terminal_failed or truncated + ): + valid_steps = [] + for step in sorted(deviation_steps): + if step not in dump_steps: + break + valid_steps.append(step) + excluded_steps = sorted(deviation_steps - set(valid_steps)) + status = ( + "completed" + if not invalid_frames + and not excluded_steps + and not terminal_failed + and not truncated + else "failed" + ) + eligible = valid_steps if (status == "completed" or salvage_prefix) else [] + reasons = [] + if invalid_frames or excluded_steps: + reasons.append("dump/deviation mismatch or invalid frame") + if terminal_failed: + reasons.append(f"dispatcher task ended in failure state {task_state!r}") + if truncated: + reasons.append( + f"trajectory ends before requested timestep {expected_last_step}" + ) + 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, + "task_state": task_state, + "expected_last_step": expected_last_step, + "reason": "; ".join(reasons), + } + except (OSError, ValueError, IndexError, AssertionError) as error: + return { + "task": task_path.name, + "status": "corrupt", + "valid_steps": [], + "excluded_steps": [], + "invalid_frames": 0, + "task_state": task_state, + "expected_last_step": expected_last_step, + "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, model_devi_merge_traj: bool = False, + filter_recovery_steps: bool = True, ): model_devi_files = glob.glob(os.path.join(task_path, "model_devi*.out")) if len(model_devi_files) > 1: @@ -3446,6 +3824,14 @@ 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)] + if filter_recovery_steps: + 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") @@ -3495,6 +3881,9 @@ def _select_by_model_devi_standard( counter["failed"] = 0 counter["accurate"] = 0 for tt in modd_system_task: + valid_steps = _recovery_valid_steps(tt) + if valid_steps is not None and not valid_steps: + continue with warnings.catch_warnings(): warnings.simplefilter("ignore") all_conf = _read_model_devi_file( @@ -3606,6 +3995,9 @@ def _select_by_model_devi_adaptive_trust_low( coll_v = [] coll_f = [] for tt in modd_system_task: + valid_steps = _recovery_valid_steps(tt) + if valid_steps is not None and not valid_steps: + continue with warnings.catch_warnings(): warnings.simplefilter("ignore") model_devi = _read_model_devi_file( diff --git a/tests/generator/test_model_devi_recovery.py b/tests/generator/test_model_devi_recovery.py new file mode 100644 index 000000000..bdde8734f --- /dev/null +++ b/tests/generator/test_model_devi_recovery.py @@ -0,0 +1,288 @@ +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, + _recovery_task_states, + _select_by_model_devi_adaptive_trust_low, + _select_by_model_devi_standard, + post_model_devi, +) + +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 +""" + +MISMATCHED_TYPE_DUMP = DUMP.replace("2 1 1 0 0 0 0 0", "2 2 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}} + ) + with self.assertRaisesRegex(ValueError, "lammps"): + _normalize_model_devi_recovery( + { + "model_devi_engine": "gromacs", + "model_devi_recovery": {"enabled": True}, + } + ) + + 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_task_report_salvages_only_contiguous_prefix(self): + with tempfile.TemporaryDirectory() as root: + task = Path(root) / "task.000.000000" + (task / "traj").mkdir(parents=True) + for step in (0, 20): + (task / "traj" / f"{step}.lammpstrj").write_text( + DUMP.replace("\n0\n", f"\n{step}\n", 1) + ) + np.savetxt( + task / "model_devi.out", + np.asarray([[step, 0.1, 0, 0, 0.2, 0, 0] for step in (0, 10, 20)]), + ) + report = _recovery_task_report(task, False, True) + self.assertEqual(report["status"], "failed") + self.assertEqual(report["valid_steps"], [0]) + self.assertEqual(report["excluded_steps"], [10, 20]) + + 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) + + def test_selection_skips_empty_allow_list_before_loading_output(self): + with tempfile.TemporaryDirectory() as root: + iteration = Path(root) / "iter.000000" + task = iteration / "01.model_devi" / "task.000.000000" + task.mkdir(parents=True) + (iteration / "exploration_report.json").write_text( + json.dumps({"tasks": [{"task": task.name, "valid_steps": []}]}) + ) + standard = _select_by_model_devi_standard( + [str(task)], 0.0, 1.0, 0.0, 1.0, None, "lammps" + ) + adaptive = _select_by_model_devi_adaptive_trust_low( + [str(task)], 1.0, 1, 0.0, 1.0, 1, 0.0 + ) + self.assertEqual(standard[-1]["candidate"], 0) + self.assertEqual(adaptive[3]["candidate"], 0) + + def test_post_and_selection_tolerate_missing_task_output(self): + with tempfile.TemporaryDirectory() as root: + previous = Path.cwd() + os.chdir(root) + self.addCleanup(os.chdir, previous) + try: + work_path = Path("iter.000000") / "01.model_devi" + valid_task = work_path / "task.000.000000" + missing_task = work_path / "task.001.000000" + (valid_task / "traj").mkdir(parents=True) + missing_task.mkdir(parents=True) + (valid_task / "traj" / "0.lammpstrj").write_text(DUMP) + np.savetxt( + valid_task / "model_devi.out", + np.asarray([[0, 0.1, 0, 0, 0.2, 0, 0]]), + ) + post_model_devi( + 0, + { + "model_devi_recovery": { + "enabled": True, + "max_failed_tasks": 1, + "max_failed_ratio": 0.5, + } + }, + {}, + ) + report = json.loads( + Path("iter.000000/exploration_report.json").read_text() + ) + self.assertEqual(report["failed_tasks"], 1) + selected = _select_by_model_devi_standard( + [str(valid_task), str(missing_task)], + 0.0, + 1.0, + 0.0, + 1.0, + None, + "lammps", + ) + self.assertEqual(selected[-1]["candidate"], 1) + finally: + os.chdir(previous) + + def test_task_report_preserves_terminal_failure_and_restart_state(self): + with tempfile.TemporaryDirectory() as root: + previous = Path.cwd() + try: + os.chdir(root) + iteration = Path("iter.000000") + work_path = iteration / "01.model_devi" + task = work_path / "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], + ] + ), + ) + (work_path / "cur_job.json").write_text( + json.dumps({"nsteps": 1000, "trj_freq": 10}) + ) + (work_path / "recovery_task_states.json").write_text( + json.dumps({"tasks": {task.name: 7}}) + ) + policy = { + "model_devi_recovery": { + "enabled": True, + "max_failed_tasks": 1, + "max_failed_ratio": 1.0, + "salvage_prefix": True, + } + } + post_model_devi(0, policy, {}) + report = json.loads((iteration / "exploration_report.json").read_text()) + self.assertEqual(report["tasks"][0]["status"], "failed") + self.assertEqual(report["tasks"][0]["valid_steps"], [0]) + self.assertIn("failure state", report["tasks"][0]["reason"]) + self.assertIn("1000", report["tasks"][0]["reason"]) + post_model_devi(0, policy, {}) + repeated = json.loads( + (iteration / "exploration_report.json").read_text() + ) + self.assertEqual(repeated["tasks"][0]["status"], "failed") + self.assertEqual(repeated["tasks"][0]["valid_steps"], [0]) + self.assertEqual(repeated["tasks"][0]["excluded_steps"], [10]) + finally: + os.chdir(previous) + + def test_task_report_uses_last_scheduled_dump_step(self): + with tempfile.TemporaryDirectory() as root: + task = Path(root) / "task.000.000000" + (task / "traj").mkdir(parents=True) + steps = [0, 30, 60, 90] + for step in steps: + (task / "traj" / f"{step}.lammpstrj").write_text( + DUMP.replace("\n0\n", f"\n{step}\n", 1) + ) + np.savetxt( + task / "model_devi.out", + np.asarray([[step, 0.1, 0, 0, 0.2, 0, 0] for step in steps]), + ) + report = _recovery_task_report( + task, False, False, task_state=5, expected_last_step=90 + ) + self.assertEqual(report["status"], "completed") + self.assertEqual(report["valid_steps"], steps) + + def test_task_report_validates_atom_mapping_across_split_files(self): + with tempfile.TemporaryDirectory() as root: + task = Path(root) / "task.000.000000" + (task / "traj").mkdir(parents=True) + (task / "traj" / "0.lammpstrj").write_text(DUMP) + (task / "traj" / "10.lammpstrj").write_text( + MISMATCHED_TYPE_DUMP.replace("\n0\n", "\n10\n", 1) + ) + 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]) + self.assertEqual(report["invalid_frames"], 1) + + def test_task_report_rejects_zero_volume_box(self): + with tempfile.TemporaryDirectory() as root: + task = Path(root) / "task.000.000000" + (task / "traj").mkdir(parents=True) + (task / "traj" / "0.lammpstrj").write_text( + DUMP.replace("0 5\n0 5\n0 5", "0 0\n0 5\n0 5") + ) + np.savetxt( + task / "model_devi.out", + np.asarray([[0, 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"], []) + self.assertEqual(report["invalid_frames"], 1) + + def test_extracts_dispatcher_task_states(self): + result = _recovery_task_states( + { + "belonging_jobs": [ + { + "job-hash": { + "job_state": 7, + "task_states": [5, 7], + "job_task_list": [ + {"task_work_path": "task.000.000000"}, + {"task_work_path": "task.001.000000"}, + ], + } + } + ] + } + ) + self.assertEqual(result, {"task.000.000000": 5, "task.001.000000": 7}) + + +if __name__ == "__main__": + unittest.main()