Source code for amorphgen.cli

"""
amorphgen.cli
--------------
Command-line interface for AmorphGen.

Common usage lives in :data:`_EXAMPLES` (shown by ``amorphgen --examples``
and at the bottom of ``amorphgen -h``) — kept in one place so the two stay
in sync. Below are only the advanced patterns not covered there.

Examples (advanced)
-------------------
Use a custom fine-tuned model:
    amorphgen POSCAR --model-path /data/InO_finetuned.model

Random generation with custom minsep:
    amorphgen --random-gen --composition "In2O3*8" --target-density 5.5 \
        --minsep In-In=2.8,In-O=1.9,O-O=2.5

Optimise with cubic cell constraint:
    amorphgen POSCAR --stages 1 --cell-filter cubic --model mace-mpa-0-medium

Batch quench from snapshots:
    amorphgen --batch-quench --snapshot-dir snapshots/ --n-runs 20

Convert structure files between formats (xyz/extxyz/vasp/cif):
    amorphgen --convert snapshots/ --format vasp -o snapshots_vasp/
    amorphgen --convert traj_frame.xyz --format cif
    amorphgen --config convert.yaml          # YAML-driven (uses convert: block)
"""

from __future__ import annotations

import argparse
import sys
import os


# Concise, task-oriented usage shown at the bottom of ``-h`` and by
# ``--examples``. Kept short on purpose; the full flag list is the rest of -h.
_EXAMPLES = """\
common usage:
  # Random generation + relax (one step)
  amorphgen --random-gen --composition "Nb2O5*40" -n 10 --relax -o out -m chgnet

  # Generate only, then optimise the structures separately
  amorphgen --random-gen --composition "SiO2*32" -n 20 -o gen
  amorphgen --batch-opt --input-dir gen/random_initial -o opt -C cubic

  # Full 7-stage melt-quench pipeline from a crystal
  amorphgen POSCAR -o mq_run

  # Analyse a folder of structures (RDF, CN, density, S(q))
  amorphgen --analyse --input-dir opt --save-plot plots --save-pdf

  # List available calculator models
  amorphgen --list-models

composition:  "Nb2O5*40" (40 formula units)  or  "Nb=80,O=200" (atom counts)
notes:        modes are mutually exclusive (pick one); on Apple Silicon prefer
              '-m chgnet' or '-d cpu' (MACE + mps hits a float64 limitation).
"""


class _HelpFormatter(argparse.ArgumentDefaultsHelpFormatter,
                     argparse.RawDescriptionHelpFormatter):
    """Show argument defaults *and* keep the epilog's literal formatting."""


class _ExamplesAction(argparse.Action):
    """``--examples``: print the common-usage block and exit (skips the full -h)."""

    def __init__(self, option_strings, dest, **kwargs):
        kwargs["nargs"] = 0
        super().__init__(option_strings, dest, **kwargs)

    def __call__(self, parser, namespace, values, option_string=None):
        print(_EXAMPLES)
        parser.exit()


def _typed(*flags, argv=None):
    """True if any of *flags* was typed on the command line (regardless of
    whether its value equals the parser default)."""
    argv = sys.argv[1:] if argv is None else argv
    typed = {tok.split("=", 1)[0] for tok in argv if tok.startswith("-")}
    return bool(typed.intersection(flags))


def _DC(section, key):
    """Pipeline-stage CLI defaults come from DEFAULT_CONFIG (single source of truth)."""
    from .configs.default_config import DEFAULT_CONFIG
    return DEFAULT_CONFIG[section][key]


def _parse_tolerance(value):
    """Parse one absolute descriptor tolerance for argparse."""
    import math

    name, separator, raw_value = value.partition("=")
    name = name.strip()
    if not separator or not name:
        raise argparse.ArgumentTypeError(
            "tolerance must be NAME=VALUE, e.g. density=0.02")
    try:
        tolerance = float(raw_value)
    except ValueError as exc:
        raise argparse.ArgumentTypeError(
            f"tolerance for '{name}' must be a finite positive number") from exc
    if not math.isfinite(tolerance) or tolerance <= 0:
        raise argparse.ArgumentTypeError(
            f"tolerance for '{name}' must be a finite positive number")
    return name, tolerance


def _convergence_options(args, config):
    """Resolve convergence settings; each CLI tolerance overrides its YAML key."""
    import math

    tolerances = dict(config.get("tolerances", {}))
    tolerances.update(args.tolerance or [])
    confidence = args.convergence_confidence
    if confidence is None:
        confidence = config.get("convergence_confidence", 0.95)
    if not math.isfinite(confidence) or not 0 < confidence < 1:
        raise ValueError("convergence confidence must be finite and between 0 and 1")
    max_structures = args.convergence_max_structures
    if max_structures is None:
        max_structures = config.get("convergence_max_structures", 1000000)
    if max_structures < 2:
        raise ValueError("convergence max structures must be at least 2")
    enabled = bool(args.convergence or config.get("convergence", False) or tolerances)
    return enabled, tolerances, confidence, max_structures


def _parse_descriptor_bounds(value):
    """Parse fixed population support bounds for sequential inference."""
    import math

    name, separator, raw = value.partition("=")
    name = name.strip()
    try:
        bounds = [float(part) for part in raw.split(",")]
    except ValueError:
        bounds = []
    if (not separator or not name or len(bounds) != 2
            or not all(math.isfinite(part) for part in bounds)
            or not bounds[0] < bounds[1]
            or not math.isfinite(bounds[1] - bounds[0])):
        raise argparse.ArgumentTypeError(
            "descriptor bounds must be NAME=LOW,HIGH with finite LOW < HIGH")
    return name, bounds


def _until_convergence_options(args, config):
    """Resolve the immutable sequential sampling contract before model setup."""
    import math

    random_cfg = config.get("random_gen", {})
    analysis = config.get("analysis", {})
    enabled = bool(args.until_converged or random_cfg.get("until_converged", False))
    if not enabled:
        if (args.descriptor_bounds or args.convergence_batch_size is not None
                or args.convergence_min_structures is not None):
            raise ValueError("descriptor bounds and sampling controls require --until-converged")
        return None
    if not args.random_gen:
        raise ValueError("--until-converged requires --random-gen")
    if not (args.relax or random_cfg.get("relax", False)):
        raise ValueError("--until-converged requires --relax")
    if config.get("engine", "ase") != "torchsim":
        raise ValueError("--until-converged requires --engine torchsim")
    if args.indices is not None:
        raise ValueError("--until-converged cannot use --indices; sampling must retain its complete prefix")
    if _typed("-n", "--n-structures"):
        raise ValueError("--until-converged uses --convergence-max-structures instead of --n-structures")
    output_format = (args.format if _typed("--format")
                     else random_cfg.get("output_format", args.format))
    if output_format != "xyz" or config.get("opt", {}).get("output_format", "xyz") != "xyz":
        raise ValueError("--until-converged requires xyz output to preserve descriptor and seed metadata")
    seed = (args.seed if args.seed is not None
            else random_cfg.get("seed", config.get("seed")))
    if isinstance(seed, bool) or not isinstance(seed, int) or seed < 0:
        raise ValueError("--until-converged requires an explicit nonnegative integer --seed")
    tolerances = dict(analysis.get("tolerances", {}))
    tolerances.update(args.tolerance or [])
    bounds = dict(analysis.get("descriptor_bounds", {}))
    bounds.update(args.descriptor_bounds or [])
    if not tolerances:
        raise ValueError("--until-converged requires at least one declared --tolerance")
    if set(bounds) != set(tolerances):
        raise ValueError("every convergence tolerance requires matching descriptor bounds, with no extra bounds")
    from ase.data import atomic_numbers
    for name, tolerance in tolerances.items():
        if not isinstance(name, str):
            raise ValueError("convergence descriptor names must be strings")
        if name not in {"density", "energy.total", "energy.per_atom"}:
            kind, separator, label = name.partition(".")
            count = {"coordination": 2, "total_coordination": 1,
                     "bond_distance": 2, "bond_angle": 3}.get(kind)
            elements = label.split("-")
            if (not separator or count is None or len(elements) != count
                    or any(element not in atomic_numbers for element in elements)):
                raise ValueError(f"Unsupported sequential descriptor {name!r}")
        if (isinstance(tolerance, bool) or not isinstance(tolerance, (int, float))
                or not math.isfinite(tolerance) or tolerance <= 0):
            raise ValueError(f"Tolerance for {name!r} must be finite and positive")
        support = bounds[name]
        if (not isinstance(support, (list, tuple)) or len(support) != 2
                or any(isinstance(x, bool) or not isinstance(x, (int, float))
                       or not math.isfinite(x) for x in support)
                or not support[0] < support[1]
                or not math.isfinite(support[1] - support[0])):
            raise ValueError(f"Descriptor bounds for {name!r} require finite lower < upper")
    targets = {name: {"bounds": bounds[name], "tolerance": tolerance, "components": 1}
               for name, tolerance in tolerances.items()}
    confidence = (args.convergence_confidence if args.convergence_confidence is not None
                  else analysis.get("convergence_confidence", .95))
    if (isinstance(confidence, bool) or not isinstance(confidence, (int, float))
            or not math.isfinite(confidence) or not 0 < confidence < 1):
        raise ValueError("convergence confidence must be finite and between 0 and 1")
    batch_size = (args.convergence_batch_size if args.convergence_batch_size is not None
                  else random_cfg.get("convergence_batch_size", 8))
    minimum = (args.convergence_min_structures if args.convergence_min_structures is not None
               else random_cfg.get("convergence_min_structures", 2))
    maximum = (args.convergence_max_structures if args.convergence_max_structures is not None
               else analysis.get("convergence_max_structures", 1000))
    for name, value, lower in (("batch size", batch_size, 1),
                               ("min structures", minimum, 2),
                               ("max structures", maximum, 2)):
        if isinstance(value, bool) or not isinstance(value, int) or value < lower:
            raise ValueError(f"convergence {name} must be an integer >= {lower}")
    if maximum < minimum:
        raise ValueError("convergence max structures must be >= min structures")
    cutoff = args.cutoff if _typed("--cutoff") else analysis.get("cutoff", args.cutoff)
    needs_cutoff = any(name not in {"density", "energy.total", "energy.per_atom"}
                       for name in targets)
    try:
        numeric_cutoff = float(cutoff)
    except (ValueError, TypeError):
        numeric_cutoff = None
    if (isinstance(cutoff, bool) or numeric_cutoff is None
            or not math.isfinite(numeric_cutoff) or numeric_cutoff <= 0):
        if needs_cutoff:
            raise ValueError("--until-converged requires an explicit positive numeric --cutoff for structural targets")
        numeric_cutoff = None
    return {"targets": targets, "batch_size": batch_size, "min_structures": minimum,
            "max_structures": maximum, "confidence": confidence, "cutoff": numeric_cutoff}


def _cutoff_window_option(args, config):
    """Resolve the positive cutoff half-window: CLI > YAML > 0.1 A."""
    import math

    window = args.cutoff_window
    if window is None:
        window = config.get("cutoff_window", 0.1)
    if (isinstance(window, bool) or not isinstance(window, (int, float))
            or not math.isfinite(window) or window <= 0):
        raise ValueError("cutoff window must be a finite positive number in A")
    return float(window)


def _collect_convergence_summaries(target, prefix, summaries):
    """Flatten selected descriptor uncertainty trees to their public names."""
    if not isinstance(summaries, dict):
        return
    if "per_structure" in summaries and "sem" in summaries:
        target[prefix] = summaries
        return
    for key, value in summaries.items():
        label = "-".join(map(str, key)) if isinstance(key, tuple) else str(key)
        _collect_convergence_summaries(target, f"{prefix}.{label}", value)


def _get_parser():
    """Build and return the argument parser (without parsing)."""
    p = argparse.ArgumentParser(
        description="AmorphGen: amorphous structure generation via melt-quench MD and random placement",
        formatter_class=_HelpFormatter,
        epilog=_EXAMPLES,
    )
    from . import __version__
    p.add_argument("--version", action="version",
                   version=f"amorphgen {__version__}")
    return _add_arguments(p)


