Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
30 changes: 30 additions & 0 deletions docs/api/session.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
189 changes: 189 additions & 0 deletions src/trade_study/_warm_start.py
Original file line number Diff line number Diff line change
@@ -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,
)
Loading
Loading