Source code for diffBloch.app.cli

"""Thin command-line entry point.

This is the orchestration / SLURM boundary: a workflow engine (e.g. Dagster or Prefect) shells out
to it, or a SLURM job runs it. Kept deliberately thin — it delegates to the library and holds no
science.
"""

from __future__ import annotations

import argparse
import logging
import sys
from pathlib import Path

import yaml
from pydantic import ValidationError

from diffBloch import __version__
from diffBloch.app.loggers import ConsoleLogger, CSVLogger, print_summary_box
from diffBloch.app.loggers.summary import SummaryLogger
from diffBloch.app.program import (
    converge_experiment,
    preprocess_experiment,
    refine_experiment,
    run_experiment,
)
from diffBloch.config import load_config, write_experiment_lock
from diffBloch.observability import Logger, MultiLogger


def _add_stage_flags(parser: argparse.ArgumentParser) -> None:
    """Add the flags shared by ``infer``, ``preprocess``, and ``refine`` (same preprocess surface)."""
    parser.add_argument("experiment_directory", help="Path to the experiment directory")
    parser.add_argument(
        "--csv", metavar="PATH", help="append per-rotation observations to a long-format CSV log"
    )
    parser.add_argument(
        "--refresh",
        action="store_true",
        help="ignore any existing preprocess checkpoints and recompute (regenerates each "
        "dataset's plan.<stem>.npz/.lock)",
    )
    parser.add_argument(
        "--no-checkpoint",
        action="store_true",
        help="neither read nor write the preprocess checkpoint (leave the experiment dir alone)",
    )
    parser.add_argument(
        "--device",
        metavar="DEVICE",
        default="cuda",
        help="run the forward solve on this torch device (default: cuda; use 'cpu' to override)",
    )
    parser.add_argument(
        "--workers",
        metavar="N",
        type=int,
        default=1,
        help="fan orientation-plan builds and per-rotation searches over N threads (default 1); "
        "cap host threads to 1 (OMP_NUM_THREADS/MKL_NUM_THREADS/TORCH_NUM_THREADS, or "
        "torch.set_num_threads(1)) or the node-sized BLAS/torch pools oversubscribe the cores",
    )
    parser.add_argument(
        "--max-batch",
        metavar="N",
        type=int,
        default=None,
        help="cap the matrix_exp propagator block to N (N,N) operators (memory only, matches the "
        "unbounded solve to machine precision); default derives a memory-safe block per beam "
        "count. Raise to fill a larger GPU, e.g. 1024 on a high-memory accelerator",
    )
    parser.add_argument(
        "--plot-thickness",
        action="store_true",
        help="save one wR2-vs-thickness PNG per rotation from the thickness grid search; ORs with "
        "preprocess.thickness.plot in experiment.yaml, so "
        "either turns it on. Defaults to '<inputs.structure's directory>/thickness_optim', "
        "override with --plot-thickness-dir",
    )
    parser.add_argument(
        "--plot-thickness-dir",
        metavar="PATH",
        default=None,
        help="override the output directory for thickness plots; only takes effect when plotting "
        "is on (--plot-thickness or preprocess.thickness.plot)",
    )