def _add_arguments(p):
    """Add all arguments to the parser and return it."""

    # ── Positional + global ───────────────────────────────────────────────────
    p.add_argument("input_file", nargs="?",
                   help="Input structure (POSCAR/xyz/cif). Pipeline mode only.")
    p.add_argument("--config", default=None, metavar="FILE",
                   help="YAML config file (CLI overrides YAML).")
    p.add_argument("-o", "--work-dir", default=None,
                   help="Output directory (auto-named per mode).")
    p.add_argument("--format", default="xyz", choices=["xyz", "vasp", "cif"],
                   help="Output format for generated/optimised structures.")
    p.add_argument("--resume", action="store_true",
                   help="Skip completed work in --random-gen / --batch-quench "
                        "/ pipeline modes.")
    p.add_argument("--examples", action=_ExamplesAction,
                   help="Show common usage examples and exit.")

    # ── Mode selectors ────────────────────────────────────────────────────────
    g_mode = p.add_argument_group(
        "modes",
        "Pick one mode (default: full melt-quench pipeline if input_file given).")
    g_mode.add_argument("--random-gen", action="store_true",
                        help="Generate random structures.")
    g_mode.add_argument("--batch-quench", action="store_true",
                        help="Quench multiple snapshots independently.")
    g_mode.add_argument("--batch-opt", action="store_true",
                        help="Optimise all structures in --input-dir.")
    g_mode.add_argument("--analyse", action="store_true",
                        help="Analyse structures (RDF, CN, angles, density).")
    g_mode.add_argument("--rank-from-log", default=None, metavar="LOG",
                        help="Parse a random-gen log and rank by total energy.")
    g_mode.add_argument("--extract-snapshots", default=None, metavar="TRAJ",
                        help="Extract N snapshots from a trajectory file "
                             "(use with --n-runs N and --select).")
    g_mode.add_argument("--mq-ensemble", action="store_true",
                        help="Full MQ-ensemble workflow: stages 1-4 from a "
                             "crystalline input, extract N snapshots from "
                             "stage 4, run stages 5-6-7 on each independently. "
                             "Use with input_file, --n-structures N.")
    g_mode.add_argument("--hybrid-ensemble", action="store_true",
                        help="Hybrid ensemble: starts from a directory of "
                             "disordered structures (e.g. --random-gen "
                             "outputs), runs stages 4-5-6-7 on each. "
                             "Use with --input-dir DIR.")
    g_mode.add_argument("--list-models", action="store_true",
                        help="Print available foundation models and exit.")
    g_mode.add_argument("--convert", default=None, metavar="PATH",
                        help="Convert one structure file or every "
                             "ASE-readable file in a directory to the format "
                             "given by --format. Output goes to --work-dir "
                             "(defaults to <PATH>_<format>/). VASP outputs "
                             "are sorted by species.")

    # ── Calculator ────────────────────────────────────────────────────────────
    g_calc = p.add_argument_group("calculator")
    model_group = g_calc.add_mutually_exclusive_group()
    model_group.add_argument("-m", "--model", default="mace-mpa-0", metavar="NAME",
                             help="Foundation model (mace-mpa-0, chgnet, sevennet, ...).")
    model_group.add_argument("--model-path", default=None, metavar="PATH",
                             help="Path to a local .model file.")
    g_calc.add_argument("-d", "--device", default="auto",
                        choices=["auto", "cuda", "cpu", "mps"],
                        help="Device.")
    g_calc.add_argument("--dtype", "--default-dtype",
                        dest="default_dtype",
                        default="auto",
                        choices=["auto", "float32", "float64"],
                        help="MLIP precision. 'auto' (default) picks the "
                             "right value per backend: float32 for CHGNet "
                             "(only option it supports) and classical "
                             "potentials; float64 for MACE and SevenNet. "
                             "Explicit 'float32' is ~2x faster MD but may "
                             "fail for backends that require float64. "
                             "(``--default-dtype`` accepted as legacy alias.)")

    # ── Optimisation (stages 1, 7; also random-gen --relax) ───────────────────
    g_opt = p.add_argument_group("optimisation")
    g_opt.add_argument("-f", "--fmax", type=float, default=0.01,
                       help="Force convergence (eV/A).")
    g_opt.add_argument("--opt-steps", type=int, default=1000,
                       help="Max optimisation steps.")
    g_opt.add_argument("-O", "--optimizer", default="LBFGS",
                       choices=["LBFGS", "FIRE", "BFGSLineSearch", "BFGS", "MDMin"],
                       help="Optimizer.")
    g_opt.add_argument("-C", "--cell-filter", default="FrechetCellFilter",
                       choices=["FrechetCellFilter", "UnitCellFilter",
                                "ExpCellFilter", "StrainFilter", "cubic", "none"],
                       help="Cell filter ('cubic' = isotropic V; 'none' = fixed "
                            "cell). Amorphous-input modes (--random-gen, "
                            "--hybrid-ensemble, --batch-opt) default to 'cubic'; "
                            "the melt-quench pipeline defaults to FrechetCellFilter.")

    # ── Melt-quench pipeline ──────────────────────────────────────────────────
    g_pipe = p.add_argument_group(
        "pipeline (melt-quench MD)",
        "Stages: 1=opt, 2=eq-premelt, 3=melt, 4=eq-high, 5=quench, "
        "6=eq-low, 7=final-opt.")
    g_pipe.add_argument("--stages", nargs="+", type=int,
                        default=[1, 2, 3, 4, 5, 6, 7], metavar="N",
                        help="Stages to run.")
    g_pipe.add_argument("--engine", choices=["ase", "torchsim"], default=None,
                        help="Relaxation engine for --batch-opt and --random-gen "
                             "--relax: 'ase' (default, one structure at a time) or "
                             "'torchsim' (all structures in one batched GPU/CPU "
                             "call; pip install \"amorphgen[torchsim]\"; MACE, "
                             "SevenNet or LJ models; no MPS).")
    g_pipe.add_argument("--batch-size", default=None, metavar="N|auto",
                        help="torch-sim engine: structures per batched chunk. "
                             "Default auto: a one-structure GPU memory probe picks "
                             "the largest chunk that fits (16 on CPU). Outputs are "
                             "written after each chunk and MD trajectories every 100 steps, "
                             "so --resume loses at most 100 steps.")
    g_pipe.add_argument("--run-index", type=int, default=None, metavar="INT",
                        help="Run index for the per-run seed stream of the MD "
                             "stages (with --seed). Normally inferred from the "
                             "run_NNNN/ directory or SLURM_ARRAY_TASK_ID; set it "
                             "when several single-structure jobs share one --seed.")
    g_pipe.add_argument("--seed", type=int, default=None, metavar="INT",
                        help="Global random seed: seeds random placement AND "
                             "the velocity initialisation / Langevin noise of "
                             "every MD stage (per stage and per run). Same seed, "
                             "same CPU and versions -> identical output; on a GPU "
                             "MLIP forces are not bit-reproducible, so trajectories "
                             "diverge after a few thousand steps.")
    g_pipe.add_argument("--timestep", type=float, default=0.5,
                        help="MD timestep in fs (applies to all MD stages).")
    # Stage 2
    g_pipe.add_argument("--eq-premelt-ensemble", default="NVT",
                        choices=["NVT", "NPT"], help="Stage 2 ensemble.")
    g_pipe.add_argument("--eq-premelt-T", type=int, default=300,
                        help="Stage 2 T (K).")
    g_pipe.add_argument("--eq-premelt-steps", type=int, default=_DC("eq_premelt", "steps"),
                        help="Stage 2 MD steps.")
    # Stage 3
    g_pipe.add_argument("--melt-ensemble", default=_DC("melt", "ensemble"),
                        choices=["NVT", "NPT"], help="Stage 3 ensemble.")
    g_pipe.add_argument("--melt-T-start", type=int, default=_DC("melt", "T_start"),
                        help="Stage 3 T_start (K).")
    g_pipe.add_argument("--melt-T-end", type=int, default=_DC("melt", "T_end"),
                        help="Stage 3 T_end (K).")
    g_pipe.add_argument("--melt-T-step", type=int, default=_DC("melt", "T_step"),
                        help="Stage 3 ramp segment (K).")
    g_pipe.add_argument("--melt-steps-per-T", type=int, default=_DC("melt", "steps_per_T"),
                        help="Stage 3 steps per segment.")
    # Stage 4
    g_pipe.add_argument("--eq-high-ensemble", default=_DC("eq_high", "ensemble"),
                        choices=["NVT", "NPT"], help="Stage 4 ensemble.")
    g_pipe.add_argument("--eq-high-T", type=int, default=_DC("eq_high", "T"),
                        help="Stage 4 T (K).")
    g_pipe.add_argument("--eq-high-steps", type=int, default=_DC("eq_high", "steps"),
                        help="Stage 4 MD steps.")
    # Stage 5
    g_pipe.add_argument("--quench-ensemble", default=_DC("quench", "ensemble"),
                        choices=["NVT", "NPT"], help="Stage 5 ensemble.")
    g_pipe.add_argument("--quench-T-start", type=int, default=_DC("quench", "T_start"),
                        help="Stage 5 T_start (K).")
    g_pipe.add_argument("--quench-T-end", type=int, default=_DC("quench", "T_end"),
                        help="Stage 5 T_end (K).")
    g_pipe.add_argument("--quench-T-step", type=int, default=_DC("quench", "T_step"),
                        help="Stage 5 ramp segment (K, negative = cooling).")
    g_pipe.add_argument("--quench-steps-per-T", type=int, default=_DC("quench", "steps_per_T"),
                        help="Stage 5 steps per segment.")
    # Stage 6
    g_pipe.add_argument("--eq-low-ensemble", default=_DC("eq_low", "ensemble"),
                        choices=["NVT", "NPT"], help="Stage 6 ensemble.")
    g_pipe.add_argument("--eq-low-T", type=int, default=_DC("eq_low", "T"),
                        help="Stage 6 T (K).")
    g_pipe.add_argument("--eq-low-steps", type=int, default=_DC("eq_low", "steps"),
                        help="Stage 6 MD steps.")

    # ── Random generation ─────────────────────────────────────────────────────
    g_rand = p.add_argument_group("random-gen", "Used with --random-gen.")
    g_rand.add_argument("--composition", default=None, metavar="SPEC",
                        help='"In2O3*16" (formula*N units) or "In=32,O=48".')
    g_rand.add_argument("-n", "--n-structures", type=int, default=1,
                        help="Number of structures.")
    g_rand.add_argument("--relax", action="store_true",
                        help="Relax each generated structure.")
    g_rand.add_argument("--until-converged", action="store_true",
                        help="Generate and relax independent torch-sim batches until all "
                             "declared descriptor precision targets pass an anytime-valid rule.")
    g_rand.add_argument("--convergence-batch-size", type=int, default=None, metavar="N",
                        help="Independent structures per sequential convergence look (default 8).")
    g_rand.add_argument("--convergence-min-structures", type=int, default=None, metavar="N",
                        help="Minimum structures before sequential stopping (default 2).")
    g_rand.add_argument("--target-density", type=float, default=None,
                        help="Target density (g/cm3); auto if omitted.")
    g_rand.add_argument("--density-scale", type=float, default=1.0,
                        metavar="FACTOR",
                        help="Multiplier on the auto-estimated density "
                             "(default 1.0). Use ~1.2 for tight-network "
                             "compositions (a-Si, a-Ge, metallic glasses) "
                             "where sphere-packing underestimates by 15-25%%. "
                             "Ignored when --target-density is set.")
    g_rand.add_argument("--minsep", default=None, metavar="SPEC",
                        help="Per-pair min separations, e.g. In-In=2.8,In-O=1.9.")
    g_rand.add_argument("--target-cn", default=None, metavar="SPEC",
                        help="Target CNs for coordination-aware placement, "
                             "e.g. Si=4,O=2 (auto-detected if omitted).")
    g_rand.add_argument("--no-sc", action="store_true",
                        help="Disable SC (Seed-Coordinate) coordination-aware "
                             "placement; use plain random rejection sampling.")
    g_rand.add_argument("--dmax", default=None, metavar="SPEC",
                        help="Per-pair bonding cutoffs (auto: minsep * dmax-factor).")
    g_rand.add_argument("--dmax-factor", type=float, default=1.5,
                        help="Auto dmax multiplier.")
    g_rand.add_argument("--cn-tolerance", type=int, default=None,
                        help="Over-coordination tolerance for "
                             "coordination-aware placement (0 or 1).")
    g_rand.add_argument("--repair-iters", type=int, default=0,
                        metavar="N",
                        help="EXPERIMENTAL: post-placement repair pass for "
                             "under-coordinated atoms; N is the max number of "
                             "single-atom relocation proposals (default 0 = "
                             "off).  Useful for tetrahedral covalent networks "
                             "(a-Si, a-Ge) where greedy placement leaves many "
                             "atoms below target CN.")
    g_rand.add_argument("--max-attempts", type=int, default=500000,
                        help="Max placement attempts per atom.")
    g_rand.add_argument("--retry-mode", default="expand",
                        choices=["expand", "reduce-minsep", "none"],
                        help="Placement-stall policy. 'expand' (default): "
                             "grow the cell 5%% per retry, minseps stay "
                             "physical — right when the density is an "
                             "estimate. 'reduce-minsep': hold the cell (and "
                             "density) FIXED and soften only non-bonded "
                             "minseps (same-element, anion-anion) 5%% per "
                             "retry — for fixed-density film / isochoric "
                             "studies where expansion would corrupt the "
                             "comparison. 'none': auto-retry OFF — nothing "
                             "is adjusted; unplaceable structures fail/skip "
                             "so both density AND minseps stay exact. Bond "
                             "minseps are never reduced in any mode.")

    # ── Batch quench ──────────────────────────────────────────────────────────
    g_bq = p.add_argument_group("batch-quench", "Used with --batch-quench.")
    g_bq.add_argument("--snapshot-dir", default="snapshots", metavar="PATH",
                      help="Directory of structures, or a single trajectory "
                           "file (.xyz/.extxyz) — file gets auto-extracted "
                           "into N uniform snapshots.")
    g_bq.add_argument("--n-runs", type=int, default=20,
                      help="Number of quench runs.")
    g_bq.add_argument("--select", default=None,
                      choices=["uniform", "last", "decorrelated"],
                      help="Snapshot selection: decorrelated for --mq-ensemble, "
                           "uniform otherwise. Decorrelated uses autocorrelation "
                           "and species diffusion to choose minimum spacing.")
    g_bq.add_argument("--burn-in-frames", type=int, default=None, metavar="N",
                      help="Discard the first N frames of the trajectory "
                           "before sampling snapshots (automatic for "
                           "decorrelated selection, zero otherwise).")
    g_bq.add_argument("--decorrelation-distance", type=float, default=None,
                      metavar="ANGSTROM",
                      help="Length scale for the species displacement correlation "
                           "proxy; default: median nearest-neighbour distance.")
    g_bq.add_argument("--batch-stages", nargs="+", type=int,
                      default=[5, 6, 7], metavar="N",
                      help="Stages to run per snapshot.")

    # ── Batch optimisation ────────────────────────────────────────────────────
    g_bo = p.add_argument_group("batch-opt", "Used with --batch-opt.")
    g_bo.add_argument("--indices", default=None, metavar="SPEC",
                      help="Only these structure indices, e.g. 80-90 or 0,5,7-9 "
                           "(inclusive). --random-gen: generate/relax only those "
                           "indices (same seeds as a full run); --batch-opt: only "
                           "files whose name ends in one of them.")
    g_bo.add_argument("--pattern", default="*.xyz", metavar="GLOB",
                      help="--batch-opt: input file glob (default *.xyz, then "
                           "*.extxyz/*.vasp/*.cif).")
    g_bo.add_argument("--input-dir", default=None, metavar="DIR",
                     help="Directory of structures to optimise "
                          "(also used by --hybrid-ensemble and --analyse). "
                          "For --random-gen output pass its random_initial/ "
                          "or random_opt/ subdirectory.")

    # ── Analysis ──────────────────────────────────────────────────────────────
    g_an = p.add_argument_group("analyse", "Used with --analyse.")
    g_an.add_argument("--screen", action="store_true",
                      help="Label candidates before analysis using YAML "
                           "analysis.screening, or default coordination, "
                           "crystal-like, close-contact and relaxation screens. "
                           "Only screens with exclude: true remove structures.")
    g_an.add_argument("--screening-output", default=None, metavar="PREFIX",
                      help="Write screening JSON and candidate/count CSV tables "
                           "(default: <work-dir>/screening). Requires screening.")
    g_an.add_argument("--convergence", action="store_true",
                      help="Report order-independent uncertainty versus ensemble "
                           "size and estimates against declared tolerances.")
    g_an.add_argument("--tolerance", action="append", type=_parse_tolerance,
                      default=None, metavar="NAME=VALUE",
                      help="Absolute uncertainty tolerance for a descriptor, e.g. "
                           "density=0.02 or coordination.Si-O=0.1. Repeat for "
                           "multiple descriptors; enables --convergence. "
                           "CLI declarations override YAML per descriptor.")
    g_an.add_argument("--convergence-confidence", type=float, default=None,
                      metavar="LEVEL",
                      help="Confidence level for convergence uncertainty "
                           "(default 0.95).")
    g_an.add_argument("--descriptor-bounds", action="append", type=_parse_descriptor_bounds,
                      default=None, metavar="NAME=LOW,HIGH",
                      help="Fixed population support for an --until-converged descriptor. "
                           "Repeat once per tolerance; CLI overrides YAML per descriptor.")
    g_an.add_argument("--convergence-max-structures", type=int, default=None,
                      metavar="N",
                      help="Maximum structures: sequential generation cap with --until-converged "
                           "(default 1000), or analysis projection cap (default 1000000).")
    g_an.add_argument("--cutoff", default="auto-rdf",
                      help="Bond cutoff. 'auto-rdf' (default): first minimum "
                           "of each partial g(r), so every pair gets its own "
                           "value. 'auto': minsep from the radii table. A "
                           "number: one value for all pairs. Per-pair "
                           "overrides keep auto-rdf for the rest: "
                           "'In-O=2.6,Zn-O=2.3'; prefix a base to change it: "
                           "'auto,In-O=2.6' or '2.4,In-O=2.6'.")
    g_an.add_argument("--cutoff-window", type=float, default=None, metavar="FLOAT",
                      help="Positive half-window in A for cutoff robustness: "
                           "report nearby contact shares and coordination at "
                           "five cutoffs across +/- this value (default 0.1).")
    g_an.add_argument("--sq", action="store_true",
                      help="Compute the total structure factor S(q) via the "
                           "direct (Debye) method: correct FSDP intensities "
                           "and S(q->inf)=1. Saved as PNG+CSV under "
                           "--save-plot. Note: the FSDP region needs a large "
                           "box (q_min = 2*pi/L; ~450+ atoms recommended).")
    from .scattering_cli import add_scattering_arguments
    add_scattering_arguments(g_an)
    g_an.add_argument("--sq-weighting", default="xray",
                      choices=["xray", "neutron", "unweighted"],
                      help="Scattering-factor weighting for --sq. Use 'xray' "
                           "to compare with X-ray diffraction (heavy elements "
                           "dominate), 'neutron' for neutron data.")
    g_an.add_argument("--sq-method", default="direct",
                      choices=["direct", "ft"],
                      help="How to compute S(q) for --sq. 'direct' (default): "
                           "Debye sum at reciprocal-lattice q-vectors -- no "
                           "real-space truncation, resolves the FSDP. 'ft': "
                           "Fourier transform of g(r) truncated at L/2 -- "
                           "smoother but damps/shifts the FSDP; use for "
                           "comparison with FT-based codes. Both are "
                           "Faber-Ziman weighted.")
    g_an.add_argument("--sq-partials", action="store_true",
                      help="With --sq (direct method): also report the "
                           "Faber-Ziman partial structure factors S_ab(q) "
                           "for every element pair. Printed as first-peak "
                           "positions, added as s_<A-B> columns to "
                           "analysis_sq.csv and drawn in "
                           "analysis_sq_partials.png. Independent of "
                           "--sq-weighting (the weighting only combines "
                           "the partials into the total).")
    g_an.add_argument("--tr", action="store_true",
                      help="Total correlation function T(r) = 4*pi*r*rho*g(r), "
                           "the curve diffraction papers plot beside S(q): the "
                           "weighted S(q) Fourier-transformed over the measured "
                           "q range. Directly comparable with published data, "
                           "unlike the unweighted g(r) of --total-rdf. Saved as "
                           "analysis_tr.png + CSV (r, weighted g(r), T(r), G(r)).")
    g_an.add_argument("--tr-qrange", nargs=2, type=float, default=(0.3, 20.0),
                      metavar=("QMIN", "QMAX"),
                      help="Integration limits for --tr in 1/A (default 0.3 20). "
                           "Set them to the experiment's own range: qmax fixes the "
                           "real-space resolution and the truncation ripple.")
    g_an.add_argument("--tr-window", default="lorch", choices=["lorch", "none"],
                      help="Window for the --tr transform: 'lorch' damps the "
                           "qmax truncation ripple (default), 'none' leaves it. "
                           "Match whichever the paper used.")
    g_an.add_argument("--tr-scan", action="store_true",
                      help="With --tr: sweep the transform choices (qmax and the "
                           "window) and print how far the first T(r) peak and its "
                           "integrated count move. The q range belongs to the "
                           "measurement, not the model, so this is the honest "
                           "error bar on a comparison.")
    g_an.add_argument("--pair-panels", action="store_true",
                      help="Also plot each element pair in its own panel: "
                           "analysis_rdf_panels.png for g(r) and, with "
                           "--sq-partials, analysis_sq_partials_panels.png "
                           "for S_ab(q). Easier to read than one axis when "
                           "there are many pairs (IGZO has ten).")
    g_an.add_argument("--sq-smooth", type=float, default=None, metavar="SIGMA_Q",
                      help="Gaussian re-binning width (1/A) for the direct "
                           "S(q), weighted by q-vectors per shell; reduces "
                           "the low-q speckle noise without moving peaks "
                           "(default 0.05; 0 = raw). Keep well below the "
                           "FSDP width (~0.3 1/A). Raw values are always "
                           "kept in the CSV as s_q_raw.")
    g_an.add_argument("--rings", nargs="?", const="auto", default=None, metavar="PAIR",
                      help="Ring statistics (shortest-ring per network edge). "
                           "Optional PAIR such as Ge-O selects the node-bridge "
                           "pair; default auto (least electronegative element as "
                           "nodes). Printed, appended to --save-report, and "
                           "written as analysis_rings.{json,csv,png} with "
                           "per-structure CSVs under --save-plot.")
    g_an.add_argument("--ring-max-size", type=int, default=None, metavar="N",
                      help="Largest ring to search in network nodes (default 12).")
    g_an.add_argument("--ring-cutoff", type=float, default=None, metavar="A",
                      help="Bond cutoff in A for rings only (default: analyser pair cutoff).")
    g_an.add_argument("--voronoi", nargs="?", const="all", default=None, metavar="ELEMENT",
                      help="Voronoi indices <n3 n4 n5 n6> for all atoms or for "
                           "ELEMENT only. Printed, appended to --save-report, and "
                           "written as analysis_voronoi.csv under --save-plot.")
    g_an.add_argument("--connectivity", action="store_true",
                      help="Polyhedral connectivity: corner/edge/face sharing "
                           "between cation-centred polyhedra and the fraction "
                           "of cations in edge-sharing pairs. Printed, appended "
                           "to --save-report, analysis_connectivity.csv under "
                           "--save-plot.")
    g_an.add_argument("--voids", action="store_true",
                      help="Sample periodic free-space clearance and accessible volume.")
    g_an.add_argument("--bond-order", action="store_true",
                      help="Steinhardt q6, Lechner-Dellago averaged qbar6, ordered "
                           "fraction and largest ordered cluster (uses "
                           "--order-cutoff or --cutoff).")
    g_an.add_argument("--order-cutoff", default=None, metavar="SPEC",
                      help="Independent neighbor cutoff for bond order and MQ melt "
                           "memory, in A (number or pair overrides; default: --cutoff).")
    g_an.add_argument("--qbar6-threshold", type=float, default=None,
                      help="Minimum qbar6 for crystal-like order (default 0.3; "
                           "calibrate for the material). Also used by --mq-ensemble.")
    g_an.add_argument("--order-min-neighbors", type=int, default=None,
                      help="Minimum shell neighbors for crystal-like order "
                           "(default 4). Also used by --mq-ensemble.")
    g_an.add_argument("--void-samples", type=int, default=None,
                      help="Random points per cell for --voids (default 10000).")
    g_an.add_argument("--void-probe-radius", type=float, default=None,
                      help="Probe radius in A for --voids (default 0).")
    g_an.add_argument("--void-probe-radii", type=float, nargs="+", default=None,
                      metavar="A", help="Probe radii in A for a free-volume curve "
                      "from the same samples (default: histogram bin edges).")
    g_an.add_argument("--void-bins", type=int, default=None,
                      help="Clearance histogram bins (default 50).")
    g_an.add_argument("--void-seed", type=int, default=None,
                      help="Reproducible void sampling seed (default 0).")
    g_an.add_argument("--oxygen-speciation", action="store_true",
                      help="Oxygen speciation by network-former coordination.")
    g_an.add_argument("--network-formers", default=None, metavar="Si,Al",
                      help="Comma-separated network formers for oxygen speciation; "
                           "defaults to the Al/B/Ge/P/Si present.")
    g_an.add_argument("--elastic", action="store_true",
                      help="Elastic tensor and Voigt/Reuss/Hill moduli from calculator stresses.")
    g_an.add_argument("--elastic-strain", type=float, default=None,
                      help="Central finite strain amplitude (default 0.005).")
    g_an.add_argument("--elastic-relax", action="store_true",
                      help="Relax internal positions at each fixed cell for --elastic.")
    g_an.add_argument("--vdos", action="store_true",
                      help="Harmonic vibrational DOS from calculator forces (6N evaluations/cell).")
    g_an.add_argument("--vdos-displacement", type=float, default=None,
                      help="Finite displacement in A (default 0.01).")
    g_an.add_argument("--vdos-sigma", type=float, default=None,
                      help="Gaussian DOS width in THz (default 0.1).")
    g_an.add_argument("--vdos-npoints", type=int, default=None,
                      help="Frequency grid points (default 400).")
    g_an.add_argument("--total-cn", action="append", default=None, metavar="SPEC",
                      help="Total first-shell coordination of one element over "
                           "several partner types, repeatable. 'O' counts every "
                           "bonded partner (O surrounded by Ga+In+Zn in IGZO); "
                           "'O:In+Ga' counts only the named partners. YAML: "
                           "total_cn: [O, 'O:In+Ga']. The report already "
                           "prints the all-bonded total automatically for "
                           "elements with more than one partner type.")
    g_an.add_argument("--check-dimers", action="store_true",
                      help="Report unphysical close contacts (O-O peroxide, "
                           "Cl-Cl, metal-metal dimers) below 0.85 x the "
                           "radii-derived minsep.")
    g_an.add_argument("--per-structure", action="store_true",
                      help="Per-structure comparison table.")
    g_an.add_argument("--save-report", default=None, metavar="FILE",
                      help="Write text report.")
    g_an.add_argument("--save-plot", default=None, metavar="DIR",
                      help="Save RDF/CN/angles plots (PNG + CSV).")
    g_an.add_argument("--save-pdf", action="store_true",
                      help="Also save plots as vector PDF.")
    g_an.add_argument("--dpi", type=int, default=None, metavar="N",
                      help="Plot DPI (default 300).")
    g_an.add_argument("--show-title", action="store_true",
                      help="Add titles to plots (off by default).")
    g_an.add_argument("--total-rdf", action="store_true",
                      help="Include total g(r) in RDF plot.")
    g_an.add_argument("--smearing", type=float, default=None, metavar="SIGMA",
                      help="RDF Gaussian smearing in A (default 0.05, ~thermal/"
                           "experimental broadening; 0 = raw histogram). "
                           "S(q) is always computed from the raw g(r).")
    g_an.add_argument("--reference", default=None, metavar="YAML",
                      help="Reference YAML; adds validation table.")

    return p


