diff --git a/CHANGELOG.md b/CHANGELOG.md index 2e46ddf..fbc16ed 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -14,6 +14,10 @@ All notable changes to this project will be documented in this file. - Opt-in `EvaluationCache` reuses grid evaluations by typed configuration, replicate namespace/id, model/scorer revision, fidelity, objective and annotation definitions; includes provenance, bypass, invalidation and conflicting-evidence checks (#154). - `preference_sweep()` reports ranking/selection stability, regret, feasible Pareto alternatives, raw-unit practical equivalence and optional paired uncertainty under explicit normalization and preference assumptions, with exportable per-design summaries (#155). +### Fixed + +- Queued, retried and imported adaptive evaluations now join NSGA-II generations and process constraints when completed; interrupted imports resume idempotently. The adaptive extra now requires Optuna >=4.5 for its public generation API (#164). + ## [0.3.0] — 2026-10-02 ### Added diff --git a/docs/api/session.md b/docs/api/session.md index 5bbfb27..64edd93 100644 --- a/docs/api/session.md +++ b/docs/api/session.md @@ -36,6 +36,9 @@ 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. +Queued, retried and imported trials join the sampler population when completed; +constraints are processed through the same tell lifecycle as sampled trials. +The adaptive extra requires Optuna 4.5 or newer for public generation assignment. To import completed observations, give both sessions the same explicit `revision="simulator-v2-scorer-v1-data-v3-fidelity-high"` and matching factor, @@ -51,7 +54,9 @@ 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 +are refused. All rows are validated before any are imported. Interrupted imports +resume their pending completion on the next identical import without adding +another trial. 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 diff --git a/pyproject.toml b/pyproject.toml index b1bafa1..36ebe0e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -68,7 +68,7 @@ design = [ "scipy>=1.10", ] adaptive = [ - "optuna>=4.0", + "optuna>=4.5", ] parallel = [ "joblib>=1.3", diff --git a/src/trade_study/_warm_start.py b/src/trade_study/_warm_start.py index 160467a..15bfeb1 100644 --- a/src/trade_study/_warm_start.py +++ b/src/trade_study/_warm_start.py @@ -179,7 +179,7 @@ def _fingerprint(trial: optuna.trial.FrozenTrial) -> str: return json.dumps( { "params": trial.params, - "values": trial.values, + "values": trial.values or trial.user_attrs.get("_import_values"), "summary": { k: trial.user_attrs.get(k) for k in ("scores", "n_reps", "standard_error") diff --git a/src/trade_study/session.py b/src/trade_study/session.py index 1cba00c..3bdc2d3 100644 --- a/src/trade_study/session.py +++ b/src/trade_study/session.py @@ -140,6 +140,7 @@ def __init__( seed=seed, constraints_func=self._constraint_values if self.constraints else None, ) + self._sampler = sampler self._study = _optuna.create_study( study_name=study_name, storage=self._storage, @@ -203,7 +204,9 @@ def ask(self, n: int = 1) -> list[tuple[int, dict[str, Any]]]: proposals = [] for _ in range(n): trial = self._study.ask() - proposals.append((trial.number, self._suggest(trial))) + config = self._suggest(trial) + self._register_generation(trial.number) + proposals.append((trial.number, config)) return proposals def tell( @@ -252,6 +255,7 @@ def tell( self._storage.set_trial_user_attr( internal, "n_reps", {k: v[2] for k, v in summary.items()} ) + self._register_generation(trial_id) self._study.tell( trial_id, [summary[o.name][0] * o.weight for o in self.observables] ) @@ -265,6 +269,10 @@ def _trial_id(self, trial_id: int) -> int: msg = f"Unknown trial id {trial_id}" raise ValueError(msg) from error + def _register_generation(self, trial_id: int) -> None: + trial = self._storage.get_trial(self._trial_id(trial_id)) + self._sampler.get_trial_generation(self._study, trial) + def trials(self, state: str | None = None) -> list[SessionTrial]: """Inspect trials in creation order without exposing storage internals. @@ -371,6 +379,7 @@ def retry( ) internal = self._storage.create_new_trial(self._study_id, template) child = self._storage.get_trial(internal) + self._register_generation(child.number) return child.number, dict(child.params) def enqueue(self, config: dict[str, Any]) -> None: @@ -412,6 +421,8 @@ def warm_start(self, source: AdaptiveSession | ResultsTable) -> list[int]: caller's assertion of matching model, scorer, data and fidelity; the library cannot establish that assertion from scores alone. """ + import optuna as _optuna + if self.revision is None: msg = "warm_start requires an explicit model/data revision" raise ValueError(msg) @@ -419,12 +430,19 @@ def warm_start(self, source: AdaptiveSession | ResultsTable) -> list[int]: templates = _import_trials( results, self.factors, self.observables, self.constraints, self._schema ) + trials = self._study.get_trials() + unfinished = { + t.user_attrs["evaluation_id"]: t.number + for t in trials + if t.state.name == "RUNNING" and "_import_values" in t.user_attrs + } existing = { t.user_attrs.get( "evaluation_id", f"{self._session_id}:{t.number}" ): _fingerprint(t) - for t in self._study.get_trials() + for t in trials if t.state.name == "COMPLETE" + or (t.state.name == "RUNNING" and "_import_values" in t.user_attrs) } pending = [] for trial in templates: @@ -433,13 +451,25 @@ def warm_start(self, source: AdaptiveSession | ResultsTable) -> list[int]: if identity in existing and existing[identity] != fingerprint: msg = f"Conflicting imported evaluation {identity!r}" raise ValueError(msg) - if identity not in existing: + if identity not in existing or identity in unfinished: + trial.user_attrs["_resume_trial"] = unfinished.pop(identity, None) pending.append(trial) existing[identity] = fingerprint imported = [] for trial in pending: - self._study.add_trial(trial) - imported.append(self._study.get_trials()[-1].number) + number = trial.user_attrs.pop("_resume_trial") + if number is None: + running = _optuna.trial.create_trial( + state=_optuna.trial.TrialState.RUNNING, + params=trial.params, + distributions=trial.distributions, + user_attrs={**trial.user_attrs, "_import_values": trial.values}, + ) + internal = self._storage.create_new_trial(self._study_id, running) + number = self._storage.get_trial(internal).number + self._register_generation(number) + self._study.tell(number, trial.values) + imported.append(number) return imported def results(self) -> ResultsTable: diff --git a/tests/test_warm_start.py b/tests/test_warm_start.py index 0ae4b4e..6f2a21f 100644 --- a/tests/test_warm_start.py +++ b/tests/test_warm_start.py @@ -196,3 +196,86 @@ def test_reopening_validates_schema_and_refuses_legacy_storage(tmp_path: Path) - _destination(legacy_path) with pytest.raises(ValueError, match="nonempty"): AdaptiveSession(_FACTORS, _OBS, revision=" ") + + +def _loaded(path: Path) -> optuna.Study: + return optuna.load_study( + study_name="trade-study-adaptive", + storage=optuna.storages.JournalStorage( + optuna.storages.journal.JournalFileBackend(str(path)) + ), + ) + + +@pytest.mark.parametrize("mode", ["queued", "retried", "imported"]) +def test_known_parameters_join_the_constrained_population( + tmp_path: Path, mode: str +) -> None: + path = tmp_path / "population.journal" + session = _destination(path) + if mode == "imported": + expected = session.warm_start(_source())[0] + else: + session.enqueue(_CONFIG) + ((expected, _config),) = session.ask() + if mode == "retried": + session.fail(expected, "worker failed") + expected, _config = session.retry(expected) + session.tell(expected, {"loss": [1, 2], "reward": [2, 4], "aux": [1, 3]}) + stored = _loaded(path) + population = optuna.samplers.NSGAIISampler().get_population(stored, 0) + assert [t.number for t in population] == [expected] + assert population[0].system_attrs["constraints"] == pytest.approx([-3]) + + +def test_imported_population_can_supply_the_next_generation(tmp_path: Path) -> None: + source = _destination() + for _ in range(50): + source.enqueue(_CONFIG) + ((trial_id, _config),) = source.ask() + source.tell(trial_id, {"loss": [1, 2], "reward": [2, 4], "aux": [1, 3]}) + path = tmp_path / "parents.journal" + destination = _destination(path) + destination.warm_start(source) + stored = _loaded(path) + sampler = optuna.samplers.NSGAIISampler( + constraints_func=lambda trial: trial.system_attrs["constraints"], + ) + assert len(sampler.get_population(stored, 0)) == 50 + assert len(sampler.get_parent_population(stored, 1)) == 50 + assert all( + t.system_attrs["constraints"] == pytest.approx([-3]) for t in stored.trials + ) + + +def test_interrupted_import_resumes_without_duplicate_trials( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + table = _source().results() + path = tmp_path / "interrupted.journal" + destination = _destination(path) + + def interrupt( + sampler: optuna.samplers.NSGAIISampler, + study: optuna.Study, + trial: optuna.trial.FrozenTrial, + ) -> int: + del sampler, study, trial + raise KeyboardInterrupt + + with monkeypatch.context() as patch: + patch.setattr(optuna.samplers.NSGAIISampler, "get_trial_generation", interrupt) + with pytest.raises(KeyboardInterrupt): + destination.warm_start(table) + reopened = _destination(path) + assert len(reopened.trials("pending")) == 1 + duplicated_input = ResultsTable( + table.configs * 2, + np.vstack([table.scores, table.scores]), + table.observable_names, + metadata=table.metadata * 2, + ) + assert reopened.warm_start(duplicated_input) == [0] + assert len(reopened.trials()) == 1 + assert len(reopened.trials("complete")) == 1 + assert reopened.warm_start(table) == []