From 7caf426fa8a4d6d10ebe62b9c9d3f73a8161f202 Mon Sep 17 00:00:00 2001 From: "Joshua C. Macdonald" Date: Sat, 3 Oct 2026 01:16:40 -0400 Subject: [PATCH] feat: warm start adaptive sessions with schema-validated observations --- CHANGELOG.md | 1 + docs/api/session.md | 30 +++++ src/trade_study/_warm_start.py | 189 +++++++++++++++++++++++++++++++ src/trade_study/session.py | 105 ++++++++++++++++- tests/test_warm_start.py | 198 +++++++++++++++++++++++++++++++++ 5 files changed, 521 insertions(+), 2 deletions(-) create mode 100644 src/trade_study/_warm_start.py create mode 100644 tests/test_warm_start.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 2a3f1a4..6cf49b9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,6 +10,7 @@ All notable changes to this project will be documented in this file. - Incremental grid checkpoints preserve completed design-point/replicate evaluations across interruptions, including parallel workers and incomplete `Study` phases. `run_grid(max_retries=...)` provides opt-in bounded retries (#151). - Grouped surrogate validation holds out whole designs or regimes alongside separate row-validation metrics. Prediction and recommendation expose observed-support diagnostics and warn on extrapolation for GP and RF (#152). - `run_sequential()` allocates bounded extra replication to unresolved comparisons and feasibility boundaries, preserves raw replicate ids, and reports budgets, stopping reasons and simultaneous finite-horizon mean intervals under explicit bounded-score assumptions (#153). +- Adaptive sessions queue known configurations and import compatible completed observations with persistent identity/provenance and duplicate protection. Versioned session schemas validate reopened journals and refuse unverifiable legacy storage (#154). ## [0.3.0] — 2026-10-02 diff --git a/docs/api/session.md b/docs/api/session.md index 832706c..5bbfb27 100644 --- a/docs/api/session.md +++ b/docs/api/session.md @@ -32,6 +32,36 @@ Duplicate `tell` or `fail` reports and transitions from terminal states are rejected. Retrying an external evaluation may repeat its side effects; the caller must make those operations safe to repeat. +Use `session.enqueue(config)` to evaluate a known complete configuration before +sampling new ones. Every call queues a distinct evaluation. Queued configurations +are visible through `trials("waiting")`; `ask()` supplies their evaluation ids. +Continuous bounds, log factors and categorical/discrete levels are validated. + +To import completed observations, give both sessions the same explicit +`revision="simulator-v2-scorer-v1-data-v3-fidelity-high"` and matching factor, +objective (including weights/directions) and constraint definitions: + +```python +new_session.warm_start(previous_session) +# Saved tables from previous_session.results() also retain the import schema: +new_session.warm_start(load_results("previous-results")) +``` + +Imports preserve raw means, standard errors, replicate counts and evaluation +provenance. They create completed trials without calling a simulator. Repeating +an import skips known evaluation identities, including after reopening and +through intermediate sessions; conflicting results for an existing identity +are refused. All rows are validated before any are imported. Serialize imports +into a destination session. A bare grid table lacks a verifiable session schema +and cannot be imported automatically. The revision is the caller's assertion of +matching simulator, scorer, data, randomness and fidelity, not evidence inferred +from the scores. + +Journals now store a versioned schema and refuse incompatible definitions on +reopen. Nonempty legacy journals without that identity are also refused: use a +new journal or study name. Legacy observations require reconstruction and +validation against their original definitions; they are not silently adopted. + ::: trade_study.AdaptiveSession ::: trade_study.SessionTrial diff --git a/src/trade_study/_warm_start.py b/src/trade_study/_warm_start.py new file mode 100644 index 0000000..160467a --- /dev/null +++ b/src/trade_study/_warm_start.py @@ -0,0 +1,189 @@ +"""Schema and observation validation for importing adaptive evaluations.""" + +from __future__ import annotations + +import json +from numbers import Real +from typing import TYPE_CHECKING, Any + +import numpy as np + +from ._checkpoint import _value +from .design import FactorType + +if TYPE_CHECKING: + import optuna + + from .design import Factor + from .protocols import Constraint, Observable, ResultsTable + + +def _session_schema( + factors: list[Factor], + observables: list[Observable], + constraints: list[Constraint], + revision: str | None, +) -> str: + return json.dumps( + { + "version": 1, + "factors": _value(factors, revision), + "observables": _value(observables, revision), + "constraints": _value(constraints, revision), + "revision": revision, + }, + sort_keys=True, + allow_nan=False, + ) + + +def _validate_config(config: dict[str, Any], factors: list[Factor]) -> dict[str, Any]: + if set(config) != {f.name for f in factors}: + msg = "Configuration must contain every session factor exactly" + raise ValueError(msg) + validated: dict[str, Any] = {} + for factor in factors: + value = config[factor.name] + if factor.factor_type == FactorType.CONTINUOUS: + value = float(value) + if ( + factor.bounds is None + or not factor.bounds[0] <= value <= factor.bounds[1] + ): + msg = f"Configuration outside bounds for {factor.name!r}" + raise ValueError(msg) + else: + if value not in (factor.levels or []): + msg = f"Configuration outside declared levels for {factor.name!r}" + raise ValueError(msg) + value = next(level for level in (factor.levels or []) if level == value) + validated[factor.name] = value + return validated + + +def _distributions( + factors: list[Factor], +) -> dict[str, optuna.distributions.BaseDistribution]: + import optuna as _optuna + + distributions: dict[str, optuna.distributions.BaseDistribution] = {} + for factor in factors: + if factor.factor_type == FactorType.CONTINUOUS and factor.bounds is not None: + distributions[factor.name] = _optuna.distributions.FloatDistribution( + *factor.bounds, + log=factor.log_scale, + ) + else: + distributions[factor.name] = _optuna.distributions.CategoricalDistribution( + factor.levels or [], + ) + return distributions + + +def _summary( + metadata: dict[str, Any], + observables: list[Observable], + constraints: list[Constraint], + values: list[float], +) -> dict[str, Any]: + scores = metadata.get("scores", {}) + counts = metadata.get("n_reps", {}) + errors = metadata.get("standard_error", {}) + required = {o.name for o in observables} | {c.observable for c in constraints} + for name in required: + count = counts.get(name, 0) + mean = scores.get(name, np.nan) + error = errors.get(name, np.nan) + valid_count = ( + isinstance(count, int) and not isinstance(count, bool) and count >= 1 + ) + invalid_error = valid_count and count > 1 and (not _finite(error) or error < 0) + if not valid_count or not _finite(mean) or invalid_error: + msg = f"Invalid imported score/count/standard error for {name!r}" + raise ValueError(msg) + if not np.allclose( + values, [scores[o.name] * o.weight for o in observables], rtol=1e-12, atol=1e-12 + ): + msg = "Imported weighted objective values disagree with raw score metadata" + raise ValueError(msg) + for constraint in constraints: + if constraint.confidence is not None and counts[constraint.observable] < 2: + msg = "Imported confidence constraints need at least two replicates" + raise ValueError(msg) + constraint.bound( + scores[constraint.observable], errors.get(constraint.observable, np.nan) + ) + return { + "scores": dict(scores), + "n_reps": dict(counts), + "standard_error": dict(errors), + } + + +def _finite(value: object) -> bool: + return isinstance(value, Real) and bool(np.isfinite(float(value))) + + +def _import_trials( + results: ResultsTable, + factors: list[Factor], + observables: list[Observable], + constraints: list[Constraint], + schema: str, +) -> list[optuna.trial.FrozenTrial]: + import optuna as _optuna + + if ( + results.observable_names != [o.name for o in observables] + or results.scores.shape != (len(results.configs), len(observables)) + or len(results.metadata) != len(results.configs) + ): + msg = "Imported results have an incompatible observable or row schema" + raise ValueError(msg) + templates = [] + for config, values, meta in zip( + results.configs, results.scores, results.metadata, strict=True + ): + identity = meta.get("evaluation_id") + if ( + meta.get("session_schema") != schema + or not isinstance(identity, str) + or not identity.strip() + ): + msg = ( + "Imported results require matching session schema " + "and evaluation provenance" + ) + raise ValueError(msg) + summary = _summary(meta, observables, constraints, values.tolist()) + summary.update({ + "evaluation_id": meta["evaluation_id"], + "provenance": [ + *meta.get("provenance", []), + {"session_id": meta.get("session_id"), "trial": meta.get("trial")}, + ], + }) + templates.append( + _optuna.trial.create_trial( + state=_optuna.trial.TrialState.COMPLETE, + values=values.tolist(), + params=_validate_config(config, factors), + distributions=_distributions(factors), + user_attrs=summary, + ) + ) + return templates + + +def _fingerprint(trial: optuna.trial.FrozenTrial) -> str: + return json.dumps( + { + "params": trial.params, + "values": trial.values, + "summary": { + k: trial.user_attrs.get(k) + for k in ("scores", "n_reps", "standard_error") + }, + }, + sort_keys=True, + ) diff --git a/src/trade_study/session.py b/src/trade_study/session.py index 63f0e0c..1cba00c 100644 --- a/src/trade_study/session.py +++ b/src/trade_study/session.py @@ -22,9 +22,11 @@ from dataclasses import dataclass from pathlib import Path from typing import TYPE_CHECKING, Any +from uuid import uuid4 import numpy as np +from ._warm_start import _fingerprint, _import_trials, _session_schema, _validate_config from .design import FactorType from .protocols import Direction, ResultsTable @@ -91,6 +93,7 @@ def __init__( seed: int = 42, path: str | Path | None = None, study_name: str = _STUDY_NAME, + revision: str | None = None, ) -> None: """Create or reopen a session. @@ -105,6 +108,12 @@ def __init__( path: Journal file for a persistent session; in-memory when ``None``. study_name: Name of the study inside the journal. + revision: Caller-managed simulator/scorer/data/fidelity revision. + Required for importing completed observations between searches. + + Raises: + ValueError: If the revision is empty or a journal's stored schema + is incompatible or absent in an existing nonempty study. """ import optuna as _optuna from optuna.storages.journal import JournalFileBackend @@ -112,6 +121,11 @@ def __init__( self.factors = factors self.observables = observables self.constraints = list(constraints or []) + if revision is not None and not revision.strip(): + msg = "revision must be a nonempty model/data revision" + raise ValueError(msg) + self.revision = revision + self._schema = _session_schema(factors, observables, self.constraints, revision) for constraint in self.constraints: _constraint_value(constraint, 0.0) if path is None: @@ -137,6 +151,14 @@ def __init__( load_if_exists=True, ) self._study_id = self._storage.get_study_id_from_name(study_name) + stored = self._study.user_attrs.get("trade_study_identity") + if stored is None and not self._study.get_trials(): + stored = {"schema": self._schema, "session_id": str(uuid4())} + self._study.set_user_attr("trade_study_identity", stored) + elif not isinstance(stored, dict) or stored.get("schema") != self._schema: + msg = "Incompatible or legacy session schema; use a new journal/study name" + raise ValueError(msg) + self._session_id = str(stored["session_id"]) def _constraint_values(self, trial: optuna.trial.FrozenTrial) -> list[float]: scores = trial.user_attrs.get("scores", {}) @@ -269,7 +291,12 @@ def trials(self, state: str | None = None) -> list[SessionTrial]: msg = f"Unknown trial state {state!r}" raise ValueError(msg) return [ - SessionTrial(t.number, dict(t.params), states[t.state.name], t.user_attrs) + SessionTrial( + t.number, + dict(t.params or t.system_attrs.get("fixed_params", {})), + states[t.state.name], + t.user_attrs, + ) for t in self._study.get_trials(deepcopy=True) if state is None or states[t.state.name] == state ] @@ -346,6 +373,75 @@ def retry( child = self._storage.get_trial(internal) return child.number, dict(child.params) + def enqueue(self, config: dict[str, Any]) -> None: + """Queue a known configuration before sampling new configurations. + + Args: + config: Complete factor configuration inside the declared domain. + + Raises: + ValueError: If the configuration's factor names or values are invalid. + + Notes: + Each call queues a distinct evaluation. Completed observations + should instead be imported with :meth:`warm_start`. + """ + if set(config) != {f.name for f in self.factors}: + msg = "Configuration must contain every session factor exactly" + raise ValueError(msg) + self._study.enqueue_trial(_validate_config(config, self.factors)) + + def warm_start(self, source: AdaptiveSession | ResultsTable) -> list[int]: + """Import compatible completed evaluations, without re-evaluating them. + + Args: + source: Session or saved/loaded table produced by session.results(). + Factor/objective/constraint definitions and explicit revision + must match. Bare grid tables have no verifiable session schema. + + Returns: + Newly imported trial ids. Repeated imports skip evaluations already + present, including after reopening and through intermediate imports. + + Raises: + ValueError: If revision/schema/provenance/observations are invalid, + or an existing evaluation id has conflicting results. + + Notes: + Serialize imports into a destination session. The revision is the + caller's assertion of matching model, scorer, data and fidelity; + the library cannot establish that assertion from scores alone. + """ + if self.revision is None: + msg = "warm_start requires an explicit model/data revision" + raise ValueError(msg) + results = source.results() if isinstance(source, AdaptiveSession) else source + templates = _import_trials( + results, self.factors, self.observables, self.constraints, self._schema + ) + existing = { + t.user_attrs.get( + "evaluation_id", f"{self._session_id}:{t.number}" + ): _fingerprint(t) + for t in self._study.get_trials() + if t.state.name == "COMPLETE" + } + pending = [] + for trial in templates: + identity = trial.user_attrs["evaluation_id"] + fingerprint = _fingerprint(trial) + if identity in existing and existing[identity] != fingerprint: + msg = f"Conflicting imported evaluation {identity!r}" + raise ValueError(msg) + if identity not in existing: + pending.append(trial) + existing[identity] = fingerprint + imported = [] + for trial in pending: + self._study.add_trial(trial) + imported.append(self._study.get_trials()[-1].number) + return imported + def results(self) -> ResultsTable: """Return completed trials. @@ -368,13 +464,18 @@ def results(self) -> ResultsTable: metadata=[ { "trial": t.number, + "session_id": self._session_id, + "session_schema": self._schema, + "evaluation_id": t.user_attrs.get( + "evaluation_id", f"{self._session_id}:{t.number}" + ), "scores": t.user_attrs.get("scores", {}), "standard_error": t.user_attrs.get("standard_error", {}), "n_reps": t.user_attrs.get("n_reps", {}), "constraints": self._constraint_values(t), **{ k: t.user_attrs[k] - for k in ("retry_of", "retry_attempt") + for k in ("retry_of", "retry_attempt", "provenance") if k in t.user_attrs }, } diff --git a/tests/test_warm_start.py b/tests/test_warm_start.py new file mode 100644 index 0000000..0ae4b4e --- /dev/null +++ b/tests/test_warm_start.py @@ -0,0 +1,198 @@ +"""Schema-validated queueing, observation imports, and persistent provenance.""" + +from __future__ import annotations + +import copy +from typing import TYPE_CHECKING + +import numpy as np +import optuna +import pytest + +from trade_study import ( + AdaptiveSession, + Constraint, + Direction, + Factor, + FactorType, + Observable, + ResultsTable, + load_results, + save_results, +) + +if TYPE_CHECKING: + from pathlib import Path + +_FACTORS = [ + Factor("x", FactorType.CONTINUOUS, bounds=(0.1, 10), log_scale=True), + Factor("mode", FactorType.CATEGORICAL, levels=["a", "b"]), + Factor("size", FactorType.DISCRETE, levels=[1, 2, 4]), +] +_OBS = [ + Observable("loss", Direction.MINIMIZE, weight=2), + Observable("reward", Direction.MAXIMIZE), +] +_CONFIG = {"x": 1.0, "mode": "b", "size": 2} +_CONSTRAINTS = [Constraint("safe", "aux", "<=", 5)] + + +def _source() -> AdaptiveSession: + session = AdaptiveSession( + _FACTORS, _OBS, constraints=_CONSTRAINTS, revision="model-v1-data-v2" + ) + session.enqueue(_CONFIG) + ((trial, _),) = session.ask() + session.tell(trial, {"loss": [1.0, 2.0], "reward": [2.0, 4.0], "aux": [1.0, 3.0]}) + return session + + +def _destination(path: Path | None = None) -> AdaptiveSession: + return AdaptiveSession( + _FACTORS, + _OBS, + constraints=_CONSTRAINTS, + revision="model-v1-data-v2", + path=path, + ) + + +def test_queued_mixed_log_configuration_and_seeded_proposals(tmp_path: Path) -> None: + source = _source() + first, second = _destination(), _destination() + for session in (first, second): + assert session.warm_start(source) == [0] + session.enqueue(_CONFIG) + assert session.trials("waiting")[0].config == _CONFIG + assert first.ask(5) == second.ask(5) + assert first.trials("pending")[0].config == _CONFIG + path = tmp_path / "queued.journal" + persisted = _destination(path) + persisted.enqueue(_CONFIG) + ((_, config),) = _destination(path).ask() + assert config == _CONFIG + + +def test_saved_observations_import_once_even_through_another_session( + tmp_path: Path, +) -> None: + source = _source() + save_results(source.results(), tmp_path / "table") + table = load_results(tmp_path / "table") + path = tmp_path / "destination.journal" + destination = _destination(path) + assert destination.warm_start(table) == [0] + np.testing.assert_array_equal(destination.results().scores, source.results().scores) + meta = destination.results().metadata[0] + assert meta["n_reps"] == {"loss": 2, "reward": 2, "aux": 2} + assert meta["scores"] == {"loss": 1.5, "reward": 3.0, "aux": 2.0} + assert ( + meta["provenance"][0]["session_id"] + == source.results().metadata[0]["session_id"] + ) + assert destination.warm_start(source) == [] + assert _destination(path).warm_start(table) == [] + intermediate = _destination() + intermediate.warm_start(source) + assert _destination(path).warm_start(intermediate) == [] + assert source.warm_start(source) == [] + + +@pytest.mark.parametrize("change", ["revision", "weight", "factor", "constraint"]) +def test_schema_mismatch_is_rejected(change: str) -> None: + source = _source() + factors = _FACTORS[:-1] if change == "factor" else _FACTORS + obs = ( + [Observable("loss", Direction.MINIMIZE), _OBS[1]] + if change == "weight" + else _OBS + ) + constraints = ( + [Constraint("safe", "aux", "<=", 4)] if change == "constraint" else _CONSTRAINTS + ) + revision = "model-v2-data-v2" if change == "revision" else "model-v1-data-v2" + destination = AdaptiveSession( + factors, obs, constraints=constraints, revision=revision + ) + with pytest.raises(ValueError, match="matching session schema"): + destination.warm_start(source) + assert destination.results().configs == [] + + +@pytest.mark.parametrize( + "change", ["schema", "count", "error", "weighted", "config", "provenance"] +) +def test_malformed_import_is_rejected_without_adding_trials(change: str) -> None: + table = copy.deepcopy(_source().results()) + if change == "schema": + table.observable_names.reverse() + elif change == "count": + table.metadata[0]["n_reps"]["loss"] = "two" + elif change == "error": + table.metadata[0]["standard_error"]["loss"] = float("nan") + elif change == "weighted": + table.scores[0, 0] = 100 + elif change == "config": + table.configs[0]["x"] = 100 + else: + del table.metadata[0]["evaluation_id"] + destination = _destination() + with pytest.raises(ValueError, match=r"Imported|imported|Configuration"): + destination.warm_start(table) + assert destination.trials() == [] + + +def test_conflicting_existing_evaluation_is_rejected() -> None: + table = _source().results() + destination = _destination() + destination.warm_start(table) + table.scores[0, 0] = 4 + table.metadata[0]["scores"]["loss"] = 2 + with pytest.raises(ValueError, match="Conflicting imported evaluation"): + destination.warm_start(table) + np.testing.assert_array_equal(destination.results().scores, [[3, 3]]) + + +def test_import_requires_revision_and_source_schema() -> None: + destination = AdaptiveSession(_FACTORS, _OBS) + with pytest.raises(ValueError, match="explicit model/data revision"): + destination.warm_start(_source()) + grid_table = ResultsTable([_CONFIG], np.array([[1.0, 2.0]]), ["loss", "reward"]) + with pytest.raises(ValueError, match="row schema"): + _destination().warm_start(grid_table) + + +@pytest.mark.parametrize( + "config", + [ + {"x": 1.0}, + {**_CONFIG, "x": 0}, + {**_CONFIG, "mode": "other"}, + {**_CONFIG, "size": 3}, + {**_CONFIG, "extra": 1}, + ], +) +def test_invalid_queued_configuration_is_rejected(config: dict[str, object]) -> None: + with pytest.raises(ValueError, match="Configuration"): + _destination().enqueue(config) + + +def test_reopening_validates_schema_and_refuses_legacy_storage(tmp_path: Path) -> None: + path = tmp_path / "schema.journal" + _destination(path).ask() + with pytest.raises(ValueError, match="session schema"): + AdaptiveSession( + _FACTORS, _OBS, constraints=_CONSTRAINTS, revision="different", path=path + ) + legacy_path = tmp_path / "legacy.journal" + legacy = optuna.create_study( + study_name="trade-study-adaptive", + storage=optuna.storages.JournalStorage( + optuna.storages.journal.JournalFileBackend(str(legacy_path)) + ), + ) + legacy.ask() + with pytest.raises(ValueError, match="legacy session schema"): + _destination(legacy_path) + with pytest.raises(ValueError, match="nonempty"): + AdaptiveSession(_FACTORS, _OBS, revision=" ")