[docs] def parse_args(): """Parse command-line arguments.""" return _get_parser().parse_args()
def _parse_composition(spec: str) -> dict[str, int]: """Parse composition from two supported formats. Format 1 (atom counts): 'In=32,O=48' -> {'In': 32, 'O': 48} Format 2 (formula * N): 'In2O3*16' -> {'In': 32, 'O': 48} 'SiO2*16' -> {'Si': 16, 'O': 32} 'Si64' -> {'Si': 64} The '*N' multiplier scales a chemical formula by N formula units. If no '*' and no '=' is present, the formula is used as-is. Raises ValueError with a helpful message on malformed input. """ from ase.data import chemical_symbols as _chem_syms spec = spec.strip() if not spec: raise ValueError("Empty composition string.") # --- Format 2: chemical formula (no '=' signs) --- if "=" not in spec: from ase import Atoms as _Atoms # Split on '*' for multiplier: e.g. 'In2O3*16' if "*" in spec: parts = spec.split("*", 1) formula = parts[0].strip() try: multiplier = int(parts[1].strip()) except ValueError: raise ValueError( f"Invalid multiplier in '{spec}'. " f"Expected 'Formula*N' (e.g. 'In2O3*16')." ) if multiplier <= 0: raise ValueError( f"Multiplier must be positive, got {multiplier}." ) else: formula = spec multiplier = 1 try: tmp = _Atoms(formula * multiplier) except Exception: raise ValueError( f"Cannot parse '{formula}' as a chemical formula.\n" f" Accepted formats:\n" f" In=32,O=48 (atom counts)\n" f" In2O3*16 (formula * N formula units)\n" f" SiO2*16 (formula * N)\n" f" Si64 (element + count)" ) syms = tmp.get_chemical_symbols() comp = {} for s in syms: comp[s] = comp.get(s, 0) + 1 total = sum(comp.values()) formula_str = tmp.get_chemical_formula(mode="hill") print(f" [Composition] {formula_str}: {comp} ({total} atoms)") return comp # --- Format 1: Element=count pairs --- comp = {} for part in spec.split(","): part = part.strip() if not part: continue if "=" not in part: raise ValueError( f"Invalid composition entry: '{part}'. " f"Expected 'Element=count' (e.g. 'Si=16,O=32').\n" f" Or use formula format: 'SiO2*16'" ) pieces = part.split("=", 1) sym = pieces[0].strip() count_str = pieces[1].strip() if not sym: raise ValueError(f"Empty element symbol in: '{part}'") if sym not in _chem_syms: raise ValueError( f"Unknown element symbol '{sym}'. " f"Check spelling (case-sensitive: 'O' not 'o', 'In' not 'IN')." ) try: count = int(count_str) except ValueError: raise ValueError( f"Invalid count for '{sym}': '{count_str}' (must be an integer)." ) if count <= 0: raise ValueError( f"Count for '{sym}' must be positive, got {count}." ) comp[sym] = count if not comp: raise ValueError(f"No valid entries in composition: '{spec}'") return comp def _classical_kwargs(override: dict) -> dict: """Extract classical_params from config override if present.""" kw = {} if override.get("classical_params"): kw["classical_params"] = override["classical_params"] return kw def _parse_minsep(spec: str) -> dict[str, float]: """Parse 'In-In=2.8,In-O=1.9,O-O=2.5' -> {'In-In': 2.8, ...}. Raises ValueError on malformed input. """ minsep = {} for part in spec.split(","): part = part.strip() if not part: continue if "=" not in part: raise ValueError( f"Invalid minsep entry: '{part}'. " f"Expected 'A-B=distance' (e.g. 'Si-O=1.6')." ) pair, val_str = part.split("=", 1) pair = pair.strip() if "-" not in pair: raise ValueError( f"Invalid pair format: '{pair}'. Expected 'A-B' (e.g. 'Si-O')." ) try: val = float(val_str.strip()) except ValueError: raise ValueError( f"Invalid minsep value for '{pair}': '{val_str.strip()}'." ) if val <= 0: raise ValueError( f"Minsep for '{pair}' must be positive, got {val}." ) minsep[pair] = val return minsep def _parse_target_cn(spec: str) -> dict[str, int]: """Parse 'Si=4,O=2' -> {'Si': 4, 'O': 2}. Raises ValueError on malformed input. """ target_cn = {} for part in spec.split(","): part = part.strip() if not part: continue if "=" not in part: raise ValueError( f"Invalid target-cn entry: '{part}'. " f"Expected 'Element=CN' (e.g. 'Si=4,O=2')." ) sym, cn_str = part.split("=", 1) sym = sym.strip() try: cn = int(cn_str.strip()) except ValueError: raise ValueError( f"Invalid CN for '{sym}': '{cn_str.strip()}' (must be an integer)." ) if cn <= 0: raise ValueError( f"CN for '{sym}' must be positive, got {cn}." ) target_cn[sym] = cn return target_cn def _parse_dmax(spec: str) -> dict[str, float]: """Parse 'Si-O=2.0,Si-Si=3.2' -> {'Si-O': 2.0, ...}. Raises ValueError on malformed input. """ dmax = {} for part in spec.split(","): part = part.strip() if not part: continue if "=" not in part: raise ValueError( f"Invalid dmax entry: '{part}'. " f"Expected 'A-B=distance' (e.g. 'Si-O=2.0')." ) pair, val_str = part.split("=", 1) pair = pair.strip() if "-" not in pair: raise ValueError( f"Invalid pair format: '{pair}'. Expected 'A-B' (e.g. 'Si-O')." ) try: val = float(val_str.strip()) except ValueError: raise ValueError( f"Invalid dmax value for '{pair}': '{val_str.strip()}'." ) if val <= 0: raise ValueError( f"Dmax for '{pair}' must be positive, got {val}." ) dmax[pair] = val return dmax def _build_override(args, parser, explicit_only: bool = False, argv: list[str] | None = None) -> dict: """ Build a config override dict from CLI args. Parameters ---------- args : argparse.Namespace parser : argparse.ArgumentParser explicit_only : bool If True, only include args the user explicitly set (for YAML mode, so defaults don't overwrite YAML values). If False, include all args with their defaults (for non-YAML mode). argv : list[str], optional Raw command-line tokens used to detect flags that were typed with a value equal to the parser default (defaults to ``sys.argv[1:]``). Returns ------- dict — config override ready to merge with DEFAULT_CONFIG or YAML. """ if explicit_only: explicit = {k for k, v in vars(args).items() if v != parser.get_default(k)} # A flag typed on the command line with a value equal to the parser # default (e.g. ``--eq-high-ensemble NVT``) is still an explicit # choice and must beat the YAML / DEFAULT_CONFIG value. argparse # cannot tell the two apart, so look at the raw argv as well. if argv is None: argv = sys.argv[1:] typed = {tok.split("=", 1)[0] for tok in argv if tok.startswith("-")} for action in parser._actions: if action.dest != "help" and typed.intersection(action.option_strings): explicit.add(action.dest) def get(key): return getattr(args, key) if key in explicit else None else: def get(key): return getattr(args, key) mapping = { "model": get("model"), "seed": get("seed"), "run_index": get("run_index"), "engine": get("engine"), "model_path": get("model_path"), "device": get("device"), "default_dtype": get("default_dtype"), "opt": { "fmax": get("fmax"), "max_steps": get("opt_steps"), "optimizer": get("optimizer"), "cell_filter": get("cell_filter"), "output_format": get("format"), }, "eq_premelt": { "ensemble": get("eq_premelt_ensemble"), "T": get("eq_premelt_T"), "steps": get("eq_premelt_steps"), "timestep": get("timestep"), }, "melt": { "ensemble": get("melt_ensemble"), "T_start": get("melt_T_start"), "T_end": get("melt_T_end"), "T_step": get("melt_T_step"), "steps_per_T": get("melt_steps_per_T"), "timestep": get("timestep"), }, "eq_high": { "ensemble": get("eq_high_ensemble"), "T": get("eq_high_T"), "steps": get("eq_high_steps"), "timestep": get("timestep"), }, "quench": { "ensemble": get("quench_ensemble"), "T_start": get("quench_T_start"), "T_end": get("quench_T_end"), "T_step": get("quench_T_step"), "steps_per_T": get("quench_steps_per_T"), "timestep": get("timestep"), }, "eq_low": { "ensemble": get("eq_low_ensemble"), "T": get("eq_low_T"), "steps": get("eq_low_steps"), "timestep": get("timestep"), }, "final_opt": { "fmax": get("fmax"), "max_steps": get("opt_steps"), "optimizer": get("optimizer"), "cell_filter": get("cell_filter"), "output_format": get("format"), }, } if explicit_only: # Strip None values so they don't overwrite YAML/defaults def _strip_none(d): out = {} for k, v in d.items(): if isinstance(v, dict): nested = _strip_none(v) if nested: out[k] = nested elif v is not None: out[k] = v return out return _strip_none(mapping) return mapping def _apply_amorphous_cubic_default(args, override, ff_default): """Default the cell filter to ``cubic`` for amorphous-input modes. ``--hybrid-ensemble`` and ``--batch-opt`` operate on disordered/amorphous structures, which are isotropic and should relax under a cubic (hydrostatic) constraint rather than the anisotropic ``FrechetCellFilter`` parser default. This mirrors the cubic default already applied to ``--random-gen``. The 7-stage melt-quench pipeline is intentionally *not* covered here, because it starts from a (possibly non-cubic) crystal that a cubic constraint would distort. An explicit ``-C`` on the CLI, or a non-default ``cell_filter`` in YAML, is preserved. Mutates and returns ``override``. """ if not (getattr(args, "hybrid_ensemble", False) or getattr(args, "batch_opt", False)): return override if _typed("-C", "--cell-filter"): return override # user set -C explicitly -> respect it if not isinstance(override, dict): return override for section in ("opt", "final_opt"): sec = override.get(section) if isinstance(sec, dict): # None/absent (YAML didn't set) or the FrechetCellFilter parser # default (no-YAML path) -> cubic; a real YAML choice is kept. if sec.get("cell_filter") in (None, ff_default): sec["cell_filter"] = "cubic" else: override[section] = {"cell_filter": "cubic"} return override # ═════════════════════════════════════════════════════════════════════════════ # Convert mode (--convert) # ═════════════════════════════════════════════════════════════════════════════ def _run_convert(args, yaml_cfg: dict | None = None) -> None: """Thin CLI wrapper around :func:`amorphgen.utils.convert.convert`. Resolves the (input_path, output_format, output_dir) tuple from CLI flags first, falling back to a ``convert:`` block in the YAML config when the CLI flag is at its default (i.e. user did not set it). """ from .utils import convert as _convert yaml_block = (yaml_cfg or {}).get("convert", {}) if yaml_cfg else {} # Detect "user explicitly set --format" by comparing to the parser default. # If the user did NOT pass --format, fall through to the YAML value. parser = _get_parser() fmt_default = parser.get_default("format") input_path = args.convert or yaml_block.get("input") if args.format and _typed("--format"): output_format = args.format else: output_format = yaml_block.get("format", args.format or "vasp") # Ignore the "melt_quench_run" fallback when YAML drives convert mode # without an explicit -o flag, so the auto "<input>_<format>/" default # in convert() can take over. work_dir = args.work_dir if work_dir == "melt_quench_run": work_dir = None output_dir = work_dir or yaml_block.get("output_dir") if not input_path: print("Error: --convert needs an input PATH (or `convert.input` in " "the YAML config).") sys.exit(1) try: _convert(input_path, output_format=output_format, output_dir=output_dir) except (FileNotFoundError, ValueError) as exc: print(f"Error: {exc}") sys.exit(1) # ═════════════════════════════════════════════════════════════════════════════ # Ensemble workflow helpers (--mq-ensemble, --hybrid-ensemble) # ═════════════════════════════════════════════════════════════════════════════ def _has_quench_outputs(work_dir, cfg, stages): """Recognise completed or interrupted downstream runs before reselection.""" from pathlib import Path root = Path(work_dir) directories = [root] + [path for path in root.glob("run_*") if path.is_dir()] patterns = [f"stage{stage}*" for stage in stages] + ["final_amorphous.*"] if any(path.is_file() for directory in directories for pattern in patterns for path in directory.glob(pattern)): return True sections = {4: "eq_high", 5: "quench", 6: "eq_low", 7: "final_opt"} names = [section[key] for stage in stages if stage in sections for section in [cfg.get(sections[stage], {})] for key in ("traj_file", "log_file", "output_xyz") if section.get(key)] return any((directory / name).is_file() for directory in directories for name in names) def _snapshot_sampling_kwargs(args, override, *, mq=False, report_path=None, quench_dir=None): """Use the actual stage-4 timestep for automatic and manual extraction.""" from .configs import DEFAULT_CONFIG from .utils import merge_config from .utils.common import TRAJ_LOG_INTERVAL cfg = merge_config(DEFAULT_CONFIG, override) select = args.select or ("decorrelated" if mq else "uniform") burn_in = args.burn_in_frames if burn_in is None and select != "decorrelated": burn_in = 0 return { "select": select, "burn_in_frames": burn_in, "timestep_fs": cfg["eq_high"]["timestep"], "frame_stride": TRAJ_LOG_INTERVAL, "decorrelation_distance": args.decorrelation_distance, "report_path": (report_path if mq or select == "decorrelated" or (report_path is not None and os.path.isfile(report_path)) else None), "resume": bool(args.resume and quench_dir is not None and _has_quench_outputs( quench_dir, cfg, (5, 6, 7) if mq else args.batch_stages)), } def _run_mq_ensemble(args, override: dict, analysis_config=None) -> None: """Full melt-quench ensemble: stages 1-4 once + N independent quenches. Output layout under args.work_dir: shared/ stages 1-4 outputs (incl. stage4_eq_traj.xyz trajectory) snapshots/ up to N decorrelated snapshots from stage 4 trajectory quench_runs/ per-snapshot stages 5-6-7 outputs (run_0000, run_0001, ...) final/ collected final amorphous structures (mq_NNNN.<fmt>) melt_memory.{json,csv,txt} original-crystal order retention at melt endpoints snapshot_sampling.{json,txt} spacing and effective independent sample count """ from .pipeline.run_pipeline import MeltQuenchPipeline from .pipeline import batch_quench from .utils import get_calculator, extract_snapshots from .pipeline.random_gen import _FORMAT_MAP from .analysis.descriptors import bond_order_options from .analysis.melt_memory import prepare_melt_memory, report_melt_memory work_dir = args.work_dir or "mq_ensemble" shared_dir = os.path.join(work_dir, "shared") snap_dir = os.path.join(work_dir, "snapshots") quench_dir = os.path.join(work_dir, "quench_runs") final_dir = os.path.join(work_dir, "final") os.makedirs(work_dir, exist_ok=True) if analysis_config is None: analysis_config = override.get("analysis", {}) order_options = bond_order_options(args, analysis_config) # Validate and resolve the original-crystal order definition before MD. memory_reference = prepare_melt_memory(args.input_file, **order_options) bar = "=" * 70 print(f"\n{bar}") print(f" AmorphGen MQ-ensemble: {args.input_file} -> " f"{args.n_structures} amorphous structures") print(f" Output: {work_dir}/") print(bar) # ── Phase 1: shared stages 1-4 (with resume) ───────────────────────────── print(f"\n[Phase 1/3] Stages 1-4 (shared) -> {shared_dir}/") pipe = MeltQuenchPipeline(args.input_file, work_dir=shared_dir, cfg_override=override) pipe.run(stages=[1, 2, 3, 4], resume=args.resume) # ── Phase 2: extract N snapshots from stage 4 trajectory ───────────────── traj = os.path.join(shared_dir, override.get("eq_high", {}).get( "traj_file", "stage4_eq_traj.xyz")) if not os.path.isfile(traj): # Backwards-compat: older runs wrote stage4_eq.xyz as the trajectory legacy = os.path.join(shared_dir, "stage4_eq.xyz") if os.path.isfile(legacy): traj = legacy else: report_melt_memory( args.input_file, shared_dir, [], work_dir, cfg_override=override, prepared=memory_reference) print(f"Error: stage 4 trajectory not found " f"({traj} or {legacy})") sys.exit(1) print(f"\n[Phase 2/3] Extracting {args.n_structures} snapshots from {traj}") try: snap_files = extract_snapshots(traj, n_snapshots=args.n_structures, output_dir=snap_dir, **_snapshot_sampling_kwargs( args, override, mq=True, report_path=os.path.join(work_dir, "snapshot_sampling.json"), quench_dir=quench_dir)) except Exception: # Preserve diagnostics from completed MD even if a trajectory cannot # be read or the requested sampling range is invalid. try: report_melt_memory( args.input_file, shared_dir, [], work_dir, cfg_override=override, prepared=memory_reference) except Exception as report_error: print(f"Warning: could not save melt-memory report: {report_error}") raise # Keep the original snapshot-extraction failure. # Save melt diagnostics before any expensive or interrupted quench run. # Use the exact extracted files, excluding stale snapshots on resume. report_melt_memory( args.input_file, shared_dir, snap_files, work_dir, cfg_override=override, prepared=memory_reference) # ── Phase 3: stages 5-6-7 per snapshot (with resume) ───────────────────── print(f"\n[Phase 3/3] Stages 5-6-7 (per snapshot) -> {quench_dir}/") calc = get_calculator( **_classical_kwargs(override), model=override.get("model", args.model), device=override.get("device", args.device), model_path=override.get("model_path", args.model_path), default_dtype=override.get("default_dtype", args.default_dtype), ) batch_quench.run( snapshot_files=snap_files, n_runs=len(snap_files), cfg_override=override, work_dir=quench_dir, stages=[5, 6, 7], calc=calc, resume=args.resume, ) _collect_ensemble_final(quench_dir, final_dir, args.format, prefix="mq", fmt_map=_FORMAT_MAP, snapshot_files=snap_files) print(f"\n{bar}") print(f" MQ ensemble complete -> {final_dir}/") print(bar) def _run_hybrid_ensemble(args, override: dict) -> None: """Hybrid ensemble: stages 4-7 per disordered input structure. Output layout under args.work_dir: quench_runs/ per-input stages 4-7 outputs (run_0000, run_0001, ...) final/ collected final amorphous structures (hybrid_NNNN.<fmt>) """ import glob as _glob from .pipeline import batch_quench from .utils import get_calculator from .pipeline.random_gen import _FORMAT_MAP, random_gen_dir_hint work_dir = args.work_dir or "hybrid_ensemble" quench_dir = os.path.join(work_dir, "quench_runs") final_dir = os.path.join(work_dir, "final") os.makedirs(work_dir, exist_ok=True) # Find input structures (any ASE-readable format) snap_files = [] for pattern in ("*.xyz", "*.extxyz", "*.vasp", "*.cif", "POSCAR*"): snap_files = sorted(_glob.glob(os.path.join(args.input_dir, pattern))) if snap_files: break if not snap_files: print(f"Error: no structure files in {args.input_dir}/ " f"(looked for *.xyz, *.extxyz, *.vasp, *.cif, POSCAR*)") hint = random_gen_dir_hint(args.input_dir) if hint: print(hint) sys.exit(1) bar = "=" * 70 print(f"\n{bar}") print(f" AmorphGen hybrid-ensemble: {len(snap_files)} structures from " f"{args.input_dir}/") print(f" Output: {work_dir}/") print(bar) if override.get('engine', 'ase') == 'torchsim': from .pipeline.batch_quench import run_torchsim # torch-sim runs NVT only. DEFAULT_CONFIG has stage 4 as NPT, so unless # the user chose an ensemble (flag or YAML) switch the MD stages to NVT # here; an explicit NPT is refused before anything starts. for key, flag in (("eq_high", "--eq-high-ensemble"), ("quench", "--quench-ensemble"), ("eq_low", "--eq-low-ensemble")): chosen = (override.get(key) or {}).get("ensemble") if chosen is None: override.setdefault(key, {})["ensemble"] = "NVT" elif str(chosen).upper() != "NVT": print(f"Error: --engine torchsim runs NVT only; {flag} {chosen} is not " f"supported. Drop the flag (NVT is used) or use --engine ase.") sys.exit(1) print(" torch-sim engine: MD stages run NVT (fixed cell)") run_torchsim(snap_files, cfg_override=override, work_dir=quench_dir, stages=[4, 5, 6, 7], resume=args.resume, batch_size=(int(args.batch_size) if getattr(args, 'batch_size', None) and str(args.batch_size).isdigit() else getattr(args, 'batch_size', None))) _collect_ensemble_final(quench_dir, final_dir, args.format, prefix='hybrid', fmt_map=_FORMAT_MAP, snapshot_files=snap_files, flat_single_run=False) print(f'\n{bar}\n Hybrid ensemble complete (torch-sim) -> {final_dir}/\n{bar}') return calc = get_calculator( **_classical_kwargs(override), model=override.get("model", args.model), device=override.get("device", args.device), model_path=override.get("model_path", args.model_path), default_dtype=override.get("default_dtype", args.default_dtype), ) batch_quench.run( snapshot_files=snap_files, n_runs=len(snap_files), cfg_override=override, work_dir=quench_dir, stages=[4, 5, 6, 7], calc=calc, resume=args.resume, ) _collect_ensemble_final(quench_dir, final_dir, args.format, prefix="hybrid", fmt_map=_FORMAT_MAP, snapshot_files=snap_files) print(f"\n{bar}") print(f" Hybrid ensemble complete -> {final_dir}/") print(bar) def _collect_ensemble_final(quench_dir: str, final_dir: str, output_format: str, prefix: str, fmt_map: dict, snapshot_files: list[str] | None = None, flat_single_run: bool = True) -> None: """Collect batch outputs, including the ASE single-input flat layout. Input filenames preserve the snapshot index of a flat output and restrict collection to the current batch. Torch-sim always uses run_NNNN directories. """ import glob as _glob from ase.io import read, write from .pipeline.batch_quench import _run_dir_name from .utils.relaxation import read_relaxation_metadata, write_relaxation_metadata if output_format not in fmt_map: print(f"Warning: unknown format '{output_format}', using 'xyz'") output_format = "xyz" ase_format, ext = fmt_map[output_format] def final_output(run_dir): for name in ("final_amorphous.xyz", "final_amorphous.extxyz"): path = os.path.join(run_dir, name) if os.path.isfile(path): return path return None flat_output = final_output(quench_dir) if flat_single_run else None if snapshot_files is None: runs = [(os.path.basename(d), final_output(d)) for d in sorted(_glob.glob(os.path.join(quench_dir, "run_*"))) if os.path.isdir(d)] if flat_output: runs = [(name, src) for name, src in runs if name != "run_0000"] runs.insert(0, ("run_0000", flat_output)) runs = [(name, src) for name, src in runs if src is not None] else: runs = [] for i, snap_file in enumerate(snapshot_files): name = _run_dir_name(snap_file, fallback_idx=i) src = (flat_output if len(snapshot_files) == 1 and flat_output else final_output(os.path.join(quench_dir, name))) if src is None: raise FileNotFoundError( f"No final batch output for {snap_file!r} in {quench_dir!r} " f"(expected {name}/final_amorphous.xyz or a flat singleton output)." ) runs.append((name, src)) if not runs: raise FileNotFoundError(f"No final batch outputs found in {quench_dir!r}.") os.makedirs(final_dir, exist_ok=True) n_collected = 0 for name, src in runs: idx = name.removeprefix("run_") dest = os.path.join(final_dir, f"{prefix}_{idx}{ext}") atoms = read(src) read_relaxation_metadata(src, atoms) if ase_format == "vasp": atoms = atoms[atoms.numbers.argsort()] write(dest, atoms, format=ase_format, sort=True) else: write(dest, atoms, format=ase_format) write_relaxation_metadata(dest, atoms) n_collected += 1 print(f" Collected {n_collected} final structures -> {final_dir}/") def _requires_calculator(args, analysis_config=None) -> bool: """Will this invocation construct a calculator? Gates the fail-fast backend check. Modes that only read, transform, or analyse geometry never need a backend and must keep working on a torch-free install. Elastic and harmonic vibrational descriptors do. """ # Calculator-free modes (checked first — they may combine with input_file) if (args.list_models or args.rank_from_log or args.convert or args.extract_snapshots): return False if args.analyse: cfg = analysis_config or {} return bool(getattr(args, "elastic", False) or getattr(args, "vdos", False) or cfg.get("elastic", False) or cfg.get("vdos", False)) # Random generation only builds a calculator when relaxing if args.random_gen: return bool(args.relax) # MD / optimisation modes always need one if (args.batch_opt or args.batch_quench or args.mq_ensemble or args.hybrid_ensemble): return True # Bare input_file = full melt-quench pipeline return args.input_file is not None
[docs] def main(): from .utils.preemption import (checkpoint_signals, stop_if_requested, PreemptionRequested) with checkpoint_signals(): try: try: result = _main() except SystemExit as exc: if exc.code in (None, 0): stop_if_requested() raise # A signal during a short stage or non-simulation command must # still prevent successful Slurm dependencies from starting. stop_if_requested() return result except PreemptionRequested as exc: print(f"[Preemption] {exc}", file=sys.stderr) raise SystemExit(exc.code) from None
def _main(): args = parse_args() # ── Smart default for --work-dir based on mode ─────────────────────────── if args.work_dir is None: if args.random_gen: from ase import Atoms comp = args.composition or "" try: symbols = [] for pair in comp.replace(" ", "").split(","): el, n = pair.split("=") symbols.extend([el] * int(n)) formula = Atoms(symbols).get_chemical_formula(mode="hill") args.work_dir = f"random_{formula}" except Exception: args.work_dir = "random_structures" elif args.batch_quench: args.work_dir = "batch_quench" elif args.batch_opt: args.work_dir = "batch_opt" elif getattr(args, "analyse", False): args.work_dir = "analysis" elif args.convert: # Leave None so the convert helper picks # "<input>_<format>/" as its automatic default. pass elif getattr(args, "mq_ensemble", False): args.work_dir = "mq_ensemble_run" elif getattr(args, "hybrid_ensemble", False): args.work_dir = "hybrid_run" elif getattr(args, "extract_snapshots", None): args.work_dir = "snapshots" elif getattr(args, "batch_quench", False): args.work_dir = "batch_quench_run" else: args.work_dir = "melt_quench_run" # ── List models ─────────────────────────────────────────────────────────── if args.list_models: from .utils import list_models list_models() sys.exit(0) # ── Build config override ───────────────────────────────────────────────── # Precedence: CLI args > YAML config > DEFAULT_CONFIG if args.config is not None: from .configs import load_yaml_config from .utils import merge_config yaml_cfg = load_yaml_config(args.config) print(f"[Config] Loaded: {args.config}") # Build override from only explicitly-set CLI args # We need the parser to detect defaults — re-parse to get it cli_override = _build_override(args, _get_parser(), explicit_only=True) # Merge: YAML first, then CLI on top override = merge_config(yaml_cfg, cli_override) else: # No YAML: still pass only what was typed, so DEFAULT_CONFIG (the # documented defaults) fills the rest exactly as in YAML mode. override = _build_override(args, _get_parser(), explicit_only=True) # --hybrid-ensemble / --batch-opt relax with the stage-7 ("final_opt") # settings; let a YAML that only has an `opt:` block drive that relax # instead of silently falling back to the defaults. if (getattr(args, "hybrid_ensemble", False) or getattr(args, "batch_opt", False)) \ and isinstance(override, dict) and isinstance(override.get("opt"), dict): merged = dict(override["opt"]); merged.update(override.get("final_opt") or {}) override["final_opt"] = merged # Amorphous-input modes default to a cubic (isotropic) cell filter. override = _apply_amorphous_cubic_default( args, override, _get_parser().get_default("cell_filter")) try: until_options = _until_convergence_options(args, override) except ValueError as exc: print(f"Error: sequential convergence: {exc}") sys.exit(1) # ── Fail fast when the requested backend is missing ────────────────────── if (args.screen or args.screening_output is not None) and not args.analyse: print("Error: --screen and --screening-output require --analyse.") sys.exit(1) # Calculator-requiring modes abort BEFORE any setup work (no work dir, no # structure loading) with a copy-pasteable install hint. Backend knowledge # lives in utils.calculators (require_backend); this is just the gate. # The same gate refuses a precision the model can't run (CHGNet + float64), # which --mq-ensemble would otherwise only hit in phase 3, after stages # 1-4 of MD. if until_options is not None or _requires_calculator(args, override.get("analysis", {})): from .utils.calculators import (require_backend, require_dtype, BackendNotInstalledError) model = override.get("model", args.model) or "mace-mpa-0" model_path = override.get("model_path", getattr(args, "model_path", None)) try: require_backend(model, model_path=model_path) require_dtype(model, override.get("default_dtype", args.default_dtype), model_path=model_path) except (BackendNotInstalledError, ValueError, NotImplementedError) as exc: print(f"Error: {exc}") sys.exit(1) # ── Rank structures from a random-gen log file ──────────────────────────── if args.rank_from_log: from .analysis.energy import rank_from_log, format_log_ranking result = rank_from_log(args.rank_from_log) print(format_log_ranking(result, logfile=args.rank_from_log)) return # ── Convert one or many structure files to a different format ──────────── yaml_convert_block = override.get("convert") if isinstance(override, dict) else None if args.convert or (yaml_convert_block and yaml_convert_block.get("input")): _run_convert(args, yaml_cfg=override if isinstance(override, dict) else None) return # ── Extract snapshots from a trajectory file ────────────────────────────── if args.extract_snapshots: from .utils import extract_snapshots out_dir = args.work_dir or "snapshots" # Snapshot count: prefer --n-structures (the unified count flag); # fall back to --n-runs for backwards compatibility. --n-structures # has default 1, --n-runs has default 20, so pick the one the user # actually changed. if _typed("-n", "--n-structures"): n_snap = args.n_structures elif _typed("--n-runs"): n_snap = args.n_runs else: n_snap = 20 # historical default extract_snapshots( args.extract_snapshots, n_snapshots=n_snap, output_dir=out_dir, output_format=args.format, **_snapshot_sampling_kwargs( args, override, report_path=os.path.join(out_dir, "snapshot_sampling.json")), ) return # ── MQ-ensemble mode ────────────────────────────────────────────────────── # Full melt-quench ensemble workflow in one command: # stages 1-4 (shared, from crystal) -> extract N snapshots -> stages 5-7 # independently per snapshot. Resume-aware at every step. if args.mq_ensemble: if args.input_file is None: print("Error: input_file (crystal structure) is required for --mq-ensemble.") sys.exit(1) _run_mq_ensemble(args, override, analysis_config=override.get("analysis", {})) return # ── Hybrid-ensemble mode ────────────────────────────────────────────────── # Take all structures in --input-dir and run stages 4-5-6-7 on each. # Useful for "AmorphGen random + chgnet quench" workflows. if args.hybrid_ensemble: if args.input_dir is None: print("Error: --input-dir is required for --hybrid-ensemble.") sys.exit(1) _run_hybrid_ensemble(args, override) return # ── Structure analysis mode ──────────────────────────────────────────────── if args.analyse: from .analysis import StructureAnalyser source = args.input_dir or args.input_file if source is None: print("Error: --input-dir or input_file is required for --analyse.") print(" Example: amorphgen --analyse --input-dir optimised_structures/") sys.exit(1) # Read analysis block from YAML config (if present) an_cfg = override.get("analysis", {}) from .analysis.screening import validate_screening_config try: screening_spec = an_cfg.get("screening", False) if args.screen and not screening_spec: screening_spec = True screening_config = validate_screening_config(screening_spec) screening_prefix = args.screening_output or an_cfg.get("screening_output") if screening_prefix and not screening_config: raise ValueError("screening_output requires --screen or analysis.screening") except (TypeError, ValueError) as exc: print(f"Error: screening: {exc}") sys.exit(1) try: cutoff_window = _cutoff_window_option(args, an_cfg) except ValueError as exc: print(f"Error: cutoff robustness: {exc}") sys.exit(1) try: convergence_enabled, tolerances, convergence_confidence, convergence_max = ( _convergence_options(args, an_cfg)) except ValueError as exc: print(f"Error: convergence analysis: {exc}") sys.exit(1) convergence_descriptors = {} # Parse cutoff: CLI > YAML > default "auto" cutoff = args.cutoff parser = _get_parser() if not _typed("--cutoff") and "cutoff" in an_cfg: cutoff = an_cfg["cutoff"] from .analysis.cutoff import parse_cutoff_spec try: cutoff = parse_cutoff_spec(cutoff) except ValueError as exc: print(f"Error: {exc}. Use a number, 'auto', 'auto-rdf', or " f"per-pair overrides such as 'In-O=2.6,Zn-O=2.3'.") sys.exit(1) sa = StructureAnalyser(source, cutoff=cutoff) screening = None if screening_config: from .analysis.screening import ( format_screening_report, write_screening_outputs, mark_screening_analysed) candidates = sa try: sa, screening = candidates.screened(screening_config) except (TypeError, ValueError) as exc: print(f"Error: screening: {exc}") sys.exit(1) screening_prefix = screening_prefix or os.path.join(args.work_dir, "screening") # Preserve every decision even if a later analysis fails. Only # successful completion below marks retained candidates analysed. write_screening_outputs(screening, screening_prefix) if sa is None: screening_text = format_screening_report(screening) print(screening_text) print(" No structures retained for analysis.") empty_report_path = args.save_report or an_cfg.get("save_report") if empty_report_path: candidates.save_report(empty_report_path, text=screening_text) return # Per-structure or grouped analysis per_structure = args.per_structure or an_cfg.get("per_structure", False) if per_structure: text = sa.per_structure_summary(cutoff_window=cutoff_window) else: text = sa.summary(cutoff_window=cutoff_window) # Dimer check: CLI flag > YAML. summary() prints itself, so print the # dimer section too; the concatenated text feeds --save-report. if args.check_dimers or an_cfg.get("check_dimers", False): from .analysis.structure import format_dimer_report dimers = sa.dimer_report() dimer_text = format_dimer_report(dimers) if convergence_enabled: for name, summary in ( ("count", dimers.get("uncertainty")), ("fraction_of_sites", dimers.get("site_fraction_uncertainty")), ("fraction_of_structures", dimers.get("structure_fraction_uncertainty"))): _collect_convergence_summaries( convergence_descriptors, f"dimers.{name}", summary) for pair, data in dimers.get("pairs", {}).items(): _collect_convergence_summaries( convergence_descriptors, f"dimers.count.{pair}", data.get("uncertainty")) _collect_convergence_summaries( convergence_descriptors, f"dimers.min_distance.{pair}", data.get("min_distance_uncertainty")) print(dimer_text) text += "\n" + dimer_text # Requested total coordinations: CLI (repeatable) > YAML list total_cn = args.total_cn if total_cn is None and "total_cn" in an_cfg: y = an_cfg["total_cn"] total_cn = [y] if isinstance(y, str) else list(y) if total_cn: from .analysis.analyser import format_total_cn tcn_text = format_total_cn(sa, total_cn) print(tcn_text) text += "\n" + tcn_text # Save report: CLI > YAML report_path = args.save_report if report_path is None and "save_report" in an_cfg: report_path = an_cfg["save_report"] if report_path: sa.save_report(report_path, text=text) # Save plots: CLI > YAML plot_dir = args.save_plot if plot_dir is None and "save_plot" in an_cfg: plot_dir = an_cfg["save_plot"] # Plot settings from YAML plot_kwargs = {"cutoff_window": cutoff_window} if "rdf_pairs" in an_cfg: plot_kwargs["rdf_pairs"] = an_cfg["rdf_pairs"] if "angle_triplets" in an_cfg: plot_kwargs["angle_triplets"] = an_cfg["angle_triplets"] if "angle_style" in an_cfg: plot_kwargs["angle_style"] = an_cfg["angle_style"] if "rmax" in an_cfg: plot_kwargs["rmax"] = an_cfg["rmax"] # Smearing: CLI > YAML > default (0.0) # smearing: CLI > YAML analyse block > DEFAULT_SMEARING; 0 = raw. from .analysis.rdf import DEFAULT_SMEARING smearing = args.smearing if smearing is None: smearing = an_cfg.get("smearing", DEFAULT_SMEARING) plot_kwargs["smearing"] = float(smearing) # Total RDF: CLI flag or YAML if args.total_rdf or an_cfg.get("total_rdf", False): plot_kwargs["show_total_rdf"] = True # Publication-quality knobs (CLI > YAML) if args.save_pdf or an_cfg.get("save_pdf", False): plot_kwargs["save_pdf"] = True if args.dpi is not None: plot_kwargs["dpi"] = args.dpi elif "dpi" in an_cfg: plot_kwargs["dpi"] = an_cfg["dpi"] if args.show_title or an_cfg.get("show_title", False): plot_kwargs["show_title"] = True if args.pair_panels or an_cfg.get("pair_panels", False): plot_kwargs["pair_panels"] = True if total_cn: plot_kwargs["total_cn"] = list(total_cn) if plot_dir: sa.plot(output_dir=plot_dir, **plot_kwargs) # S(q): CLI flag > YAML (direct method, Faber-Ziman normalised) sq_result = None tr = None if (args.sq or an_cfg.get("sq", False) or args.experiment_sq or an_cfg.get("experiment_sq")): sq_weighting = args.sq_weighting if (not _typed("--sq-weighting") and "sq_weighting" in an_cfg): sq_weighting = an_cfg["sq_weighting"] L_min = min(min(a.cell.lengths()) for a in sa.atoms_list) q_min = 2 * 3.141592653589793 / L_min sq_method = args.sq_method if not _typed("--sq-method") and "sq_method" in an_cfg: sq_method = an_cfg["sq_method"] sq_qmax = (args.sq_qmax if _typed("--sq-qmax") else an_cfg.get("sq_qmax", args.sq_qmax)) sq_nq = (args.sq_nq if _typed("--sq-nq") else an_cfg.get("sq_nq", args.sq_nq)) import math if (not math.isfinite(sq_qmax) or sq_qmax <= 0.1 or isinstance(sq_nq, bool) or not isinstance(sq_nq, int) or sq_nq < 2): print("Error: --sq-qmax must exceed 0.1 and --sq-nq must be at least 2.") sys.exit(1) print(f"\n S(q): {sq_method} method, {sq_weighting} weighting " f"(q_min = 2pi/L = {q_min:.2f} A^-1)") if sq_method == "ft": print(" Note: FT of g(r) is truncated at r = L/2 " f"({L_min/2:.1f} A); the FSDP is damped/shifted. " "Use --sq-method direct for the FSDP.") elif L_min < 15.0: print(" Warning: cell < 15 A — the FSDP region " "(~1-2 A^-1) is under-resolved at this box size; " "use ~450+ atom boxes for a quantitative S(q).") # sq_smooth: CLI > YAML analyse block > DEFAULT_SQ_SMOOTH; 0 = raw. from .analysis.rdf import DEFAULT_SQ_SMOOTH sq_smooth = args.sq_smooth if sq_smooth is None: sq_smooth = float(an_cfg.get("sq_smooth", DEFAULT_SQ_SMOOTH)) sq_partials = bool(args.sq_partials or an_cfg.get("sq_partials", False)) if sq_method == "ft": sq_result = sa.structure_factor(weighting=sq_weighting, qmax=sq_qmax, nq=sq_nq) if sq_partials: print(" Note: --sq-partials needs the direct method; " "partials skipped for --sq-method ft.") else: sq_result = sa.structure_factor_direct(weighting=sq_weighting, qmax=sq_qmax, nq=sq_nq, sigma_q=sq_smooth, partials=sq_partials) if sq_smooth > 0: print(f" S(q) re-binned with sigma_q = {sq_smooth:.2f} " f"1/A (n_per_bin-weighted); raw values kept in CSV") if sq_partials: import numpy as _np _q = _np.asarray(sq_result["q"], dtype=float) print(" Faber-Ziman partials S_ab(q), first peak below 3 A^-1:") for pair, s_ab in sq_result["partials"].items(): s_ab = _np.asarray(s_ab, dtype=float) m = (_q > q_min) & (_q < 3.0) & ~_np.isnan(s_ab) if m.any(): k = _np.argmax(_np.where(m, s_ab, -_np.inf)) print(f" {pair:<8s} q = {_q[k]:.2f} A^-1, " f"S = {s_ab[k]:.2f}") sq_result["calculation"] = { "method": sq_method, "weighting": sq_weighting, "qmax": sq_qmax, "nq": sq_nq, "sigma_q": sq_smooth if sq_method == "direct" else 0.0, "normalization": "Faber-Ziman", } if sq_method == "ft": from .analysis.rdf import _shared_rmax sq_result["calculation"]["rmax"] = _shared_rmax(sa.atoms_list, None) if convergence_enabled: _collect_convergence_summaries( convergence_descriptors, "sq.total", sq_result.get("uncertainty")) _collect_convergence_summaries( convergence_descriptors, "sq", sq_result.get("partials_uncertainty")) if plot_dir: from .analysis.plotting import plot_sq plot_sq(sq_result, output_dir=plot_dir, dpi=plot_kwargs.get("dpi", 300), save_pdf=plot_kwargs.get("save_pdf", False), weighting=sq_weighting, method=sq_method, show_title=plot_kwargs.get("show_title", False), pair_panels=plot_kwargs.get("pair_panels", False)) else: print(" (pass --save-plot DIR to write the S(q) PNG + CSV)") # T(r): CLI flag > YAML key, weighted like --sq if (args.tr or an_cfg.get("tr", False) or args.experiment_tr or an_cfg.get("experiment_tr")): tr_w = args.sq_weighting if not _typed("--sq-weighting") and "sq_weighting" in an_cfg: tr_w = an_cfg["sq_weighting"] qlo, qhi = (an_cfg.get("tr_qrange", args.tr_qrange) if not _typed("--tr-qrange") else args.tr_qrange) win = args.tr_window if _typed("--tr-window") else an_cfg.get("tr_window", args.tr_window) print(f"\n T(r): {tr_w} weighting, q = {qlo}-{qhi} 1/A, " f"{win} window") try: tr = sa.total_correlation(weighting=tr_w, qmin=qlo, qmax=qhi, window=None if win == "none" else win) except ValueError as exc: print(f" T(r) skipped: {exc}") else: tr["calculation"] = { "method": "direct", "weighting": tr_w, "qmin": qlo, "qmax": qhi, "window": win, "nq": 400, "nr": 600, "rmax": 10.0, "sigma_q": 0.05, "normalization": "T(r) = 4*pi*r*rho*g(r)", } if convergence_enabled: _collect_convergence_summaries( convergence_descriptors, "tr", tr.get("curve_uncertainty")) from .analysis.rdf import coordination_from_Tr, first_Tr_peak pk, r_lo, r_hi = first_Tr_peak(tr) if pk is None: tr_line = (f" T(r) ({tr_w}, q = {qlo}-{qhi} 1/A, {win} window): " f"no resolved first peak; " f"rho = {tr['rho']:.4f} atoms/A^3") else: n_first = coordination_from_Tr(tr, r_lo, r_hi) tr_line = (f" T(r) ({tr_w}, q = {qlo}-{qhi} 1/A, {win} window): " f"first peak at r = {pk:.2f} A " f"({r_lo:.2f}-{r_hi:.2f} A, weighted count {n_first:.2f}); " f"rho = {tr['rho']:.4f} atoms/A^3") print(tr_line) # the report file is already written by this point (rings does # the same), so append rather than adding to `text` if report_path: with open(report_path, "a") as rf: rf.write("\n" + tr_line + "\n") if args.tr_scan or an_cfg.get("tr_scan", False): from .analysis.rdf import scan_Tr_qmax, format_Tr_scan rows = scan_Tr_qmax(sa.atoms_list, weighting=tr_w, qmin=qlo) scan_text = format_Tr_scan(rows) print(scan_text) if report_path: with open(report_path, "a") as rf: rf.write("\n" + scan_text + "\n") if plot_dir: from .analysis.plotting import plot_tr plot_tr(tr, output_dir=plot_dir, dpi=plot_kwargs.get("dpi", 300), save_pdf=plot_kwargs.get("save_pdf", False), show_title=plot_kwargs.get("show_title", False)) else: print(" (pass --save-plot DIR to write the T(r) PNG + CSV)") from .scattering_cli import run_scattering_comparisons try: run_scattering_comparisons( sa, args, an_cfg, _typed, sq=sq_result, tr=tr, plot_dir=plot_dir, report_path=report_path, dpi=plot_kwargs.get("dpi", 300), save_pdf=plot_kwargs.get("save_pdf", False)) except (ValueError, OSError, KeyError) as exc: print(f"Error: scattering comparison: {exc}") sys.exit(1) # Ring statistics and Voronoi indices: CLI flag > YAML key. # (YAML: rings: true | "Ge-O"; voronoi: true | "Ge"; the older # ring_bond_pair / voronoi_element keys are still honoured.) rings_opt = args.rings if rings_opt is None: y = an_cfg.get("rings", an_cfg.get("ring_bond_pair")) if y is True: rings_opt = "auto" elif isinstance(y, (list, tuple)): rings_opt = "-".join(y) elif isinstance(y, str): rings_opt = y if rings_opt: pair = None if rings_opt == "auto" else tuple(rings_opt.split("-")) ring_max_size = (args.ring_max_size if args.ring_max_size is not None else an_cfg.get("ring_max_size", 12)) ring_cutoff = (args.ring_cutoff if args.ring_cutoff is not None else an_cfg.get("ring_cutoff")) try: rings = sa.ring_statistics(bond_pair=pair, cutoff=ring_cutoff, max_ring=ring_max_size) except ValueError as exc: print(f"Error: ring analysis: {exc}") sys.exit(1) if sa._file_list: rings["structure_files"] = [str(f) for f in sa._file_list] if convergence_enabled: _collect_convergence_summaries( convergence_descriptors, "rings", rings.get("uncertainty")) from .analysis.descriptors import format_descriptor, save_descriptor ring_text = format_descriptor("rings", rings) print(ring_text) if report_path: with open(report_path, "a") as rf: rf.write("\n" + ring_text + "\n") if plot_dir: save_descriptor("rings", rings, plot_dir, dpi=plot_kwargs.get("dpi", 300), save_pdf=plot_kwargs.get("save_pdf", False), show_title=plot_kwargs.get("show_title", False)) if args.connectivity or an_cfg.get("connectivity", False): from .analysis.structure import format_connectivity_report conn = sa.polyhedral_connectivity() if convergence_enabled: _collect_convergence_summaries( convergence_descriptors, "connectivity.edge_or_face_percent", conn.get("uncertainty")) for name, key in (("link_percent", "link_percent_uncertainty"), ("n_links", "n_links_uncertainty"), ("fraction_of_sites", "site_fraction_uncertainty"), ("fraction_of_structures", "structure_fraction_uncertainty"), ("face_percent", "face_percent_uncertainty")): _collect_convergence_summaries( convergence_descriptors, f"connectivity.{name}", conn.get(key)) for element, data in conn.get("per_species", {}).items(): _collect_convergence_summaries( convergence_descriptors, f"connectivity.{element}", data.get("uncertainty")) conn_text = format_connectivity_report(conn) print(conn_text) if report_path: with open(report_path, "a") as rf: rf.write("\n" + conn_text + "\n") if plot_dir and "error" not in conn: os.makedirs(plot_dir, exist_ok=True) with open(os.path.join(plot_dir, "analysis_connectivity.csv"), "w") as fh: fh.write("quantity,value\n") for k in ("corner", "edge", "face"): fh.write(f"link_percent_{k},{conn['link_percent'][k]:.4f}\n") fh.write(f"cation_edge_or_face_percent,{conn['cation_edge_or_face_percent']:.4f}\n") for i, v in enumerate(conn["per_structure_edge_percent"]): fh.write(f"structure_{i}_edge_or_face_percent,{v:.4f}\n") print(f" Saved: {os.path.join(plot_dir, 'analysis_connectivity.csv')}") vor_opt = args.voronoi if vor_opt is None: y = an_cfg.get("voronoi", an_cfg.get("voronoi_element")) if y is True: vor_opt = "all" elif isinstance(y, str): vor_opt = y if vor_opt: elem = None if vor_opt == "all" else vor_opt vor = sa.voronoi(element=elem) if convergence_enabled: _collect_convergence_summaries( convergence_descriptors, "voronoi", vor.get("uncertainty")) lines = [f"\n Voronoi indices <n3 n4 n5 n6> ({elem or 'all atoms'}): " f"{vor['total_atoms']} atoms, mean faces = {vor['mean_faces']:.2f}"] for idx, count, pct in vor["top_10"]: lines.append(f" {str(idx):16s} {count:6d} ({pct:5.1f}%)") vor_text = "\n".join(lines) print(vor_text) if report_path: with open(report_path, "a") as rf: rf.write("\n" + vor_text + "\n") if plot_dir: os.makedirs(plot_dir, exist_ok=True) with open(os.path.join(plot_dir, "analysis_voronoi.csv"), "w") as fh: fh.write("voronoi_index,count,percent\n") for idx, count, pct in vor["top_10"]: fh.write(f"\"{idx}\",{count},{pct:.4f}\n") print(f" Saved: {os.path.join(plot_dir, 'analysis_voronoi.csv')}") # Optional structural, mechanical and vibrational descriptors. from .analysis.descriptors import run_descriptor_analysis try: descriptor_results = run_descriptor_analysis( sa, args, an_cfg, override, plot_dir=plot_dir, report_path=report_path, plot_kwargs=plot_kwargs) except (ValueError, RuntimeError, NotImplementedError) as exc: print(f"Error: descriptor analysis: {exc}") sys.exit(1) if an_cfg.get("energy_ranking", False): er = sa.energy_ranking() if convergence_enabled: summaries = er.get("uncertainty", {}) _collect_convergence_summaries( convergence_descriptors, "energy.total", summaries.get("energy")) _collect_convergence_summaries( convergence_descriptors, "energy.per_atom", summaries.get("energy_per_atom")) if er.get("best_energy") is not None: print(f"\n Energy ranking:") print(f" Best: {er['best_energy']:.4f} eV/atom") print(f" Worst: {er['worst_energy']:.4f} eV/atom") print(f" Spread: {er['spread']:.4f} eV/atom") elif er.get("warning") or er.get("error"): print(f"\n Energy ranking: {er.get('warning') or er.get('error')}") if convergence_enabled: from ase.data import atomic_numbers from .analysis.convergence_output import ( format_convergence_report, save_convergence_report) for name, result in (descriptor_results or {}).items(): _collect_convergence_summaries( convergence_descriptors, name, result.get("uncertainty")) rdf_names = {name for name in tolerances if name.startswith("rdf.")} rdf_names.update(f"rdf.{pair}" for pair in (an_cfg.get("rdf_pairs") or [])) if args.total_rdf or an_cfg.get("total_rdf", False): rdf_names.add("rdf.total") angle_names = {name for name in tolerances if name.startswith("angle_distribution.")} angle_names.update(f"angle_distribution.{triplet}" for triplet in (an_cfg.get("angle_triplets") or [])) try: for name in sorted(rdf_names): pair = name.removeprefix("rdf.") symbols = pair.split("-") if pair != "total" and (len(symbols) != 2 or any( symbol not in atomic_numbers for symbol in symbols)): raise ValueError( f"Unknown RDF descriptor '{name}'; expected rdf.total " "or rdf.Element-Element") result = sa.rdf( pair=None if pair == "total" else pair, rmax=an_cfg.get("rmax"), sigma=float(smearing), confidence=convergence_confidence, n_bootstrap=0) convergence_descriptors[name] = result["uncertainty"] for name in sorted(angle_names): triplet = name.removeprefix("angle_distribution.") symbols = triplet.split("-") if len(symbols) != 3 or any( symbol not in atomic_numbers for symbol in symbols): raise ValueError( f"Unknown angle descriptor '{name}'; expected " "angle_distribution.Element-Element-Element") distributions = sa.angle_distribution( triplet, confidence=convergence_confidence, n_bootstrap=0) if triplet in distributions: convergence_descriptors[name] = distributions[triplet]["uncertainty"] convergence = sa.convergence_report( tolerances, descriptors=convergence_descriptors, confidence=convergence_confidence, max_structures=convergence_max) except (ValueError, TypeError) as exc: print(f"Error: convergence analysis: {exc}") sys.exit(1) convergence_text = format_convergence_report(convergence) print(convergence_text) if report_path: with open(report_path, "a") as handle: handle.write("\n" + convergence_text + "\n") if plot_dir: save_convergence_report( convergence, plot_dir, dpi=plot_kwargs.get("dpi", 300), save_pdf=plot_kwargs.get("save_pdf", False)) # Validation against literature reference YAML ref_path = args.reference or an_cfg.get("reference") if ref_path: from .analysis.validate import (validate_against_reference, format_validation_report) import yaml with open(ref_path) as f: reference = yaml.safe_load(f) v_result = validate_against_reference(sa, reference) v_text = format_validation_report(v_result) print(v_text) if report_path: with open(report_path, "a") as rf: rf.write("\n" + v_text + "\n") if screening is not None: mark_screening_analysed(screening, screening["retained_indices"]) write_screening_outputs(screening, screening_prefix) screening_text = format_screening_report(screening) print(screening_text) if report_path: with open(report_path, "a", encoding="utf-8") as handle: handle.write("\n" + screening_text + "\n") return # ── Random generation mode ──────────────────────────────────────────────── if args.random_gen: from .pipeline.random_gen import _batch_random_unlocked from .utils.run_lock import run_lock from .utils import get_calculator # Read random_gen block from YAML config (if present) rg_cfg = override.get("random_gen", {}) # Composition: CLI > YAML > error if args.composition is not None: composition = _parse_composition(args.composition) elif "composition" in rg_cfg: composition = rg_cfg["composition"] else: print("Error: --composition is required for --random-gen mode.") print(" Examples:") print(' --composition "In2O3*16" (formula * N units = 80 atoms)') print(" --composition In=32,O=48 (explicit atom counts)") print(" Or in YAML: random_gen: { composition: {In: 32, O: 48} }") sys.exit(1) # n_structures: CLI > YAML > default (10) parser = _get_parser() n_structures = args.n_structures if not _typed("-n", "--n-structures") and "n_structures" in rg_cfg: n_structures = rg_cfg["n_structures"] # target_density: CLI > YAML > None target_density = args.target_density if target_density is None and "target_density" in rg_cfg: target_density = rg_cfg["target_density"] # density_scale: CLI > YAML > 1.0 (parser default) density_scale = args.density_scale if not _typed("--density-scale") \ and "density_scale" in rg_cfg: density_scale = rg_cfg["density_scale"] # output_format: CLI > YAML > default output_format = args.format if not _typed("--format") and "output_format" in rg_cfg: output_format = rg_cfg["output_format"] # minsep: CLI > YAML > None (auto-generated) minsep = None if args.minsep is not None: minsep = _parse_minsep(args.minsep) elif "minsep" in rg_cfg: minsep = rg_cfg["minsep"] # target_cn: --no-sc > CLI > YAML > auto target_cn = None if args.no_sc: target_cn = {} # empty dict disables coordination-aware placement elif args.target_cn is not None: target_cn = _parse_target_cn(args.target_cn) elif "target_cn" in rg_cfg: target_cn = rg_cfg["target_cn"] # dmax: CLI > YAML > None (auto-generated if target_cn set) dmax_dict = None if args.dmax is not None: dmax_dict = _parse_dmax(args.dmax) elif "dmax" in rg_cfg: dmax_dict = rg_cfg["dmax"] # cn_tolerance: CLI > YAML > auto (from composition) cn_tolerance = args.cn_tolerance if cn_tolerance is None and "cn_tolerance" in rg_cfg: cn_tolerance = rg_cfg["cn_tolerance"] # If still None, generate_random will auto-detect from composition # dmax_factor: CLI > YAML > default (1.5) if args.dmax_factor == 1.5 and "dmax_factor" in rg_cfg: args.dmax_factor = rg_cfg["dmax_factor"] # cell_filter: CLI > YAML random_gen > explicit YAML opt > "cubic" # Random gen produces cubic cells, so default to cubic (not # FrechetCellFilter). A --relax is an optimisation, so an explicit # cell_filter under the YAML ``opt:`` block is honoured too -- the # shipped example_classical.yaml sets ``opt: cell_filter: none`` # (classical potentials have no stress), and ignoring it made the # stress guard tell users to set a value their YAML already had. ff_default = parser.get_default("cell_filter") cell_filter = args.cell_filter if cell_filter == ff_default: opt_cf = None if isinstance(override, dict) and isinstance(override.get("opt"), dict): opt_cf = override["opt"].get("cell_filter") if "cell_filter" in rg_cfg: cell_filter = rg_cfg["cell_filter"] elif opt_cf not in (None, ff_default): # Explicitly set by the user (the no-YAML path puts the # parser default here, which is not a user choice). cell_filter = opt_cf else: cell_filter = "cubic" # relax: CLI flag or YAML (default: no relaxation) do_relax = args.relax if not do_relax and "relax" in rg_cfg: do_relax = rg_cfg["relax"] # The merged opt block already gives explicit CLI flags precedence # over YAML. Resolve once so both relaxation engines use the same # settings, retaining the random-gen defaults for omitted values. opt_cfg = override.get("opt", {}) or {} fmax = opt_cfg.get("fmax", 0.05) max_relax_steps = opt_cfg.get("max_steps", args.opt_steps) optimizer = opt_cfg.get("optimizer", args.optimizer) use_torchsim = do_relax and override.get("engine", "ase") == "torchsim" relax_settings = None if do_relax: relax_settings = { "model": override.get("model", args.model), "model_path": override.get("model_path", args.model_path), "device": override.get("device", args.device), "default_dtype": override.get("default_dtype", args.default_dtype), "classical_params": override.get("classical_params"), "engine": "torchsim" if use_torchsim else "ase", } if use_torchsim: relax_settings["torchsim"] = { "pressure_tol_gpa": opt_cfg.get("pressure_tol_gpa", 0.02), "output_format": opt_cfg.get("output_format", "xyz"), "batch_size": args.batch_size or opt_cfg.get("batch_size") or "auto", } if until_options is not None: from .pipeline.until_converged import run_until_converged sequential_override = dict(override) sequential_override["opt"] = { **opt_cfg, "fmax": fmax, "max_steps": max_relax_steps, "optimizer": optimizer, "cell_filter": cell_filter, } if args.batch_size is not None: sequential_override["opt"]["batch_size"] = ( int(args.batch_size) if str(args.batch_size).isdigit() else args.batch_size) generation = { "target_density": target_density, "density_scale": density_scale, "minsep": minsep, "max_attempts_per_atom": args.max_attempts, "target_cn": target_cn, "dmax": dmax_dict, "cn_tolerance": cn_tolerance, "dmax_factor": args.dmax_factor, "repair_iters": args.repair_iters, "retry_mode": args.retry_mode, "seed": (args.seed if args.seed is not None else rg_cfg.get("seed", override.get("seed"))), } try: result = run_until_converged( composition, args.work_dir, **until_options, generation=generation, cfg_override=sequential_override, resume=args.resume) except (ValueError, RuntimeError, OSError, ImportError) as exc: print(f"Error: sequential convergence: {exc}") sys.exit(1) print(f"Sequential convergence: {result['status']} " f"after {result['n_structures']} structures.") if result["status"] == "max_structures_reached": sys.exit(2) if result["status"] != "converged": sys.exit(1) return # Placement and the optional separate torch-sim phase share ownership # of the whole output tree, including resume metadata validation. with run_lock(args.work_dir): calc = None if do_relax and not use_torchsim: calc = get_calculator( **_classical_kwargs(override), model=override.get("model", args.model), device=override.get("device", args.device), model_path=override.get("model_path", args.model_path), default_dtype=override.get("default_dtype", args.default_dtype), ) files = _batch_random_unlocked( composition=composition, n_structures=n_structures, output_dir=args.work_dir, output_format=output_format, relax=do_relax and not use_torchsim, calc=calc, resume_settings=relax_settings, safety=override.get("safety"), repulsive_core=override.get("repulsive_core"), fmax=fmax, max_relax_steps=max_relax_steps, optimizer=optimizer, cell_filter=cell_filter, target_density=target_density, density_scale=density_scale, minsep=minsep, max_attempts_per_atom=args.max_attempts, target_cn=target_cn, dmax=dmax_dict, cn_tolerance=cn_tolerance, dmax_factor=args.dmax_factor, repair_iters=args.repair_iters, retry_mode=args.retry_mode, indices=args.indices, seed=(args.seed if args.seed is not None else rg_cfg.get("seed", override.get("seed"))), resume=args.resume, ) if use_torchsim: from .pipeline.opt_cell import batch_optimize batch_optimize(input_dir=os.path.join(args.work_dir, "random_initial"), output_dir=os.path.join(args.work_dir, "random_opt"), cfg_override=override, calc=None, engine="torchsim", fmax=fmax, max_steps=max_relax_steps, cell_filter=cell_filter, optimizer=optimizer, resume=args.resume, batch_size=(int(args.batch_size) if args.batch_size and str(args.batch_size).isdigit() else args.batch_size), indices=args.indices) return # ── Batch optimisation mode ────────────────────────────────────────────── if args.batch_opt: if args.input_dir is None: print("Error: --input-dir is required for --batch-opt mode.") print(" Example: amorphgen --batch-opt --input-dir random_Ga2O3/random_initial/") sys.exit(1) from .pipeline.opt_cell import batch_optimize from .utils import get_calculator # batch_optimize returns [] only when nothing matched (it has printed # why); exit non-zero so a script or job chain doesn't carry on. if override.get("engine", "ase") == "torchsim": paths = batch_optimize(input_dir=args.input_dir, output_dir=args.work_dir, cfg_override=override, calc=None, engine="torchsim", resume=args.resume, batch_size=(int(args.batch_size) if args.batch_size and str(args.batch_size).isdigit() else args.batch_size), pattern=args.pattern, indices=args.indices) if not paths: sys.exit(1) return calc = get_calculator( **_classical_kwargs(override), model=override.get("model", args.model), device=override.get("device", args.device), model_path=override.get("model_path", args.model_path), default_dtype=override.get("default_dtype", args.default_dtype), ) paths = batch_optimize( input_dir=args.input_dir, output_dir=args.work_dir, cfg_override=override, calc=calc, pattern=args.pattern, indices=args.indices, ) if not paths: sys.exit(1) return # ── Batch quench mode ───────────────────────────────────────────────────── if args.batch_quench: from .pipeline import batch_quench from .utils import get_calculator import glob snap_source = args.snapshot_dir # Polymorphic input: if --snapshot-dir is a single trajectory file # (e.g. shared/stage4_eq.xyz from a Stage 4 run), extract --n-runs # selected frames into a 'snapshots_extracted/' subdir of # the work directory and use that as the snapshot source. # Putting the extracted dir inside work_dir avoids race conditions # when several array tasks point at files in the same source dir. extracted_files = None if os.path.isfile(snap_source): from .utils import extract_snapshots os.makedirs(args.work_dir, exist_ok=True) extracted_dir = os.path.join(args.work_dir, "snapshots_extracted") print(f"[batch-quench] '{snap_source}' is a file — extracting " f"up to {args.n_runs} snapshots to {extracted_dir}/") extracted_files = extract_snapshots( snap_source, n_snapshots=args.n_runs, output_dir=extracted_dir, **_snapshot_sampling_kwargs( args, override, report_path=os.path.join(args.work_dir, "snapshot_sampling.json"), quench_dir=args.work_dir)) snap_source = extracted_dir elif args.select == "decorrelated": raise ValueError("--select decorrelated requires a trajectory file as --snapshot-dir.") # Accept any ASE-readable structure format. extxyz/xyz are the # original use case (snapshots from MD trajectory); vasp/cif/POSCAR # let users feed in pre-relaxed structures from --random-gen or DFT. snap_files: list = extracted_files or [] if extracted_files is None: for pattern in ("*.xyz", "*.extxyz", "*.vasp", "*.cif", "POSCAR*"): snap_files = sorted(glob.glob( os.path.join(snap_source, pattern))) if snap_files: break if not snap_files: from .pipeline.random_gen import random_gen_dir_hint print(f"Error: no snapshot files found in {snap_source}/ " f"(looked for *.xyz, *.extxyz, *.vasp, *.cif, POSCAR*)") hint = random_gen_dir_hint(snap_source) if hint: print(hint) sys.exit(1) calc = get_calculator( **_classical_kwargs(override), model=override.get("model", args.model), device=override.get("device", args.device), model_path=override.get("model_path", args.model_path), default_dtype=override.get("default_dtype", args.default_dtype), ) batch_quench.run( snapshot_files=snap_files, n_runs=args.n_runs, select="uniform" if extracted_files is not None else (args.select or "uniform"), cfg_override=override, work_dir=args.work_dir, stages=args.batch_stages, calc=calc, resume=args.resume, ) return # ── Standard pipeline mode ──────────────────────────────────────────────── if args.input_file is None: print("Error: input_file is required for melt-quench pipeline mode.") print(" Usage: amorphgen POSCAR [--model NAME] [--stages 1 2 3 4 5 6 7]") print(" Other modes that don't need an input file:") print(' amorphgen --random-gen --composition "SiO2*16"') print(" amorphgen --batch-opt --input-dir structures/") print(" amorphgen --analyse structures/") print(" Run 'amorphgen --help' for full options.") sys.exit(1) # For optimisation-only (--stages 1 or --stages 7), call opt_cell directly # so that output filenames are derived from the input file name. if args.stages == [1] or args.stages == [7]: from .pipeline.opt_cell import run as opt_run from .utils import get_calculator os.makedirs(args.work_dir, exist_ok=True) orig_dir = os.getcwd() os.chdir(args.work_dir) try: input_path = os.path.join(orig_dir, args.input_file) calc = get_calculator( **_classical_kwargs(override), model=override.get("model", args.model), device=override.get("device", args.device), model_path=override.get("model_path", args.model_path), default_dtype=override.get("default_dtype", args.default_dtype), ) stage_key = "opt" if args.stages == [1] else "final_opt" opt_run(input_path, cfg_override=override, calc=calc, stage_key=stage_key) finally: os.chdir(orig_dir) return from .pipeline.run_pipeline import MeltQuenchPipeline pipe = MeltQuenchPipeline( input_file=args.input_file, work_dir=args.work_dir, cfg_override=override, ) pipe.run(stages=args.stages, resume=args.resume) if __name__ == "__main__": main()