[docs] def main(argv: list[str] | None = None) -> int: """Parse arguments and dispatch to a subcommand. Returns a process exit code.""" parser = argparse.ArgumentParser( prog="diffbloch", description="Differentiable Bloch-wave electron-diffraction structure refinement", formatter_class=argparse.RawDescriptionHelpFormatter, epilog=( "a typical workflow:\n" " diffbloch validate experiment/experiment.yaml\n" " diffbloch converge experiment/\n" " diffbloch preprocess experiment/\n" " diffbloch refine experiment/\n" "\n" "Each command has its own flags; see e.g. 'diffbloch refine --help'. " "reproducibility/experiment.lock is created automatically on first run; " "use 'diffbloch lock-experiment --force experiment/' to refresh it after inputs " "intentionally change." ), ) parser.add_argument("--version", action="version", version=f"diffbloch {__version__}") parser.add_argument("--debug", action="store_true", help="show full tracebacks on error") sub = parser.add_subparsers(dest="command") # Registered in typical workflow order so the --help listing reads as the pipeline. p_validate = sub.add_parser("validate", help="Validate an experiment.yaml and report") p_validate.add_argument("config", help="Path to experiment.yaml") p_lock_experiment = sub.add_parser( "lock-experiment", help="Create reproducibility/experiment.lock from the current input files (other commands " "create it automatically on first run; existing locks require --force)", ) p_lock_experiment.add_argument( "--force", action="store_true", help="rewrite an existing experiment.lock after an intentional input change; this " "invalidates existing plan and refinement locks", ) p_lock_experiment.add_argument("experiment_directory", help="Path to the experiment directory") p_converge = sub.add_parser( "converge", help="Test convergence of g_max, sg_max, and rocking-curve tilt steps" ) p_converge.add_argument("experiment_directory", help="Path to the experiment directory") p_converge.add_argument( "--device", metavar="DEVICE", default="cuda", help="run the convergence simulations on this torch device (default: cuda)", ) p_converge.add_argument( "--orientations", metavar="N", type=int, default=1, help="use the first N orientations for convergence testing (default: 1)", ) p_converge.add_argument( "--csv", metavar="PATH", help="append per-trial observations to a long-format CSV log" ) p_preprocess = sub.add_parser( "preprocess", help="Settle the coupled preprocess Plan and write the checkpoint (no score)" ) _add_stage_flags(p_preprocess) p_infer = sub.add_parser("infer", help="Score every rotation of an experiment") _add_stage_flags(p_infer) p_refine = sub.add_parser( "refine", help="Gradient-refine the structure against the data (reuses the checkpoint)" ) _add_stage_flags(p_refine) p_refine.add_argument( "--verbose-refinement", action="store_true", help="also report per-rotation wR2/R_obs/diffraction-loss every step, not just the epoch " "mean (n_orientations x louder; a diagnosis tool, off by default)", ) p_refine.add_argument( "--profile", action="store_true", help="log per-phase wall time (structure factors, each rotation's solve, backward, " "optimizer step) via stdlib diagnostics logging; forces a CUDA sync per measured block " "(real overhead) so use only to diagnose one run, not routinely", ) p_refine.add_argument( "--no-checkpoint-activations", action="store_true", help="do not gradient-checkpoint each per-orientation/per-segment solve; trades a full " "forward recompute on backward for higher peak memory (gradients are unaffected either " "way) -- try this if backward is much slower than forward and you have memory headroom", ) args = parser.parse_args(argv) if args.command == "validate": # Expected user errors get a concise stderr message + nonzero exit; tracebacks are reserved # for --debug. This is the CLI/orchestration boundary, not a place to leak Python internals. try: cfg = load_config(args.config) except (FileNotFoundError, yaml.YAMLError, ValidationError) as exc: if args.debug: raise print(f"error: {args.config}: {exc}", file=sys.stderr) return 1 print(f"OK: experiment '{cfg.name}' validated.") return 0 if args.command == "lock-experiment": lock_path = ( Path(args.experiment_directory) / "reproducibility" / "experiment.lock" ).resolve() replacing = lock_path.exists() try: experiment_lock = write_experiment_lock(args.experiment_directory, force=args.force) except (FileExistsError, FileNotFoundError, ValidationError, yaml.YAMLError) as exc: if args.debug: raise print(f"error: {exc}", file=sys.stderr) return 1 refs = ( [entry.ref for entry in experiment_lock.experimental_data] if isinstance(experiment_lock.experimental_data, list) else [experiment_lock.experimental_data.ref] ) print(f"{'rewrote' if replacing else 'wrote'} {lock_path}") if replacing: print( "warning: if this accepts changed input bytes, existing plan and refinement locks " "are invalid and must be regenerated" ) print(f" - {'Structure':<20} {experiment_lock.structure.ref}") for ref in refs: print(f" - {'Experimental data':<20} {ref}") return 0 if args.command == "infer": logging.basicConfig( level=logging.INFO, format="%(asctime)s %(message)s", datefmt="%H:%M:%S" ) try: run_experiment( args.experiment_directory, logger=_build_logger(csv=args.csv), checkpoint=not args.no_checkpoint, refresh=args.refresh, device=args.device, workers=args.workers, max_batch=args.max_batch, plot_thickness=args.plot_thickness, plot_thickness_dir=args.plot_thickness_dir, ) except (FileNotFoundError, ValueError, ValidationError, yaml.YAMLError) as exc: if args.debug: raise print(f"error: {exc}", file=sys.stderr) return 1 return 0 if args.command == "preprocess": logging.basicConfig( level=logging.INFO, format="%(asctime)s %(message)s", datefmt="%H:%M:%S" ) try: plan = preprocess_experiment( args.experiment_directory, logger=_build_logger(csv=args.csv), checkpoint=not args.no_checkpoint, refresh=args.refresh, device=args.device, workers=args.workers, max_batch=args.max_batch, plot_thickness=args.plot_thickness, plot_thickness_dir=args.plot_thickness_dir, ) except (FileNotFoundError, ValueError, ValidationError, yaml.YAMLError) as exc: if args.debug: raise print(f"error: {exc}", file=sys.stderr) return 1 # ConsoleLogger printed "PREPROCESS COMPLETE" the moment preprocessing settled -- the same # box a refine/infer run gets, from the same sink. print() print("Pipeline") for index, record in enumerate(plan.provenance, start=1): print(f" {index:>2}. {record.name.replace('_', ' ').title()}") print() # List the checkpoint pairs actually on disk (one per dataset; none under --no-checkpoint). reproducibility_dir = Path(args.experiment_directory) / "reproducibility" checkpoints = ( sorted(reproducibility_dir.glob("plan.*.npz")) if reproducibility_dir.is_dir() else [] ) if checkpoints: print("Output files") for npz in checkpoints: print(f" • {'Plan':<20} {npz.resolve()}") lock = npz.with_suffix(".lock") if lock.exists(): print(f" • {'Plan Lock':<20} {lock.resolve()}") return 0 if args.command == "refine": logging.basicConfig( level=logging.INFO, format="%(asctime)s %(message)s", datefmt="%H:%M:%S" ) try: # The written summary is one more sink on the run's event stream, chosen here beside # the console/CSV ones rather than by refine_experiment: an API caller composes it (or # not) for themselves instead of having a file appear as a side effect of refining. report_path = (Path(args.experiment_directory) / "refinement_report.txt").resolve() refine_sinks: tuple[Logger, ...] = ( _build_logger(csv=args.csv), SummaryLogger(report_path), ) refine_experiment( args.experiment_directory, logger=MultiLogger(refine_sinks), checkpoint=not args.no_checkpoint, refresh=args.refresh, device=args.device, workers=args.workers, max_batch=args.max_batch, verbose=args.verbose_refinement, profile=args.profile, checkpoint_activations=not args.no_checkpoint_activations, plot_thickness=args.plot_thickness, plot_thickness_dir=args.plot_thickness_dir, ) except (FileNotFoundError, ValueError, ValidationError, yaml.YAMLError) as exc: if args.debug: raise print(f"error: {exc}", file=sys.stderr) return 1 # ConsoleLogger printed "REFINEMENT COMPLETE" and the artifact list off the run's terminal # event. The report is the one output this file chose the location of, so it is also the # one line this file still prints. print(f" • {'Refinement Report':<20} {report_path}") return 0 if args.command == "converge": logging.basicConfig( level=logging.INFO, format="%(asctime)s %(message)s", datefmt="%H:%M:%S" ) try: settled = converge_experiment( args.experiment_directory, logger=_build_logger(csv=args.csv), device=args.device, n_orientations=args.orientations, ) except (FileNotFoundError, ValueError, ValidationError, yaml.YAMLError) as exc: if args.debug: raise print(f"error: {exc}", file=sys.stderr) return 1 print() print_summary_box( "CONVERGENCE COMPLETE", ( ("g_max", f"{settled.g_max:g}"), ("sg_max", f"{settled.sg_max:g}"), ("Tilt steps", str(settled.tilt_steps)), ), ) return 0 parser.print_help() return 0
def _build_logger(*, csv: str | None, per_rotation: bool = True) -> Logger: """Combine the observation sinks every command gets: the console, plus a CSV log if asked. The single place a CLI run's sinks are assembled, so a newly observable phase is rendered by teaching :class:`~diffBloch.app.loggers.ConsoleLogger` its event rather than by wiring another sink into each command. ``per_rotation`` opts the console into the settled per-rotation stream. """ console = ConsoleLogger(per_rotation=per_rotation) if csv is None: return console return MultiLogger((console, CSVLogger(Path(csv)))) if __name__ == "__main__": sys.exit(main())