From 37c22a689f61de57be102fd892890f770d85bac7 Mon Sep 17 00:00:00 2001 From: Simone Vadilonga Date: Tue, 11 Aug 2026 08:52:13 +0200 Subject: [PATCH 01/10] Extract the shared Ax trial loop from the measurement-fit optimizer Co-Authored-By: Claude Sonnet 5 --- src/grax_opt/dynamic.py | 208 ++++++++++++++------------------ src/grax_opt/loop.py | 253 +++++++++++++++++++++++++++++++++++++++ src/grax_opt/optimize.py | 81 ++++++++++--- 3 files changed, 412 insertions(+), 130 deletions(-) create mode 100644 src/grax_opt/loop.py diff --git a/src/grax_opt/dynamic.py b/src/grax_opt/dynamic.py index 4aa57c6..c839024 100644 --- a/src/grax_opt/dynamic.py +++ b/src/grax_opt/dynamic.py @@ -16,15 +16,14 @@ from .evaluation import normalize_evaluation_selection from .data import MeasurementData, load_measurement_data from .objective import build_evaluation_measurement, simulate_efficiency_curve +from .loop import TrialEvaluation, TrialLoopState, run_ax_trial_loop from .optimize import ( OptimizationResult, TrialRecord, - _complete_ax_trial, + _atomic_write_text, _describe_optimizer_compute_context, _evaluate_candidate_batch, _import_ax_client as _import_ax_client_from_optimize, - _import_data_required_exception, - _import_max_parallelism_exception, _import_objective_properties, _patch_torch_fork_rng_for_cpu_only, _resolve_optimizer_backend, @@ -555,7 +554,7 @@ def _write_measurement_fit_result_json( payload["grazing_angle_deg"] = float(config.grazing_angle_deg) else: payload["cff"] = float(config.cff) - output_path.write_text(json.dumps(payload, indent=2), encoding="utf-8") + _atomic_write_text(output_path, json.dumps(payload, indent=2)) def _persist_measurement_fit_optimizer_artifacts( @@ -669,20 +668,9 @@ def optimize_to_measurements( measurement = load_measurement_data(config.measurement_path) evaluation_measurement = build_evaluation_measurement(config, measurement) ax_client = _create_ax_client_for_measurement_fit_config(config) - max_parallelism_exception = _import_max_parallelism_exception() - data_required_exception = _import_data_required_exception() - - trial_records: list[TrialRecord] = [] - best_loss = float("inf") - best_parameters: dict[str, float] = {} - best_grating_parameters: dict[str, object] = {} - best_solver_parameters: dict[str, float | None] = {} - early_stop_reason: str | None = None - stopped_early = False - completed_trials = 0 - trial_index_cursor = 0 - no_improvement_trials = 0 - optimizer_resolved_max_workers = _resolve_simulation_max_workers(config.max_workers) + + state = TrialLoopState() + state.resolved_max_workers = _resolve_simulation_max_workers(config.max_workers) build_grating_fn = lambda trial_parameters: config.build_grating( resolve_measurement_fit_trial_parameters(config, trial_parameters) @@ -692,135 +680,123 @@ def optimize_to_measurements( resolve_measurement_fit_trial_parameters(config, trial_parameters), ) - while trial_index_cursor < config.total_trials: - candidates: list[tuple[int, dict[str, float]]] = [] - while len(candidates) < config.batch_size and trial_index_cursor < config.total_trials: - try: - raw_parameters, trial_index = ax_client.get_next_trial() - except Exception as error: - if max_parallelism_exception is not None and isinstance(error, max_parallelism_exception): - break - if data_required_exception is not None and isinstance(error, data_required_exception): - break - raise - parameters = {name: float(value) for name, value in raw_parameters.items()} - candidates.append((int(trial_index), parameters)) - trial_index_cursor += 1 - - if not candidates: - break - - evaluated_candidates = _evaluate_candidate_batch( - candidates, + def evaluate_candidates(candidates) -> list[TrialEvaluation]: + """Evaluate one batch of measurement-fit candidates. + + Args: + candidates: Candidate ``(trial_index, parameters)`` pairs. + + Returns: + One evaluation per candidate, ordered by trial index. + """ + + evaluated = _evaluate_candidate_batch( + list(candidates), config=config, measurement=measurement, backend_effective=backend_effective, build_grating_fn=build_grating_fn, resolve_solver_parameters_fn=resolve_solver_parameters_fn, ) - for trial_index, parameters, loss, trial_resolved_max_workers in evaluated_candidates: - optimizer_resolved_max_workers = int(trial_resolved_max_workers) - _complete_ax_trial( - ax_client=ax_client, - config=config, - trial_index=trial_index, + return [ + TrialEvaluation( + trial_index=int(trial_index), + parameters=dict(parameters), loss=float(loss), + resolved_max_workers=int(trial_resolved_max_workers), ) - trial_records.append( - TrialRecord( - trial_index=int(trial_index), - loss=float(loss), - parameters=dict(parameters), - ) - ) - completed_trials += 1 - current_grating_parameters = resolve_measurement_fit_trial_parameters(config, parameters) - current_solver_parameters = _resolve_measurement_fit_solver_parameters( - config, - current_grating_parameters, + for trial_index, parameters, loss, trial_resolved_max_workers in evaluated + ] + + def on_trial_completed(*, evaluation: TrialEvaluation, state: TrialLoopState, improved: bool) -> None: + """Refresh derived best-fit state and rewrite optimizer artifacts. + + Args: + evaluation: Evaluation for the trial that just completed. + state: Mutable loop state to update. + improved: Whether this trial produced a new best loss. + """ + + if improved: + state.best_grating_parameters = dict( + resolve_measurement_fit_trial_parameters(config, evaluation.parameters) ) - if loss < best_loss: - best_loss = float(loss) - best_parameters = dict(parameters) - best_grating_parameters = dict(current_grating_parameters) - best_solver_parameters = dict(current_solver_parameters) - no_improvement_trials = 0 - else: - no_improvement_trials += 1 - - _persist_measurement_fit_optimizer_artifacts( - config=config, - evaluation_measurement=evaluation_measurement, - best_parameters=best_parameters, - best_grating_parameters=best_grating_parameters, - best_solver_parameters=best_solver_parameters, - best_loss=best_loss, - trial_records=trial_records, - stopped_early=stopped_early, - completed_trials=completed_trials, - early_stop_reason=early_stop_reason, - backend_requested=config.backend, - backend_effective=backend_effective, - optimizer_requested_max_workers=config.max_workers, - optimizer_resolved_max_workers=optimizer_resolved_max_workers, + state.best_solver_parameters = dict( + _resolve_measurement_fit_solver_parameters( + config, + state.best_grating_parameters, + ) ) + _persist_measurement_fit_optimizer_artifacts( + config=config, + evaluation_measurement=evaluation_measurement, + best_parameters=state.best_parameters, + best_grating_parameters=state.best_grating_parameters, + best_solver_parameters=state.best_solver_parameters, + best_loss=state.best_loss, + trial_records=state.trial_records, + stopped_early=state.stopped_early, + completed_trials=state.completed_trials, + early_stop_reason=state.early_stop_reason, + backend_requested=config.backend, + backend_effective=backend_effective, + optimizer_requested_max_workers=config.max_workers, + optimizer_resolved_max_workers=state.resolved_max_workers, + ) - if ( - config.enable_early_stopping - and completed_trials >= config.early_stopping_warmup_trials - and no_improvement_trials >= config.early_stopping_patience - ): - stopped_early = True - early_stop_reason = ( - "Early stopping triggered after " - f"{no_improvement_trials} non-improving trials." - ) - break - if stopped_early: - break + run_ax_trial_loop( + ax_client=ax_client, + config=config, + state=state, + evaluate_candidates=evaluate_candidates, + on_trial_completed=on_trial_completed, + ) - if not trial_records: + if not state.trial_records: raise RuntimeError("Optimization produced no completed trials.") - if not best_parameters: - best_parameters = dict(trial_records[-1].parameters) - best_grating_parameters = resolve_measurement_fit_trial_parameters(config, best_parameters) - best_solver_parameters = _resolve_measurement_fit_solver_parameters( + if not state.best_parameters: + state.best_parameters = dict(state.trial_records[-1].parameters) + state.best_grating_parameters = resolve_measurement_fit_trial_parameters( + config, + state.best_parameters, + ) + state.best_solver_parameters = _resolve_measurement_fit_solver_parameters( config, - best_grating_parameters, + state.best_grating_parameters, ) - best_loss = float(trial_records[-1].loss) + state.best_loss = float(state.trial_records[-1].loss) result_paths = _persist_measurement_fit_optimizer_artifacts( config=config, evaluation_measurement=evaluation_measurement, - best_parameters=best_parameters, - best_grating_parameters=best_grating_parameters, - best_solver_parameters=best_solver_parameters, - best_loss=best_loss, - trial_records=trial_records, - stopped_early=stopped_early, - completed_trials=completed_trials, - early_stop_reason=early_stop_reason, + best_parameters=state.best_parameters, + best_grating_parameters=state.best_grating_parameters, + best_solver_parameters=state.best_solver_parameters, + best_loss=state.best_loss, + trial_records=state.trial_records, + stopped_early=state.stopped_early, + completed_trials=state.completed_trials, + early_stop_reason=state.early_stop_reason, backend_requested=config.backend, backend_effective=backend_effective, optimizer_requested_max_workers=config.max_workers, - optimizer_resolved_max_workers=optimizer_resolved_max_workers, + optimizer_resolved_max_workers=state.resolved_max_workers, ) return OptimizationResult( - best_parameters=best_parameters, - best_grating_parameters=best_grating_parameters, - best_loss=best_loss, + best_parameters=state.best_parameters, + best_grating_parameters=state.best_grating_parameters, + best_loss=state.best_loss, measurement_path=config.measurement_path, result_json_path=result_paths[0], trial_history_csv_path=result_paths[1], best_fit_plot_path=result_paths[2], loss_history_plot_path=result_paths[3], - trial_records=trial_records, - stopped_early=stopped_early, - completed_trials=completed_trials, - early_stop_reason=early_stop_reason, + trial_records=state.trial_records, + stopped_early=state.stopped_early, + completed_trials=state.completed_trials, + early_stop_reason=state.early_stop_reason, ) diff --git a/src/grax_opt/loop.py b/src/grax_opt/loop.py new file mode 100644 index 0000000..6288c82 --- /dev/null +++ b/src/grax_opt/loop.py @@ -0,0 +1,253 @@ +"""Shared Ax trial-loop runtime for the measurement-fit optimizers.""" + +from __future__ import annotations + +import math +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass, field +from typing import Any + +from .optimize import ( + TrialRecord, + _complete_ax_trial, + _import_data_required_exception, + _import_max_parallelism_exception, +) + + +@dataclass(frozen=True) +class TrialEvaluation: + """Outcome of evaluating one optimizer candidate. + + Attributes: + trial_index: Ax trial index the candidate was generated for. + parameters: Free parameter values evaluated for the candidate. + loss: Objective value produced by the evaluation. + resolved_max_workers: Worker count the evaluation actually used. + extras: Mode-specific payload such as per-measurement losses and + cached simulated curves. + """ + + trial_index: int + parameters: dict[str, float] + loss: float + resolved_max_workers: int + extras: dict[str, Any] = field(default_factory=dict) + + +@dataclass +class TrialLoopState: + """Mutable run state shared across optimizer trial iterations. + + Attributes: + trial_records: Completed trial records in completion order. + best_loss: Best objective value observed so far. + best_parameters: Free parameters for the best trial. + best_grating_parameters: Resolved grating parameters for the best trial. + best_solver_parameters: Resolved solver parameters for the best trial. + best_extras: Mode-specific payload captured from the best trial. + completed_trials: Number of trials successfully evaluated. + trial_index_cursor: Number of candidates drawn from Ax so far. + no_improvement_trials: Consecutive trials without a significant improvement. + stopped_early: Whether early stopping ended the run. + early_stop_reason: Human-readable early-stopping reason, or ``None``. + resolved_max_workers: Worker count reported by the most recent trial. + """ + + trial_records: list[TrialRecord] = field(default_factory=list) + best_loss: float = float("inf") + best_parameters: dict[str, float] = field(default_factory=dict) + best_grating_parameters: dict[str, object] = field(default_factory=dict) + best_solver_parameters: dict[str, float | None] = field(default_factory=dict) + best_extras: dict[str, Any] = field(default_factory=dict) + completed_trials: int = 0 + trial_index_cursor: int = 0 + no_improvement_trials: int = 0 + stopped_early: bool = False + early_stop_reason: str | None = None + resolved_max_workers: int = 1 + + +def is_significant_improvement( + previous_best_loss: float, + loss: float, + minimum_relative_improvement: float, +) -> bool: + """Return whether a loss improves on the previous best by enough to reset patience. + + Args: + previous_best_loss: Best objective value before this trial. + loss: Objective value produced by this trial. + minimum_relative_improvement: Minimum relative gain that counts as progress. + + Returns: + ``True`` when the improvement should reset the early-stopping counter. + """ + + if not loss < previous_best_loss: + return False + if not math.isfinite(previous_best_loss): + return True + if minimum_relative_improvement <= 0.0: + return True + denominator = abs(previous_best_loss) + if denominator == 0.0: + return True + return ((previous_best_loss - loss) / denominator) >= minimum_relative_improvement + + +def _collect_candidates( + *, + ax_client: Any, + state: TrialLoopState, + total_trials: int, + batch_size: int, + max_parallelism_exception: type[BaseException] | None, + data_required_exception: type[BaseException] | None, +) -> list[tuple[int, dict[str, float]]]: + """Draw up to one batch of candidates from the Ax client. + + Args: + ax_client: Ax client generating the candidates. + state: Mutable loop state whose cursor is advanced per draw. + total_trials: Cumulative trial budget for the run. + batch_size: Maximum number of candidates to draw at once. + max_parallelism_exception: Ax max-parallelism exception type, if importable. + data_required_exception: Ax data-required exception type, if importable. + + Returns: + The drawn candidates as ``(trial_index, parameters)`` pairs. + + Raises: + Exception: Any error from Ax that is not a known back-pressure signal. + """ + + candidates: list[tuple[int, dict[str, float]]] = [] + while len(candidates) < batch_size and state.trial_index_cursor < total_trials: + try: + raw_parameters, trial_index = ax_client.get_next_trial() + except Exception as error: + if max_parallelism_exception is not None and isinstance( + error, max_parallelism_exception + ): + break + if data_required_exception is not None and isinstance(error, data_required_exception): + break + raise + parameters = {name: float(value) for name, value in raw_parameters.items()} + candidates.append((int(trial_index), parameters)) + state.trial_index_cursor += 1 + return candidates + + +def run_ax_trial_loop( + *, + ax_client: Any, + config: Any, + state: TrialLoopState, + evaluate_candidates: Callable[ + [Sequence[tuple[int, Mapping[str, float]]]], list[TrialEvaluation] + ], + on_trial_completed: Callable[..., None], +) -> None: + """Run the Ax ask-and-tell loop shared by the measurement-fit optimizers. + + The loop owns candidate generation, Ax completion, best-so-far tracking, and + early stopping. Mode-specific evaluation and artifact persistence are supplied + by the ``evaluate_candidates`` and ``on_trial_completed`` callbacks. + + Args: + ax_client: Ax client used to generate and complete trials. + config: Configuration supplying ``total_trials``, ``batch_size``, + ``objective_name``, ``objective_sem``, and the early-stopping settings. + state: Mutable loop state, pre-populated when resuming a run. + evaluate_candidates: Callable evaluating one batch of candidates. + on_trial_completed: Callable invoked after each completed trial with + ``evaluation``, ``state``, and ``improved`` keyword arguments. + """ + + max_parallelism_exception = _import_max_parallelism_exception() + data_required_exception = _import_data_required_exception() + total_trials = int(config.total_trials) + minimum_relative_improvement = float( + getattr(config, "early_stopping_min_relative_improvement", 0.0) + ) + + while state.trial_index_cursor < total_trials: + candidates = _collect_candidates( + ax_client=ax_client, + state=state, + total_trials=total_trials, + batch_size=int(config.batch_size), + max_parallelism_exception=max_parallelism_exception, + data_required_exception=data_required_exception, + ) + if not candidates: + break + + for evaluation in evaluate_candidates(candidates): + state.resolved_max_workers = int(evaluation.resolved_max_workers) + _complete_ax_trial( + ax_client=ax_client, + config=config, + trial_index=evaluation.trial_index, + loss=float(evaluation.loss), + ) + state.trial_records.append( + TrialRecord( + trial_index=int(evaluation.trial_index), + loss=float(evaluation.loss), + parameters=dict(evaluation.parameters), + extras=_numeric_extras(evaluation.extras), + ) + ) + state.completed_trials += 1 + + previous_best_loss = state.best_loss + improved = bool(evaluation.loss < previous_best_loss) + if improved: + state.best_loss = float(evaluation.loss) + state.best_parameters = dict(evaluation.parameters) + state.best_extras = dict(evaluation.extras) + if is_significant_improvement( + previous_best_loss, + float(evaluation.loss), + minimum_relative_improvement, + ): + state.no_improvement_trials = 0 + else: + state.no_improvement_trials += 1 + + on_trial_completed(evaluation=evaluation, state=state, improved=improved) + + if ( + config.enable_early_stopping + and state.completed_trials >= config.early_stopping_warmup_trials + and state.no_improvement_trials >= config.early_stopping_patience + ): + state.stopped_early = True + state.early_stop_reason = ( + "Early stopping triggered after " + f"{state.no_improvement_trials} non-improving trials." + ) + break + if state.stopped_early: + break + + +def _numeric_extras(extras: Mapping[str, Any]) -> dict[str, float]: + """Return only the scalar entries of a trial extras payload. + + Args: + extras: Mode-specific payload attached to a trial evaluation. + + Returns: + The subset of entries that are finite-castable scalars, for CSV output. + """ + + numeric: dict[str, float] = {} + for name, value in extras.items(): + if isinstance(value, bool) or not isinstance(value, (int, float)): + continue + numeric[str(name)] = float(value) + return numeric diff --git a/src/grax_opt/optimize.py b/src/grax_opt/optimize.py index b8b4218..2fc7418 100644 --- a/src/grax_opt/optimize.py +++ b/src/grax_opt/optimize.py @@ -5,10 +5,12 @@ import concurrent.futures import csv import inspect +import io import os import platform +import tempfile import warnings -from dataclasses import dataclass +from dataclasses import dataclass, field from pathlib import Path from typing import Any @@ -247,11 +249,19 @@ def _evaluate_candidate_batch( @dataclass(frozen=True) class TrialRecord: - """Summary of one completed Ax trial.""" + """Summary of one completed Ax trial. + + Attributes: + trial_index: Ax trial index for the completed trial. + loss: Objective value observed for the trial. + parameters: Free parameter values evaluated for the trial. + extras: Optional scalar diagnostics such as per-measurement losses. + """ trial_index: int loss: float parameters: dict[str, float] + extras: dict[str, float] = field(default_factory=dict) @dataclass(frozen=True) @@ -333,11 +343,46 @@ def _import_data_required_exception(): return DataRequiredError +def _atomic_write_text(output_path: Path, text: str) -> None: + """Write text to a path atomically so a crash cannot truncate the file. + + Args: + output_path: Destination path to replace. + text: File contents to write. + """ + + output_path.parent.mkdir(parents=True, exist_ok=True) + handle = tempfile.NamedTemporaryFile( + "w", + encoding="utf-8", + newline="", + dir=str(output_path.parent), + prefix=f".{output_path.name}.", + suffix=".tmp", + delete=False, + ) + temporary_path = Path(handle.name) + try: + with handle: + handle.write(text) + handle.flush() + os.fsync(handle.fileno()) + os.replace(temporary_path, output_path) + except BaseException: + temporary_path.unlink(missing_ok=True) + raise + + def _write_trial_history_csv( trial_records: list[TrialRecord], output_path: Path, ) -> None: - """Write per-trial optimization history to CSV.""" + """Write per-trial optimization history to CSV. + + Args: + trial_records: Completed trial records to serialize. + output_path: Destination CSV path. + """ parameter_names: list[str] = [] for record in trial_records: @@ -345,17 +390,25 @@ def _write_trial_history_csv( if name not in parameter_names: parameter_names.append(name) - with output_path.open("w", newline="", encoding="utf-8") as handle: - writer = csv.writer(handle) - writer.writerow(["trial_index", "loss", *parameter_names]) - for record in trial_records: - writer.writerow( - [ - record.trial_index, - record.loss, - *[record.parameters.get(name, "") for name in parameter_names], - ] - ) + extra_names: list[str] = [] + for record in trial_records: + for name in record.extras: + if name not in extra_names: + extra_names.append(name) + + buffer = io.StringIO() + writer = csv.writer(buffer) + writer.writerow(["trial_index", "loss", *parameter_names, *extra_names]) + for record in trial_records: + writer.writerow( + [ + record.trial_index, + record.loss, + *[record.parameters.get(name, "") for name in parameter_names], + *[record.extras.get(name, "") for name in extra_names], + ] + ) + _atomic_write_text(output_path, buffer.getvalue()) def json_safe_grating_parameters(parameters: dict[str, object]) -> dict[str, object]: From 7d017cc2b616bd7be9256c9208432223c1e2ec55 Mon Sep 17 00:00:00 2001 From: Simone Vadilonga Date: Tue, 11 Aug 2026 08:58:03 +0200 Subject: [PATCH 02/10] Add joint multi-angle measurement fitting to grax_opt Fits one parameter set against several measured curves recorded at different grazing angles, each keeping its own energy grid. Simulated efficiencies are reassembled by CaseExecutionResult.index because the parallel batch runner yields in completion order. Co-Authored-By: Claude Sonnet 5 --- src/grax_opt/__init__.py | 29 +- src/grax_opt/joint.py | 918 ++++++++++++++++++++++++++++++++++++++ src/grax_opt/objective.py | 237 ++++++++++ 3 files changed, 1183 insertions(+), 1 deletion(-) create mode 100644 src/grax_opt/joint.py diff --git a/src/grax_opt/__init__.py b/src/grax_opt/__init__.py index a6a1e34..01e0b50 100644 --- a/src/grax_opt/__init__.py +++ b/src/grax_opt/__init__.py @@ -8,7 +8,13 @@ optimize_to_measurements as _optimize_to_measurements, resolve_measurement_fit_trial_parameters, ) -from .objective import build_evaluation_measurement, evaluate_trial +from .joint import ( + AngleMeasurementSpec, + JointMeasurementFitConfig, + JointOptimizationResult, + optimize_to_joint_measurements as _optimize_to_joint_measurements, +) +from .objective import build_evaluation_measurement, evaluate_trial, reduce_joint_losses from .optimize import OptimizationResult, TrialRecord, json_safe_grating_parameters def optimize_to_measurements(config): @@ -24,7 +30,26 @@ def optimize_to_measurements(config): return _optimize_to_measurements(config) +def optimize_to_joint_measurements(config): + """Fit one parameter set jointly against several measured curves. + + Each measurement is recorded at its own fixed grazing angle and keeps its + own energy grid. + + Args: + config: Spec mapping describing the joint optimization run. + + Returns: + JointOptimizationResult: Result bundle with persisted artifacts. + """ + + return _optimize_to_joint_measurements(config) + + __all__ = [ + "AngleMeasurementSpec", + "JointMeasurementFitConfig", + "JointOptimizationResult", "MeasurementFitConfig", "MeasurementData", "OptimizationResult", @@ -36,6 +61,8 @@ def optimize_to_measurements(config): "json_safe_grating_parameters", "load_measurement_data", "sample_measurement_data", + "optimize_to_joint_measurements", "optimize_to_measurements", + "reduce_joint_losses", "resolve_measurement_fit_trial_parameters", ] diff --git a/src/grax_opt/joint.py b/src/grax_opt/joint.py new file mode 100644 index 0000000..9e644c9 --- /dev/null +++ b/src/grax_opt/joint.py @@ -0,0 +1,918 @@ +"""Joint multi-angle measurement-fit optimization. + +Fits one parameter set simultaneously against several measured curves recorded +at different grazing angles. Each measurement keeps its own energy grid, and the +per-measurement losses are combined into a single joint objective. +""" + +from __future__ import annotations + +import csv +import inspect +import io +import json +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +import matplotlib.pyplot as plt +import numpy as np + +from grax.simulation import _resolve_max_workers as _resolve_simulation_max_workers + +from .config import ParameterBounds +from .data import load_measurement_data, sample_measurement_data +from .dynamic import ( + _default_solver_parameter_resolver, + _normalize_constraints, + _normalize_parameter_bounds, + _validate_constraint_graph, + build_free_parameter_names, + build_measurement_fit_ax_parameters, + resolve_measurement_fit_trial_parameters, +) +from .loop import TrialEvaluation, TrialLoopState, run_ax_trial_loop +from .objective import evaluate_joint_trial_with_metadata +from .optimize import ( + TrialRecord, + _atomic_write_text, + _describe_optimizer_compute_context, + _import_ax_client, + _import_objective_properties, + _patch_torch_fork_rng_for_cpu_only, + _resolve_optimizer_backend, + _save_loss_history_plot, + _write_trial_history_csv, + json_safe_grating_parameters, +) + +JOINT_LOSS_REDUCTIONS = frozenset({"mean", "sum", "pooled", "weighted"}) + + +@dataclass(frozen=True) +class AngleMeasurementSpec: + """One measured curve at a fixed grazing angle. + + Attributes: + grazing_angle_deg: Fixed grazing angle for this measurement. + measurement_path: Path to the measured two-column dataset. + evaluation_energies_ev: Energies used for evaluation. When empty, the + measurement's own energy grid is used. + measurement_efficiency: Optional measured efficiencies to use directly + instead of interpolating the file. Supply this when the values were + prepared upstream, for example by smoothing or downsampling. + weight: Relative weight used by the ``"weighted"`` joint reduction. + label: Identifier used in artifacts. Defaults to the grazing angle. + """ + + grazing_angle_deg: float + measurement_path: Path + evaluation_energies_ev: list[float] = field(default_factory=list) + measurement_efficiency: list[float] | None = None + weight: float = 1.0 + label: str | None = None + + def __post_init__(self) -> None: + """Normalize paths and validate the measurement definition.""" + + object.__setattr__(self, "measurement_path", Path(self.measurement_path)) + object.__setattr__( + self, + "evaluation_energies_ev", + [float(energy_ev) for energy_ev in self.evaluation_energies_ev], + ) + if self.measurement_efficiency is not None: + object.__setattr__( + self, + "measurement_efficiency", + [float(value) for value in self.measurement_efficiency], + ) + if self.label is None: + object.__setattr__(self, "label", f"alpha{float(self.grazing_angle_deg):g}deg") + + if float(self.grazing_angle_deg) <= 0.0: + raise ValueError("grazing_angle_deg must be > 0.") + if any(energy_ev <= 0.0 for energy_ev in self.evaluation_energies_ev): + raise ValueError("evaluation_energies_ev values must be > 0.") + if not float(self.weight) > 0.0: + raise ValueError("weight must be > 0.") + if self.measurement_efficiency is not None: + if len(self.measurement_efficiency) == 0: + raise ValueError("measurement_efficiency must not be empty when provided.") + if ( + len(self.evaluation_energies_ev) > 0 + and len(self.measurement_efficiency) != len(self.evaluation_energies_ev) + ): + raise ValueError( + "measurement_efficiency must have the same length as evaluation_energies_ev." + ) + + @classmethod + def from_mapping(cls, mapping: Mapping[str, object]) -> AngleMeasurementSpec: + """Build a measurement spec from a plain mapping. + + Args: + mapping: Spec mapping describing one angle's measurement. + + Returns: + The normalized measurement spec. + + Raises: + ValueError: If required keys are missing or unexpected keys remain. + """ + + spec = dict(mapping) + if "grazing_angle_deg" not in spec: + raise ValueError("Each joint measurement requires 'grazing_angle_deg'.") + if "measurement_path" not in spec: + raise ValueError("Each joint measurement requires 'measurement_path'.") + instance = cls( + grazing_angle_deg=float(spec.pop("grazing_angle_deg")), + measurement_path=Path(spec.pop("measurement_path")), # type: ignore[arg-type] + evaluation_energies_ev=list(spec.pop("evaluation_energies_ev", []) or []), + measurement_efficiency=( + None + if spec.get("measurement_efficiency") is None + else list(spec.pop("measurement_efficiency")) # type: ignore[arg-type] + ), + weight=float(spec.pop("weight", 1.0)), + label=spec.pop("label", None), # type: ignore[arg-type] + ) + spec.pop("measurement_efficiency", None) + if spec: + raise ValueError(f"Unexpected joint measurement keys: {sorted(spec)}") + return instance + + +@dataclass(frozen=True) +class JointAngleMeasurement: + """A measurement spec resolved onto its evaluation grid. + + Attributes: + label: Identifier used in artifacts. + grazing_angle_deg: Fixed grazing angle for this measurement. + measurement_path: Path the measurement was loaded from. + evaluation_energies_ev: Energies used for evaluation. + evaluation_efficiency: Measured efficiencies on the evaluation grid. + weight: Relative weight for the ``"weighted"`` joint reduction. + """ + + label: str + grazing_angle_deg: float + measurement_path: Path + evaluation_energies_ev: np.ndarray + evaluation_efficiency: np.ndarray + weight: float + + +def prepare_joint_measurements( + measurements: Sequence[AngleMeasurementSpec], +) -> list[JointAngleMeasurement]: + """Resolve every measurement spec onto its evaluation grid. + + Args: + measurements: Normalized measurement specs. + + Returns: + One resolved measurement per spec, in the same order. + + Raises: + ValueError: If supplied efficiencies do not match the resolved grid. + """ + + prepared: list[JointAngleMeasurement] = [] + for spec in measurements: + measurement = load_measurement_data(spec.measurement_path) + if len(spec.evaluation_energies_ev) > 0: + energies = np.asarray(spec.evaluation_energies_ev, dtype=float) + if spec.measurement_efficiency is None: + sampled = sample_measurement_data(measurement, energies) + efficiencies = np.asarray(sampled.efficiency, dtype=float) + else: + efficiencies = np.asarray(spec.measurement_efficiency, dtype=float) + else: + energies = np.asarray(measurement.energy_ev, dtype=float) + if spec.measurement_efficiency is None: + efficiencies = np.asarray(measurement.efficiency, dtype=float) + else: + efficiencies = np.asarray(spec.measurement_efficiency, dtype=float) + if efficiencies.shape != energies.shape: + raise ValueError( + "measurement_efficiency must have the same length as the measurement " + f"energy grid for {spec.label!r}." + ) + prepared.append( + JointAngleMeasurement( + label=str(spec.label), + grazing_angle_deg=float(spec.grazing_angle_deg), + measurement_path=spec.measurement_path, + evaluation_energies_ev=energies, + evaluation_efficiency=efficiencies, + weight=float(spec.weight), + ) + ) + return prepared + + +@dataclass(frozen=True) +class JointMeasurementFitConfig: + """Configuration for a joint multi-angle measurement fit. + + Attributes: + build_grating: Callable building the grating from resolved parameters. + parameter_bounds: Mapping of parameter names to lower/upper bounds. + output_dir: Directory where optimizer artifacts are written. + measurements: Per-angle measurement specs to fit jointly. + diffraction_order: Diffraction order selected for evaluation. + fourier_orders: Fourier orders used by the solver. + roughness_sigma_nm: Optional roughness passed to the solver. + validate_physical_results: Whether to validate simulated results. + total_trials: Cumulative trial budget across resumed runs. + batch_size: Number of candidates generated per Ax batch. + random_seed: Optional Ax random seed. + equality_constraints: Mapping of tied parameter names to their source. + objective_name: Ax objective name. + experiment_name: Ax experiment name. + failure_penalty: Loss reported for failed trials. + objective_sem: Standard error reported with each observation. + enable_early_stopping: Whether early stopping is active. + early_stopping_patience: Non-improving trials tolerated before stopping. + early_stopping_min_relative_improvement: Minimum relative gain counted + as an improvement. + early_stopping_warmup_trials: Trials completed before stopping applies. + joint_loss_reduction: How per-measurement losses are combined. + save_best_fit_plot: Whether to write the multi-panel best-fit plot. + save_loss_plot: Whether to write the loss-history plot. + save_comparison_csv: Whether to write the long-form comparison CSV. + backend: Requested RCWA backend. + max_workers: Trial-level worker count for the batch runner. + solver_parameter_resolver: Optional solver-parameter hook. + resume: Whether to resume from a previous checkpoint. + checkpoint_dir: Checkpoint directory. Defaults to ``output_dir/checkpoint``. + checkpoint_interval: Trials between checkpoint flushes. + """ + + build_grating: Any + parameter_bounds: Mapping[str, ParameterBounds | Sequence[float]] + output_dir: Path + measurements: Sequence[AngleMeasurementSpec | Mapping[str, object]] = field( + default_factory=list + ) + diffraction_order: int = 1 + fourier_orders: int = 20 + roughness_sigma_nm: float | None = None + validate_physical_results: bool = True + total_trials: int = 20 + batch_size: int = 1 + random_seed: int | None = None + equality_constraints: Mapping[str, str] = field(default_factory=dict) + objective_name: str = "joint_loss" + experiment_name: str = "joint_measurement_fit" + failure_penalty: float = 1.0e6 + objective_sem: float = 1.0e-6 + enable_early_stopping: bool = False + early_stopping_patience: int = 8 + early_stopping_min_relative_improvement: float = 5.0e-3 + early_stopping_warmup_trials: int = 8 + joint_loss_reduction: str = "mean" + save_best_fit_plot: bool = True + save_loss_plot: bool = True + save_comparison_csv: bool = True + backend: str = "auto" + max_workers: int | str | None = None + solver_parameter_resolver: Any = None + resume: bool = False + checkpoint_dir: Path | None = None + checkpoint_interval: int = 1 + + def __post_init__(self) -> None: + """Normalize paths, bounds, measurements, and validate settings.""" + + object.__setattr__(self, "output_dir", Path(self.output_dir)) + object.__setattr__( + self, + "parameter_bounds", + _normalize_parameter_bounds(self.parameter_bounds), + ) + object.__setattr__( + self, + "equality_constraints", + _normalize_constraints(self.equality_constraints), + ) + object.__setattr__( + self, + "measurements", + [ + item + if isinstance(item, AngleMeasurementSpec) + else AngleMeasurementSpec.from_mapping(item) + for item in self.measurements + ], + ) + if self.checkpoint_dir is not None: + object.__setattr__(self, "checkpoint_dir", Path(self.checkpoint_dir)) + + if not callable(self.build_grating): + raise ValueError("build_grating must be callable.") + if len(self.measurements) == 0: + raise ValueError("measurements must be provided and non-empty.") + labels = [str(spec.label) for spec in self.measurements] + if len(set(labels)) != len(labels): + raise ValueError("measurements must have unique labels.") + if self.diffraction_order <= 0: + raise ValueError("diffraction_order must be > 0.") + if self.fourier_orders <= 0: + raise ValueError("fourier_orders must be > 0.") + if self.roughness_sigma_nm is not None and self.roughness_sigma_nm < 0.0: + raise ValueError("roughness_sigma_nm must be >= 0 when provided.") + if self.total_trials <= 0: + raise ValueError("total_trials must be > 0.") + if self.batch_size <= 0: + raise ValueError("batch_size must be > 0.") + resolved_max_workers = _resolve_simulation_max_workers(self.max_workers) + if resolved_max_workers > 1 and self.batch_size > 1: + raise ValueError( + "batch_size > 1 cannot be combined with optimizer max_workers > 1. " + "Use trial-level multiprocessing or candidate batching, but not both." + ) + if self.failure_penalty <= 0.0: + raise ValueError("failure_penalty must be > 0.") + if not np.isfinite(self.objective_sem) or self.objective_sem <= 0.0: + raise ValueError("objective_sem must be finite and > 0.") + if self.early_stopping_patience <= 0: + raise ValueError("early_stopping_patience must be > 0.") + if ( + not np.isfinite(self.early_stopping_min_relative_improvement) + or self.early_stopping_min_relative_improvement < 0.0 + ): + raise ValueError( + "early_stopping_min_relative_improvement must be finite and >= 0." + ) + if self.early_stopping_warmup_trials < 0: + raise ValueError("early_stopping_warmup_trials must be >= 0.") + if self.joint_loss_reduction not in JOINT_LOSS_REDUCTIONS: + raise ValueError( + "joint_loss_reduction must be one of 'mean', 'sum', 'pooled', or 'weighted'." + ) + if self.backend not in {"auto", "numba", "numpy"}: + raise ValueError("backend must be one of 'auto', 'numba', or 'numpy'.") + if self.checkpoint_interval <= 0: + raise ValueError("checkpoint_interval must be > 0.") + if self.solver_parameter_resolver is not None and not callable( + self.solver_parameter_resolver + ): + raise ValueError("solver_parameter_resolver must be callable when provided.") + + _validate_constraint_graph( + parameter_bounds=self.parameter_bounds, + equality_constraints=self.equality_constraints, + ) + if not build_free_parameter_names(self): + raise ValueError("At least one parameter must remain free for optimization.") + + @classmethod + def from_mapping(cls, mapping: Mapping[str, object]) -> JointMeasurementFitConfig: + """Build a joint configuration from a plain spec mapping. + + Args: + mapping: Spec mapping describing the joint optimization run. + + Returns: + The normalized joint configuration. + + Raises: + ValueError: If required keys are missing or unexpected keys remain. + """ + + config = dict(mapping) + if "build_grating" not in config: + raise ValueError("Joint measurement-fit spec requires 'build_grating'.") + parameter_bounds = config.pop("parameter_bounds", None) + if parameter_bounds is None: + raise ValueError("Joint measurement-fit spec requires 'parameter_bounds'.") + measurements = config.pop("measurements", None) + if measurements is None: + raise ValueError("Joint measurement-fit spec requires 'measurements'.") + equality_constraints = config.pop("equality_constraints", None) or {} + + instance = cls( + build_grating=config.pop("build_grating"), + parameter_bounds=parameter_bounds, # type: ignore[arg-type] + output_dir=config.pop("output_dir"), # type: ignore[arg-type] + measurements=list(measurements), # type: ignore[arg-type] + diffraction_order=int(config.pop("diffraction_order", 1)), + fourier_orders=int(config.pop("fourier_orders", 20)), + roughness_sigma_nm=config.pop("roughness_sigma_nm", None), # type: ignore[arg-type] + validate_physical_results=bool(config.pop("validate_physical_results", True)), + total_trials=int(config.pop("total_trials", 20)), + batch_size=int(config.pop("batch_size", 1)), + random_seed=config.pop("random_seed", None), # type: ignore[arg-type] + equality_constraints=equality_constraints, # type: ignore[arg-type] + objective_name=str(config.pop("objective_name", "joint_loss")), + experiment_name=str(config.pop("experiment_name", "joint_measurement_fit")), + failure_penalty=float(config.pop("failure_penalty", 1.0e6)), + objective_sem=float(config.pop("objective_sem", 1.0e-6)), + enable_early_stopping=bool(config.pop("enable_early_stopping", False)), + early_stopping_patience=int(config.pop("early_stopping_patience", 8)), + early_stopping_min_relative_improvement=float( + config.pop("early_stopping_min_relative_improvement", 5.0e-3) + ), + early_stopping_warmup_trials=int(config.pop("early_stopping_warmup_trials", 8)), + joint_loss_reduction=str(config.pop("joint_loss_reduction", "mean")), + save_best_fit_plot=bool(config.pop("save_best_fit_plot", True)), + save_loss_plot=bool(config.pop("save_loss_plot", True)), + save_comparison_csv=bool(config.pop("save_comparison_csv", True)), + backend=str(config.pop("backend", "auto")), + max_workers=config.pop("max_workers", None), # type: ignore[arg-type] + solver_parameter_resolver=config.pop("solver_parameter_resolver", None), + resume=bool(config.pop("resume", False)), + checkpoint_dir=config.pop("checkpoint_dir", None), # type: ignore[arg-type] + checkpoint_interval=int(config.pop("checkpoint_interval", 1)), + ) + if config: + raise ValueError(f"Unexpected joint measurement-fit spec keys: {sorted(config)}") + return instance + + +@dataclass(frozen=True) +class JointOptimizationResult: + """Result bundle returned by the joint multi-angle optimizer. + + Attributes: + best_parameters: Best free parameters returned by Ax. + best_grating_parameters: Best resolved grating parameters. + best_loss: Best joint objective value found. + per_measurement_best_losses: Per-measurement losses for the best trial. + measurements: Resolved measurements used for the fit. + result_json_path: Path to the persisted JSON summary. + trial_history_csv_path: Path to the per-trial history CSV. + best_fit_plot_path: Path to the multi-panel best-fit plot, if written. + loss_history_plot_path: Path to the loss-history plot, if written. + comparison_csv_path: Path to the long-form comparison CSV, if written. + trial_records: Per-trial records including per-measurement losses. + stopped_early: Whether early stopping ended the run. + completed_trials: Number of trials successfully evaluated. + early_stop_reason: Human-readable early-stopping reason, or ``None``. + """ + + best_parameters: dict[str, float] + best_grating_parameters: dict[str, object] + best_loss: float + per_measurement_best_losses: dict[str, float] + measurements: list[JointAngleMeasurement] + result_json_path: Path + trial_history_csv_path: Path + best_fit_plot_path: Path | None + loss_history_plot_path: Path | None + comparison_csv_path: Path | None + trial_records: list[TrialRecord] + stopped_early: bool + completed_trials: int + early_stop_reason: str | None + + +def _resolve_joint_solver_parameters( + config: JointMeasurementFitConfig, + resolved_parameters: Mapping[str, float], +) -> dict[str, float | None]: + """Resolve solver parameters for one joint trial. + + Args: + config: Joint optimization configuration. + resolved_parameters: Fully expanded grating parameters. + + Returns: + Solver parameters for the trial. + """ + + if config.solver_parameter_resolver is not None: + return dict(config.solver_parameter_resolver(resolved_parameters)) + return _default_solver_parameter_resolver(resolved_parameters) + + +def _create_ax_client_for_joint_config(config: JointMeasurementFitConfig) -> Any: + """Create and configure an Ax client for a joint optimization run. + + Args: + config: Joint optimization configuration. + + Returns: + A configured Ax client with the experiment created. + + Raises: + RuntimeError: If the installed Ax client API is unsupported. + """ + + _patch_torch_fork_rng_for_cpu_only() + print(_describe_optimizer_compute_context()) + ax_client_cls = _import_ax_client() + + client_kwargs: dict[str, object] = {} + client_signature = inspect.signature(ax_client_cls) + if config.random_seed is not None and "random_seed" in client_signature.parameters: + client_kwargs["random_seed"] = config.random_seed + ax_client = ax_client_cls(**client_kwargs) + + create_signature = inspect.signature(ax_client.create_experiment) + create_kwargs: dict[str, object] = { + "parameters": build_measurement_fit_ax_parameters(config), + "name": config.experiment_name, + } + if "objective_name" in create_signature.parameters: + create_kwargs["objective_name"] = config.objective_name + if "minimize" in create_signature.parameters: + create_kwargs["minimize"] = True + elif "objectives" in create_signature.parameters: + objective_properties = _import_objective_properties() + create_kwargs["objectives"] = {config.objective_name: objective_properties(minimize=True)} + else: + raise RuntimeError("Unsupported Ax client create_experiment signature.") + ax_client.create_experiment(**create_kwargs) + return ax_client + + +def _save_joint_best_fit_plot( + *, + measurements: Sequence[JointAngleMeasurement], + simulated_by_label: Mapping[str, np.ndarray], + output_path: Path, +) -> None: + """Save one measurement-vs-simulation panel per angle. + + Args: + measurements: Resolved measurements used for the fit. + simulated_by_label: Best-fit simulated curves keyed by label. + output_path: Destination image path. + """ + + figure, axes = plt.subplots( + len(measurements), + 1, + figsize=(10, 4.0 * len(measurements)), + squeeze=False, + ) + for axis, measurement in zip(axes[:, 0], measurements, strict=True): + axis.plot( + measurement.evaluation_energies_ev, + measurement.evaluation_efficiency, + "o-", + linewidth=1.0, + label="Measurement", + ) + simulated = simulated_by_label.get(measurement.label) + if simulated is not None: + axis.plot( + measurement.evaluation_energies_ev, + simulated, + "s-", + linewidth=1.0, + label="Best fit", + ) + axis.set_xlabel("Photon Energy (eV)") + axis.set_ylabel("Diffraction Efficiency") + axis.set_title( + f"Joint Best Fit at Grazing Angle = {measurement.grazing_angle_deg:g} deg" + ) + axis.grid(True, alpha=0.3) + axis.legend(loc="best") + figure.tight_layout() + figure.savefig(output_path, dpi=150, bbox_inches="tight") + plt.close(figure) + + +def _write_joint_comparison_csv( + *, + measurements: Sequence[JointAngleMeasurement], + simulated_by_label: Mapping[str, np.ndarray], + diffraction_order: int, + output_path: Path, +) -> None: + """Write a long-form measured-versus-simulated comparison table. + + Args: + measurements: Resolved measurements used for the fit. + simulated_by_label: Best-fit simulated curves keyed by label. + diffraction_order: Diffraction order the fit was evaluated at. + output_path: Destination CSV path. + """ + + buffer = io.StringIO() + writer = csv.writer(buffer) + writer.writerow( + [ + "label", + "grazing_angle_deg", + "energy_ev", + "measured_efficiency", + "simulated_efficiency", + "diffraction_order", + ] + ) + for measurement in measurements: + simulated = simulated_by_label.get(measurement.label) + for point_index, energy_ev in enumerate(measurement.evaluation_energies_ev): + writer.writerow( + [ + measurement.label, + measurement.grazing_angle_deg, + float(energy_ev), + float(measurement.evaluation_efficiency[point_index]), + "" if simulated is None else float(simulated[point_index]), + int(diffraction_order), + ] + ) + _atomic_write_text(output_path, buffer.getvalue()) + + +def _write_joint_result_json( + *, + config: JointMeasurementFitConfig, + measurements: Sequence[JointAngleMeasurement], + state: TrialLoopState, + backend_requested: str, + backend_effective: str, + output_path: Path, +) -> None: + """Write the joint optimizer JSON summary. + + Args: + config: Joint optimization configuration. + measurements: Resolved measurements used for the fit. + state: Loop state carrying the best-so-far results. + backend_requested: Backend requested by the caller. + backend_effective: Backend actually used. + output_path: Destination JSON path. + """ + + payload: dict[str, object] = { + "optimization_mode": "joint_measurement_fit", + "experiment_name": config.experiment_name, + "objective_name": config.objective_name, + "joint_loss_reduction": config.joint_loss_reduction, + "measurements": [ + { + "label": measurement.label, + "grazing_angle_deg": measurement.grazing_angle_deg, + "measurement_path": str(measurement.measurement_path), + "evaluation_energies_ev": [ + float(energy_ev) for energy_ev in measurement.evaluation_energies_ev + ], + "point_count": int(len(measurement.evaluation_energies_ev)), + "weight": measurement.weight, + } + for measurement in measurements + ], + "parameter_bounds": { + name: [bounds.lower, bounds.upper] for name, bounds in config.parameter_bounds.items() + }, + "equality_constraints": dict(config.equality_constraints), + "best_loss": state.best_loss, + "per_measurement_best_losses": dict(state.best_extras.get("per_measurement_losses", {})), + "best_parameters": dict(state.best_parameters), + "best_grating_parameters": json_safe_grating_parameters(state.best_grating_parameters), + "best_solver_parameters": dict(state.best_solver_parameters), + "stopped_early": state.stopped_early, + "completed_trials": state.completed_trials, + "early_stop_reason": state.early_stop_reason, + "backend_requested": backend_requested, + "backend_effective": backend_effective, + "optimizer_execution_strategy": "trial_batch_runner", + "optimizer_requested_max_workers": config.max_workers, + "optimizer_resolved_max_workers": state.resolved_max_workers, + "diffraction_order": config.diffraction_order, + "fourier_orders": config.fourier_orders, + } + _atomic_write_text(output_path, json.dumps(payload, indent=2)) + + +def _persist_joint_optimizer_artifacts( + *, + config: JointMeasurementFitConfig, + measurements: Sequence[JointAngleMeasurement], + state: TrialLoopState, + backend_requested: str, + backend_effective: str, + write_heavy_artifacts: bool, +) -> tuple[Path, Path, Path | None, Path | None, Path | None]: + """Persist joint optimizer artifacts and return their paths. + + Args: + config: Joint optimization configuration. + measurements: Resolved measurements used for the fit. + state: Loop state carrying the best-so-far results. + backend_requested: Backend requested by the caller. + backend_effective: Backend actually used. + write_heavy_artifacts: Whether to rewrite plots and the comparison CSV. + + Returns: + Paths to the JSON summary, trial-history CSV, best-fit plot, loss-history + plot, and comparison CSV. Optional paths are ``None`` when disabled. + """ + + config.output_dir.mkdir(parents=True, exist_ok=True) + result_json_path = config.output_dir / "best_result.json" + trial_history_csv_path = config.output_dir / "trial_history.csv" + best_fit_plot_path = config.output_dir / "best_fit.png" if config.save_best_fit_plot else None + loss_history_plot_path = ( + config.output_dir / "optimization_loss_history.png" if config.save_loss_plot else None + ) + comparison_csv_path = ( + config.output_dir / "best_fit_comparison.csv" if config.save_comparison_csv else None + ) + + _write_joint_result_json( + config=config, + measurements=measurements, + state=state, + backend_requested=backend_requested, + backend_effective=backend_effective, + output_path=result_json_path, + ) + _write_trial_history_csv(state.trial_records, trial_history_csv_path) + + simulated_by_label = state.best_extras.get("simulated_by_label", {}) + if write_heavy_artifacts: + if best_fit_plot_path is not None and simulated_by_label: + _save_joint_best_fit_plot( + measurements=measurements, + simulated_by_label=simulated_by_label, + output_path=best_fit_plot_path, + ) + if comparison_csv_path is not None and simulated_by_label: + _write_joint_comparison_csv( + measurements=measurements, + simulated_by_label=simulated_by_label, + diffraction_order=config.diffraction_order, + output_path=comparison_csv_path, + ) + if loss_history_plot_path is not None and state.trial_records: + _save_loss_history_plot( + trial_records=state.trial_records, + output_path=loss_history_plot_path, + stopped_early=state.stopped_early, + ) + + return ( + result_json_path, + trial_history_csv_path, + best_fit_plot_path, + loss_history_plot_path, + comparison_csv_path, + ) + + +def optimize_to_joint_measurements( + config: JointMeasurementFitConfig | Mapping[str, object], +) -> JointOptimizationResult: + """Fit one parameter set jointly against several measured curves. + + Each measurement is recorded at its own fixed grazing angle and keeps its own + energy grid. Every trial evaluates all measurements in a single batch, and the + per-measurement losses are combined using ``joint_loss_reduction``. + + Args: + config: Joint configuration, or a spec mapping describing the run. + Required keys: ``build_grating``, ``parameter_bounds``, + ``output_dir``, and ``measurements``. + + Returns: + JointOptimizationResult: Result bundle with persisted artifact paths. + + Raises: + RuntimeError: If the optimization produced no completed trials. + """ + + if not isinstance(config, JointMeasurementFitConfig): + config = JointMeasurementFitConfig.from_mapping(config) + + backend_effective = _resolve_optimizer_backend(config.backend) + measurements = prepare_joint_measurements(config.measurements) + ax_client = _create_ax_client_for_joint_config(config) + + state = TrialLoopState() + state.resolved_max_workers = _resolve_simulation_max_workers(config.max_workers) + + build_grating_fn = lambda trial_parameters: config.build_grating( + resolve_measurement_fit_trial_parameters(config, trial_parameters) + ) + resolve_solver_parameters_fn = lambda trial_parameters: _resolve_joint_solver_parameters( + config, + resolve_measurement_fit_trial_parameters(config, trial_parameters), + ) + + def evaluate_candidates(candidates) -> list[TrialEvaluation]: + """Evaluate one batch of joint candidates. + + Args: + candidates: Candidate ``(trial_index, parameters)`` pairs. + + Returns: + One evaluation per candidate. + """ + + evaluations: list[TrialEvaluation] = [] + for trial_index, parameters in candidates: + ( + joint_loss, + per_measurement_losses, + simulated_by_label, + resolved_max_workers, + ) = evaluate_joint_trial_with_metadata( + config, + parameters, + measurements, + backend=backend_effective, + build_grating_fn=build_grating_fn, + resolve_solver_parameters_fn=resolve_solver_parameters_fn, + ) + extras: dict[str, Any] = { + "per_measurement_losses": per_measurement_losses, + "simulated_by_label": simulated_by_label, + } + for label, loss in per_measurement_losses.items(): + extras[f"loss_{label}"] = float(loss) + evaluations.append( + TrialEvaluation( + trial_index=int(trial_index), + parameters=dict(parameters), + loss=float(joint_loss), + resolved_max_workers=int(resolved_max_workers), + extras=extras, + ) + ) + return evaluations + + def on_trial_completed( + *, evaluation: TrialEvaluation, state: TrialLoopState, improved: bool + ) -> None: + """Refresh derived best state and rewrite joint artifacts. + + Args: + evaluation: Evaluation for the trial that just completed. + state: Mutable loop state to update. + improved: Whether this trial produced a new best joint loss. + """ + + if improved: + state.best_grating_parameters = dict( + resolve_measurement_fit_trial_parameters(config, evaluation.parameters) + ) + state.best_solver_parameters = dict( + _resolve_joint_solver_parameters(config, state.best_grating_parameters) + ) + _persist_joint_optimizer_artifacts( + config=config, + measurements=measurements, + state=state, + backend_requested=config.backend, + backend_effective=backend_effective, + write_heavy_artifacts=improved, + ) + + run_ax_trial_loop( + ax_client=ax_client, + config=config, + state=state, + evaluate_candidates=evaluate_candidates, + on_trial_completed=on_trial_completed, + ) + + if not state.trial_records: + raise RuntimeError("Joint optimization produced no completed trials.") + + if not state.best_parameters: + state.best_parameters = dict(state.trial_records[-1].parameters) + state.best_grating_parameters = dict( + resolve_measurement_fit_trial_parameters(config, state.best_parameters) + ) + state.best_solver_parameters = dict( + _resolve_joint_solver_parameters(config, state.best_grating_parameters) + ) + state.best_loss = float(state.trial_records[-1].loss) + + result_paths = _persist_joint_optimizer_artifacts( + config=config, + measurements=measurements, + state=state, + backend_requested=config.backend, + backend_effective=backend_effective, + write_heavy_artifacts=True, + ) + + return JointOptimizationResult( + best_parameters=state.best_parameters, + best_grating_parameters=state.best_grating_parameters, + best_loss=state.best_loss, + per_measurement_best_losses=dict(state.best_extras.get("per_measurement_losses", {})), + measurements=measurements, + result_json_path=result_paths[0], + trial_history_csv_path=result_paths[1], + best_fit_plot_path=result_paths[2], + loss_history_plot_path=result_paths[3], + comparison_csv_path=result_paths[4], + trial_records=state.trial_records, + stopped_early=state.stopped_early, + completed_trials=state.completed_trials, + early_stop_reason=state.early_stop_reason, + ) diff --git a/src/grax_opt/objective.py b/src/grax_opt/objective.py index 25df687..083f10e 100644 --- a/src/grax_opt/objective.py +++ b/src/grax_opt/objective.py @@ -2,7 +2,9 @@ from __future__ import annotations +from collections.abc import Sequence from typing import Any, Callable, Dict, Mapping, Optional +import logging import warnings import numpy as np @@ -13,6 +15,8 @@ from .data import MeasurementData, sample_measurement_data from .evaluation import build_evaluation_cases +module_logger = logging.getLogger(__name__) + LossFunction = Callable[[np.ndarray, np.ndarray], float] BuildGratingFunction = Callable[[Mapping[str, float]], object] ResolveSolverParametersFunction = Callable[[Mapping[str, float]], Dict[str, Optional[float]]] @@ -200,6 +204,239 @@ def simulate_efficiency_curve( return efficiencies +def reduce_joint_losses( + per_measurement_losses: Mapping[str, float], + *, + reduction: str = "mean", + weights: Mapping[str, float] | None = None, + point_counts: Mapping[str, int] | None = None, +) -> float: + """Combine per-measurement losses into one joint objective value. + + Args: + per_measurement_losses: Loss for each measurement, keyed by label. + reduction: One of ``"mean"``, ``"sum"``, ``"pooled"``, or ``"weighted"``. + weights: Explicit per-measurement weights, required for ``"weighted"``. + point_counts: Evaluation point count per measurement, required for + ``"pooled"``. + + Returns: + The reduced joint loss. + + Raises: + ValueError: If the losses are empty, the reduction is unknown, or the + inputs required by the chosen reduction are missing. + """ + + if len(per_measurement_losses) == 0: + raise ValueError("per_measurement_losses must not be empty.") + + labels = list(per_measurement_losses) + losses = np.asarray([float(per_measurement_losses[label]) for label in labels], dtype=float) + + if reduction == "sum": + return float(np.sum(losses)) + if reduction == "mean": + weight_values = np.ones(len(labels), dtype=float) + elif reduction == "pooled": + if point_counts is None: + raise ValueError("point_counts is required for the 'pooled' reduction.") + weight_values = np.asarray( + [float(point_counts[label]) for label in labels], + dtype=float, + ) + elif reduction == "weighted": + if weights is None: + raise ValueError("weights is required for the 'weighted' reduction.") + weight_values = np.asarray([float(weights[label]) for label in labels], dtype=float) + else: + raise ValueError( + "joint_loss_reduction must be one of 'mean', 'sum', 'pooled', or 'weighted'." + ) + + weight_total = float(np.sum(weight_values)) + if weight_total <= 0.0: + raise ValueError("Joint loss reduction weights must sum to a positive value.") + return float(np.sum(weight_values * losses) / weight_total) + + +def simulate_joint_efficiency_curves_with_metadata( + config: Any, + trial_parameters: Mapping[str, float], + joint_measurements: Sequence[Any], + *, + backend: str, + build_grating_fn: BuildGratingFunction | None = None, + resolve_solver_parameters_fn: ResolveSolverParametersFunction | None = None, +) -> tuple[dict[str, np.ndarray], int]: + """Simulate efficiency curves for every measurement in one flat batch. + + All measurements are evaluated in a single :class:`BatchSimulationRunner` + batch so trial-level ``max_workers`` parallelizes across angles as well as + energies. Results are reassembled by ``result.index`` because the parallel + runner yields in completion order rather than input order. + + Args: + config: Joint optimization configuration describing the simulation setup. + trial_parameters: Ax trial parameters for the current candidate. + joint_measurements: Prepared per-angle measurements to evaluate. + backend: RCWA backend to use for the simulation. + build_grating_fn: Hook that builds a grating from the trial parameters. + resolve_solver_parameters_fn: Hook that resolves solver parameters. + + Returns: + A mapping of measurement label to simulated efficiencies, and the + runner's resolved worker count. + + Raises: + RuntimeError: If the required build hooks are missing. + _BatchCaseFailure: If any case fails or no result is returned for a case. + """ + + if build_grating_fn is None: + raise RuntimeError("build_grating_fn is required for joint optimizer execution.") + if resolve_solver_parameters_fn is None: + raise RuntimeError( + "resolve_solver_parameters_fn is required for joint optimizer execution." + ) + grating = build_grating_fn(trial_parameters) + solver_parameters = resolve_solver_parameters_fn(trial_parameters) + + cases: list[dict[str, object]] = [] + case_slots: list[tuple[str, int]] = [] + for joint_measurement in joint_measurements: + label = str(joint_measurement.label) + for point_index, energy_ev in enumerate(joint_measurement.evaluation_energies_ev): + cases.append( + { + "case_id": f"trial_eval_{label}_{point_index}", + "grating": grating, + "energy_ev": float(energy_ev), + "grazing_angle_deg": float(joint_measurement.grazing_angle_deg), + "diffraction_order": int(config.diffraction_order), + "fourier_orders": int(config.fourier_orders), + "roughness_sigma_nm": solver_parameters["roughness_sigma_nm"], + } + ) + case_slots.append((label, point_index)) + + runner = BatchSimulationRunner( + default_diffraction_order=int(config.diffraction_order), + default_fourier_orders=int(config.fourier_orders), + max_workers=getattr(config, "max_workers", None), + validate_physical_results=bool(config.validate_physical_results), + backend=backend, + ) + + simulated: dict[str, np.ndarray] = { + str(joint_measurement.label): np.full( + len(joint_measurement.evaluation_energies_ev), + np.nan, + dtype=float, + ) + for joint_measurement in joint_measurements + } + filled = np.zeros(len(cases), dtype=bool) + for result in runner.run_cases(cases): + if result.status != "ok": + raise _BatchCaseFailure(result.case_id, result.status, int(runner.resolved_max_workers)) + flat_index = int(result.index) + label, point_index = case_slots[flat_index] + simulated[label][point_index] = float(result.selected_efficiency) + filled[flat_index] = True + + if not bool(filled.all()): + missing_index = int(np.argmin(filled)) + raise _BatchCaseFailure( + str(cases[missing_index]["case_id"]), + "missing", + int(runner.resolved_max_workers), + ) + + return simulated, int(runner.resolved_max_workers) + + +def evaluate_joint_trial_with_metadata( + config: Any, + trial_parameters: Mapping[str, float], + joint_measurements: Sequence[Any], + *, + loss_function: LossFunction | None = None, + backend: str, + build_grating_fn: BuildGratingFunction | None = None, + resolve_solver_parameters_fn: ResolveSolverParametersFunction | None = None, +) -> tuple[float, dict[str, float], dict[str, np.ndarray], int]: + """Evaluate one joint multi-angle trial. + + Args: + config: Joint optimization configuration describing the simulation setup. + trial_parameters: Ax trial parameters for the current candidate. + joint_measurements: Prepared per-angle measurements to evaluate. + loss_function: Optional custom per-measurement loss function. + backend: RCWA backend to use for the simulation. + build_grating_fn: Hook that builds a grating from the trial parameters. + resolve_solver_parameters_fn: Hook that resolves solver parameters. + + Returns: + The joint loss, the per-measurement losses, the simulated curves, and + the resolved worker count. On failure the penalty is reported for the + joint loss and every measurement, with empty simulated curves. + """ + + _warn_if_numpy_backend_requested(backend, stacklevel=2) + selected_loss_function = loss_function or mean_squared_error + labels = [str(joint_measurement.label) for joint_measurement in joint_measurements] + + try: + simulated, resolved_max_workers = simulate_joint_efficiency_curves_with_metadata( + config, + trial_parameters, + joint_measurements, + backend=backend, + build_grating_fn=build_grating_fn, + resolve_solver_parameters_fn=resolve_solver_parameters_fn, + ) + except _BatchCaseFailure as error: + module_logger.warning( + "Joint optimizer trial penalized: %s (resolved_max_workers=%s).", + error, + error.resolved_max_workers, + ) + penalty = float(config.failure_penalty) + return penalty, {label: penalty for label in labels}, {}, int(error.resolved_max_workers) + except Exception as error: + module_logger.warning( + "Joint optimizer trial penalized by %s: %s.", + type(error).__name__, + error, + ) + penalty = float(config.failure_penalty) + return penalty, {label: penalty for label in labels}, {}, _trial_max_workers(config) + + per_measurement_losses = { + str(joint_measurement.label): float( + selected_loss_function( + np.asarray(joint_measurement.evaluation_efficiency, dtype=float), + simulated[str(joint_measurement.label)], + ) + ) + for joint_measurement in joint_measurements + } + joint_loss = reduce_joint_losses( + per_measurement_losses, + reduction=str(config.joint_loss_reduction), + weights={ + str(joint_measurement.label): float(joint_measurement.weight) + for joint_measurement in joint_measurements + }, + point_counts={ + str(joint_measurement.label): len(joint_measurement.evaluation_energies_ev) + for joint_measurement in joint_measurements + }, + ) + return joint_loss, per_measurement_losses, simulated, int(resolved_max_workers) + + def evaluate_trial( config: Any, trial_parameters: Mapping[str, float], From da177fbcf04758555e98d140ffdc80e637d4b328 Mon Sep 17 00:00:00 2001 From: Simone Vadilonga Date: Tue, 11 Aug 2026 09:03:20 +0200 Subject: [PATCH 03/10] Add checkpoint and resume support to the measurement-fit optimizers Persists the Ax client snapshot alongside optimizer run state and an append-only trial log, so an interrupted run can continue and total_trials can be raised to extend a finished one. A problem fingerprint refuses to resume into a changed search space. Co-Authored-By: Claude Sonnet 5 --- src/grax_opt/checkpoint.py | 779 +++++++++++++++++++++++++++++++++++++ src/grax_opt/dynamic.py | 51 ++- src/grax_opt/joint.py | 36 +- 3 files changed, 849 insertions(+), 17 deletions(-) create mode 100644 src/grax_opt/checkpoint.py diff --git a/src/grax_opt/checkpoint.py b/src/grax_opt/checkpoint.py new file mode 100644 index 0000000..153dffc --- /dev/null +++ b/src/grax_opt/checkpoint.py @@ -0,0 +1,779 @@ +"""Checkpoint and resume support for the measurement-fit optimizers. + +A checkpoint directory holds three files: + +``ax_client_snapshot.json`` + The Ax client state, so a resumed run keeps its surrogate model instead of + restarting the generation strategy from scratch. +``optimizer_state.json`` + Run state Ax does not own: best-so-far results, counters, timing metadata, + and the problem fingerprint guarding against resuming into a changed search + space. +``trial_records.jsonl`` + Append-only per-trial history. A crash costs at most the final line. +""" + +from __future__ import annotations + +import hashlib +import json +import logging +import os +import tempfile +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from datetime import datetime +from pathlib import Path +from typing import Any + +from .optimize import TrialRecord + +module_logger = logging.getLogger(__name__) + +CHECKPOINT_SCHEMA_VERSION = 1 + + +@dataclass(frozen=True) +class OptimizerCheckpointPaths: + """Filesystem layout for one optimizer checkpoint directory. + + Attributes: + checkpoint_dir: Directory holding the checkpoint files. + ax_snapshot_path: Path to the serialized Ax client state. + state_path: Path to the optimizer run-state JSON. + trial_records_path: Path to the append-only trial-record log. + """ + + checkpoint_dir: Path + ax_snapshot_path: Path + state_path: Path + trial_records_path: Path + + @classmethod + def for_config(cls, config: Any) -> OptimizerCheckpointPaths: + """Resolve the checkpoint layout for a configuration. + + Args: + config: Optimizer configuration with ``output_dir`` and an optional + ``checkpoint_dir``. + + Returns: + The resolved checkpoint paths. + """ + + checkpoint_dir = getattr(config, "checkpoint_dir", None) + if checkpoint_dir is None: + checkpoint_dir = Path(config.output_dir) / "checkpoint" + checkpoint_dir = Path(checkpoint_dir) + return cls( + checkpoint_dir=checkpoint_dir, + ax_snapshot_path=checkpoint_dir / "ax_client_snapshot.json", + state_path=checkpoint_dir / "optimizer_state.json", + trial_records_path=checkpoint_dir / "trial_records.jsonl", + ) + + def exists(self) -> bool: + """Return whether a usable checkpoint is present. + + Returns: + ``True`` when both the Ax snapshot and the run state exist. + """ + + return self.ax_snapshot_path.is_file() and self.state_path.is_file() + + +def _atomic_write_json(output_path: Path, payload: Mapping[str, Any]) -> None: + """Write JSON atomically so a crash cannot truncate the file. + + Args: + output_path: Destination path to replace. + payload: JSON-serializable payload. + """ + + output_path.parent.mkdir(parents=True, exist_ok=True) + handle = tempfile.NamedTemporaryFile( + "w", + encoding="utf-8", + dir=str(output_path.parent), + prefix=f".{output_path.name}.", + suffix=".tmp", + delete=False, + ) + temporary_path = Path(handle.name) + try: + with handle: + json.dump(payload, handle, indent=2) + handle.flush() + os.fsync(handle.fileno()) + os.replace(temporary_path, output_path) + except BaseException: + temporary_path.unlink(missing_ok=True) + raise + + +def _file_content_hash(path: Path) -> str: + """Return a stable content hash for a measurement file. + + Args: + path: File to hash. + + Returns: + The hex digest, or a sentinel when the file is unreadable. + """ + + try: + return hashlib.sha256(path.read_bytes()).hexdigest() + except OSError: + return "unreadable" + + +def build_problem_fingerprint(config: Any) -> dict[str, Any]: + """Describe the parts of a configuration that a resume may not change. + + Settings that only affect run length or performance -- ``total_trials``, + ``batch_size``, ``max_workers``, ``backend``, artifact flags, and the + early-stopping settings -- are deliberately excluded so they can be tuned + between runs. + + Args: + config: Optimizer configuration to fingerprint. + + Returns: + A JSON-serializable description of the optimization problem. + """ + + fingerprint: dict[str, Any] = { + "parameter_bounds": { + name: [float(bounds.lower), float(bounds.upper)] + for name, bounds in sorted(config.parameter_bounds.items()) + }, + "equality_constraints": dict(sorted(config.equality_constraints.items())), + "objective_name": str(config.objective_name), + "diffraction_order": int(config.diffraction_order), + "fourier_orders": int(config.fourier_orders), + "roughness_sigma_nm": config.roughness_sigma_nm, + "failure_penalty": float(config.failure_penalty), + } + + measurements = getattr(config, "measurements", None) + if measurements is not None: + fingerprint["joint_loss_reduction"] = str(config.joint_loss_reduction) + fingerprint["measurements"] = [ + { + "label": str(spec.label), + "grazing_angle_deg": float(spec.grazing_angle_deg), + "measurement_path": str(Path(spec.measurement_path).resolve()), + "content_hash": _file_content_hash(Path(spec.measurement_path)), + "evaluation_energies_ev": [ + float(energy_ev) for energy_ev in spec.evaluation_energies_ev + ], + "weight": float(spec.weight), + } + for spec in measurements + ] + return fingerprint + + measurement_path = Path(config.measurement_path) + fingerprint["measurement_path"] = str(measurement_path.resolve()) + fingerprint["content_hash"] = _file_content_hash(measurement_path) + fingerprint["angle_mode"] = str(config.angle_mode) + fingerprint["grazing_angle_deg"] = float(config.grazing_angle_deg) + fingerprint["cff"] = float(config.cff) + fingerprint["evaluation_energies_ev"] = [ + float(energy_ev) for energy_ev in config.evaluation_energies_ev + ] + fingerprint["evaluation_grazing_angles_deg"] = [ + float(angle) for angle in config.evaluation_grazing_angles_deg + ] + return fingerprint + + +def fingerprint_hash(fingerprint: Mapping[str, Any]) -> str: + """Return a stable hash for a problem fingerprint. + + Args: + fingerprint: Fingerprint produced by :func:`build_problem_fingerprint`. + + Returns: + The hex digest of the canonical JSON encoding. + """ + + canonical = json.dumps(fingerprint, sort_keys=True, default=str) + return hashlib.sha256(canonical.encode("utf-8")).hexdigest() + + +def _describe_fingerprint_differences( + stored: Mapping[str, Any], + current: Mapping[str, Any], +) -> list[str]: + """List the fingerprint keys that differ between two runs. + + Args: + stored: Fingerprint recorded in the checkpoint. + current: Fingerprint of the configuration being run now. + + Returns: + Sorted dotted key paths that differ. + """ + + differences: list[str] = [] + for key in sorted(set(stored) | set(current)): + stored_value = stored.get(key) + current_value = current.get(key) + if stored_value == current_value: + continue + if isinstance(stored_value, dict) and isinstance(current_value, dict): + differences.extend( + f"{key}.{nested}" + for nested in _describe_fingerprint_differences(stored_value, current_value) + ) + continue + differences.append(key) + return differences + + +def verify_fingerprint( + *, + stored_fingerprint: Mapping[str, Any], + current_fingerprint: Mapping[str, Any], +) -> None: + """Refuse to resume when the optimization problem itself changed. + + Args: + stored_fingerprint: Fingerprint recorded in the checkpoint. + current_fingerprint: Fingerprint of the configuration being run now. + + Raises: + ValueError: If the fingerprints differ, naming the changed keys. + """ + + if stored_fingerprint == current_fingerprint: + return + differences = _describe_fingerprint_differences(stored_fingerprint, current_fingerprint) + raise ValueError( + "Cannot resume: the checkpoint was created for a different optimization problem " + f"(changed: {', '.join(differences) or 'unknown'}). Use a different checkpoint_dir " + "or set resume=False to start a new run." + ) + + +def append_trial_record(handle: Any, record: TrialRecord) -> None: + """Append one completed trial to the trial-record log. + + Args: + handle: Open append-mode file handle. + record: Completed trial to serialize. + """ + + payload = { + "trial_index": int(record.trial_index), + "loss": float(record.loss), + "parameters": {name: float(value) for name, value in record.parameters.items()}, + "extras": {name: float(value) for name, value in record.extras.items()}, + } + handle.write(json.dumps(payload) + "\n") + + +def load_trial_records(trial_records_path: Path) -> list[TrialRecord]: + """Load completed trial records, tolerating a torn final line. + + Args: + trial_records_path: Path to the append-only trial-record log. + + Returns: + The recoverable trial records in file order. + """ + + if not trial_records_path.is_file(): + return [] + records: list[TrialRecord] = [] + with trial_records_path.open("r", encoding="utf-8") as handle: + for line in handle: + stripped = line.strip() + if not stripped: + continue + try: + payload = json.loads(stripped) + records.append( + TrialRecord( + trial_index=int(payload["trial_index"]), + loss=float(payload["loss"]), + parameters={ + str(name): float(value) + for name, value in payload.get("parameters", {}).items() + }, + extras={ + str(name): float(value) + for name, value in payload.get("extras", {}).items() + }, + ) + ) + except (json.JSONDecodeError, KeyError, TypeError, ValueError): + module_logger.warning("Ignoring malformed trial record during resume.") + return records + + +def load_checkpoint_state(state_path: Path) -> dict[str, Any]: + """Read the optimizer run-state JSON. + + Args: + state_path: Path to the run-state file. + + Returns: + The decoded run state. + + Raises: + ValueError: If the file is missing or cannot be decoded. + """ + + try: + return json.loads(state_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as error: + raise ValueError( + f"Cannot resume: the optimizer state at {state_path} is missing or unreadable " + f"({error}). Use a different checkpoint_dir or set resume=False to start a new run." + ) from error + + +def write_checkpoint_state( + *, + paths: OptimizerCheckpointPaths, + payload: Mapping[str, Any], +) -> None: + """Persist the optimizer run state atomically. + + Args: + paths: Resolved checkpoint paths. + payload: Run state to persist. + """ + + _atomic_write_json(paths.state_path, payload) + + +def save_ax_client_snapshot(ax_client: Any, snapshot_path: Path) -> None: + """Persist the Ax client state atomically. + + Args: + ax_client: Ax client to serialize. + snapshot_path: Destination snapshot path. + """ + + snapshot_path.parent.mkdir(parents=True, exist_ok=True) + temporary_path = snapshot_path.with_name(f".{snapshot_path.name}.tmp") + try: + ax_client.save_to_json_file(filepath=str(temporary_path)) + os.replace(temporary_path, snapshot_path) + except BaseException: + temporary_path.unlink(missing_ok=True) + raise + + +def load_ax_client_snapshot(snapshot_path: Path, *, recorded_ax_version: str | None = None) -> Any: + """Restore an Ax client from a snapshot. + + Args: + snapshot_path: Path to the serialized Ax client state. + recorded_ax_version: Ax version recorded when the snapshot was written. + + Returns: + The restored Ax client. + + Raises: + ValueError: If the snapshot cannot be loaded. + """ + + from .optimize import _import_ax_client + + ax_client_cls = _import_ax_client() + installed_ax_version = _installed_ax_version() + if recorded_ax_version is not None and recorded_ax_version != installed_ax_version: + module_logger.warning( + "Resuming a checkpoint written with ax %s using ax %s.", + recorded_ax_version, + installed_ax_version, + ) + try: + return ax_client_cls.load_from_json_file(filepath=str(snapshot_path)) + except Exception as error: + raise ValueError( + f"Cannot resume: the Ax snapshot at {snapshot_path} could not be loaded " + f"({type(error).__name__}: {error}). It was written with ax " + f"{recorded_ax_version or 'unknown'} and this environment has ax " + f"{installed_ax_version}. Use a different checkpoint_dir or set resume=False " + "to start a new run." + ) from error + + +def _installed_ax_version() -> str: + """Return the installed Ax version. + + Returns: + The version string, or ``"unknown"`` when it cannot be determined. + """ + + try: + import ax + except ImportError: + return "unknown" + return str(getattr(ax, "__version__", "unknown")) + + +def ax_trial_count(ax_client: Any) -> int | None: + """Return how many trials the Ax client has already issued. + + Args: + ax_client: Ax client to inspect. + + Returns: + The trial count, or ``None`` when it cannot be determined. + """ + + experiment = getattr(ax_client, "experiment", None) + if experiment is None: + return None + try: + return int(len(experiment.trials)) + except (AttributeError, TypeError): + return None + + +def build_checkpoint_state_payload( + *, + config: Any, + state: Any, + fingerprint: Mapping[str, Any], + previous_state: Mapping[str, Any] | None, + backend_requested: str, + backend_effective: str, + best_extras_payload: Mapping[str, Any], + elapsed_seconds: float, +) -> dict[str, Any]: + """Assemble the optimizer run-state payload. + + Timing metadata follows the batch-sweep convention: ``created`` is preserved + across resumes and elapsed time accumulates. + + Args: + config: Optimizer configuration for the current run. + state: Trial-loop state to persist. + fingerprint: Problem fingerprint for the current configuration. + previous_state: Run state loaded at resume, if any. + backend_requested: Backend requested by the caller. + backend_effective: Backend actually used. + best_extras_payload: JSON-safe extras captured from the best trial. + elapsed_seconds: Wall time spent in the current run. + + Returns: + The run state to persist. + """ + + now_iso = datetime.now().isoformat() + previous = dict(previous_state or {}) + total_trials_history = list(previous.get("total_trials_history", [])) + total_trials_history.append(int(config.total_trials)) + cumulative_elapsed = float(previous.get("cumulative_elapsed_seconds", 0.0)) + float( + elapsed_seconds + ) + + return { + "schema_version": CHECKPOINT_SCHEMA_VERSION, + "ax_version": _installed_ax_version(), + "fingerprint": dict(fingerprint), + "fingerprint_hash": fingerprint_hash(fingerprint), + "best_loss": float(state.best_loss), + "best_parameters": dict(state.best_parameters), + "best_extras": dict(best_extras_payload), + "completed_trials": int(state.completed_trials), + "trial_index_cursor": int(state.trial_index_cursor), + "no_improvement_trials": int(state.no_improvement_trials), + "stopped_early": bool(state.stopped_early), + "early_stop_reason": state.early_stop_reason, + "total_trials_history": total_trials_history, + "random_seed": config.random_seed, + "backend_requested": backend_requested, + "backend_effective": backend_effective, + "optimizer_resolved_max_workers": int(state.resolved_max_workers), + "created": previous.get("created", now_iso), + "current_run_started": previous.get("current_run_started_marker", now_iso), + "last_updated": now_iso, + "cumulative_elapsed_seconds": cumulative_elapsed, + "last_run_elapsed_seconds": float(elapsed_seconds), + "run_count": int(previous.get("run_count", 0)) + 1, + } + + +def json_safe_extras(extras: Mapping[str, Any]) -> dict[str, Any]: + """Convert a trial extras payload into JSON-serializable values. + + Args: + extras: Mode-specific payload attached to a trial evaluation. + + Returns: + The payload with arrays converted to lists. + """ + + safe: dict[str, Any] = {} + for name, value in extras.items(): + if hasattr(value, "tolist"): + safe[str(name)] = value.tolist() + elif isinstance(value, Mapping): + safe[str(name)] = { + str(inner_name): ( + inner_value.tolist() if hasattr(inner_value, "tolist") else inner_value + ) + for inner_name, inner_value in value.items() + } + else: + safe[str(name)] = value + return safe + + +def restore_best_extras(payload: Mapping[str, Any]) -> dict[str, Any]: + """Rebuild a best-trial extras payload loaded from JSON. + + Args: + payload: Extras payload as stored in the checkpoint. + + Returns: + The payload with simulated curves restored as arrays. + """ + + import numpy as np + + restored = dict(payload) + simulated = restored.get("simulated_by_label") + if isinstance(simulated, Mapping): + restored["simulated_by_label"] = { + str(label): np.asarray(values, dtype=float) for label, values in simulated.items() + } + return restored + + +class OptimizerCheckpointSession: + """Persists and restores optimizer progress for one run. + + Checkpoints are always written so an interrupted run can be continued later, + but they are only read back when the configuration sets ``resume=True``. + + Attributes: + paths: Resolved checkpoint file layout. + enabled: Whether checkpoint writing is active. + resumed: Whether state was restored from an existing checkpoint. + previous_state: Run state loaded at resume, if any. + """ + + def __init__( + self, + *, + config: Any, + backend_requested: str, + backend_effective: str, + ) -> None: + """Initialize a checkpoint session for one optimizer run. + + Args: + config: Optimizer configuration for the run. + backend_requested: Backend requested by the caller. + backend_effective: Backend actually used. + """ + + self._config = config + self._backend_requested = backend_requested + self._backend_effective = backend_effective + self.paths = OptimizerCheckpointPaths.for_config(config) + self.enabled = True + self.resumed = False + self.previous_state: dict[str, Any] | None = None + self._fingerprint = build_problem_fingerprint(config) + self._handle: Any = None + self._since_flush = 0 + self._started_monotonic = 0.0 + + def restore_or_create_ax_client(self, create_ax_client: Any, state: Any) -> Any: + """Restore the Ax client from a checkpoint, or create a fresh one. + + Args: + create_ax_client: Callable creating a new Ax client for the config. + state: Trial-loop state to populate when resuming. + + Returns: + The Ax client to drive the run. + + Raises: + ValueError: If a partial or mismatched checkpoint is found. + """ + + if not bool(getattr(self._config, "resume", False)): + return create_ax_client(self._config) + + if not self.paths.exists(): + if self.paths.ax_snapshot_path.is_file() or self.paths.state_path.is_file(): + raise ValueError( + f"Cannot resume: the checkpoint at {self.paths.checkpoint_dir} is " + "incomplete. Use a different checkpoint_dir or set resume=False to " + "start a new run." + ) + module_logger.info( + "No checkpoint found at %s; starting a new optimization run.", + self.paths.checkpoint_dir, + ) + return create_ax_client(self._config) + + checkpoint_state = load_checkpoint_state(self.paths.state_path) + verify_fingerprint( + stored_fingerprint=checkpoint_state.get("fingerprint", {}), + current_fingerprint=self._fingerprint, + ) + if checkpoint_state.get("random_seed") != self._config.random_seed: + module_logger.warning( + "Resuming with random_seed=%s but the checkpoint recorded %s; " + "already-generated trials are unaffected.", + self._config.random_seed, + checkpoint_state.get("random_seed"), + ) + + ax_client = load_ax_client_snapshot( + self.paths.ax_snapshot_path, + recorded_ax_version=checkpoint_state.get("ax_version"), + ) + restore_trial_loop_state( + state=state, + checkpoint_state=checkpoint_state, + trial_records=load_trial_records(self.paths.trial_records_path), + ax_client=ax_client, + ) + state.best_extras = restore_best_extras(checkpoint_state.get("best_extras", {})) + self.resumed = True + self.previous_state = checkpoint_state + module_logger.info( + "Resumed optimization from %s with %s completed trials.", + self.paths.checkpoint_dir, + state.completed_trials, + ) + return ax_client + + def __enter__(self) -> OptimizerCheckpointSession: + """Open the trial-record log for appending. + + Returns: + This session. + """ + + import time + + self._started_monotonic = time.perf_counter() + if self.enabled: + self.paths.checkpoint_dir.mkdir(parents=True, exist_ok=True) + self._handle = self.paths.trial_records_path.open("a", encoding="utf-8") + return self + + def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None: + """Flush and close the trial-record log. + + Args: + exc_type: Exception type, if the block raised. + exc_value: Exception value, if the block raised. + traceback: Traceback, if the block raised. + """ + + if self._handle is not None: + self._handle.flush() + self._handle.close() + self._handle = None + + def record_trial(self, *, state: Any, ax_client: Any) -> None: + """Append the newest trial and periodically persist full state. + + Args: + state: Trial-loop state after the trial completed. + ax_client: Ax client to snapshot. + """ + + if not self.enabled or self._handle is None or not state.trial_records: + return + append_trial_record(self._handle, state.trial_records[-1]) + self._since_flush += 1 + if self._since_flush >= int(self._config.checkpoint_interval): + self._handle.flush() + self._since_flush = 0 + self.persist(state=state, ax_client=ax_client) + + def persist(self, *, state: Any, ax_client: Any) -> None: + """Write the Ax snapshot and optimizer run state atomically. + + Args: + state: Trial-loop state to persist. + ax_client: Ax client to snapshot. + """ + + if not self.enabled: + return + import time + + if self._handle is not None: + self._handle.flush() + try: + save_ax_client_snapshot(ax_client, self.paths.ax_snapshot_path) + except Exception as error: + module_logger.warning( + "Could not write the Ax checkpoint snapshot: %s: %s.", + type(error).__name__, + error, + ) + return + write_checkpoint_state( + paths=self.paths, + payload=build_checkpoint_state_payload( + config=self._config, + state=state, + fingerprint=self._fingerprint, + previous_state=self.previous_state, + backend_requested=self._backend_requested, + backend_effective=self._backend_effective, + best_extras_payload=json_safe_extras(state.best_extras), + elapsed_seconds=time.perf_counter() - self._started_monotonic, + ), + ) + + +def restore_trial_loop_state( + *, + state: Any, + checkpoint_state: Mapping[str, Any], + trial_records: Sequence[TrialRecord], + ax_client: Any, +) -> None: + """Restore in-memory loop state from a checkpoint. + + Ax is authoritative for how many candidates have been issued, so the cursor + is taken from the client and reconciled against the recovered trial records. + + Args: + state: Trial-loop state to populate in place. + checkpoint_state: Decoded optimizer run state. + trial_records: Trial records recovered from the log. + ax_client: Restored Ax client. + """ + + state.trial_records = list(trial_records) + state.completed_trials = len(state.trial_records) + state.best_loss = float(checkpoint_state.get("best_loss", float("inf"))) + state.best_parameters = dict(checkpoint_state.get("best_parameters", {})) + state.no_improvement_trials = int(checkpoint_state.get("no_improvement_trials", 0)) + state.resolved_max_workers = int( + checkpoint_state.get("optimizer_resolved_max_workers", state.resolved_max_workers) + ) + + persisted_cursor = int(checkpoint_state.get("trial_index_cursor", state.completed_trials)) + client_cursor = ax_trial_count(ax_client) + if client_cursor is None: + state.trial_index_cursor = persisted_cursor + return + if client_cursor != state.completed_trials: + module_logger.warning( + "Checkpoint drift on resume: Ax reports %s issued trials but %s trial records " + "were recovered. Continuing from the larger count.", + client_cursor, + state.completed_trials, + ) + state.trial_index_cursor = max(client_cursor, persisted_cursor) diff --git a/src/grax_opt/dynamic.py b/src/grax_opt/dynamic.py index c839024..6d532d4 100644 --- a/src/grax_opt/dynamic.py +++ b/src/grax_opt/dynamic.py @@ -16,6 +16,7 @@ from .evaluation import normalize_evaluation_selection from .data import MeasurementData, load_measurement_data from .objective import build_evaluation_measurement, simulate_efficiency_curve +from .checkpoint import OptimizerCheckpointSession from .loop import TrialEvaluation, TrialLoopState, run_ax_trial_loop from .optimize import ( OptimizationResult, @@ -178,6 +179,10 @@ class MeasurementFitConfig: objective evaluation through ``BatchSimulationRunner``. solver_parameter_resolver: Optional callable that resolves solver parameters from the fully expanded parameter dictionary. + resume: Whether to continue a previous run from its checkpoint. + ``total_trials`` is cumulative, so raising it extends the run. + checkpoint_dir: Checkpoint directory. Defaults to ``output_dir/checkpoint``. + checkpoint_interval: Number of trials between checkpoint flushes. """ build_grating: BuildGratingFunction @@ -210,12 +215,17 @@ class MeasurementFitConfig: evaluation_grazing_angles_deg: list[float] = field(default_factory=list) max_workers: int | str | None = None solver_parameter_resolver: ResolveSolverParametersFunction | None = None + resume: bool = False + checkpoint_dir: Path | None = None + checkpoint_interval: int = 1 def __post_init__(self) -> None: """Normalize paths, bounds, and validation settings.""" object.__setattr__(self, "measurement_path", Path(self.measurement_path)) object.__setattr__(self, "output_dir", Path(self.output_dir)) + if self.checkpoint_dir is not None: + object.__setattr__(self, "checkpoint_dir", Path(self.checkpoint_dir)) object.__setattr__( self, "parameter_bounds", @@ -245,6 +255,8 @@ def __post_init__(self) -> None: raise ValueError("total_trials must be > 0.") if self.batch_size <= 0: raise ValueError("batch_size must be > 0.") + if self.checkpoint_interval <= 0: + raise ValueError("checkpoint_interval must be > 0.") resolved_max_workers = _resolve_simulation_max_workers(self.max_workers) if resolved_max_workers > 1 and self.batch_size > 1: raise ValueError( @@ -348,6 +360,9 @@ def from_mapping(cls, mapping: Mapping[str, object]) -> "MeasurementFitConfig": ), max_workers=config.pop("max_workers", None), solver_parameter_resolver=config.pop("solver_parameter_resolver", None), + resume=bool(config.pop("resume", False)), + checkpoint_dir=config.pop("checkpoint_dir", None), + checkpoint_interval=int(config.pop("checkpoint_interval", 1)), ) if config: raise ValueError(f"Unexpected measurement-fit spec keys: {sorted(config)}") @@ -667,11 +682,26 @@ def optimize_to_measurements( backend_effective = _resolve_optimizer_backend(config.backend) measurement = load_measurement_data(config.measurement_path) evaluation_measurement = build_evaluation_measurement(config, measurement) - ax_client = _create_ax_client_for_measurement_fit_config(config) - state = TrialLoopState() state.resolved_max_workers = _resolve_simulation_max_workers(config.max_workers) + checkpoint = OptimizerCheckpointSession( + config=config, + backend_requested=config.backend, + backend_effective=backend_effective, + ) + ax_client = checkpoint.restore_or_create_ax_client( + lambda run_config: _create_ax_client_for_measurement_fit_config(run_config), + state, + ) + if checkpoint.resumed and state.best_parameters: + state.best_grating_parameters = dict( + resolve_measurement_fit_trial_parameters(config, state.best_parameters) + ) + state.best_solver_parameters = dict( + _resolve_measurement_fit_solver_parameters(config, state.best_grating_parameters) + ) + build_grating_fn = lambda trial_parameters: config.build_grating( resolve_measurement_fit_trial_parameters(config, trial_parameters) ) @@ -743,14 +773,17 @@ def on_trial_completed(*, evaluation: TrialEvaluation, state: TrialLoopState, im optimizer_requested_max_workers=config.max_workers, optimizer_resolved_max_workers=state.resolved_max_workers, ) + checkpoint.record_trial(state=state, ax_client=ax_client) - run_ax_trial_loop( - ax_client=ax_client, - config=config, - state=state, - evaluate_candidates=evaluate_candidates, - on_trial_completed=on_trial_completed, - ) + with checkpoint: + run_ax_trial_loop( + ax_client=ax_client, + config=config, + state=state, + evaluate_candidates=evaluate_candidates, + on_trial_completed=on_trial_completed, + ) + checkpoint.persist(state=state, ax_client=ax_client) if not state.trial_records: raise RuntimeError("Optimization produced no completed trials.") diff --git a/src/grax_opt/joint.py b/src/grax_opt/joint.py index 9e644c9..2e8b21d 100644 --- a/src/grax_opt/joint.py +++ b/src/grax_opt/joint.py @@ -21,6 +21,7 @@ from grax.simulation import _resolve_max_workers as _resolve_simulation_max_workers +from .checkpoint import OptimizerCheckpointSession from .config import ParameterBounds from .data import load_measurement_data, sample_measurement_data from .dynamic import ( @@ -788,11 +789,27 @@ def optimize_to_joint_measurements( backend_effective = _resolve_optimizer_backend(config.backend) measurements = prepare_joint_measurements(config.measurements) - ax_client = _create_ax_client_for_joint_config(config) state = TrialLoopState() state.resolved_max_workers = _resolve_simulation_max_workers(config.max_workers) + checkpoint = OptimizerCheckpointSession( + config=config, + backend_requested=config.backend, + backend_effective=backend_effective, + ) + ax_client = checkpoint.restore_or_create_ax_client( + lambda run_config: _create_ax_client_for_joint_config(run_config), + state, + ) + if checkpoint.resumed and state.best_parameters: + state.best_grating_parameters = dict( + resolve_measurement_fit_trial_parameters(config, state.best_parameters) + ) + state.best_solver_parameters = dict( + _resolve_joint_solver_parameters(config, state.best_grating_parameters) + ) + build_grating_fn = lambda trial_parameters: config.build_grating( resolve_measurement_fit_trial_parameters(config, trial_parameters) ) @@ -869,14 +886,17 @@ def on_trial_completed( backend_effective=backend_effective, write_heavy_artifacts=improved, ) + checkpoint.record_trial(state=state, ax_client=ax_client) - run_ax_trial_loop( - ax_client=ax_client, - config=config, - state=state, - evaluate_candidates=evaluate_candidates, - on_trial_completed=on_trial_completed, - ) + with checkpoint: + run_ax_trial_loop( + ax_client=ax_client, + config=config, + state=state, + evaluate_candidates=evaluate_candidates, + on_trial_completed=on_trial_completed, + ) + checkpoint.persist(state=state, ax_client=ax_client) if not state.trial_records: raise RuntimeError("Joint optimization produced no completed trials.") From 16c0f4ea8f148c6a0aad808c856f3177c648ed69 Mon Sep 17 00:00:00 2001 From: Simone Vadilonga Date: Tue, 11 Aug 2026 09:05:53 +0200 Subject: [PATCH 04/10] Cache best-fit curves instead of re-simulating on every trial The best-fit plot now reuses the winning trial's simulated curve, and plots are rewritten only when the best improves. Penalized trials log the failing case or exception instead of failing silently. Co-Authored-By: Claude Sonnet 5 --- src/grax_opt/dynamic.py | 75 +++++++++++++++++++++++++++++-------- src/grax_opt/objective.py | 79 +++++++++++++++++++++++++++++++++++---- src/grax_opt/optimize.py | 31 ++++++++++++--- 3 files changed, 156 insertions(+), 29 deletions(-) diff --git a/src/grax_opt/dynamic.py b/src/grax_opt/dynamic.py index 6d532d4..9213982 100644 --- a/src/grax_opt/dynamic.py +++ b/src/grax_opt/dynamic.py @@ -588,8 +588,34 @@ def _persist_measurement_fit_optimizer_artifacts( backend_effective: str, optimizer_requested_max_workers: int | str | None, optimizer_resolved_max_workers: int, + best_simulated_efficiency: np.ndarray | None = None, + write_heavy_artifacts: bool = True, ) -> tuple[Path, Path, Path | None, Path | None]: - """Persist measurement-fit optimizer artifacts and return their paths.""" + """Persist measurement-fit optimizer artifacts and return their paths. + + Args: + config: Measurement-fit configuration for the run. + evaluation_measurement: Measurement sampled onto the evaluation grid. + best_parameters: Best free parameters found so far. + best_grating_parameters: Best resolved grating parameters. + best_solver_parameters: Best resolved solver parameters. + best_loss: Best objective value found so far. + trial_records: Completed trial records. + stopped_early: Whether early stopping ended the run. + completed_trials: Number of trials successfully evaluated. + early_stop_reason: Human-readable early-stopping reason, or ``None``. + backend_requested: Backend requested by the caller. + backend_effective: Backend actually used. + optimizer_requested_max_workers: Worker count requested by the caller. + optimizer_resolved_max_workers: Worker count actually used. + best_simulated_efficiency: Cached best-fit curve. When ``None`` the + curve is re-simulated so the plot can still be written. + write_heavy_artifacts: Whether to rewrite the plots. + + Returns: + Paths to the JSON summary, trial-history CSV, best-fit plot, and + loss-history plot. Optional paths are ``None`` when disabled. + """ config.output_dir.mkdir(parents=True, exist_ok=True) result_json_path = config.output_dir / "best_result.json" @@ -618,26 +644,28 @@ def _persist_measurement_fit_optimizer_artifacts( output_path=result_json_path, ) _write_trial_history_csv(trial_records, trial_history_csv_path) - if best_fit_plot_path is not None: - simulated_efficiency = simulate_efficiency_curve( - config, - best_parameters, - evaluation_measurement, - backend=backend_effective, - build_grating_fn=lambda trial_parameters: config.build_grating( - resolve_measurement_fit_trial_parameters(config, trial_parameters) - ), - resolve_solver_parameters_fn=lambda trial_parameters: _resolve_measurement_fit_solver_parameters( + if best_fit_plot_path is not None and write_heavy_artifacts and best_parameters: + simulated_efficiency = best_simulated_efficiency + if simulated_efficiency is None: + simulated_efficiency = simulate_efficiency_curve( config, - resolve_measurement_fit_trial_parameters(config, trial_parameters), - ), - ) + best_parameters, + evaluation_measurement, + backend=backend_effective, + build_grating_fn=lambda trial_parameters: config.build_grating( + resolve_measurement_fit_trial_parameters(config, trial_parameters) + ), + resolve_solver_parameters_fn=lambda trial_parameters: _resolve_measurement_fit_solver_parameters( + config, + resolve_measurement_fit_trial_parameters(config, trial_parameters), + ), + ) _save_best_fit_plot( measurement=evaluation_measurement, simulated_efficiency=simulated_efficiency, output_path=best_fit_plot_path, ) - if loss_history_plot_path is not None: + if loss_history_plot_path is not None and write_heavy_artifacts and trial_records: _save_loss_history_plot( trial_records=trial_records, output_path=loss_history_plot_path, @@ -734,8 +762,19 @@ def evaluate_candidates(candidates) -> list[TrialEvaluation]: parameters=dict(parameters), loss=float(loss), resolved_max_workers=int(trial_resolved_max_workers), + extras=( + {} + if simulated_efficiency is None + else {"simulated_efficiency": simulated_efficiency} + ), ) - for trial_index, parameters, loss, trial_resolved_max_workers in evaluated + for ( + trial_index, + parameters, + loss, + trial_resolved_max_workers, + simulated_efficiency, + ) in evaluated ] def on_trial_completed(*, evaluation: TrialEvaluation, state: TrialLoopState, improved: bool) -> None: @@ -772,6 +811,8 @@ def on_trial_completed(*, evaluation: TrialEvaluation, state: TrialLoopState, im backend_effective=backend_effective, optimizer_requested_max_workers=config.max_workers, optimizer_resolved_max_workers=state.resolved_max_workers, + best_simulated_efficiency=state.best_extras.get("simulated_efficiency"), + write_heavy_artifacts=improved, ) checkpoint.record_trial(state=state, ax_client=ax_client) @@ -815,6 +856,8 @@ def on_trial_completed(*, evaluation: TrialEvaluation, state: TrialLoopState, im backend_effective=backend_effective, optimizer_requested_max_workers=config.max_workers, optimizer_resolved_max_workers=state.resolved_max_workers, + best_simulated_efficiency=state.best_extras.get("simulated_efficiency"), + write_heavy_artifacts=True, ) return OptimizationResult( diff --git a/src/grax_opt/objective.py b/src/grax_opt/objective.py index 083f10e..47dcf68 100644 --- a/src/grax_opt/objective.py +++ b/src/grax_opt/objective.py @@ -489,7 +489,7 @@ def evaluate_trial( return float(selected_loss_function(evaluation_measurement.efficiency, simulated_efficiency)) -def evaluate_trial_with_metadata( +def evaluate_trial_curve_with_metadata( config: Any, trial_parameters: Mapping[str, float], measurement: MeasurementData, @@ -498,10 +498,27 @@ def evaluate_trial_with_metadata( backend: str, build_grating_fn: BuildGratingFunction | None = None, resolve_solver_parameters_fn: ResolveSolverParametersFunction | None = None, -) -> tuple[float, int]: - """Evaluate one Ax trial and return loss plus resolved worker count.""" +) -> tuple[float, int, np.ndarray | None]: + """Evaluate one Ax trial and also return its simulated curve. - _warn_if_numpy_backend_requested(backend, stacklevel=2) + Returning the curve lets callers plot the best fit without re-running the + simulation afterwards. + + Args: + config: Optimization configuration describing the simulation setup. + trial_parameters: Ax trial parameters for the current candidate. + measurement: Energy grid and target efficiencies used for evaluation. + loss_function: Optional custom loss function. + backend: RCWA backend to use for the simulation. + build_grating_fn: Hook that builds a grating from the trial parameters. + resolve_solver_parameters_fn: Hook that resolves solver parameters. + + Returns: + The loss, the resolved worker count, and the simulated efficiencies. + The curve is ``None`` when the trial was penalized. + """ + + _warn_if_numpy_backend_requested(backend, stacklevel=3) selected_loss_function = loss_function or mean_squared_error evaluation_measurement = build_evaluation_measurement(config, measurement) try: @@ -514,10 +531,58 @@ def evaluate_trial_with_metadata( resolve_solver_parameters_fn=resolve_solver_parameters_fn, ) except _BatchCaseFailure as error: - return float(config.failure_penalty), int(error.resolved_max_workers) - except Exception: - return float(config.failure_penalty), _trial_max_workers(config) + module_logger.warning( + "Optimizer trial penalized: %s (resolved_max_workers=%s).", + error, + error.resolved_max_workers, + ) + return float(config.failure_penalty), int(error.resolved_max_workers), None + except Exception as error: + module_logger.warning( + "Optimizer trial penalized by %s: %s.", + type(error).__name__, + error, + ) + return float(config.failure_penalty), _trial_max_workers(config), None return ( float(selected_loss_function(evaluation_measurement.efficiency, simulated_efficiency)), int(resolved_max_workers), + simulated_efficiency, + ) + + +def evaluate_trial_with_metadata( + config: Any, + trial_parameters: Mapping[str, float], + measurement: MeasurementData, + *, + loss_function: LossFunction | None = None, + backend: str, + build_grating_fn: BuildGratingFunction | None = None, + resolve_solver_parameters_fn: ResolveSolverParametersFunction | None = None, +) -> tuple[float, int]: + """Evaluate one Ax trial and return loss plus resolved worker count. + + Args: + config: Optimization configuration describing the simulation setup. + trial_parameters: Ax trial parameters for the current candidate. + measurement: Energy grid and target efficiencies used for evaluation. + loss_function: Optional custom loss function. + backend: RCWA backend to use for the simulation. + build_grating_fn: Hook that builds a grating from the trial parameters. + resolve_solver_parameters_fn: Hook that resolves solver parameters. + + Returns: + The loss and the resolved worker count. + """ + + loss, resolved_max_workers, _simulated_efficiency = evaluate_trial_curve_with_metadata( + config, + trial_parameters, + measurement, + loss_function=loss_function, + backend=backend, + build_grating_fn=build_grating_fn, + resolve_solver_parameters_fn=resolve_solver_parameters_fn, ) + return loss, resolved_max_workers diff --git a/src/grax_opt/optimize.py b/src/grax_opt/optimize.py index 2fc7418..18da53d 100644 --- a/src/grax_opt/optimize.py +++ b/src/grax_opt/optimize.py @@ -20,7 +20,7 @@ from grax.materials import material_label from .data import MeasurementData -from .objective import evaluate_trial_with_metadata +from .objective import evaluate_trial_curve_with_metadata def _is_cuda_usable() -> bool: @@ -181,8 +181,21 @@ def _evaluate_candidate_worker( backend_effective: str, build_grating_fn=None, resolve_solver_parameters_fn=None, -) -> tuple[int, dict[str, float], float, int]: - """Evaluate one optimizer candidate and return trial index, params, and loss.""" +) -> tuple[int, dict[str, float], float, int, Any]: + """Evaluate one optimizer candidate. + + Args: + candidate: Candidate ``(trial_index, parameters)`` pair. + config: Optimization configuration describing the simulation setup. + measurement: Measurement data used for evaluation. + backend_effective: RCWA backend to use for the simulation. + build_grating_fn: Optional grating-build hook. + resolve_solver_parameters_fn: Optional solver-parameter hook. + + Returns: + The trial index, parameters, loss, resolved worker count, and the + simulated curve for the candidate. + """ trial_index, parameters = candidate evaluate_kwargs: dict[str, object] = { @@ -192,13 +205,19 @@ def _evaluate_candidate_worker( evaluate_kwargs["build_grating_fn"] = build_grating_fn if resolve_solver_parameters_fn is not None: evaluate_kwargs["resolve_solver_parameters_fn"] = resolve_solver_parameters_fn - loss, resolved_max_workers = evaluate_trial_with_metadata( + loss, resolved_max_workers, simulated_efficiency = evaluate_trial_curve_with_metadata( config, parameters, measurement, **evaluate_kwargs, ) - return int(trial_index), dict(parameters), float(loss), int(resolved_max_workers) + return ( + int(trial_index), + dict(parameters), + float(loss), + int(resolved_max_workers), + simulated_efficiency, + ) def _evaluate_candidate_batch( @@ -209,7 +228,7 @@ def _evaluate_candidate_batch( backend_effective: str, build_grating_fn=None, resolve_solver_parameters_fn=None, -) -> list[tuple[int, dict[str, float], float, int]]: +) -> list[tuple[int, dict[str, float], float, int, Any]]: """Evaluate a candidate batch, optionally in parallel.""" if len(candidates) <= 1: From bf71d7d9356b176cf9a600eb4c5e8684871e60c3 Mon Sep 17 00:00:00 2001 From: Simone Vadilonga Date: Tue, 11 Aug 2026 09:10:58 +0200 Subject: [PATCH 05/10] Add tests for joint multi-angle fitting and optimizer resume Covers out-of-order result reassembly, joint loss reductions, resuming with an extended total_trials, and the fingerprint and corruption guards. Co-Authored-By: Claude Sonnet 5 --- tests/smoke/test_grax_opt_ax_snapshot.py | 63 ++++ tests/unit/test_grax_opt_joint.py | 412 +++++++++++++++++++++++ tests/unit/test_grax_opt_resume.py | 343 +++++++++++++++++++ 3 files changed, 818 insertions(+) create mode 100644 tests/smoke/test_grax_opt_ax_snapshot.py create mode 100644 tests/unit/test_grax_opt_joint.py create mode 100644 tests/unit/test_grax_opt_resume.py diff --git a/tests/smoke/test_grax_opt_ax_snapshot.py b/tests/smoke/test_grax_opt_ax_snapshot.py new file mode 100644 index 0000000..972a10f --- /dev/null +++ b/tests/smoke/test_grax_opt_ax_snapshot.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +pytest.importorskip("ax", reason="Ax is only installed with the 'opt' extra.") + +from grax_opt.checkpoint import ( # noqa: E402 + ax_trial_count, + load_ax_client_snapshot, + save_ax_client_snapshot, +) +from grax_opt.optimize import _import_ax_client, _import_objective_properties # noqa: E402 + + +def _build_ax_client(): + ax_client_cls = _import_ax_client() + objective_properties = _import_objective_properties() + ax_client = ax_client_cls(random_seed=7) + ax_client.create_experiment( + name="snapshot_smoke", + parameters=[ + {"name": "x", "type": "range", "bounds": [0.0, 1.0], "value_type": "float"} + ], + objectives={"loss": objective_properties(minimize=True)}, + ) + return ax_client + + +def _run_trials(ax_client, count: int) -> None: + for _ in range(count): + parameters, trial_index = ax_client.get_next_trial() + loss = float(parameters["x"]) ** 2 + ax_client.complete_trial(trial_index=trial_index, raw_data={"loss": (loss, 1.0e-6)}) + + +def test_ax_client_snapshot_round_trip_continues_the_experiment(tmp_path: Path) -> None: + ax_client = _build_ax_client() + _run_trials(ax_client, 3) + snapshot_path = tmp_path / "ax_client_snapshot.json" + + save_ax_client_snapshot(ax_client, snapshot_path) + restored = load_ax_client_snapshot(snapshot_path) + + assert ax_trial_count(restored) == 3 + + _run_trials(restored, 2) + assert ax_trial_count(restored) == 5 + + _parameters, trial_index = restored.get_next_trial() + assert trial_index == 5 + + +def test_save_ax_client_snapshot_leaves_no_temporary_file(tmp_path: Path) -> None: + ax_client = _build_ax_client() + _run_trials(ax_client, 1) + snapshot_path = tmp_path / "ax_client_snapshot.json" + + save_ax_client_snapshot(ax_client, snapshot_path) + + assert snapshot_path.is_file() + assert [path.name for path in tmp_path.iterdir() if ".tmp" in path.name] == [] diff --git a/tests/unit/test_grax_opt_joint.py b/tests/unit/test_grax_opt_joint.py new file mode 100644 index 0000000..c450534 --- /dev/null +++ b/tests/unit/test_grax_opt_joint.py @@ -0,0 +1,412 @@ +from __future__ import annotations + +import csv +import json +from pathlib import Path +from types import SimpleNamespace + +import numpy as np +import pytest + +from grax_opt import joint as joint_module +from grax_opt import objective as objective_module +from grax_opt.joint import ( + AngleMeasurementSpec, + JointAngleMeasurement, + JointMeasurementFitConfig, + optimize_to_joint_measurements, +) +from grax_opt.objective import ( + evaluate_joint_trial_with_metadata, + reduce_joint_losses, + simulate_joint_efficiency_curves_with_metadata, +) + + +class _FakeRunner: + resolved_max_workers = 2 + + batch_sizes: list[int] = [] + result_order: list[int] | None = None + efficiencies: dict[int, float] | None = None + status_by_index: dict[int, str] | None = None + omit_indices: set[int] = set() + + def __init__(self, **_kwargs: object) -> None: + pass + + def run_cases(self, cases, metadata=None): + type(self).batch_sizes.append(len(cases)) + order = type(self).result_order or list(range(len(cases))) + for index in order: + if index in type(self).omit_indices: + continue + status = (type(self).status_by_index or {}).get(index, "ok") + efficiency = (type(self).efficiencies or {}).get(index, 0.25) + yield SimpleNamespace( + index=index, + case_id=cases[index]["case_id"], + status=status, + selected_efficiency=efficiency, + ) + + +class _FakeAxClient: + def __init__(self, start_index: int = 0) -> None: + self.next_index = start_index + self.completed: list[int] = [] + + def create_experiment(self, **_kwargs: object) -> None: + return None + + def get_next_trial(self): + trial_index = self.next_index + self.next_index += 1 + return {"depth_nm": 5.0 + trial_index}, trial_index + + def complete_trial(self, trial_index, raw_data=None, data=None) -> None: + self.completed.append(int(trial_index)) + + def save_to_json_file(self, filepath: str) -> None: + Path(filepath).write_text(json.dumps({"next_index": self.next_index}), encoding="utf-8") + + +def _reset_fake_runner() -> None: + _FakeRunner.batch_sizes = [] + _FakeRunner.result_order = None + _FakeRunner.efficiencies = None + _FakeRunner.status_by_index = None + _FakeRunner.omit_indices = set() + + +def _write_measurement(path: Path, rows: list[tuple[float, float]]) -> Path: + path.write_text( + "".join(f"{energy_ev} {efficiency}\n" for energy_ev, efficiency in rows), + encoding="utf-8", + ) + return path + + +def _joint_measurements() -> list[JointAngleMeasurement]: + return [ + JointAngleMeasurement( + label="a1", + grazing_angle_deg=1.0, + measurement_path=Path("a1.dat"), + evaluation_energies_ev=np.array([100.0, 200.0]), + evaluation_efficiency=np.array([0.1, 0.2]), + weight=1.0, + ), + JointAngleMeasurement( + label="a2", + grazing_angle_deg=2.0, + measurement_path=Path("a2.dat"), + evaluation_energies_ev=np.array([300.0, 400.0]), + evaluation_efficiency=np.array([0.3, 0.4]), + weight=1.0, + ), + ] + + +def _joint_eval_config() -> SimpleNamespace: + return SimpleNamespace( + diffraction_order=1, + fourier_orders=5, + max_workers=2, + validate_physical_results=True, + failure_penalty=1.0e6, + joint_loss_reduction="mean", + ) + + +def _joint_spec(tmp_path: Path, **overrides: object) -> dict[str, object]: + first_path = _write_measurement(tmp_path / "m1.dat", [(100.0, 0.2), (200.0, 0.3)]) + second_path = _write_measurement(tmp_path / "m2.dat", [(100.0, 0.4), (200.0, 0.5)]) + spec: dict[str, object] = { + "build_grating": lambda _parameters: SimpleNamespace(period_lpermm=2000.0), + "parameter_bounds": {"depth_nm": (1.0, 20.0)}, + "output_dir": tmp_path / "out", + "measurements": [ + {"grazing_angle_deg": 1.0, "measurement_path": first_path}, + {"grazing_angle_deg": 2.0, "measurement_path": second_path}, + ], + "total_trials": 3, + "save_loss_plot": False, + } + spec.update(overrides) + return spec + + +def test_angle_measurement_spec_rejects_non_positive_angle(tmp_path: Path) -> None: + with pytest.raises(ValueError, match="grazing_angle_deg must be > 0"): + AngleMeasurementSpec(grazing_angle_deg=0.0, measurement_path=tmp_path / "m.dat") + + +def test_angle_measurement_spec_defaults_label_from_angle(tmp_path: Path) -> None: + spec = AngleMeasurementSpec(grazing_angle_deg=2.5, measurement_path=tmp_path / "m.dat") + + assert spec.label == "alpha2.5deg" + + +def test_angle_measurement_spec_rejects_mismatched_efficiency_length(tmp_path: Path) -> None: + with pytest.raises(ValueError, match="same length as evaluation_energies_ev"): + AngleMeasurementSpec( + grazing_angle_deg=1.0, + measurement_path=tmp_path / "m.dat", + evaluation_energies_ev=[100.0, 200.0], + measurement_efficiency=[0.1], + ) + + +def test_joint_config_rejects_empty_measurements(tmp_path: Path) -> None: + with pytest.raises(ValueError, match="measurements must be provided and non-empty"): + JointMeasurementFitConfig( + build_grating=lambda _parameters: None, + parameter_bounds={"depth_nm": (1.0, 2.0)}, + output_dir=tmp_path / "out", + measurements=[], + ) + + +def test_joint_config_rejects_duplicate_measurement_labels(tmp_path: Path) -> None: + with pytest.raises(ValueError, match="unique labels"): + JointMeasurementFitConfig( + build_grating=lambda _parameters: None, + parameter_bounds={"depth_nm": (1.0, 2.0)}, + output_dir=tmp_path / "out", + measurements=[ + {"grazing_angle_deg": 1.0, "measurement_path": tmp_path / "m.dat"}, + {"grazing_angle_deg": 1.0, "measurement_path": tmp_path / "other.dat"}, + ], + ) + + +def test_joint_config_rejects_unknown_reduction(tmp_path: Path) -> None: + with pytest.raises(ValueError, match="joint_loss_reduction must be one of"): + JointMeasurementFitConfig( + build_grating=lambda _parameters: None, + parameter_bounds={"depth_nm": (1.0, 2.0)}, + output_dir=tmp_path / "out", + measurements=[{"grazing_angle_deg": 1.0, "measurement_path": tmp_path / "m.dat"}], + joint_loss_reduction="median", + ) + + +def test_joint_config_from_mapping_rejects_unexpected_keys(tmp_path: Path) -> None: + with pytest.raises(ValueError, match="Unexpected joint measurement-fit spec keys"): + JointMeasurementFitConfig.from_mapping( + { + "build_grating": lambda _parameters: None, + "parameter_bounds": {"depth_nm": (1.0, 2.0)}, + "output_dir": tmp_path / "out", + "measurements": [ + {"grazing_angle_deg": 1.0, "measurement_path": tmp_path / "m.dat"} + ], + "not_a_real_key": 1, + } + ) + + +def test_reduce_joint_losses_supports_mean_sum_pooled_and_weighted() -> None: + losses = {"a": 1.0, "b": 3.0} + + assert reduce_joint_losses(losses) == pytest.approx(2.0) + assert reduce_joint_losses(losses, reduction="sum") == pytest.approx(4.0) + assert reduce_joint_losses( + losses, + reduction="pooled", + point_counts={"a": 1, "b": 3}, + ) == pytest.approx(2.5) + assert reduce_joint_losses( + losses, + reduction="weighted", + weights={"a": 3.0, "b": 1.0}, + ) == pytest.approx(1.5) + + +def test_reduce_joint_losses_rejects_unknown_reduction() -> None: + with pytest.raises(ValueError, match="joint_loss_reduction must be one of"): + reduce_joint_losses({"a": 1.0}, reduction="median") + + +def test_simulate_joint_curves_reassembles_out_of_order_results( + monkeypatch: pytest.MonkeyPatch, +) -> None: + _reset_fake_runner() + _FakeRunner.result_order = [3, 1, 0, 2] + _FakeRunner.efficiencies = {0: 0.11, 1: 0.22, 2: 0.33, 3: 0.44} + monkeypatch.setattr(objective_module, "BatchSimulationRunner", _FakeRunner) + + simulated, resolved_max_workers = simulate_joint_efficiency_curves_with_metadata( + _joint_eval_config(), + {"depth_nm": 5.0}, + _joint_measurements(), + backend="numba", + build_grating_fn=lambda _parameters: SimpleNamespace(period_lpermm=2000.0), + resolve_solver_parameters_fn=lambda _parameters: {"roughness_sigma_nm": None}, + ) + + assert np.allclose(simulated["a1"], np.array([0.11, 0.22])) + assert np.allclose(simulated["a2"], np.array([0.33, 0.44])) + assert resolved_max_workers == 2 + + +def test_simulate_joint_curves_builds_one_flat_batch(monkeypatch: pytest.MonkeyPatch) -> None: + _reset_fake_runner() + monkeypatch.setattr(objective_module, "BatchSimulationRunner", _FakeRunner) + + simulate_joint_efficiency_curves_with_metadata( + _joint_eval_config(), + {"depth_nm": 5.0}, + _joint_measurements(), + backend="numba", + build_grating_fn=lambda _parameters: SimpleNamespace(period_lpermm=2000.0), + resolve_solver_parameters_fn=lambda _parameters: {"roughness_sigma_nm": None}, + ) + + assert _FakeRunner.batch_sizes == [4] + + +def test_simulate_joint_curves_raises_when_results_are_incomplete( + monkeypatch: pytest.MonkeyPatch, +) -> None: + _reset_fake_runner() + _FakeRunner.omit_indices = {2} + monkeypatch.setattr(objective_module, "BatchSimulationRunner", _FakeRunner) + + with pytest.raises(objective_module._BatchCaseFailure): + simulate_joint_efficiency_curves_with_metadata( + _joint_eval_config(), + {"depth_nm": 5.0}, + _joint_measurements(), + backend="numba", + build_grating_fn=lambda _parameters: SimpleNamespace(period_lpermm=2000.0), + resolve_solver_parameters_fn=lambda _parameters: {"roughness_sigma_nm": None}, + ) + + +def test_evaluate_joint_trial_returns_failure_penalty_when_case_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + _reset_fake_runner() + _FakeRunner.status_by_index = {1: "error"} + monkeypatch.setattr(objective_module, "BatchSimulationRunner", _FakeRunner) + + joint_loss, per_measurement_losses, simulated, _workers = ( + evaluate_joint_trial_with_metadata( + _joint_eval_config(), + {"depth_nm": 5.0}, + _joint_measurements(), + backend="numba", + build_grating_fn=lambda _parameters: SimpleNamespace(period_lpermm=2000.0), + resolve_solver_parameters_fn=lambda _parameters: {"roughness_sigma_nm": None}, + ) + ) + + assert joint_loss == pytest.approx(1.0e6) + assert per_measurement_losses == {"a1": 1.0e6, "a2": 1.0e6} + assert simulated == {} + + +def test_evaluate_joint_trial_averages_per_measurement_losses( + monkeypatch: pytest.MonkeyPatch, +) -> None: + _reset_fake_runner() + monkeypatch.setattr(objective_module, "BatchSimulationRunner", _FakeRunner) + + joint_loss, per_measurement_losses, _simulated, _workers = ( + evaluate_joint_trial_with_metadata( + _joint_eval_config(), + {"depth_nm": 5.0}, + _joint_measurements(), + backend="numba", + build_grating_fn=lambda _parameters: SimpleNamespace(period_lpermm=2000.0), + resolve_solver_parameters_fn=lambda _parameters: {"roughness_sigma_nm": None}, + ) + ) + + expected_a1 = float(np.mean((np.array([0.25, 0.25]) - np.array([0.1, 0.2])) ** 2)) + expected_a2 = float(np.mean((np.array([0.25, 0.25]) - np.array([0.3, 0.4])) ** 2)) + assert per_measurement_losses["a1"] == pytest.approx(expected_a1) + assert per_measurement_losses["a2"] == pytest.approx(expected_a2) + assert joint_loss == pytest.approx((expected_a1 + expected_a2) / 2.0) + + +def test_optimize_to_joint_measurements_writes_artifacts( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _reset_fake_runner() + monkeypatch.setattr(objective_module, "BatchSimulationRunner", _FakeRunner) + monkeypatch.setattr( + joint_module, + "_create_ax_client_for_joint_config", + lambda _config: _FakeAxClient(), + ) + + result = optimize_to_joint_measurements(_joint_spec(tmp_path)) + + assert result.completed_trials == 3 + payload = json.loads(result.result_json_path.read_text(encoding="utf-8")) + assert payload["optimization_mode"] == "joint_measurement_fit" + assert set(payload["per_measurement_best_losses"]) == {"alpha1deg", "alpha2deg"} + assert payload["joint_loss_reduction"] == "mean" + + rows = list(csv.reader(result.trial_history_csv_path.open(encoding="utf-8"))) + assert rows[0] == ["trial_index", "loss", "depth_nm", "loss_alpha1deg", "loss_alpha2deg"] + assert len(rows) - 1 == 3 + + comparison_rows = list(csv.reader(result.comparison_csv_path.open(encoding="utf-8"))) + assert len(comparison_rows) - 1 == 4 + + +def test_optimize_to_joint_measurements_builds_one_flat_batch_per_trial( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _reset_fake_runner() + monkeypatch.setattr(objective_module, "BatchSimulationRunner", _FakeRunner) + monkeypatch.setattr( + joint_module, + "_create_ax_client_for_joint_config", + lambda _config: _FakeAxClient(), + ) + + optimize_to_joint_measurements(_joint_spec(tmp_path, total_trials=4)) + + assert _FakeRunner.batch_sizes == [4, 4, 4, 4] + + +def test_optimize_to_joint_measurements_uses_supplied_measurement_efficiency( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _reset_fake_runner() + monkeypatch.setattr(objective_module, "BatchSimulationRunner", _FakeRunner) + monkeypatch.setattr( + joint_module, + "_create_ax_client_for_joint_config", + lambda _config: _FakeAxClient(), + ) + measurement_path = _write_measurement(tmp_path / "m.dat", [(100.0, 0.2), (200.0, 0.3)]) + + result = optimize_to_joint_measurements( + { + "build_grating": lambda _parameters: SimpleNamespace(period_lpermm=2000.0), + "parameter_bounds": {"depth_nm": (1.0, 20.0)}, + "output_dir": tmp_path / "out", + "measurements": [ + { + "grazing_angle_deg": 1.0, + "measurement_path": measurement_path, + "evaluation_energies_ev": [100.0, 200.0], + "measurement_efficiency": [0.25, 0.25], + } + ], + "total_trials": 1, + "save_loss_plot": False, + "save_best_fit_plot": False, + } + ) + + assert result.best_loss == pytest.approx(0.0) diff --git a/tests/unit/test_grax_opt_resume.py b/tests/unit/test_grax_opt_resume.py new file mode 100644 index 0000000..10ee968 --- /dev/null +++ b/tests/unit/test_grax_opt_resume.py @@ -0,0 +1,343 @@ +from __future__ import annotations + +import csv +import json +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from grax_opt import checkpoint as checkpoint_module +from grax_opt import dynamic as dynamic_module +from grax_opt import joint as joint_module +from grax_opt import objective as objective_module +from grax_opt.checkpoint import OptimizerCheckpointPaths +from grax_opt.joint import optimize_to_joint_measurements +from grax_opt.loop import is_significant_improvement + + +class _FakeRunner: + resolved_max_workers = 1 + + call_count = 0 + + def __init__(self, **_kwargs: object) -> None: + pass + + def run_cases(self, cases, metadata=None): + type(self).call_count += 1 + for index, case in enumerate(cases): + yield SimpleNamespace( + index=index, + case_id=case["case_id"], + status="ok", + selected_efficiency=0.3, + ) + + +class _FakeAxClient: + def __init__(self, start_index: int = 0) -> None: + self.next_index = start_index + self.completed: list[int] = [] + + def create_experiment(self, **_kwargs: object) -> None: + return None + + def get_next_trial(self): + trial_index = self.next_index + self.next_index += 1 + return {"depth_nm": 5.0 + trial_index}, trial_index + + def complete_trial(self, trial_index, raw_data=None, data=None) -> None: + self.completed.append(int(trial_index)) + + def save_to_json_file(self, filepath: str) -> None: + Path(filepath).write_text(json.dumps({"next_index": self.next_index}), encoding="utf-8") + + @property + def experiment(self) -> SimpleNamespace: + return SimpleNamespace(trials={index: None for index in range(self.next_index)}) + + +def _install_fakes(monkeypatch: pytest.MonkeyPatch) -> None: + _FakeRunner.call_count = 0 + monkeypatch.setattr(objective_module, "BatchSimulationRunner", _FakeRunner) + monkeypatch.setattr( + dynamic_module, + "_create_ax_client_for_measurement_fit_config", + lambda _config: _FakeAxClient(), + ) + monkeypatch.setattr( + joint_module, + "_create_ax_client_for_joint_config", + lambda _config: _FakeAxClient(), + ) + monkeypatch.setattr( + checkpoint_module, + "load_ax_client_snapshot", + lambda snapshot_path, recorded_ax_version=None: _FakeAxClient( + start_index=int(json.loads(Path(snapshot_path).read_text(encoding="utf-8"))["next_index"]) + ), + ) + + +def _write_measurement(path: Path, rows: list[tuple[float, float]]) -> Path: + path.write_text( + "".join(f"{energy_ev} {efficiency}\n" for energy_ev, efficiency in rows), + encoding="utf-8", + ) + return path + + +def _spec(tmp_path: Path, **overrides: object) -> dict[str, object]: + measurement_path = tmp_path / "m.dat" + if not measurement_path.is_file(): + _write_measurement(measurement_path, [(100.0, 0.2), (200.0, 0.3), (300.0, 0.4)]) + spec: dict[str, object] = { + "build_grating": lambda _parameters: SimpleNamespace(period_lpermm=2000.0), + "parameter_bounds": {"depth_nm": (1.0, 20.0)}, + "measurement_path": measurement_path, + "output_dir": tmp_path / "out", + "evaluation_energies_ev": [100.0, 200.0, 300.0], + "total_trials": 5, + "save_best_fit_plot": False, + "save_loss_plot": False, + } + spec.update(overrides) + return spec + + +def test_is_significant_improvement_requires_minimum_relative_gain() -> None: + assert is_significant_improvement(float("inf"), 1.0, 0.5) is True + assert is_significant_improvement(1.0, 0.4, 0.5) is True + assert is_significant_improvement(1.0, 0.9, 0.5) is False + assert is_significant_improvement(1.0, 2.0, 0.5) is False + assert is_significant_improvement(1.0, 0.999, 0.0) is True + + +def test_resume_defaults_checkpoint_dir_to_output_dir(tmp_path: Path) -> None: + paths = OptimizerCheckpointPaths.for_config( + SimpleNamespace(output_dir=tmp_path / "out", checkpoint_dir=None) + ) + + assert paths.checkpoint_dir == tmp_path / "out" / "checkpoint" + assert paths.ax_snapshot_path.name == "ax_client_snapshot.json" + assert paths.state_path.name == "optimizer_state.json" + assert paths.trial_records_path.name == "trial_records.jsonl" + + +def test_run_writes_checkpoint_files(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + _install_fakes(monkeypatch) + + dynamic_module.optimize_to_measurements(_spec(tmp_path)) + + checkpoint_dir = tmp_path / "out" / "checkpoint" + assert sorted(path.name for path in checkpoint_dir.iterdir()) == [ + "ax_client_snapshot.json", + "optimizer_state.json", + "trial_records.jsonl", + ] + + +def test_resume_with_missing_checkpoint_starts_fresh( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fakes(monkeypatch) + + result = dynamic_module.optimize_to_measurements(_spec(tmp_path, resume=True)) + + assert result.completed_trials == 5 + + +def test_resume_restores_trials_and_extends_total_trials( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fakes(monkeypatch) + + first = dynamic_module.optimize_to_measurements(_spec(tmp_path, total_trials=5)) + calls_after_first_run = _FakeRunner.call_count + assert first.completed_trials == 5 + assert calls_after_first_run == 5 + + second = dynamic_module.optimize_to_measurements( + _spec(tmp_path, total_trials=12, resume=True) + ) + + assert second.completed_trials == 12 + assert _FakeRunner.call_count - calls_after_first_run == 7 + + rows = list(csv.reader(second.trial_history_csv_path.open(encoding="utf-8"))) + assert len(rows) - 1 == 12 + assert [row[0] for row in rows[1:]] == [str(index) for index in range(12)] + + +def test_resume_does_not_rerun_when_total_trials_already_reached( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fakes(monkeypatch) + dynamic_module.optimize_to_measurements(_spec(tmp_path, total_trials=5)) + calls_after_first_run = _FakeRunner.call_count + + result = dynamic_module.optimize_to_measurements(_spec(tmp_path, total_trials=5, resume=True)) + + assert _FakeRunner.call_count == calls_after_first_run + assert result.completed_trials == 5 + assert result.result_json_path.is_file() + + +def test_resume_rejects_changed_parameter_bounds( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fakes(monkeypatch) + dynamic_module.optimize_to_measurements(_spec(tmp_path)) + + with pytest.raises(ValueError, match="different optimization problem"): + dynamic_module.optimize_to_measurements( + _spec( + tmp_path, + total_trials=9, + resume=True, + parameter_bounds={"depth_nm": (1.0, 99.0)}, + ) + ) + + +def test_resume_rejects_changed_measurement_content( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fakes(monkeypatch) + dynamic_module.optimize_to_measurements(_spec(tmp_path)) + _write_measurement(tmp_path / "m.dat", [(100.0, 0.9), (200.0, 0.8), (300.0, 0.7)]) + + with pytest.raises(ValueError, match="different optimization problem"): + dynamic_module.optimize_to_measurements(_spec(tmp_path, total_trials=9, resume=True)) + + +def test_resume_ignores_torn_trial_record_line( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fakes(monkeypatch) + dynamic_module.optimize_to_measurements(_spec(tmp_path, total_trials=5)) + trial_records_path = tmp_path / "out" / "checkpoint" / "trial_records.jsonl" + trial_records_path.write_text( + trial_records_path.read_text(encoding="utf-8") + '{"trial_index": 9, "loss": ', + encoding="utf-8", + ) + + result = dynamic_module.optimize_to_measurements(_spec(tmp_path, total_trials=6, resume=True)) + + assert result.completed_trials == 6 + + +def test_resume_raises_on_corrupt_ax_snapshot( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fakes(monkeypatch) + dynamic_module.optimize_to_measurements(_spec(tmp_path)) + monkeypatch.setattr( + checkpoint_module, + "load_ax_client_snapshot", + lambda snapshot_path, recorded_ax_version=None: (_ for _ in ()).throw( + ValueError("Cannot resume: the Ax snapshot could not be loaded.") + ), + ) + + with pytest.raises(ValueError, match="Cannot resume"): + dynamic_module.optimize_to_measurements(_spec(tmp_path, total_trials=9, resume=True)) + + +def test_resume_raises_on_incomplete_checkpoint( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fakes(monkeypatch) + dynamic_module.optimize_to_measurements(_spec(tmp_path)) + (tmp_path / "out" / "checkpoint" / "optimizer_state.json").unlink() + + with pytest.raises(ValueError, match="incomplete"): + dynamic_module.optimize_to_measurements(_spec(tmp_path, total_trials=9, resume=True)) + + +def test_checkpoint_writes_leave_no_temporary_files( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fakes(monkeypatch) + + dynamic_module.optimize_to_measurements(_spec(tmp_path)) + + checkpoint_dir = tmp_path / "out" / "checkpoint" + assert [path.name for path in checkpoint_dir.iterdir() if ".tmp" in path.name] == [] + assert [path.name for path in (tmp_path / "out").iterdir() if ".tmp" in path.name] == [] + + +def test_resume_preserves_best_across_runs( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fakes(monkeypatch) + first = dynamic_module.optimize_to_measurements(_spec(tmp_path, total_trials=3)) + + second = dynamic_module.optimize_to_measurements(_spec(tmp_path, total_trials=6, resume=True)) + + assert second.best_loss == pytest.approx(first.best_loss) + assert second.best_parameters == first.best_parameters + + +def test_resume_accumulates_elapsed_seconds( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fakes(monkeypatch) + dynamic_module.optimize_to_measurements(_spec(tmp_path, total_trials=3)) + state_path = tmp_path / "out" / "checkpoint" / "optimizer_state.json" + first_state = json.loads(state_path.read_text(encoding="utf-8")) + + dynamic_module.optimize_to_measurements(_spec(tmp_path, total_trials=6, resume=True)) + second_state = json.loads(state_path.read_text(encoding="utf-8")) + + assert second_state["run_count"] == first_state["run_count"] + 1 + assert second_state["created"] == first_state["created"] + assert second_state["total_trials_history"] == [3, 6] + assert ( + second_state["cumulative_elapsed_seconds"] >= first_state["cumulative_elapsed_seconds"] + ) + + +def test_optimize_to_joint_measurements_resumes_and_extends( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fakes(monkeypatch) + first_path = _write_measurement(tmp_path / "m1.dat", [(100.0, 0.2), (200.0, 0.3)]) + second_path = _write_measurement(tmp_path / "m2.dat", [(100.0, 0.4), (200.0, 0.5)]) + + def joint_spec(total_trials: int, resume: bool) -> dict[str, object]: + return { + "build_grating": lambda _parameters: SimpleNamespace(period_lpermm=2000.0), + "parameter_bounds": {"depth_nm": (1.0, 20.0)}, + "output_dir": tmp_path / "out", + "measurements": [ + {"grazing_angle_deg": 1.0, "measurement_path": first_path}, + {"grazing_angle_deg": 2.0, "measurement_path": second_path}, + ], + "total_trials": total_trials, + "resume": resume, + "save_best_fit_plot": False, + "save_loss_plot": False, + } + + optimize_to_joint_measurements(joint_spec(4, False)) + calls_after_first_run = _FakeRunner.call_count + + result = optimize_to_joint_measurements(joint_spec(10, True)) + + assert result.completed_trials == 10 + assert _FakeRunner.call_count - calls_after_first_run == 6 From 9dd7c44021e5abe407b5af1697b976a84d7a81c7 Mon Sep 17 00:00:00 2001 From: Simone Vadilonga Date: Tue, 11 Aug 2026 09:13:32 +0200 Subject: [PATCH 06/10] Document joint multi-angle fitting and optimizer resume Co-Authored-By: Claude Sonnet 5 --- CHANGELOG.md | 5 + docs/api/optimization.md | 21 ++++- docs/developer/module-guide.md | 22 +++++ docs/tutorials/optimizer-joint-angles.md | 115 +++++++++++++++++++++++ docs/tutorials/optimizer-resume.md | 95 +++++++++++++++++++ docs/tutorials/optimizer.md | 21 +++++ 6 files changed, 278 insertions(+), 1 deletion(-) create mode 100644 docs/tutorials/optimizer-joint-angles.md create mode 100644 docs/tutorials/optimizer-resume.md diff --git a/CHANGELOG.md b/CHANGELOG.md index 090c43c..65d50ed 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,11 @@ # Changelog ## Unreleased +- Added `grax_opt.optimize_to_joint_measurements` for fitting one parameter set jointly against several measured curves taken at different grazing angles, each with its own measurement file and energy grid, with a configurable `joint_loss_reduction` (`mean`, `sum`, `pooled`, or `weighted`). Each trial evaluates every angle in a single `BatchSimulationRunner` batch so trial-level `max_workers` parallelizes across angles as well as energies. +- Added checkpoint and resume support to `optimize_to_measurements` and `optimize_to_joint_measurements` through the new `resume`, `checkpoint_dir`, and `checkpoint_interval` spec keys. The Ax client snapshot is persisted alongside optimizer run state and an append-only trial log, so a resumed run keeps its surrogate model and `total_trials` can be raised to extend a finished run. A problem fingerprint refuses to resume into a changed search space, naming the settings that differ. +- Fixed joint multi-angle evaluation assigning simulated efficiencies to the wrong angle/energy slot when `BatchSimulationRunner` returned results in completion order rather than input order; results are now reassembled by `CaseExecutionResult.index`. +- Extracted the shared Ax trial loop into `grax_opt.loop.run_ax_trial_loop`, so both optimizer entrypoints share candidate generation, best-so-far tracking, and early stopping. `early_stopping_min_relative_improvement` is now honored instead of being validated and ignored. +- Optimizer artifacts are now written atomically, the best-fit plot reuses the winning trial's cached simulated curve instead of re-simulating it after every trial, and a penalized trial logs the failing case or exception instead of failing silently. - `random-interface` roughness is now a correlated Gaussian random field (Gaussian autocorrelation) instead of per-sample white noise. `grax.RoughnessSpec` gains `correlation_length_nm`: the lateral autocorrelation length in nanometers, defaulting to one tenth of the grating period (`0.0` reproduces the previous white-noise interface). This produces physically smooth interfaces; long correlation lengths (much larger than the grating period) wash out geometrically and are better modelled by the `debye-waller` kind. - Added per-layer roughness: `LayerSpec` accepts `roughness_sigma_nm`, and `SingleLayerStack`/`MultilayerStack`/`CustomStack` expose per-layer/per-material roughness arguments (plus a `substrate_roughness_sigma_nm` for the substrate boundary). Each value sets the roughness of that interface, falling back to the grating-level `RoughnessSpec.sigma_nm` when unset. `random-interface` perturbs each interface with its own sigma; per-layer `debye-waller` sigmas combine in quadrature into an effective damping. - Web UI: the grating form now has a roughness sigma field for the substrate, each coating layer/material, and the top cap (persisted with the grating, schema v2). The run form has a roughness-kind dropdown (None / Debye-Waller / Random interface) that selects the model applied at run time. diff --git a/docs/api/optimization.md b/docs/api/optimization.md index 4fb9002..4b30766 100644 --- a/docs/api/optimization.md +++ b/docs/api/optimization.md @@ -12,10 +12,27 @@ This page exposes the primary public optimizer API for end users. .. autofunction:: grax_opt.optimize_to_measurements ``` -## Result type +## Joint multi-angle fitting + +Fits one parameter set against several measured curves recorded at different +grazing angles, each with its own energy grid. + +```{eval-rst} +.. autofunction:: grax_opt.optimize_to_joint_measurements + +.. autoclass:: grax_opt.AngleMeasurementSpec + +.. autoclass:: grax_opt.JointMeasurementFitConfig + +.. autofunction:: grax_opt.reduce_joint_losses +``` + +## Result types ```{eval-rst} .. autoclass:: grax_opt.OptimizationResult + +.. autoclass:: grax_opt.JointOptimizationResult ``` ## See tutorials @@ -23,3 +40,5 @@ This page exposes the primary public optimizer API for end users. - [Optimizer setup guide](../tutorials/optimizer.md) - [Laminar Grating](../tutorials/optimizer-laminar-fit.md) - [Blazed Grating](../tutorials/optimizer-blazed-fit.md) +- [Joint multi-angle fits](../tutorials/optimizer-joint-angles.md) +- [Resume an optimizer run](../tutorials/optimizer-resume.md) diff --git a/docs/developer/module-guide.md b/docs/developer/module-guide.md index 216f916..d14a2c6 100644 --- a/docs/developer/module-guide.md +++ b/docs/developer/module-guide.md @@ -34,6 +34,28 @@ not part of the first-class user documentation set for the core simulation package. Keep optimization-specific user material separate unless the core docs need to mention installation of the optional `opt` extra. +- `config.py`: `ParameterBounds`, the shared bounds value type +- `data.py`: measurement loading and interpolation onto an evaluation grid +- `evaluation.py`: normalization of evaluation energies/angles for the + single-measurement optimizer +- `objective.py`: trial evaluation, loss functions, joint multi-angle + evaluation, and the failure-penalty path +- `optimize.py`: Ax imports and version sniffing, candidate batching, trial + records, plotting, and atomic artifact writes +- `loop.py`: the Ax ask-and-tell loop shared by both optimizer entrypoints, + including best-so-far tracking and early stopping +- `dynamic.py`: `MeasurementFitConfig` and `optimize_to_measurements`, the + single-measurement fit +- `joint.py`: `JointMeasurementFitConfig` and `optimize_to_joint_measurements`, + the multi-angle fit +- `checkpoint.py`: checkpoint layout, problem fingerprinting, and the + resume session used by both entrypoints + +Both optimizer entrypoints share `loop.py` and `checkpoint.py`; mode-specific +work is limited to config validation, candidate evaluation, and artifact +persistence. Add new optimizer modes the same way rather than duplicating the +trial loop. + ## Where to make changes Add new grating shapes in `gratings.py` when they need to participate in the diff --git a/docs/tutorials/optimizer-joint-angles.md b/docs/tutorials/optimizer-joint-angles.md new file mode 100644 index 0000000..4bb7e4e --- /dev/null +++ b/docs/tutorials/optimizer-joint-angles.md @@ -0,0 +1,115 @@ +# Joint multi-angle fits + +`grax_opt.optimize_to_joint_measurements` fits **one** parameter set against +**several** measured curves recorded at different grazing angles. Use it when a +single geometry has to explain all of your measurements at once, rather than +fitting each angle separately and comparing the results afterwards. + +This is a separate entrypoint from +{doc}`optimize_to_measurements `, which fits a single measurement. + +## Basic setup + +```python +from grax_opt import optimize_to_joint_measurements + +result = optimize_to_joint_measurements({ + "build_grating": build_candidate_grating, + "parameter_bounds": { + "width_to_period_ratio": (0.45, 0.80), + "depth_nm": (4.5, 6.5), + "wall_angle_deg": (1.0, 40.0), + }, + "output_dir": "results/joint_fit", + "measurements": [ + {"grazing_angle_deg": 1.0, "measurement_path": "data/alpha1.dat"}, + {"grazing_angle_deg": 2.0, "measurement_path": "data/alpha2.dat"}, + {"grazing_angle_deg": 4.0, "measurement_path": "data/alpha4.dat"}, + ], + "diffraction_order": 1, + "fourier_orders": 15, + "total_trials": 200, + "max_workers": "auto", +}) + +print(result.best_loss, result.per_measurement_best_losses) +``` + +Every angle keeps its **own** energy grid. The grids do not have to match in +range or in length. + +## Measurement keys + +Each entry in `measurements` accepts: + +- `grazing_angle_deg` (required): the fixed angle for this curve. +- `measurement_path` (required): the measured two-column dataset. +- `evaluation_energies_ev`: energies to evaluate at. When omitted, the file's + own energy grid is used. When given, the measured curve is interpolated onto + it. +- `measurement_efficiency`: measured efficiencies to use **directly** instead of + interpolating the file, as described under "Pre-prepared measurements" below. +- `weight`: relative weight for the `"weighted"` reduction. Defaults to `1.0`. +- `label`: identifier used in artifacts. Defaults to `alphadeg`. + +## Combining the per-angle losses + +Each angle contributes its own mean squared error. `joint_loss_reduction` +controls how those are combined into the single value the optimizer minimizes: + +| Reduction | Joint loss | Use when | +| --- | --- | --- | +| `"mean"` (default) | mean of the per-angle MSEs | every angle should count equally | +| `"sum"` | sum of the per-angle MSEs | same ranking as `mean`, larger magnitude | +| `"pooled"` | MSE over all points pooled together | grids have **different point counts** and every measured point should count equally | +| `"weighted"` | weighted mean using each `weight` | some angles are more trustworthy than others | + +`"mean"` and `"pooled"` differ only when the angles have unequal point counts. +If you exclude absorption edges or downsample one angle more than another, +prefer `"pooled"` — otherwise the angle with fewer points is weighted more +heavily per measured point. + +## Pre-prepared measurements + +`grax_opt` deliberately contains no smoothing, downsampling, or edge-exclusion +logic. When you preprocess the measured curves yourself, pass the resulting +values through `measurement_efficiency` so the optimizer fits exactly the +numbers you prepared: + +```python +energies, efficiencies = my_own_preprocessing(raw_path) + +spec = { + "grazing_angle_deg": 2.0, + "measurement_path": raw_path, # kept for provenance + "evaluation_energies_ev": energies, + "measurement_efficiency": efficiencies, +} +``` + +When `measurement_efficiency` is omitted, the measured curve is interpolated +from the file onto `evaluation_energies_ev`. + +## Parallelism + +Each trial evaluates **all** angles in a single batch. With +`max_workers="auto"`, the worker pool parallelizes across angles and energies +together, so a three-angle fit keeps the pool busy far better than three +separate single-angle fits would. + +Because a batch runner may return results out of completion order, simulated +efficiencies are reassembled by result index rather than by arrival order. + +## Written artifacts + +- `best_result.json`: best fit, `per_measurement_best_losses`, per-angle + metadata, and run metadata. +- `trial_history.csv`: per-trial history with one `loss_