Source code for amorphgen.pipeline.run_pipeline

"""
amorphgen.pipeline.run_pipeline
--------------------------------
Orchestrates the full melt-and-quench pipeline.

Usage
-----
    from amorphgen import MeltQuenchPipeline

    # Default MACE model
    pipe = MeltQuenchPipeline(
        input_file="POSCAR",
        work_dir="my_run",
        cfg_override={"model": "mace-mpa-0", "device": "cuda"},
    )
    final_atoms = pipe.run()

    )

    # CHGNet
    pipe = MeltQuenchPipeline(
        input_file="POSCAR",
        cfg_override={"model": "chgnet"},
    )

    # Custom fine-tuned MACE model
    pipe = MeltQuenchPipeline(
        input_file="POSCAR",
        cfg_override={"model_path": "/data/models/InO_finetuned.model"},
    )

Resuming from a checkpoint
--------------------------
    pipe.run(stages=[5, 6, 7], input_file="stage4_eq.xyz")
"""

import os
import time
import platform
import sys
import hashlib
from copy import deepcopy
from datetime import datetime

from ase.io import read

from . import opt_cell, melt_cell, equilibrate, quench, final_opt
from ..utils import get_calculator, merge_config
from ..configs import DEFAULT_CONFIG
from .manifest import RunManifest, _json_value
from ..utils.run_lock import run_lock


def _calculator_parameters(calc, seen=None):
    """Capture known calculator and wrapper parameters without runtime state."""
    seen = set() if seen is None else seen
    if id(calc) in seen:
        return None
    seen.add(id(calc))
    parameters = {
        key: _json_value(getattr(calc, key, None))
        for key in ("parameters", "pair_params", "charges", "cutoff",
                    "coulomb_method", "alpha", "coulomb", "repulsive_core_config")
    }
    base = getattr(calc, "base_calculator", None)
    if base is not None:
        parameters["base_calculator"] = _calculator_parameters(base, seen)
    return parameters


def _resume_config_value(value):
    """Include identities of live reference calculators in nested settings."""
    from ase.calculators.calculator import Calculator
    from ..utils.run_provenance import calculator_provenance

    if isinstance(value, dict):
        return {str(k): _resume_config_value(v) for k, v in value.items()}
    if isinstance(value, (list, tuple)):
        return [_resume_config_value(v) for v in value]
    if isinstance(value, Calculator):
        model = calculator_provenance({}, value, injected=True)["model"]
        model.pop("hash_unavailable_reason", None)
        return {"model": model, "parameters": _calculator_parameters(value)}
    return _json_value(value)


