"""The default experiment runner the ``run infer`` CLI exposes.
:func:`run_experiment` encodes the default recipe as one ordered ``Plan -> Plan`` pipeline:
plan-shaping (``build_orientation_plans``, selecting coupled SOLVE beams from ``g_max``/``sg_max``
and building rocking geometry plus its reduction) followed by
the config-enabled parameter fitting stages (orientation, then thickness), then ``run_inference``
evaluates it -- so a caller with an experiment directory gets a full result in one call. It is a
*convenience*, not the only path: every step is ordinary public API, so a Python user who wants a
different composition composes their own ``pipeline([...])`` with ``from_experiment`` +
``run_inference`` directly. The CLI stays thin by delegating here and holds no science.
**Checkpoint / resume.** The preprocess (the coupled fit especially) is the run's expensive phase,
so each dataset's settled ``Plan`` is checkpointed to ``<experiment_dir>/reproducibility/``
(``plan.<stem>.npz`` + ``plan.<stem>.lock`` per ``inputs.exp_data`` entry, see
:func:`_reproducibility_dir` and :func:`~diffBloch.config.schema.dataset_checkpoint_stem`) and
reused when still valid. The lock binds the checkpoint to *per-dataset* axes -- the structure and
that dataset's input bytes, the experiment's authoritative PETS cell, the dataset-scoped config
projection, its file-local ignored rotations, the software version, and the *recipe* (the Plan's
own provenance) -- so it is reused only when all match, and never restaled by changes to *other*
datasets in a pooled experiment unless the first dataset's authoritative cell changes. The match is
a longest-prefix over the recipe: an identical recipe is a full **reuse**; a recipe that *appends*
steps **resumes** from the snapshot and runs only the suffix (the append-only increment). Any run
that computes -- fresh, stale-recompute, ``refresh``, or resume -- regenerates the ``.npz`` +
``.lock`` pair, so the lock always describes the on-disk checkpoint. With ``inputs.multi_dataset``
the settled per-dataset plans are pooled in memory (:func:`~diffBloch.preprocess.pool`) just before
refinement; the pooled plan itself is never checkpointed. Diagnostics about load/resume go through
stdlib ``logging`` (not the domain-observation ``logger``).
"""
from __future__ import annotations
import logging
from dataclasses import dataclass, replace
from pathlib import Path
from typing import cast
import gemmi
import numpy as np
import torch
from pydantic import ValidationError
from diffBloch.config import (
CellParameters,
ExperimentConfig,
InputLock,
PreprocessLock,
RecipeStep,
RefinementLock,
artifact_hash_for,
code_version,
dataset_config_digest,
input_lock_for,
load_experiment,
preprocess_lock_status,
read_preprocess_lock,
refinement_config_digest,
sha256_file,
write_preprocess_lock,
write_refinement_lock,
)
from diffBloch.config.schema import dataset_checkpoint_stem
from diffBloch.core.crystal import cell_matrix_from_parameters, cell_volume
from diffBloch.core.products import MosaicSmoothed
from diffBloch.engine import (
ApparentThicknessNN,
ModelRefinementResult,
RefinementEngine,
ThicknessBounds,
build_refinement_model,
build_refinement_problem,
run_refinement_model,
)
from diffBloch.engine.plan import OrientationPlanLike
from diffBloch.io import (
ExperimentalRecord,
StructureRecord,
read_experimental_data_with_diagnostics,
read_structure,
read_structure_with_diagnostics,
)
from diffBloch.observability import (
NULL_LOGGER,
DeviceSelected,
Logger,
MultiLogger,
PreprocessCompleted,
RefinedRotationMetrics,
RefinementOutputsWritten,
)
from diffBloch.params import Device, constrain
from diffBloch.preprocess import (
OPAQUE,
ConvergenceTest,
ConvergenceTolerance,
Plan,
PlanStep,
build_orientation_plans,
fork,
optimize_orientation,
optimize_thickness,
pipeline,
plan_is_readable,
pool,
read_plan,
resolve_recipe,
run_inference,
select_beams,
setup_datasets,
step_records,
validation_mask,
write_plan,
)
from diffBloch.preprocess.driver import ConvergenceState, run_convergence
from diffBloch.preprocess.experiment import RefinementSetup
from diffBloch.preprocess.inference import InferenceResult
from diffBloch.preprocess.plan import unique_hkl_count
from diffBloch.preprocess.scoring import build_engine
from diffBloch.specs import (
ApparentThicknessNetwork,
IntegrationGeometry,
ScoredHklSelection,
TrialCoupling,
)
__all__ = [
"converge_experiment",
"preprocess_experiment",
"refine_experiment",
"run_experiment",
]
_log = logging.getLogger(__name__)
_REFINEMENT_LOCK = "refinement.lock"
_REPRODUCIBILITY_DIRNAME = "reproducibility"
def _plan_npz_name(dataset_ref: str) -> str:
"""This dataset's checkpoint filename, ``plan.<stem>.npz`` -- keyed on the file, not its
position in ``inputs.exp_data``, so reordering or inserting datasets moves nothing."""
return f"plan.{dataset_checkpoint_stem(dataset_ref)}.npz"
def _plan_lock_name(dataset_ref: str) -> str:
"""This dataset's checkpoint-lock filename, ``plan.<stem>.lock`` (see :func:`_plan_npz_name`)."""
return f"plan.{dataset_checkpoint_stem(dataset_ref)}.lock"
def _exp_data_refs(cfg: ExperimentConfig) -> tuple[str, ...]:
"""The experiment's dataset refs in pooled order -- ``inputs.exp_data`` as always-a-tuple."""
if isinstance(cfg.inputs.exp_data, list):
return tuple(cfg.inputs.exp_data)
return (cfg.inputs.exp_data,)
def _read_structure(
root: Path, cfg: ExperimentConfig, *, logger: Logger = NULL_LOGGER
) -> StructureRecord:
path = root / cfg.inputs.structure
parsed = read_structure_with_diagnostics(path, load_hydrogens=cfg.inputs.load_hydrogens)
for diagnostic in parsed.diagnostics:
logger.report(diagnostic)
return parsed.record
def _read_experimental_data(
root: Path, cfg: ExperimentConfig, *, logger: Logger = NULL_LOGGER
) -> tuple[ExperimentalRecord, ...]:
"""Read every ``inputs.exp_data`` file, single or pooled alike -- always a tuple.
Order matches ``inputs.exp_data``, which is also the order
:func:`~diffBloch.preprocess.setup_datasets` seeds datasets in and
:func:`~diffBloch.preprocess.pool` pools rotation indices in -- callers that need a flat array
aligned to the pooled ``rotation_index`` space can rely on that order without recomputing it.
"""
records: list[ExperimentalRecord] = []
for ref in _exp_data_refs(cfg):
parsed = read_experimental_data_with_diagnostics(root / ref)
for diagnostic in parsed.diagnostics:
logger.report(diagnostic)
records.append(parsed.record)
return tuple(records)
def _reproducibility_dir(root: Path) -> Path:
"""The ``root/reproducibility`` subdirectory for generated bookkeeping, created on first use.
Keeps ``root`` itself down to the handful of files someone actually reads (``experiment.yaml``,
the input structure/experimental-data files, ``refined_structure.cif``,
``refinement_report.txt``, the per-dataset thickness-NN shape plots) -- everything else, including
``experiment.lock`` (input-identity verification, not a generated output, but bookkeeping all
the same), lives here instead: preprocess/refinement checkpoints and their locks, and the raw
parameter/component ``.npz`` snapshots.
"""
directory = root / _REPRODUCIBILITY_DIRNAME
directory.mkdir(parents=True, exist_ok=True)
return directory
def _select_device(device: Device | None, *, logger: Logger = NULL_LOGGER) -> Device:
"""Resolve the execution device, falling back from CUDA to CPU when this host lacks CUDA."""
requested = "cpu" if device is None else str(device)
cuda_available = torch.cuda.is_available()
selected: Device = "cpu" if device is None else device
if torch.device(requested).type == "cuda" and not cuda_available:
selected = "cpu"
selected_name = str(selected)
logger.report(
DeviceSelected(
requested=requested,
selected=selected_name,
cuda_available=cuda_available,
)
)
return selected
# Above this unit-cell volume the coupled orientation search runs its coarse pass with the
# gather integrity checks skipped (the large-cell fast path); at or below it the search stays the
# exact, fully-validated path. The eigensolve is O(N^3) in the beam count
# and N grows with cell volume, so skipping the O(N^2) integrity checks only pays off for large
# cells; a small cell (quartz ~113 A^3) fits in seconds and gains nothing from skipping them.
# A deliberately wide heuristic gap separates the known regimes -- quartz ~113 vs zeolites
# ~1861 A^3 -- so classification is unambiguous; it is not a sharp physical boundary. Kept a code
# constant, NOT a config field: a config field would enter config_digest and restale the committed
# quartz checkpoint. It only selects whether the *search* validates its gathers, and the terminal
# always re-scores the found orientation, so the threshold never affects the reported score's
# fidelity.
_LARGE_CELL_THRESHOLD_A3 = 1000.0
[docs]
def converge_experiment(
experiment_dir: str | Path,
*,
logger: Logger = NULL_LOGGER,
device: Device = "cuda",
n_orientations: int = 1,
) -> ConvergenceState:
"""Run the standard numerical-convergence test for an experiment.
Starting ``g_max`` prefers each PETS file's own ``dstarmax`` (its processing-resolution cutoff,
:attr:`~diffBloch.io.record.ExperimentalRecord.dstar_max`) when present; a PETS version that
doesn't record it (or a file where the tag is otherwise absent) falls back to the experiment's
configured ``blochwave.g_max``, the previous manual-only behaviour. ``sg_max`` and rocking-curve
tilt steps still start from the configured simulation settings. The sweep then follows the
defaults owned by :class:`ConvergenceTest` and :class:`ConvergenceTolerance`, returning the
smallest settled values found.
With ``inputs.multi_dataset`` the sweep runs **per dataset** -- each file with its own
integration geometry, its own ``dstar_max``-or-fallback starting point, and its first
``n_orientations`` rotations -- and the returned state is the elementwise maximum over the
per-dataset settled states: the tightest single setting adequate for every pooled file. The
per-dataset settled values remain visible through the convergence logger events.
"""
root = Path(experiment_dir)
device = _select_device(device, logger=logger)
cfg, _lock = load_experiment(root)
structure = _read_structure(root, cfg, logger=logger)
records = _read_experimental_data(root, cfg, logger=logger)
refinement_setup, datasets = setup_datasets(structure, records, cfg)
refinement = replace(refinement_setup, params=refinement_setup.params.to(device))
if n_orientations < 1:
raise ValueError("n_orientations must be >= 1")
simulation = cfg.blochwave.to_policy()
settled_states: list[ConvergenceState] = []
for dataset_index, (record, dataset) in enumerate(zip(records, datasets, strict=True)):
if n_orientations > len(dataset.plan.orientations):
raise ValueError(
f"n_orientations={n_orientations} exceeds dataset {dataset_index}'s "
f"{len(dataset.plan.orientations)} orientations"
)
selected = replace(
dataset.plan,
orientations=dataset.plan.orientations[:n_orientations],
)
rocking = cfg.blochwave.to_rocking_curve(dataset.integration)
# Per dataset: a file without a dstar_max tag anchors at the configured g_max rather than
# inheriting another file's processing resolution.
starting_g_max = record.dstar_max if record.dstar_max is not None else simulation.g_max
plan = pipeline(
[
select_beams(cfg.blochwave.to_beam_selection(dataset.integration)),
build_orientation_plans(),
]
)(selected)
_plan, settled = run_convergence(
plan,
ConvergenceState(
g_max=starting_g_max,
sg_max=simulation.sg_max,
tilt_steps=rocking.sampling,
),
ConvergenceTest(),
rocking,
simulation,
refinement,
ConvergenceTolerance(),
method=cfg.blochwave.solver,
logger=logger,
)
settled_states.append(settled)
return ConvergenceState(
g_max=max(state.g_max for state in settled_states),
sg_max=max(state.sg_max for state in settled_states),
tilt_steps=max(state.tilt_steps for state in settled_states),
)
[docs]
def preprocess_experiment(
experiment_dir: str | Path,
*,
logger: Logger = NULL_LOGGER,
checkpoint: bool = True,
refresh: bool = False,
device: Device | None = "cuda",
workers: int = 1,
max_batch: int | None = None,
plot_thickness: bool = False,
plot_thickness_dir: str | Path | None = None,
) -> Plan:
"""Load and preprocess the experiment at ``experiment_dir``, returning the settled ``Plan``.
The preprocess half of :func:`run_experiment` with the terminal scoring stripped off: it loads
``experiment.yaml`` (verifying the input lock), reads the structure + experimental data, builds the
geometry via ``from_experiment``, and runs the default integrated recipe -- returning the
settled coupled :class:`~diffBloch.preprocess.plan.Plan` (fitted orientations, tilt-segment
couplings, pinned scored sets). This is the entry point for callers who want *only* the
calibrated Plan (to checkpoint it, or to drive their own downstream refinement) without paying
for the terminal inference pass.
The recipe includes **per-trial beam coupling** (the fit
re-derives the SOLVE union + SCORED set at every trial orientation). Its coupling policy,
orientation-search bounds, thickness grid, and whether
hydrogens are loaded all come from config (``blochwave`` / ``preprocess.orientation``
/ ``preprocess.thickness`` / ``inputs.load_hydrogens``). A caller wanting a different
composition (e.g. the cheaper tilt-independent fit) composes their own ``pipeline([...])`` with
the public steps.
``checkpoint`` (default ``True``) reuses/resumes a valid ``plan.npz`` in the experiment dir and
writes a fresh one after computing; ``refresh`` forces a full recompute (ignoring any snapshot)
while still regenerating the checkpoint. ``checkpoint=False`` neither reads nor writes.
``device`` (default ``"cuda"``) runs the forward solve on the selected accelerator for the
coupled preprocess fits. If CUDA is requested on a host without CUDA, the app falls back to CPU
and emits a device-selection event. The preprocess *geometry* (beam selection, coupling unions)
stays CPU-side numpy; only the eigensolve -- the O(N^3) cost -- moves to the device (the fits
move the seed params, so ``fgb`` and every per-trial score co-locate there). Device is
execution-only: it does not enter the checkpoint lock, so a committed CPU checkpoint is still
reused when a run moves to GPU.
``workers`` (default 1, sequential) fans the per-rotation orientation search over a thread pool
(rotations are independent). On a GPU run the per-trial cost is host-bound around a small
eigensolve, so overlapping rotations across cores is the main wall-clock lever (a small worker
count is usually the sweet spot; gains flatten as the solves serialise on one GPU stream). Like
``device`` it is execution-only -- the results are identical to a sequential run and it does not
enter the checkpoint lock. **Cap host threads to 1 when using it** --
``OMP_NUM_THREADS``/``MKL_NUM_THREADS``/``TORCH_NUM_THREADS`` (or ``torch.set_num_threads(1)``);
in a pod torch/BLAS size their pools from the *node* core count, not the cgroup limit, so an
uncapped run oversubscribes the cores the workers need (capping alone can dominate the speedup,
before any parallelism).
``max_batch`` (default ``None``) caps the ``matrix_exp`` propagator block for the coupled fits.
``None`` lets each solve derive a memory-safe block from its beam count; raise it to fill a
larger accelerator's memory budget. Execution-only (memory, bit-for-bit to
machine precision), out of the checkpoint lock like ``device``/``workers``. See
:func:`~diffBloch.engine.build_engine`.
``plot_thickness`` (default ``False``) ORs with ``cfg.preprocess.thickness.plot`` -- either
turns on one wR2-vs-thickness PNG per rotation from ``optimize_thickness``'s grid search, saved
under ``plot_thickness_dir`` (default ``<inputs.structure's directory>/thickness_optim``).
Execution-only like ``device``/``workers``, out of the checkpoint lock.
"""
root = Path(experiment_dir)
device = _select_device(device, logger=logger)
cfg, _lock = load_experiment(root)
return _preprocess(
root,
cfg,
logger=logger,
checkpoint=checkpoint,
refresh=refresh,
device=device,
workers=workers,
max_batch=max_batch,
plot_thickness=plot_thickness,
plot_thickness_dir=plot_thickness_dir,
).plan
[docs]
def run_experiment(
experiment_dir: str | Path,
*,
logger: Logger = NULL_LOGGER,
checkpoint: bool = True,
refresh: bool = False,
device: Device | None = "cuda",
workers: int = 1,
max_batch: int | None = None,
plot_thickness: bool = False,
plot_thickness_dir: str | Path | None = None,
) -> InferenceResult:
"""Load, preprocess, and score every rotation of the experiment at ``experiment_dir``.
:func:`preprocess_experiment` followed by the terminal forward model: it settles the coupled
``Plan`` (see that function for the recipe, ``checkpoint``/``refresh``,
``device``/``workers``, and ``plot_thickness``/``plot_thickness_dir``
semantics -- all shared), then evaluates every rotation with
``run_inference`` -- emitting
per-rotation observations to ``logger`` (the null default discards them). Returns the
:class:`~diffBloch.preprocess.inference.InferenceResult` (per-rotation ``R_obs`` + aggregate).
``device`` also runs the terminal eigensolve on the accelerator.
"""
root = Path(experiment_dir)
device = _select_device(device, logger=logger)
cfg, _lock = load_experiment(root)
outcome = _preprocess(
root,
cfg,
logger=logger,
checkpoint=checkpoint,
refresh=refresh,
device=device,
workers=workers,
max_batch=max_batch,
plot_thickness=plot_thickness,
plot_thickness_dir=plot_thickness_dir,
)
return run_inference(
outcome.plan,
outcome.refinement,
method=cfg.blochwave.solver,
device=device,
max_batch=max_batch,
logger=logger,
absorption=cfg.blochwave.to_absorption(),
)
[docs]
def refine_experiment(
experiment_dir: str | Path,
*,
logger: Logger = NULL_LOGGER,
checkpoint: bool = True,
refresh: bool = False,
device: Device | None = "cuda",
workers: int = 1,
max_batch: int | None = None,
verbose: bool = False,
profile: bool = False,
checkpoint_activations: bool = True,
plot_thickness: bool = False,
plot_thickness_dir: str | Path | None = None,
) -> ModelRefinementResult:
"""Settle the coupled ``Plan`` and gradient-refine the structure against the observed data.
:func:`preprocess_experiment` for the geometry (checkpoint reuse for free -- see it for the
recipe and ``checkpoint``/``refresh``/``device``/``workers``/
``plot_thickness``/``plot_thickness_dir`` semantics), then run the
**default** single-stage refinement on that settled ``Plan``. This is the boring config-knobs
path: the residual (:meth:`~diffBloch.config.schema.LossMetricsConfig.to_loss`), the trainable
selection (:meth:`~diffBloch.config.schema.TrainableConfig.to_spec`), and
the optimizer/step budget all come from ``experiment.yaml``. It composes no hard constraints or
penalties -- scientific
composition (hydrogen riding, freeze-H, penalties, multi-stage) is a Python/API concern, built
with :func:`~diffBloch.engine.build_refinement_model`,
:func:`~diffBloch.engine.build_refinement_problem`, and
:func:`~diffBloch.engine.with_hydrogen_riding`, then run via ``run_refinement_model``. The
:class:`~diffBloch.engine.RefinementProblem` here is pure optimization-definition data; the
imperative loop lives in ``run_refinement_model``. Returns the
:class:`~diffBloch.engine.ModelRefinementResult` (per-step losses + best snapshot);
:func:`_write_refinement_outputs` persists the best structure/parameters/summary to
``experiment_dir`` alongside a ``refinement.lock`` binding them to the settled ``Plan``
(``plan.lock``), the refinement-determining config, and the code version that produced them --
the refinement-stage counterpart to the preprocess checkpoint's own lock.
``device`` places the refinement solve on the accelerator: the seed params move there and the
forward co-locates onto them (as in the preprocess fits).
``verbose`` ("verbose refinement") reports one per-rotation wR2/R_obs/diffraction-loss line per
step in addition to the epoch mean; see :func:`~diffBloch.engine.run_refinement_model`.
Execution-only, like ``logger``.
``profile`` logs per-phase wall time (structure factors, each rotation's solve, backward,
optimizer step) via stdlib diagnostics logging; see :func:`~diffBloch.engine.run_refinement_model`
and :func:`~diffBloch.preprocess.scoring.build_engine`. Execution-only and off by default.
``checkpoint_activations`` (default ``True``) trades peak memory for one extra forward
recompute per solve on backward; disabling it removes that recompute at the cost of retaining
solve intermediates until backward. Execution-only -- gradients are unaffected. See
:class:`~diffBloch.engine.RefinementEngine`.
"""
root = Path(experiment_dir)
device = _select_device(device, logger=logger)
cfg, _lock = load_experiment(root)
outcome = _preprocess(
root,
cfg,
logger=logger,
checkpoint=checkpoint,
refresh=refresh,
device=device,
workers=workers,
max_batch=max_batch,
plot_thickness=plot_thickness,
plot_thickness_dir=plot_thickness_dir,
)
refinement, integrations, prepared = outcome.refinement, outcome.integrations, outcome.plan
validation_rotation_indices = outcome.validation_rotation_indices
plan_lock_sha256s = outcome.plan_lock_sha256s
# `engine` covers every rotation (train + validation) -- reporting always scores the whole
# experiment, e.g. the thickness-NN shape table below evaluates the trained curve at
# validation angles it never saw, which is the point. Only the *training* engine, built
# separately, excludes validation_rotation_indices from the gradient objective.
engine = build_engine(
prepared,
refinement,
loss=cfg.loss_metrics.to_loss(),
method=cfg.blochwave.solver,
max_batch=max_batch,
absorption=cfg.blochwave.to_absorption(),
profile=profile,
checkpoint_activations=checkpoint_activations,
)
# Pre-existing inconsistency, recorded not fixed: PerOrientationThickness keys its per-rotation
# lookup on the *positional* enumerate index the engine passes to forward_context, while
# ApparentThicknessNN keys on pattern.rotation_index -- the two disagree whenever an engine is
# built over a subset like the train/validation partitions below.
train_engine = engine
selection_engine = None
if validation_rotation_indices:
train_only = tuple(
op
for op in prepared.orientations
if op.pattern.rotation_index not in validation_rotation_indices
)
validation_only = tuple(
op
for op in prepared.orientations
if op.pattern.rotation_index in validation_rotation_indices
)
train_engine = build_engine(
replace(prepared, orientations=train_only),
refinement,
loss=cfg.loss_metrics.to_loss(),
method=cfg.blochwave.solver,
max_batch=max_batch,
absorption=cfg.blochwave.to_absorption(),
profile=profile,
checkpoint_activations=checkpoint_activations,
)
selection_engine = build_engine(
replace(prepared, orientations=validation_only),
refinement,
loss=cfg.loss_metrics.to_loss(),
method=cfg.blochwave.solver,
max_batch=max_batch,
absorption=cfg.blochwave.to_absorption(),
profile=profile,
checkpoint_activations=checkpoint_activations,
)
logger.report(cfg.to_declaration(integrations))
initial = refinement.params if device is None else refinement.params.to(device)
thickness_spec = cfg.refinement.thickness_nn.to_spec()
thickness_nns: tuple[ApparentThicknessNN, ...] = ()
raw_alphas: np.ndarray | None = None
if thickness_spec.enabled:
records = _read_experimental_data(root, cfg)
raw_alphas = np.concatenate([np.asarray(record.alphas) for record in records])
thickness_nns = _thickness_networks(cfg, records, thickness_spec)
model = build_refinement_model(
initial=initial,
components=thickness_nns,
component_params={
thickness_nn.key: thickness_nn.initial_params(
dtype=initial.asu_positions.dtype,
device=initial.asu_positions.device,
)
for thickness_nn in thickness_nns
},
)
else:
model = build_refinement_model(initial=initial)
problem = build_refinement_problem()
result = run_refinement_model(
train_engine,
model,
problem,
trainable=cfg.refinement.trainable.to_spec(),
steps=cfg.refinement.steps,
optimizer=cfg.refinement.optimizer.name,
lr=cfg.refinement.optimizer.lr,
logger=logger,
verbose=verbose,
profile=profile,
selection_engine=selection_engine,
)
result = _write_refinement_outputs(
root, cfg, refinement, result, plan_lock_sha256s=plan_lock_sha256s
)
# The settled-result events, then the terminal one: SummaryLogger writes the file on the latter.
_report_refinement_outcome(
logger,
engine,
result,
validation_rotation_indices=validation_rotation_indices,
thickness_nns=thickness_nns,
raw_alphas=raw_alphas,
)
return result
def _report_refinement_outcome(
logger: Logger,
engine: RefinementEngine,
result: ModelRefinementResult,
*,
validation_rotation_indices: frozenset[int],
thickness_nns: tuple[ApparentThicknessNN, ...],
raw_alphas: np.ndarray | None,
) -> None:
"""Emit the settled-result events the refinement loop itself cannot produce.
``run_refinement_model`` only ever sees the *training* engine, so the final per-rotation scores
(which cover held-out rotations too) and the trained thickness curve have to be emitted here,
where the reporting engine and the split are both in scope. :class:`RefinementOutputsWritten`
goes last and is the run's terminal event -- a sink that must write exactly once, after
everything else, acts on it. Each row already knows its dataset ref (``RotationMetrics.dataset``,
read off the rotation's own ``pattern``), so the per-dataset breakdown in the report costs
nothing here.
"""
for row in engine.per_rotation_metrics(result.best_model):
logger.report(
RefinedRotationMetrics(
rotation_index=row.rotation_index,
wr2=row.wr2,
r_obs=row.r_obs,
n_matched=row.n_matched,
is_validation=row.rotation_index in validation_rotation_indices,
dataset=row.dataset,
)
)
if raw_alphas is not None:
# Each network filters the pooled orientations to its own rotation_range itself.
for thickness_nn in thickness_nns:
logger.report(
thickness_nn.profile(
result.best_model.component_params[thickness_nn.key],
engine.orientations,
raw_alphas,
)
)
logger.report(
RefinementOutputsWritten(
structure=result.artifacts["refined_structure"], artifacts=result.artifacts
)
)
def _thickness_networks(
cfg: ExperimentConfig,
records: tuple[ExperimentalRecord, ...],
spec: ApparentThicknessNetwork,
) -> tuple[ApparentThicknessNN, ...]:
"""One thickness network per dataset, scoped to its pooled rotation-index range.
Ranges follow the cumulative pre-ignore rotation counts in ``inputs.exp_data`` order --
exactly how :func:`~diffBloch.preprocess.pool` numbers the pooled ``rotation_index`` space --
so the networks partition that space and composition finds exactly one thickness per
orientation. This offset arithmetic looks like the dataset attribution a rotation now carries on
its own ``pattern.dataset``, but it is a different quantity and cannot be replaced by it:
``ApparentThicknessNN`` indexes ``normalized_alphas`` by ``rotation_index - start`` and requires
the range to be exactly the alphas wide, so the range must span the dataset's *pre-ignore* block.
Grouping the settled plan by dataset label would yield the narrower observed span and misalign
every alpha lookup. Alphas are normalized independently per dataset so overlapping tilt ranges do
not share one thickness-vs-alpha curve. A single dataset is simply the N=1 case.
"""
bounds = ThicknessBounds(spec.min_thickness, spec.max_thickness)
offsets = np.cumsum([0, *(record.n_rotations for record in records)])
return tuple(
ApparentThicknessNN(
bounds=bounds,
normalized_alphas=_normalized_pets_alphas(np.asarray(record.alphas)),
key=f"apparent_thickness[{ref}]",
form=spec.form,
sample_thickness=spec.sample_thickness,
num_samples=spec.num_samples,
init_seed=spec.init_seed,
rotation_range=(int(start), int(end)),
label=ref,
)
for ref, record, start, end in zip(
_exp_data_refs(cfg), records, offsets[:-1], offsets[1:], strict=True
)
)
def _normalized_pets_alphas(alphas: np.ndarray) -> tuple[float, ...]:
"""Legacy MinMaxScaler ``[-1, 1]`` normalization for the PETS alpha coordinate."""
values = np.asarray(alphas, dtype=np.float64)
if values.ndim != 1 or values.size == 0 or not np.all(np.isfinite(values)):
raise ValueError("PETS alphas must be a finite non-empty 1-D array")
minimum = float(values.min())
span = float(values.max() - minimum)
if span == 0.0:
return tuple(-1.0 for _ in values)
normalized = -1.0 + 2.0 * (values - minimum) / span
return tuple(float(value) for value in normalized)
def _write_refinement_outputs(
root: Path,
cfg: ExperimentConfig,
refinement: RefinementSetup,
result: ModelRefinementResult,
*,
plan_lock_sha256s: tuple[str, ...] | None,
) -> ModelRefinementResult:
"""Persist the best structure and raw parameter/component snapshots.
``refined_structure.cif`` lands at ``root`` (the human-facing output); the raw ``.npz``
snapshots and the checkpoint/refinement locks go under :func:`_reproducibility_dir` --
bookkeeping nobody reads directly. See :class:`~diffBloch.app.loggers.summary.SummaryLogger`
for the actual human-readable summary (``refinement_report.txt``), which supersedes the
machine-readable ``refinement_summary.json`` this function used to also write.
``plan_lock_sha256s`` are the hashes of the plan locks *this run* verified or wrote
(``exp_data`` order), from :func:`_preprocess`. ``refinement.lock`` chains to exactly those;
``None`` (the run didn't checkpoint) skips the lock and the plan artifact entries entirely --
a leftover lock file on disk from some earlier run is not this run's provenance.
"""
reproducibility_dir = _reproducibility_dir(root)
structure_path = (root / "refined_structure.cif").resolve()
params_path = (reproducibility_dir / "refined_parameters.npz").resolve()
components_path = (reproducibility_dir / "refined_components.npz").resolve()
state = constrain(result.best_params, refinement.spec)
source_path = root / cfg.inputs.structure
document = gemmi.cif.read_file(str(source_path))
block = document.sole_block()
structure = read_structure(source_path, load_hydrogens=cfg.inputs.load_hydrogens)
positions = state.positions.detach().cpu().numpy()
occupancies = state.occupancies.detach().cpu().numpy()
uij_star = state.uij_star.detach().cpu().numpy()
reciprocal_basis = refinement.spec.reciprocal_basis
assert reciprocal_basis is not None
reciprocal = reciprocal_basis.detach().cpu().numpy()
reciprocal_lengths = np.linalg.norm(reciprocal, axis=1)
reciprocal_metric = reciprocal @ reciprocal.T
# refinement.cell_parameters is the authoritative cell reciprocal_basis was actually derived
# from (PETS's, not necessarily the structure CIF's own). The written header must match it, or
# the file would declare a cell inconsistent with the ADP/position convention refined under.
if refinement.cell_parameters is not None:
cell_tags = (
("_cell_length_a", refinement.cell_parameters[0]),
("_cell_length_b", refinement.cell_parameters[1]),
("_cell_length_c", refinement.cell_parameters[2]),
("_cell_angle_alpha", refinement.cell_parameters[3]),
("_cell_angle_beta", refinement.cell_parameters[4]),
("_cell_angle_gamma", refinement.cell_parameters[5]),
)
for tag, value in cell_tags:
block.set_pair(tag, f"{float(value):.6f}")
if block.find_pair("_cell_volume") is not None:
authoritative_unit_cell = cell_matrix_from_parameters(refinement.cell_parameters)
block.set_pair("_cell_volume", f"{cell_volume(authoritative_unit_cell):.5f}")
# The *effective* ADP kind (post inputs.isotropic_displacements_only override), not
# structure.adp.kind (the raw CIF classification): an atom the override force-converted to
# Uiso must be written back as Uiso, with its stale _atom_site_aniso_* row stripped below --
# otherwise re-reading this file classifies it back to Uani by the aniso row's mere presence
# (see io.cif._adp_for_site) and silently loses the override.
effective_kind = refinement.spec.adp_kind
assert effective_kind is not None
atom_loop = block.find_loop("_atom_site_label").get_loop()
tags = list(atom_loop.tags)
label_column = tags.index("_atom_site_label")
atom_rows = {
atom_loop.values[row * atom_loop.width() + label_column]: row
for row in range(atom_loop.length())
}
for index, label in enumerate(structure.labels):
row = atom_rows[label]
updates = {
"_atom_site_fract_x": positions[index, 0],
"_atom_site_fract_y": positions[index, 1],
"_atom_site_fract_z": positions[index, 2],
"_atom_site_occupancy": occupancies[index],
}
if effective_kind[index] == "Uiso":
updates["_atom_site_U_iso_or_equiv"] = np.sum(
uij_star[index] * reciprocal_metric
) / np.sum(reciprocal_metric * reciprocal_metric)
for tag, value in updates.items():
if tag in tags:
column = tags.index(tag)
atom_loop[row, column] = f"{float(value):.10g}"
aniso_column = block.find_loop("_atom_site_aniso_label")
if aniso_column:
aniso_loop = aniso_column.get_loop()
aniso_tags = list(aniso_loop.tags)
label_column = aniso_tags.index("_atom_site_aniso_label")
aniso_rows = {
aniso_loop.values[row * aniso_loop.width() + label_column]: row
for row in range(aniso_loop.length())
}
components = {
"_atom_site_aniso_U_11": (0, 0),
"_atom_site_aniso_U_22": (1, 1),
"_atom_site_aniso_U_33": (2, 2),
"_atom_site_aniso_U_12": (0, 1),
"_atom_site_aniso_U_13": (0, 2),
"_atom_site_aniso_U_23": (1, 2),
}
scale = reciprocal_lengths[:, None] * reciprocal_lengths[None, :]
stale_aniso_labels: list[str] = []
for index, label in enumerate(structure.labels):
if label not in aniso_rows:
continue
if effective_kind[index] != "Uani":
# Was Uani in the CIF but the override forces it to Uiso: strip the row rather than
# leaving it untouched (stale) or refreshing it with new anisotropic-looking values.
stale_aniso_labels.append(label)
continue
row = aniso_rows[label]
uij_cif = uij_star[index] / scale
for tag, (i, j) in components.items():
if tag in aniso_tags:
column = aniso_tags.index(tag)
aniso_loop[row, column] = f"{float(uij_cif[i, j]):.10g}"
if stale_aniso_labels:
aniso_table = block.find(list(aniso_loop.tags))
for label in stale_aniso_labels:
aniso_table.remove_row(aniso_table.find_row(label).row_index)
document.write_file(str(structure_path))
params = result.best_params
empty = np.empty((0,), dtype=np.float64)
np.savez_compressed(
str(params_path),
asu_positions=params.asu_positions.detach().cpu().numpy(),
uij_raw=empty if params.uij_raw is None else params.uij_raw.detach().cpu().numpy(),
u_iso_raw=(empty if params.u_iso_raw is None else params.u_iso_raw.detach().cpu().numpy()),
occupancy_raw=(
empty if params.occupancy_raw is None else params.occupancy_raw.detach().cpu().numpy()
),
)
component_arrays = {
f"{component_key}__{parameter_name}": parameter.detach().cpu().numpy()
for component_key, parameters in result.best_model.component_params.items()
for parameter_name, parameter in parameters.items()
}
if component_arrays:
np.savez_compressed(str(components_path), **component_arrays) # type: ignore[arg-type]
artifacts: dict[str, str] = {
"refined_structure": str(structure_path),
"refined_parameters": str(params_path),
}
if component_arrays:
artifacts["refined_components"] = str(components_path)
# refinement.lock chains only to the plan locks THIS run verified or wrote. A lock file merely
# present on disk (a --no-checkpoint run over leftovers, an opaque recipe) is not this run's
# provenance and must not be chained to.
if plan_lock_sha256s is not None:
for ref in _exp_data_refs(cfg):
stem = dataset_checkpoint_stem(ref)
artifacts[f"plan_{stem}"] = str((reproducibility_dir / _plan_npz_name(ref)).resolve())
artifacts[f"plan_lock_{stem}"] = str(
(reproducibility_dir / _plan_lock_name(ref)).resolve()
)
lock_path = (reproducibility_dir / _REFINEMENT_LOCK).resolve()
write_refinement_lock(
lock_path,
RefinementLock(
plan_lock_sha256s=list(plan_lock_sha256s),
refinement_config_digest=refinement_config_digest(cfg),
code_version=code_version(),
refined_structure=artifact_hash_for(structure_path, root=root.resolve()),
refined_parameters=artifact_hash_for(params_path, root=root.resolve()),
),
)
artifacts["refinement_lock"] = str(lock_path)
return replace(result, artifacts=artifacts)
@dataclass(frozen=True)
class PreprocessOutcome:
"""Everything :func:`_preprocess` settles, named rather than positional.
A plain tuple return made every field addition churn all three entry points at once, since each
had to restructure its unpacking to keep the underscore-prefixed elements it ignores. Fields:
``refinement`` the structure-side :class:`~diffBloch.preprocess.experiment.RefinementSetup`,
``integrations`` the per-dataset :class:`~diffBloch.specs.IntegrationGeometry` in
``inputs.exp_data`` order, ``plan`` the pooled settled ``Plan``,
``validation_rotation_indices`` the held-out pooled indices (empty when ``train_test`` is off),
and ``plan_lock_sha256s`` the locks this run verified or wrote (``None`` when it didn't
checkpoint). See :func:`_preprocess` for what each one means in full.
"""
refinement: RefinementSetup
integrations: tuple[IntegrationGeometry, ...]
plan: Plan
validation_rotation_indices: frozenset[int]
plan_lock_sha256s: tuple[str, ...] | None
def _preprocess(
root: Path,
cfg: ExperimentConfig,
*,
logger: Logger,
checkpoint: bool,
refresh: bool,
device: Device | None,
workers: int,
max_batch: int | None,
plot_thickness: bool = False,
plot_thickness_dir: str | Path | None = None,
) -> PreprocessOutcome:
"""Shared spine of the public entry points: read inputs, run the recipe per dataset, pool.
Runs :func:`~diffBloch.preprocess.setup_datasets` over every ``inputs.exp_data`` file, then --
sequentially, one dataset at a time -- that dataset's recipe under its own checkpoint
(:func:`_prepare`), and finally :func:`~diffBloch.preprocess.pool`\\ s the settled plans onto
the pooled ``rotation_index`` space. Returns the structure ``RefinementSetup`` (which the
terminal ``run_inference`` needs), the per-dataset ``IntegrationGeometry``\\ s in
``inputs.exp_data`` order (PETS-derived integration semiangles, for reporting), the pooled
settled ``Plan`` (still over *every* non-ignored rotation -- the train/validation split never
restricts preprocessing, only :func:`refine_experiment`'s gradient stage), the
``rotation_index`` set ``cfg.refinement.split`` holds out as validation (empty when
``train_test`` is off; the mask runs over the pooled pre-ignore index space, so membership is
stable under ignore edits), and the sha256s of the plan locks this run verified or wrote (in
``exp_data`` order) -- ``None`` when the run didn't checkpoint, so ``refinement.lock`` never
chains to a leftover lock this run never validated. Hydrogen sites are loaded per
``inputs.load_hydrogens``. Dataset attribution is *not* returned alongside: each rotation
already carries its own ``pattern.dataset`` from :func:`~diffBloch.preprocess.setup_datasets`,
which survives the pooled renumbering.
``plot_thickness`` (API/CLI) ORs with ``cfg.preprocess.thickness.plot`` -- either can turn
plotting on. ``plot_thickness_dir`` overrides the default output directory,
``<inputs.structure's directory>/thickness_optim``, when given. A multi-dataset experiment gets
one subdirectory per dataset under it (named by :func:`~diffBloch.config.schema.
dataset_checkpoint_stem`), so two datasets sharing a rotation index never overwrite each other's
PNG. Both are execution-only (they only decide whether/where a PNG gets written, never the
fitted ``Plan``) -- see :func:`~diffBloch.config.manifest.dataset_config_digest`.
"""
structure = _read_structure(root, cfg, logger=logger)
records = _read_experimental_data(root, cfg, logger=logger)
refs = _exp_data_refs(cfg)
refinement_setup, datasets = setup_datasets(structure, records, cfg)
cell = records[0].cell_parameters
authoritative_cell: CellParameters = (
float(cell[0]),
float(cell[1]),
float(cell[2]),
float(cell[3]),
float(cell[4]),
float(cell[5]),
)
effective_plot_dir = (
(
Path(plot_thickness_dir)
if plot_thickness_dir is not None
else (root / cfg.inputs.structure).parent / "thickness_optim"
)
if plot_thickness or cfg.preprocess.thickness.plot
else None
)
structure_lock = input_lock_for(root / cfg.inputs.structure, ref=cfg.inputs.structure)
prepared_plans: list[Plan] = []
lock_sha256s: list[str | None] = []
for i, (ref, dataset) in enumerate(zip(refs, datasets, strict=True)):
_log.info("preprocessing dataset %r (%d/%d)", ref, i + 1, len(refs))
dataset_logger = logger
if effective_plot_dir is not None:
from diffBloch.app.loggers.plotting import ThicknessPlotLogger
dataset_logger = MultiLogger(
(
logger,
ThicknessPlotLogger(effective_plot_dir / dataset_checkpoint_stem(ref)),
)
)
steps = _recipe_steps(
cfg,
refinement_setup,
dataset.integration,
dataset.mosaicity,
dataset_logger,
device=device,
workers=workers,
max_batch=max_batch,
)
prepared, lock_sha256 = _prepare(
dataset.plan,
steps,
root=root,
cfg=cfg,
dataset_ref=ref,
ignored_rotations=dataset.ignored_rotations,
authoritative_cell=authoritative_cell,
structure_lock=structure_lock,
dataset_lock=input_lock_for(root / ref, ref=ref),
checkpoint=checkpoint,
refresh=refresh,
logger=dataset_logger,
)
prepared_plans.append(prepared)
lock_sha256s.append(lock_sha256)
if checkpoint:
_prune_stale_dataset_checkpoints(_reproducibility_dir(root), refs)
offsets: list[int] = []
offset = 0
for dataset in datasets:
offsets.append(offset)
offset += dataset.n_rotations
pooled = pool(prepared_plans, offsets=offsets)
mask = validation_mask(offset, cfg.refinement.split)
validation_rotation_indices = frozenset(
op.pattern.rotation_index for op in pooled.orientations if mask[op.pattern.rotation_index]
)
integrations = tuple(dataset.integration for dataset in datasets)
plan_lock_sha256s = (
None
if any(sha is None for sha in lock_sha256s)
else tuple(sha for sha in lock_sha256s if sha is not None)
)
built = cast(tuple[OrientationPlanLike, ...], pooled.orientations)
logger.report(
PreprocessCompleted(
n_rotations=len(built),
n_stages=len(pooled.provenance),
total_hkl=unique_hkl_count(op.pattern.hkl for op in built),
matched_hkl=unique_hkl_count(op.alignment.hkl for op in built),
steps=tuple((record.name, record.params) for record in pooled.provenance),
)
)
return PreprocessOutcome(
refinement=refinement_setup,
integrations=integrations,
plan=pooled,
validation_rotation_indices=validation_rotation_indices,
plan_lock_sha256s=plan_lock_sha256s,
)
def _prune_stale_dataset_checkpoints(reproducibility_dir: Path, refs: tuple[str, ...]) -> None:
"""Unlink ``plan.<stem>.{npz,lock}`` files whose stem left ``inputs.exp_data``.
A dataset removed (or renamed) in config would otherwise leave a live-looking checkpoint pair
behind. A bare legacy ``plan.npz``/``plan.lock`` has no stem segment, matches neither glob
below, and lingers harmlessly.
"""
keep = {dataset_checkpoint_stem(ref) for ref in refs}
for pattern, suffix in (("plan.*.npz", ".npz"), ("plan.*.lock", ".lock")):
for path in reproducibility_dir.glob(pattern):
stem = path.name[len("plan.") : -len(suffix)]
if stem not in keep:
path.unlink()
_log.info("pruned stale dataset checkpoint file %s", path.name)
def _recipe_steps(
cfg: ExperimentConfig,
refinement: RefinementSetup,
integration: IntegrationGeometry,
mosaicity: MosaicSmoothed | None,
logger: Logger,
*,
device: Device | None = None,
workers: int = 1,
max_batch: int | None = None,
) -> list[PlanStep]:
"""The default recipe as an inspectable step list (its provenance keys the lock).
``preprocess.optimize_orientation`` and ``preprocess.optimize_thickness`` select which fitting
stages join the fixed recipe order. When enabled, the orientation fit runs under per-trial
coupling (:func:`_trial_coupling`). The tilt-independent fit is not offered here; compose it
directly if needed.
The orientation fit is a :func:`~diffBloch.preprocess.fork` on unit-cell volume: a **large
cell** (> ``_LARGE_CELL_THRESHOLD_A3``) skips the per-trial gather integrity checks
(``validate=False``, made sound by ``optimize_orientation``'s coupled coverage guard); a
**small cell** takes the exact, fully-validated path. The predicate
reads only the pipeline-invariant grid, so :func:`~diffBloch.preprocess.resolve_recipe` compiles
the fork to a flat branch before the lock sees it (the branch is fixed per experiment).
``validate`` is execution-only and stays out of the step identity, so the resolved branch keys
the same either way -- the committed quartz checkpoint is untouched. The thickness fit has no
such split and always runs the same path regardless of cell size.
Both ``device`` and ``workers`` are execution-only -- neither alters the recipe identity.
``device`` is threaded to *both* fits (so the coupled eigensolve runs on the same accelerator as
the terminal); ``workers`` fans both the independent initial rotation-plan builds and the
orientation searches over threads.
``max_batch`` (also execution-only) is threaded to both fits: it caps the ``matrix_exp``
propagator block so a wide coupled segment x the thickness grid can't materialize the whole
propagator at once; ``None`` picks a memory-safe block.
"""
search = cfg.preprocess.orientation.to_search()
thickness_grid = cfg.preprocess.thickness.to_grid()
coupling = _trial_coupling(cfg, integration)
def orientation_fit(*, validate: bool) -> PlanStep:
return optimize_orientation(
refinement,
search,
method=cfg.blochwave.solver,
coupling=coupling,
validate=validate,
device=device,
max_batch=max_batch,
workers=workers,
logger=logger, # per-rotation fit progress (the run's long phase)
absorption=cfg.blochwave.to_absorption(),
scores=cfg.loss_metrics.to_scores(),
residual=cfg.loss_metrics.residual,
)
def thickness_fit() -> PlanStep:
return optimize_thickness(
refinement,
thickness_grid,
method=cfg.blochwave.solver,
device=device,
max_batch=max_batch,
logger=logger, # per-rotation thickness-fit progress (the memory-heavy tail phase)
absorption=cfg.blochwave.to_absorption(),
scores=cfg.loss_metrics.to_scores(),
residual=cfg.loss_metrics.residual,
)
steps: list[PlanStep] = []
steps.append(
build_orientation_plans(
cfg.blochwave.to_rocking_curve(integration),
mosaicity,
coupling=cfg.blochwave.to_policy(),
scoring_selection=cfg.blochwave.to_beam_selection(integration),
workers=workers,
)
)
def orientation_step() -> PlanStep:
return fork(
lambda grid: grid.cell_volume > _LARGE_CELL_THRESHOLD_A3,
when_true=[orientation_fit(validate=False)],
when_false=[orientation_fit(validate=True)],
)
fitting_steps: list[PlanStep] = []
if cfg.preprocess.stage_order == "thickness_first":
if cfg.preprocess.optimize_thickness:
fitting_steps.append(thickness_fit())
if cfg.preprocess.optimize_orientation:
fitting_steps.append(orientation_step())
else:
if cfg.preprocess.optimize_orientation:
fitting_steps.append(orientation_step())
if cfg.preprocess.optimize_thickness:
fitting_steps.append(thickness_fit())
steps.extend(fitting_steps)
return steps
def _trial_coupling(cfg: ExperimentConfig, integration: IntegrationGeometry) -> TrialCoupling:
"""Assemble the per-trial beam-union policy from the top-level Bloch-wave config.
The SCORED set reuses the same Klar window as ``select_beams`` and the config's solve cutoff
(``g_max``) as the scoring-resolution cap.
"""
return TrialCoupling(
policy=cfg.blochwave.to_policy(),
scored=ScoredHklSelection(
klar=cfg.blochwave.to_beam_selection(integration),
g_max=cfg.blochwave.g_max,
),
)
def _prepare(
base: Plan,
steps: list[PlanStep],
*,
root: Path,
cfg: ExperimentConfig,
dataset_ref: str,
ignored_rotations: tuple[int, ...],
authoritative_cell: CellParameters,
structure_lock: InputLock,
dataset_lock: InputLock,
checkpoint: bool,
refresh: bool,
logger: Logger = NULL_LOGGER,
) -> tuple[Plan, str | None]:
"""Run the preprocess ``steps`` on one dataset's ``base``, reusing/resuming its checkpoint.
The checkpoint pair is this dataset's own ``plan.<stem>.npz`` + ``plan.<stem>.lock``
(:func:`_plan_npz_name`), and the lock identity is per-dataset (``structure_lock`` +
``dataset_lock`` input bytes, the authoritative PETS cell, the dataset-scoped config digest,
the file-local ``ignored_rotations``) -- see :class:`~diffBloch.config.PreprocessLock`.
``logger`` streams a per-step plan summary as the recipe runs (see :func:`pipeline`); it fires
only when steps actually execute (a fresh or resumed run, not a full-reuse load).
Returns the settled plan plus the sha256 of the lock this run *verified or wrote* -- the
provenance ``refinement.lock`` may chain to. ``None`` when the run didn't checkpoint
(``checkpoint=False`` or an opaque recipe): a lock file merely *present* on disk from some
earlier run is not this run's provenance, and must not be chained to.
"""
# Compile any `fork` away against the base grid (invariant across every step), so the recipe the
# lock keys on is a flat, fork-free step list -- the fork's shape is static by construction, so
# its branch is fixed here, before running.
steps = list(resolve_recipe(steps, base.structure_factor_grid))
records = step_records(steps)
# A recipe with an unrecorded (opaque) step cannot be safely identified -> never checkpoint it.
can_checkpoint = checkpoint and OPAQUE not in records
recipe = [RecipeStep(name=r.name, params=r.params) for r in records]
reproducibility_dir = _reproducibility_dir(root)
npz = reproducibility_dir / _plan_npz_name(dataset_ref)
lock_path = reproducibility_dir / _plan_lock_name(dataset_ref)
if can_checkpoint and not refresh and lock_path.exists() and npz.exists():
lock = _read_lock_or_none(lock_path)
status = (
"stale"
if lock is None or not plan_is_readable(npz)
else preprocess_lock_status(
lock,
structure=structure_lock,
experimental_data=dataset_lock,
authoritative_cell=authoritative_cell,
ignored_rotations=ignored_rotations,
config_digest=dataset_config_digest(cfg, exp_data=dataset_ref),
code_version=code_version(),
recipe=recipe,
plan_path=npz,
root=root,
)
)
if status == "reuse":
_log.info("loaded preprocess checkpoint (full reuse) from %s", npz)
return read_plan(npz), sha256_file(lock_path)
if status == "resume":
assert lock is not None
snapshot = read_plan(npz)
k = len(lock.recipe)
_log.info(
"resumed preprocess checkpoint after %d step(s); running %s",
k,
[r.name for r in records[k:]],
)
result = pipeline(steps[k:], logger=logger)(snapshot)
_write_checkpoint(
result,
recipe,
root=root,
cfg=cfg,
dataset_ref=dataset_ref,
ignored_rotations=ignored_rotations,
authoritative_cell=authoritative_cell,
structure_lock=structure_lock,
dataset_lock=dataset_lock,
npz=npz,
lock_path=lock_path,
)
return result, sha256_file(lock_path)
result = pipeline(steps, logger=logger)(base)
if can_checkpoint:
_write_checkpoint(
result,
recipe,
root=root,
cfg=cfg,
dataset_ref=dataset_ref,
ignored_rotations=ignored_rotations,
authoritative_cell=authoritative_cell,
structure_lock=structure_lock,
dataset_lock=dataset_lock,
npz=npz,
lock_path=lock_path,
)
return result, sha256_file(lock_path)
return result, None
def _read_lock_or_none(lock_path: Path) -> PreprocessLock | None:
"""Read a checkpoint lock, treating an unparsable one as absent (stale), not an error.
A lock that no longer parses (hand-edited, truncated, or written by an incompatible build) means
the checkpoint cannot be trusted -- recompute rather than crash.
"""
try:
return read_preprocess_lock(lock_path)
except ValidationError:
_log.info("checkpoint lock %s did not parse; treating as stale", lock_path.name)
return None
def _write_checkpoint(
plan: Plan,
recipe: list[RecipeStep],
*,
root: Path,
cfg: ExperimentConfig,
dataset_ref: str,
ignored_rotations: tuple[int, ...],
authoritative_cell: CellParameters,
structure_lock: InputLock,
dataset_lock: InputLock,
npz: Path,
lock_path: Path,
) -> None:
"""Write the dataset's ``.npz`` + regenerate its lock so the lock always describes the npz."""
write_plan(plan, npz)
lock = PreprocessLock(
structure=structure_lock,
experimental_data=dataset_lock,
authoritative_cell=authoritative_cell,
ignored_rotations=ignored_rotations,
config_digest=dataset_config_digest(cfg, exp_data=dataset_ref),
code_version=code_version(),
recipe=recipe,
plan=artifact_hash_for(npz, root=root),
)
write_preprocess_lock(lock_path, lock)
_log.info("wrote preprocess checkpoint %s + %s", npz.name, lock_path.name)