[docs] class MeltQuenchPipeline: """ End-to-end melt-and-quench pipeline for amorphous structure generation. Parameters ---------- input_file : str Path to the starting crystalline structure (any ASE-readable format). work_dir : str Directory where all output files are written. Created if absent. cfg_override : dict, optional Any keys in DEFAULT_CONFIG to override, including: * ``"model"`` : foundation model short name (any backend) * ``"model_path"`` : path to local .model file (overrides model) * ``"device"`` : ``"cuda"`` or ``"cpu"`` * ``"eq_premelt"`` : ``{"ensemble": "NVT" or "NPT", ...}`` * ``"melt"`` : ``{"ensemble": "NVT" or "NPT", ...}`` * ``"quench"`` : ``{"ensemble": "NVT" or "NPT", ...}`` * ``"eq_high"`` : ``{"ensemble": "NVT" or "NPT", ...}`` * ``"eq_low"`` : ``{"ensemble": "NVT" or "NPT", ...}`` share_calc : bool If True, one calculator is shared across all stages. calc : ASE calculator, optional A pre-built ASE calculator to use for every stage, bypassing the ``get_calculator()`` backend factory. Lets you drive the full pipeline with any ASE calculator (a fine-tuned model, a custom potential, an external code). NPT and cell-filter stages still require the calculator to provide a stress tensor. """ STAGE_NAMES = { 1: "Structure optimisation (crystalline)", 2: "Pre-melt equilibration (300 K)", 3: "Melt (heat ramp)", 4: "High-T equilibration", 5: "Quench (cooling ramp)", 6: "Low-T equilibration", 7: "Final optimisation (amorphous)", } # Default checkpoint files written by each stage. STAGE_CHECKPOINTS = { 1: "stage1_opt.xyz", 2: "stage2_eq.xyz", 3: "stage3_melted.xyz", 4: "stage4_eq.xyz", 5: "stage5_quenched.xyz", 6: "stage6_eq.xyz", 7: "stage7_opt.xyz", } def __init__(self, input_file: str, work_dir: str = "melt_quench_run", cfg_override: dict | None = None, share_calc: bool = True, calc=None): self.input_file = input_file self.work_dir = work_dir self.cfg = merge_config(DEFAULT_CONFIG, cfg_override) self.share_calc = share_calc # An explicit ASE calculator supplied here is used for every stage, # bypassing the get_calculator() factory. This lets callers drive the # full pipeline with any ASE calculator (a fine-tuned model, a custom # potential, an external code) rather than only the built-in backends. # A stress-less calculator will still be rejected by the NPT / cell- # filter guards where a stress tensor is required. self._injected_calc = calc self._calc = calc self._calc_provenance = None self._calc_config = None os.makedirs(work_dir, exist_ok=True) # Handle legacy "mace_model" key → "model" if self.cfg.get("mace_model") and not self.cfg.get("model"): self.cfg["model"] = self.cfg["mace_model"] # ───────────────────────────────────────────────────────────────────────── def _get_calc(self): """Build or return the shared calculator. An injected calculator (passed to the constructor) is always reused as-is; otherwise one is built from the config via get_calculator(). """ if self._injected_calc is not None: return self._injected_calc calculator_config = self._calculator_config() if self._calc is None or not self.share_calc or self._calc_config != calculator_config: from ..utils.common import resolve_device device = resolve_device(self.cfg.get("device", "cuda")) calc_kwargs = {} if self.cfg.get("classical_params"): calc_kwargs["classical_params"] = self.cfg["classical_params"] self._calc = get_calculator( model=self.cfg.get("model", "mace-mpa-0"), device=device, model_path=self.cfg.get("model_path"), default_dtype=self.cfg.get("default_dtype", "auto"), **calc_kwargs, ) from ..utils.run_provenance import calculator_provenance # Keep the identity of the weights actually loaded, even if the # checkpoint file is replaced before this calculator is reused. self._calc_provenance = calculator_provenance(self.cfg, self._calc) self._calc_config = calculator_config return self._calc def _calculator_config(self): return _json_value({key: self.cfg.get(key) for key in ( "model", "model_path", "device", "default_dtype", "classical_params", )}) # ───────────────────────────────────────────────────────────────────────── def _find_resume_point(self, stages: list[int], previous_stages=None) -> tuple[list[int], str]: """ Scan work_dir for completed stage checkpoints and return the remaining stages and the input file to resume from. Parameters ---------- stages : list of int The originally requested stages. Returns ------- remaining_stages : list of int Stages that still need to run. resume_input : str Path to the checkpoint file to resume from, or the original input_file if no checkpoints are found. """ resume_input = self.input_file remaining_stages = list(stages) checkpoints = self._stage_checkpoints() # Walk stages in order; if a checkpoint exists, advance past it for s in stages: if previous_stages is not None and previous_stages.get(s) not in ("completed", "skipped"): break checkpoint = checkpoints.get(s) if checkpoint is None: break checkpoint_path = os.path.join(self.work_dir, checkpoint) if os.path.isfile(checkpoint_path): # A truncated output is not a completed checkpoint. try: read(checkpoint_path) except Exception: break resume_input = checkpoint_path remaining_stages.remove(s) else: break # first missing checkpoint → start here return remaining_stages, resume_input def _stage_checkpoints(self): sections = {1: "opt", 2: "eq_premelt", 3: "melt", 4: "eq_high", 5: "quench", 6: "eq_low", 7: "final_opt"} checkpoints = {} for stage, section in sections.items(): cfg = dict(self.cfg.get("opt", {})) if stage == 7 else {} cfg.update(self.cfg.get(section, {})) checkpoints[stage] = cfg.get("output_xyz", self.STAGE_CHECKPOINTS[stage]) return checkpoints def _clear_stale_trajectory(self, stage): """Reset fresh-stage frame files before recording the stage as started. Stage setup can fail before its output writer truncates a trajectory. Clearing here prevents that failure from adopting an earlier run's frames. """ defaults = {2: ("eq_premelt", "stage2_eq_traj.xyz"), 3: ("melt", "stage3_melt_traj.xyz"), 4: ("eq_high", "stage4_eq_traj.xyz"), 5: ("quench", "stage5_quench_traj.xyz"), 6: ("eq_low", "stage6_eq_traj.xyz")} if stage not in defaults: return section, default = defaults[stage] paths = {self.cfg[section].get("traj_file", default)} if stage in (3, 5): paths.add("stage3_melt.xyz" if stage == 3 else "stage5_quench.xyz") for path in paths: if os.path.isfile(path): os.unlink(path) def _resume_settings(self, stages, input_file): """Snapshot the requested protocol before any output can be reused.""" from ..utils.common import run_index_for from ..utils.run_provenance import calculator_provenance input_path = os.path.abspath(input_file) digest = None if os.path.isfile(input_path): hasher = hashlib.sha256() with open(input_path, "rb") as stream: for chunk in iter(lambda: stream.read(1024 * 1024), b""): hasher.update(chunk) digest = hasher.hexdigest() provenance = calculator_provenance( self.cfg, self._injected_calc, injected=self._injected_calc is not None, ) model = provenance["model"] identity = {key: model[key] for key in ( "name", "path", "sha256", "hash_source", "calculator_class", )} # A shared calculator may still hold older weights after its file # was replaced. Keep requested and loaded file identities separate. requested_model = dict(identity) if (self._injected_calc is None and self.share_calc and self._calc_config == self._calculator_config() and self._calc_provenance is not None): loaded = self._calc_provenance["model"] if loaded["hash_source"] == "file": for key in ("name", "path", "sha256", "hash_source"): identity[key] = loaded[key] if self._injected_calc is not None: # ASE parameters cover classical/custom calculators that do not # expose model weights. AmorphGen pair potentials keep their # parameters as attributes rather than ASE's parameters dict. identity["parameters"] = _calculator_parameters(self._injected_calc) orig_dir = os.getcwd() try: os.chdir(self.work_dir) seed_index = run_index_for(self.cfg) finally: os.chdir(orig_dir) return { "version": 1, "config": _resume_config_value(self.cfg), "stages": list(stages), "seed_index": seed_index, "input": {"path": input_path, "sha256": digest}, "calculator": identity, "requested_model": requested_model, } # ─────────────────────────────────────────────────────────────────────────
[docs] def run(self, stages: list[int] | None = None, input_file: str | None = None, resume: bool = False) -> object: """ Execute the pipeline. Parameters ---------- stages : list of int, optional Which stages to run (default: all, i.e. ``[1, 2, 3, 4, 5, 6, 7]``). input_file : str, optional Override the input structure file (useful for resuming from a mid-pipeline checkpoint). resume : bool If True, verify the saved input and settings, then scan work_dir for valid checkpoints of stages recorded as completed. Automatically determines which stage to resume from and which input file to use. Returns ------- ase.Atoms The final optimised amorphous structure. Notes ----- ``run_manifest.json`` records configuration, calculator provenance, and stage outcomes before and during execution. Each accepted invocation adds an attempt, retaining earlier attempts when resuming or rerunning. """ with run_lock(self.work_dir): return self._run_locked(stages, input_file, resume)
def _run_locked(self, stages, input_file, resume): from ..utils.run_provenance import calculator_provenance if stages is None: stages = [1, 2, 3, 4, 5, 6, 7] stages = list(stages) if input_file is None: input_file = self.input_file settings = self._resume_settings(stages, input_file) manifest = RunManifest(self.work_dir, input_file, self.cfg, stages, self.STAGE_NAMES, resume, settings) try: manifest.attempt.update(calculator_provenance( self.cfg, self._injected_calc, injected=self._injected_calc is not None, )) manifest.attempt["seed_index"] = settings["seed_index"] manifest.save() atoms = self._run(stages, input_file, resume, manifest) except BaseException as exc: status = "interrupted" if isinstance(exc, (KeyboardInterrupt, SystemExit)) else "failed" try: manifest.finish(status, exc) except Exception as write_error: # Preserve the simulation failure if disk writes also fail. print(f"Could not update {manifest.path}: {write_error}", file=sys.stderr) raise else: manifest.finish("completed") return atoms def _run(self, stages, input_file, resume, manifest): from ..utils.run_provenance import calculator_provenance if resume: remaining, resume_input = self._find_resume_point(stages, manifest.previous_stages) skipped = [s for s in stages if s not in remaining] checkpoints = self._stage_checkpoints() manifest.skip_stages(skipped, checkpoints) if not remaining and stages: print(f" All stages already completed in {self.work_dir}/") final_checkpoint = os.path.join( self.work_dir, checkpoints[stages[-1]], ) manifest.attempt["input_file"] = os.path.abspath(final_checkpoint) manifest.save() return read(final_checkpoint) if remaining != stages: print(f" Resuming: skipping completed stages {skipped}") print(f" Starting from stage {remaining[0]} " f"(input: {os.path.basename(resume_input)})") input_file = resume_input stages = remaining manifest.attempt["input_file"] = os.path.abspath(input_file) manifest.save() atoms = read(input_file) calc = self._get_calc() atoms.calc = calc if self._injected_calc is None: provenance = deepcopy(self._calc_provenance) if provenance is None: provenance = calculator_provenance(self.cfg, calc) provenance["precision"]["requested"] = self.cfg.get("default_dtype", "auto") provenance["device"]["requested"] = self.cfg.get("device", "auto") manifest.attempt.update(provenance) manifest.save() model_name = self.cfg.get("model", "mace-mpa-0") model_path = self.cfg.get("model_path") model_display = model_path if model_path else model_name device = self.cfg.get("device", "cuda") n_atoms = len(atoms) formula = atoms.get_chemical_formula(mode="hill") bar = "=" * 65 print(f"\n{bar}") from .. import __version__ print(f" AmorphGen v{__version__} - Melt-and-Quench Pipeline") from ..utils.common import compute_density_gcm3 density = compute_density_gcm3(atoms) print(f" Model: {model_display}") print(f" Device: {device}") print(f" Input: {input_file}") print(f" System: {formula} ({n_atoms} atoms)") print(f" Density: {density:.2f} g/cm3") print(f" Stages: {stages}") print(f" Output: {self.work_dir}/") print(f"{bar}\n") orig_dir = os.getcwd() os.chdir(self.work_dir) t0 = time.time() stage_timings = [] try: for s in stages: from ..utils.preemption import stop_if_requested stop_if_requested() name = self.STAGE_NAMES.get(s, f"Stage {s}") print(f"\n{'-' * 65}") print(f" Stage {s}: {name}") print(f"{'-' * 65}\n") t_stage = time.time() stage_resume = resume and s == stages[0] and s in manifest.previous_stages if not stage_resume: self._clear_stale_trajectory(s) manifest.start_stage(s) # MD stages get the resume flag for FRAME-level resume: an # interrupted stage picks up from the last frame of its # stage trajectory (stages that never started have no # trajectory and run fresh). Optimisation stages (1, 7) # restart whole — LBFGS/FIRE state is not checkpointed. if s == 1: atoms = opt_cell.run(atoms, self.cfg, calc) elif s == 2: atoms = equilibrate.run(atoms, self.cfg, calc, stage="premelt", resume=stage_resume) elif s == 3: atoms = melt_cell.run(atoms, self.cfg, calc, resume=stage_resume) elif s == 4: atoms = equilibrate.run(atoms, self.cfg, calc, stage="high", resume=stage_resume) elif s == 5: atoms = quench.run(atoms, self.cfg, calc, resume=stage_resume) elif s == 6: atoms = equilibrate.run(atoms, self.cfg, calc, stage="low", resume=stage_resume) elif s == 7: atoms = final_opt.run(atoms, self.cfg, calc) else: print(f" WARNING: Unknown stage {s} - skipping.") manifest.finish_stage("skipped") continue dt = time.time() - t_stage stage_timings.append((s, name, dt)) manifest.finish_stage() d = compute_density_gcm3(atoms) print(f" [Stage {s} completed in {dt:.1f} s ({dt/60:.1f} min) " f"| density={d:.2f} g/cm3]") finally: os.chdir(orig_dir) elapsed = time.time() - t0 # Print summary print(f"\n{bar}") print(f" Pipeline complete ({elapsed / 60:.1f} min)") print(f" Output directory: {self.work_dir}/") print(f"{bar}\n") # Write pipeline summary log self._write_summary_log( input_file, formula, n_atoms, model_display, device, stages, stage_timings, elapsed ) return atoms def _write_summary_log(self, input_file, formula, n_atoms, model_display, device, stages, stage_timings, total_elapsed): """Write a pipeline_summary.log file with timing and config.""" logfile = os.path.join(self.work_dir, "pipeline_summary.log") bar = "=" * 65 with open(logfile, "w") as f: f.write(f"{bar}\n") f.write(f" AmorphGen — Pipeline Summary\n") f.write(f" Date: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n") f.write(f" Platform: {platform.platform()}\n") f.write(f" Python: {platform.python_version()}\n") f.write(f"{bar}\n\n") f.write(f" System: {formula} ({n_atoms} atoms)\n") f.write(f" Model: {model_display}\n") f.write(f" Device: {device}\n") f.write(f" Input: {input_file}\n") f.write(f" Output: {self.work_dir}/\n") f.write(f" Stages: {stages}\n\n") # Stage timings f.write(f" {'Stage':<45} {'Time (s)':>10} {'Time (min)':>12}\n") f.write(f" {'-'*69}\n") for s, name, dt in stage_timings: f.write(f" {s}. {name:<42} {dt:>10.1f} {dt/60:>11.1f}\n") f.write(f" {'-'*69}\n") f.write(f" {'Total':<45} {total_elapsed:>10.1f} " f"{total_elapsed/60:>11.1f}\n\n") # Per-atom timing if n_atoms > 0: f.write(f" Per-atom total: {total_elapsed/n_atoms:.2f} s/atom\n") for s, name, dt in stage_timings: f.write(f" Per-atom stage {s}: {dt/n_atoms:.2f} s/atom\n") f.write(f"\n") # Key config parameters f.write(f" Configuration:\n") for key in ["opt", "eq_premelt", "melt", "eq_high", "quench", "eq_low"]: if key in self.cfg: f.write(f" {key}:\n") for k, v in self.cfg[key].items(): f.write(f" {k}: {v}\n") f.write(f"\n{bar}\n") print(f" Summary log: {logfile}")