tunarag-python 0.2.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- tunarag/__init__.py +166 -0
- tunarag/cache.py +326 -0
- tunarag/config.py +257 -0
- tunarag/contracts.py +88 -0
- tunarag/dataset.py +481 -0
- tunarag/domain.py +83 -0
- tunarag/engine.py +979 -0
- tunarag/errors.py +179 -0
- tunarag/evaluators.py +186 -0
- tunarag/integrations/__init__.py +25 -0
- tunarag/integrations/mlflow.py +155 -0
- tunarag/integrations/runnables.py +270 -0
- tunarag/objective.py +75 -0
- tunarag/py.typed +1 -0
- tunarag/result.py +309 -0
- tunarag/retry.py +70 -0
- tunarag/search.py +286 -0
- tunarag/serialization.py +78 -0
- tunarag/stopping.py +193 -0
- tunarag/store.py +941 -0
- tunarag/synthetic.py +257 -0
- tunarag_python-0.2.1.dist-info/METADATA +1164 -0
- tunarag_python-0.2.1.dist-info/RECORD +24 -0
- tunarag_python-0.2.1.dist-info/WHEEL +4 -0
tunarag/engine.py
ADDED
|
@@ -0,0 +1,979 @@
|
|
|
1
|
+
"""Async optimization engine with retries, timeouts, caching, and sync wrapper."""
|
|
2
|
+
|
|
3
|
+
import asyncio
|
|
4
|
+
import json
|
|
5
|
+
import sqlite3
|
|
6
|
+
import time
|
|
7
|
+
from collections import defaultdict
|
|
8
|
+
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
|
9
|
+
from dataclasses import dataclass
|
|
10
|
+
from statistics import fmean
|
|
11
|
+
from typing import Any, TypeVar, cast
|
|
12
|
+
|
|
13
|
+
from .cache import CacheKey, CacheLookupStatus, CachePrivacy
|
|
14
|
+
from .config import CachePolicy, Settings
|
|
15
|
+
from .contracts import (
|
|
16
|
+
Cache,
|
|
17
|
+
Dataset,
|
|
18
|
+
Evaluator,
|
|
19
|
+
EventCallback,
|
|
20
|
+
Fingerprintable,
|
|
21
|
+
RAGAdapter,
|
|
22
|
+
ReplayableSearchStrategy,
|
|
23
|
+
SearchStrategy,
|
|
24
|
+
)
|
|
25
|
+
from .domain import Candidate, MetricValue, UsageRecord, ValueStatus
|
|
26
|
+
from .errors import ResumeError, SearchSpaceError, SyncInAsyncContextError
|
|
27
|
+
from .objective import Objective, ObjectiveResult
|
|
28
|
+
from .result import OptimizationResult, TrialSummary, UsageSummary
|
|
29
|
+
from .retry import RetryPolicy
|
|
30
|
+
from .serialization import content_hash
|
|
31
|
+
from .stopping import MaxTrials, StopCondition, StopContext
|
|
32
|
+
from .store import AttemptStatus, SQLiteStore, StudyRecord, StudyStatus, TrialRecord, TrialStatus
|
|
33
|
+
|
|
34
|
+
_OBJECTIVE_METRIC = "tunarag.objective"
|
|
35
|
+
_CACHE_PAYLOAD_VERSION = 1
|
|
36
|
+
_T = TypeVar("_T")
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
@dataclass(slots=True)
|
|
40
|
+
class _TrialFailure(Exception):
|
|
41
|
+
stage: str
|
|
42
|
+
cause_type: str
|
|
43
|
+
usage: tuple[UsageRecord, ...]
|
|
44
|
+
retryable: bool = False
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
@dataclass(frozen=True, slots=True)
|
|
48
|
+
class _SampleResult:
|
|
49
|
+
metrics: tuple[MetricValue, ...]
|
|
50
|
+
usage: tuple[UsageRecord, ...]
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class Optimizer:
|
|
54
|
+
"""Coordinate V0.1 optimization trials through public contracts."""
|
|
55
|
+
|
|
56
|
+
def __init__(
|
|
57
|
+
self,
|
|
58
|
+
*,
|
|
59
|
+
adapter: RAGAdapter,
|
|
60
|
+
strategy: SearchStrategy,
|
|
61
|
+
evaluators: Sequence[Evaluator],
|
|
62
|
+
objective: Objective,
|
|
63
|
+
store: SQLiteStore | None = None,
|
|
64
|
+
cache: Cache | None = None,
|
|
65
|
+
settings: Settings | None = None,
|
|
66
|
+
stop_conditions: Sequence[StopCondition] | None = None,
|
|
67
|
+
retry_policy: RetryPolicy | None = None,
|
|
68
|
+
callbacks: Sequence[EventCallback] = (),
|
|
69
|
+
sleeper: Callable[[float], Awaitable[None]] = asyncio.sleep,
|
|
70
|
+
clock: Callable[[], float] = time.monotonic,
|
|
71
|
+
duplicate_suggestion_limit: int = 100,
|
|
72
|
+
) -> None:
|
|
73
|
+
if not evaluators:
|
|
74
|
+
raise ValueError("optimizer requires at least one evaluator")
|
|
75
|
+
if duplicate_suggestion_limit < 1:
|
|
76
|
+
raise ValueError("duplicate suggestion limit must be positive")
|
|
77
|
+
self._adapter = adapter
|
|
78
|
+
self._strategy = strategy
|
|
79
|
+
self._evaluators = tuple(evaluators)
|
|
80
|
+
self._objective = objective
|
|
81
|
+
self._store = SQLiteStore() if store is None else store
|
|
82
|
+
self._owns_store = store is None
|
|
83
|
+
self._cache = cache
|
|
84
|
+
self._settings = Settings.resolve(environ={}) if settings is None else settings
|
|
85
|
+
self._stop_conditions = (
|
|
86
|
+
(MaxTrials(self._settings.max_trials),)
|
|
87
|
+
if stop_conditions is None
|
|
88
|
+
else tuple(stop_conditions)
|
|
89
|
+
)
|
|
90
|
+
if not self._stop_conditions:
|
|
91
|
+
raise ValueError("optimizer requires at least one stop condition")
|
|
92
|
+
self._retry_policy = retry_policy or RetryPolicy(
|
|
93
|
+
max_retries=self._settings.max_retries,
|
|
94
|
+
sample_timeout_seconds=self._settings.sample_timeout_seconds,
|
|
95
|
+
trial_timeout_seconds=self._settings.trial_timeout_seconds,
|
|
96
|
+
)
|
|
97
|
+
self._callbacks = tuple(callbacks)
|
|
98
|
+
self._event_cursors: dict[str, int] = {}
|
|
99
|
+
self._event_locks: dict[str, asyncio.Lock] = {}
|
|
100
|
+
self._sleeper = sleeper
|
|
101
|
+
self._clock = clock
|
|
102
|
+
self._duplicate_suggestion_limit = duplicate_suggestion_limit
|
|
103
|
+
|
|
104
|
+
@property
|
|
105
|
+
def store(self) -> SQLiteStore:
|
|
106
|
+
"""Return the experiment store used by this optimizer."""
|
|
107
|
+
|
|
108
|
+
return self._store
|
|
109
|
+
|
|
110
|
+
async def aoptimize(self, dataset: Dataset, *, study_name: str) -> OptimizationResult:
|
|
111
|
+
"""Run a new study and return terminal results."""
|
|
112
|
+
|
|
113
|
+
self._validate_dataset(dataset)
|
|
114
|
+
study = self._store.create_study(study_name, self._study_configuration(dataset))
|
|
115
|
+
study = self._store.transition_study(study.id, StudyStatus.RUNNING)
|
|
116
|
+
self._event_cursors[study.id] = -1
|
|
117
|
+
self._event_locks[study.id] = asyncio.Lock()
|
|
118
|
+
await self._dispatch_events(study.id)
|
|
119
|
+
return await self._continue_study(dataset, study, [], [])
|
|
120
|
+
|
|
121
|
+
async def aresume(self, dataset: Dataset, *, study_id: str) -> OptimizationResult:
|
|
122
|
+
"""Resume a compatible durable study after replaying committed successes."""
|
|
123
|
+
|
|
124
|
+
self._validate_dataset(dataset)
|
|
125
|
+
try:
|
|
126
|
+
study = self._store.get_study(study_id)
|
|
127
|
+
except KeyError as error:
|
|
128
|
+
raise ResumeError(
|
|
129
|
+
"cannot resume an unknown study",
|
|
130
|
+
stage="resume.load",
|
|
131
|
+
details={"study_id": study_id},
|
|
132
|
+
) from error
|
|
133
|
+
expected_hash = content_hash(self._study_configuration(dataset), namespace="study:v1")
|
|
134
|
+
if study.config_hash != expected_hash:
|
|
135
|
+
raise ResumeError(
|
|
136
|
+
"study configuration is incompatible with the current optimizer",
|
|
137
|
+
stage="resume.validate",
|
|
138
|
+
details={"study_id": study_id},
|
|
139
|
+
)
|
|
140
|
+
if study.status in {StudyStatus.FAILED, StudyStatus.CANCELLED}:
|
|
141
|
+
raise ResumeError(
|
|
142
|
+
"failed or cancelled studies are terminal and cannot resume",
|
|
143
|
+
stage="resume.validate",
|
|
144
|
+
details={"study_id": study_id, "status": study.status.value},
|
|
145
|
+
)
|
|
146
|
+
|
|
147
|
+
self._event_cursors[study.id] = self._latest_event_sequence(study.id)
|
|
148
|
+
self._event_locks[study.id] = asyncio.Lock()
|
|
149
|
+
summaries = self._load_terminal_summaries(study.id)
|
|
150
|
+
if study.status is StudyStatus.COMPLETED:
|
|
151
|
+
await self._replay_strategy(study.id, summaries)
|
|
152
|
+
return _build_result(
|
|
153
|
+
study.id, study.status, self._persisted_stop_reason(study.id), summaries
|
|
154
|
+
)
|
|
155
|
+
if study.status is StudyStatus.CREATED:
|
|
156
|
+
study = self._store.transition_study(study.id, StudyStatus.RUNNING)
|
|
157
|
+
|
|
158
|
+
abandoned = self._store.abandon_expired_attempts(study_id=study.id)
|
|
159
|
+
recoverable_ids = {trial.id for trial in self._store.list_recoverable_trials(study.id)}
|
|
160
|
+
unfinished = [
|
|
161
|
+
trial
|
|
162
|
+
for trial in self._store.list_trials(study.id)
|
|
163
|
+
if trial.status is TrialStatus.PENDING or trial.id in recoverable_ids
|
|
164
|
+
]
|
|
165
|
+
active = [
|
|
166
|
+
trial.id
|
|
167
|
+
for trial in self._store.list_trials(study.id)
|
|
168
|
+
if trial.status is TrialStatus.RUNNING and trial.id not in recoverable_ids
|
|
169
|
+
]
|
|
170
|
+
if active:
|
|
171
|
+
await self._dispatch_events(study.id)
|
|
172
|
+
raise ResumeError(
|
|
173
|
+
"study has attempts with active recovery leases",
|
|
174
|
+
stage="resume.recover",
|
|
175
|
+
retryable=True,
|
|
176
|
+
details={"study_id": study.id, "trial_ids": active},
|
|
177
|
+
)
|
|
178
|
+
await self._replay_strategy(study.id, summaries)
|
|
179
|
+
self._store.append_event(
|
|
180
|
+
study.id,
|
|
181
|
+
"study.resumed",
|
|
182
|
+
{
|
|
183
|
+
"terminal_trials": len(summaries),
|
|
184
|
+
"unfinished_trials": len(unfinished),
|
|
185
|
+
"abandoned_attempts": len(abandoned),
|
|
186
|
+
},
|
|
187
|
+
)
|
|
188
|
+
await self._dispatch_events(study.id)
|
|
189
|
+
return await self._continue_study(dataset, study, summaries, unfinished)
|
|
190
|
+
|
|
191
|
+
async def _continue_study(
|
|
192
|
+
self,
|
|
193
|
+
dataset: Dataset,
|
|
194
|
+
study: StudyRecord,
|
|
195
|
+
summaries: list[TrialSummary],
|
|
196
|
+
unfinished: Sequence[TrialRecord],
|
|
197
|
+
) -> OptimizationResult:
|
|
198
|
+
started_at = self._clock()
|
|
199
|
+
pending = list(unfinished)
|
|
200
|
+
stop_reason = "completed"
|
|
201
|
+
search_exhausted = False
|
|
202
|
+
try:
|
|
203
|
+
while True:
|
|
204
|
+
recovering = bool(pending)
|
|
205
|
+
capacity = (
|
|
206
|
+
self._settings.concurrency if recovering else self._batch_capacity(summaries)
|
|
207
|
+
)
|
|
208
|
+
if capacity == 0:
|
|
209
|
+
stop_reason = self._matching_stop_reason(summaries, started_at)
|
|
210
|
+
break
|
|
211
|
+
batch: list[tuple[Candidate, TrialRecord]] = []
|
|
212
|
+
while pending and len(batch) < capacity:
|
|
213
|
+
trial = pending.pop(0)
|
|
214
|
+
batch.append((_candidate_from_trial(trial), trial))
|
|
215
|
+
while not recovering and len(batch) < capacity and not search_exhausted:
|
|
216
|
+
stop_reason = self._matching_stop_reason(summaries, started_at)
|
|
217
|
+
if stop_reason:
|
|
218
|
+
break
|
|
219
|
+
reservation = await self._reserve_next_trial(study.id)
|
|
220
|
+
if reservation is None:
|
|
221
|
+
search_exhausted = True
|
|
222
|
+
self._store.append_event(study.id, "study.search_exhausted", {})
|
|
223
|
+
break
|
|
224
|
+
batch.append(reservation)
|
|
225
|
+
await self._dispatch_events(study.id)
|
|
226
|
+
if not batch:
|
|
227
|
+
if search_exhausted:
|
|
228
|
+
stop_reason = "search_exhausted"
|
|
229
|
+
break
|
|
230
|
+
completed = await asyncio.gather(
|
|
231
|
+
*(self._run_trial(dataset, candidate, trial) for candidate, trial in batch)
|
|
232
|
+
)
|
|
233
|
+
for summary in sorted(completed, key=lambda item: item.sequence):
|
|
234
|
+
summaries.append(summary)
|
|
235
|
+
if summary.status is TrialStatus.SUCCEEDED:
|
|
236
|
+
await self._observe_success(summary)
|
|
237
|
+
self._store.append_event(study.id, "study.stopping", {"reason": stop_reason})
|
|
238
|
+
study = self._store.transition_study(study.id, StudyStatus.COMPLETED)
|
|
239
|
+
await self._dispatch_events(study.id)
|
|
240
|
+
except asyncio.CancelledError:
|
|
241
|
+
current = self._store.get_study(study.id)
|
|
242
|
+
if current.status is StudyStatus.RUNNING:
|
|
243
|
+
self._store.transition_study(study.id, StudyStatus.CANCELLED)
|
|
244
|
+
raise
|
|
245
|
+
except BaseException:
|
|
246
|
+
current = self._store.get_study(study.id)
|
|
247
|
+
if current.status is StudyStatus.RUNNING:
|
|
248
|
+
self._store.transition_study(study.id, StudyStatus.FAILED)
|
|
249
|
+
raise
|
|
250
|
+
return _build_result(study.id, study.status, stop_reason, summaries)
|
|
251
|
+
|
|
252
|
+
def optimize(self, dataset: Dataset, *, study_name: str) -> OptimizationResult:
|
|
253
|
+
"""Run `aoptimize` when no event loop is active."""
|
|
254
|
+
|
|
255
|
+
try:
|
|
256
|
+
asyncio.get_running_loop()
|
|
257
|
+
except RuntimeError:
|
|
258
|
+
return asyncio.run(self.aoptimize(dataset, study_name=study_name))
|
|
259
|
+
raise SyncInAsyncContextError()
|
|
260
|
+
|
|
261
|
+
def resume(self, dataset: Dataset, *, study_id: str) -> OptimizationResult:
|
|
262
|
+
"""Run `aresume` when no event loop is active."""
|
|
263
|
+
|
|
264
|
+
try:
|
|
265
|
+
asyncio.get_running_loop()
|
|
266
|
+
except RuntimeError:
|
|
267
|
+
return asyncio.run(self.aresume(dataset, study_id=study_id))
|
|
268
|
+
raise SyncInAsyncContextError()
|
|
269
|
+
|
|
270
|
+
def close(self) -> None:
|
|
271
|
+
"""Close an internally created store; caller-owned resources remain open."""
|
|
272
|
+
|
|
273
|
+
if self._owns_store:
|
|
274
|
+
self._store.close()
|
|
275
|
+
|
|
276
|
+
def _latest_event_sequence(self, study_id: str) -> int:
|
|
277
|
+
events = self._store.list_events(study_id)
|
|
278
|
+
return -1 if not events else events[-1].sequence
|
|
279
|
+
|
|
280
|
+
async def _dispatch_events(self, study_id: str) -> None:
|
|
281
|
+
lock = self._event_locks.setdefault(study_id, asyncio.Lock())
|
|
282
|
+
async with lock:
|
|
283
|
+
cursor = self._event_cursors.get(study_id, -1)
|
|
284
|
+
events = self._store.list_events(study_id, after_sequence=cursor)
|
|
285
|
+
for event in events:
|
|
286
|
+
self._event_cursors[study_id] = event.sequence
|
|
287
|
+
if event.type == "callback.failed":
|
|
288
|
+
continue
|
|
289
|
+
for index, callback in enumerate(self._callbacks):
|
|
290
|
+
try:
|
|
291
|
+
await callback.on_event(event)
|
|
292
|
+
except Exception as error:
|
|
293
|
+
self._store.append_event(
|
|
294
|
+
study_id,
|
|
295
|
+
"callback.failed",
|
|
296
|
+
{
|
|
297
|
+
"source_sequence": event.sequence,
|
|
298
|
+
"callback_index": index,
|
|
299
|
+
"callback": _type_name(callback),
|
|
300
|
+
"cause": type(error).__name__,
|
|
301
|
+
},
|
|
302
|
+
)
|
|
303
|
+
self._event_cursors[study_id] = self._latest_event_sequence(study_id)
|
|
304
|
+
|
|
305
|
+
def _batch_capacity(self, summaries: Sequence[TrialSummary]) -> int:
|
|
306
|
+
capacity = self._settings.concurrency
|
|
307
|
+
for condition in self._stop_conditions:
|
|
308
|
+
if isinstance(condition, MaxTrials):
|
|
309
|
+
capacity = min(capacity, max(0, condition.limit - len(summaries)))
|
|
310
|
+
return capacity
|
|
311
|
+
|
|
312
|
+
@staticmethod
|
|
313
|
+
def _validate_dataset(dataset: Dataset) -> None:
|
|
314
|
+
if not dataset.version.strip():
|
|
315
|
+
raise ValueError("dataset version must not be empty")
|
|
316
|
+
|
|
317
|
+
def _load_terminal_summaries(self, study_id: str) -> list[TrialSummary]:
|
|
318
|
+
cached_trial_ids = {
|
|
319
|
+
event.payload.get("trial_id")
|
|
320
|
+
for event in self._store.list_events(study_id)
|
|
321
|
+
if event.type == "trial.cache_hit"
|
|
322
|
+
}
|
|
323
|
+
summaries: list[TrialSummary] = []
|
|
324
|
+
for trial in self._store.list_trials(study_id):
|
|
325
|
+
if trial.status in {TrialStatus.PENDING, TrialStatus.RUNNING}:
|
|
326
|
+
continue
|
|
327
|
+
persisted_metrics = self._store.list_metrics(trial.id)
|
|
328
|
+
objective_values = [
|
|
329
|
+
metric for metric in persisted_metrics if metric.name == _OBJECTIVE_METRIC
|
|
330
|
+
]
|
|
331
|
+
metrics = tuple(
|
|
332
|
+
metric for metric in persisted_metrics if metric.name != _OBJECTIVE_METRIC
|
|
333
|
+
)
|
|
334
|
+
objective: ObjectiveResult | None = None
|
|
335
|
+
if trial.status is TrialStatus.SUCCEEDED:
|
|
336
|
+
if len(objective_values) != 1 or objective_values[0].value is None:
|
|
337
|
+
raise ResumeError(
|
|
338
|
+
"successful trial is missing its committed objective",
|
|
339
|
+
stage="resume.replay",
|
|
340
|
+
details={"study_id": study_id, "trial_id": trial.id},
|
|
341
|
+
)
|
|
342
|
+
objective = ObjectiveResult(objective_values[0].value, objective_values[0].status)
|
|
343
|
+
usage = tuple(item.usage for item in self._store.list_usage(trial.id))
|
|
344
|
+
summaries.append(
|
|
345
|
+
TrialSummary(
|
|
346
|
+
id=trial.id,
|
|
347
|
+
sequence=trial.sequence,
|
|
348
|
+
candidate=_candidate_from_trial(trial),
|
|
349
|
+
status=trial.status,
|
|
350
|
+
metrics=metrics,
|
|
351
|
+
usage=usage,
|
|
352
|
+
objective=objective,
|
|
353
|
+
attempts=len(self._store.list_attempts(trial.id)),
|
|
354
|
+
cached=trial.id in cached_trial_ids,
|
|
355
|
+
error=trial.error,
|
|
356
|
+
)
|
|
357
|
+
)
|
|
358
|
+
return summaries
|
|
359
|
+
|
|
360
|
+
async def _replay_strategy(self, study_id: str, summaries: Sequence[TrialSummary]) -> None:
|
|
361
|
+
by_id = {summary.id: summary for summary in summaries}
|
|
362
|
+
current_trial_id: str | None = None
|
|
363
|
+
try:
|
|
364
|
+
for trial in self._store.list_trials(study_id):
|
|
365
|
+
current_trial_id = trial.id
|
|
366
|
+
summary = by_id.get(trial.id)
|
|
367
|
+
if isinstance(self._strategy, ReplayableSearchStrategy):
|
|
368
|
+
await self._strategy.replay(
|
|
369
|
+
_candidate_from_trial(trial),
|
|
370
|
+
() if summary is None else summary.metrics,
|
|
371
|
+
() if summary is None else summary.usage,
|
|
372
|
+
)
|
|
373
|
+
elif summary is not None and summary.status is TrialStatus.SUCCEEDED:
|
|
374
|
+
await self._strategy.observe(summary.candidate, summary.metrics, summary.usage)
|
|
375
|
+
except Exception as error:
|
|
376
|
+
raise ResumeError(
|
|
377
|
+
"search strategy failed while replaying reserved trials",
|
|
378
|
+
stage="resume.replay",
|
|
379
|
+
details={"cause": type(error).__name__, "trial_id": current_trial_id},
|
|
380
|
+
) from error
|
|
381
|
+
|
|
382
|
+
async def _observe_success(self, summary: TrialSummary) -> None:
|
|
383
|
+
try:
|
|
384
|
+
await self._strategy.observe(summary.candidate, summary.metrics, summary.usage)
|
|
385
|
+
except Exception as error:
|
|
386
|
+
raise SearchSpaceError(
|
|
387
|
+
"search strategy failed while observing a committed trial",
|
|
388
|
+
stage="strategy.observe",
|
|
389
|
+
details={"cause": type(error).__name__, "trial_id": summary.id},
|
|
390
|
+
) from error
|
|
391
|
+
|
|
392
|
+
def _persisted_stop_reason(self, study_id: str) -> str:
|
|
393
|
+
reasons = [
|
|
394
|
+
event.payload.get("reason")
|
|
395
|
+
for event in self._store.list_events(study_id)
|
|
396
|
+
if event.type == "study.stopping"
|
|
397
|
+
]
|
|
398
|
+
return reasons[-1] if reasons and isinstance(reasons[-1], str) else "already_completed"
|
|
399
|
+
|
|
400
|
+
async def _reserve_next_trial(self, study_id: str) -> tuple[Candidate, TrialRecord] | None:
|
|
401
|
+
for _ in range(self._duplicate_suggestion_limit):
|
|
402
|
+
try:
|
|
403
|
+
candidate = await self._strategy.suggest()
|
|
404
|
+
except Exception as error:
|
|
405
|
+
raise SearchSpaceError(
|
|
406
|
+
"search strategy failed while suggesting a candidate",
|
|
407
|
+
stage="strategy.suggest",
|
|
408
|
+
details={"cause": type(error).__name__},
|
|
409
|
+
) from error
|
|
410
|
+
if not isinstance(candidate, Candidate):
|
|
411
|
+
raise SearchSpaceError(
|
|
412
|
+
"search strategy returned a non-Candidate value",
|
|
413
|
+
stage="strategy.suggest",
|
|
414
|
+
details={"type": type(candidate).__name__},
|
|
415
|
+
)
|
|
416
|
+
try:
|
|
417
|
+
return candidate, self._store.reserve_trial(study_id, candidate)
|
|
418
|
+
except sqlite3.IntegrityError as error:
|
|
419
|
+
if "candidate_hash" not in str(error):
|
|
420
|
+
raise
|
|
421
|
+
return None
|
|
422
|
+
|
|
423
|
+
async def _run_trial(
|
|
424
|
+
self, dataset: Dataset, candidate: Candidate, trial: TrialRecord
|
|
425
|
+
) -> TrialSummary:
|
|
426
|
+
cache_key = self._cache_key(dataset, candidate)
|
|
427
|
+
cached_metrics = self._read_cache(trial, cache_key)
|
|
428
|
+
accumulated_usage: list[UsageRecord] = []
|
|
429
|
+
while True:
|
|
430
|
+
attempt = self._store.start_attempt(
|
|
431
|
+
trial.id, lease_seconds=self._retry_policy.lease_seconds
|
|
432
|
+
)
|
|
433
|
+
await self._dispatch_events(trial.study_id)
|
|
434
|
+
try:
|
|
435
|
+
if cached_metrics is None:
|
|
436
|
+
metrics, usage = await _wait_for(
|
|
437
|
+
self._execute_candidate(dataset, candidate),
|
|
438
|
+
self._retry_policy.trial_timeout_seconds,
|
|
439
|
+
)
|
|
440
|
+
else:
|
|
441
|
+
metrics, usage = cached_metrics, ()
|
|
442
|
+
objective = self._evaluate_objective(metrics, usage)
|
|
443
|
+
accumulated_usage.extend(usage)
|
|
444
|
+
objective_metric = MetricValue(
|
|
445
|
+
_OBJECTIVE_METRIC, cast(float, objective.score), objective.status
|
|
446
|
+
)
|
|
447
|
+
self._store.commit_trial_result(
|
|
448
|
+
trial.id,
|
|
449
|
+
status=TrialStatus.SUCCEEDED,
|
|
450
|
+
attempt_id=attempt.id,
|
|
451
|
+
metrics=(*metrics, objective_metric),
|
|
452
|
+
usage=usage,
|
|
453
|
+
)
|
|
454
|
+
cached = cached_metrics is not None
|
|
455
|
+
if cached:
|
|
456
|
+
self._store.append_event(
|
|
457
|
+
trial.study_id, "trial.cache_hit", {"trial_id": trial.id}
|
|
458
|
+
)
|
|
459
|
+
else:
|
|
460
|
+
self._write_cache(trial, cache_key, metrics)
|
|
461
|
+
await self._dispatch_events(trial.study_id)
|
|
462
|
+
return TrialSummary(
|
|
463
|
+
id=trial.id,
|
|
464
|
+
sequence=trial.sequence,
|
|
465
|
+
candidate=candidate,
|
|
466
|
+
status=TrialStatus.SUCCEEDED,
|
|
467
|
+
metrics=metrics,
|
|
468
|
+
usage=tuple(accumulated_usage),
|
|
469
|
+
objective=objective,
|
|
470
|
+
attempts=attempt.number + 1,
|
|
471
|
+
cached=cached,
|
|
472
|
+
)
|
|
473
|
+
except asyncio.CancelledError:
|
|
474
|
+
self._store.commit_trial_result(
|
|
475
|
+
trial.id,
|
|
476
|
+
status=TrialStatus.CANCELLED,
|
|
477
|
+
attempt_id=attempt.id,
|
|
478
|
+
error="trial cancelled",
|
|
479
|
+
)
|
|
480
|
+
raise
|
|
481
|
+
except _TrialFailure as failure:
|
|
482
|
+
accumulated_usage.extend(failure.usage)
|
|
483
|
+
failure_message = f"{failure.stage} failed: {failure.cause_type}"
|
|
484
|
+
if failure.retryable and attempt.number < self._retry_policy.max_retries:
|
|
485
|
+
self._store.finish_attempt(
|
|
486
|
+
attempt.id,
|
|
487
|
+
AttemptStatus.FAILED,
|
|
488
|
+
usage=failure.usage,
|
|
489
|
+
error=failure_message,
|
|
490
|
+
)
|
|
491
|
+
delay = self._retry_policy.delay_before_retry(attempt.number)
|
|
492
|
+
self._store.append_event(
|
|
493
|
+
trial.study_id,
|
|
494
|
+
"trial.retry_scheduled",
|
|
495
|
+
{
|
|
496
|
+
"trial_id": trial.id,
|
|
497
|
+
"attempt": attempt.number,
|
|
498
|
+
"delay_seconds": delay,
|
|
499
|
+
},
|
|
500
|
+
)
|
|
501
|
+
await self._dispatch_events(trial.study_id)
|
|
502
|
+
try:
|
|
503
|
+
await self._sleeper(delay)
|
|
504
|
+
except asyncio.CancelledError:
|
|
505
|
+
self._store.transition_trial(trial.id, TrialStatus.CANCELLED)
|
|
506
|
+
raise
|
|
507
|
+
cached_metrics = None
|
|
508
|
+
continue
|
|
509
|
+
self._store.commit_trial_result(
|
|
510
|
+
trial.id,
|
|
511
|
+
status=TrialStatus.FAILED,
|
|
512
|
+
attempt_id=attempt.id,
|
|
513
|
+
usage=failure.usage,
|
|
514
|
+
error=failure_message,
|
|
515
|
+
)
|
|
516
|
+
await self._dispatch_events(trial.study_id)
|
|
517
|
+
return TrialSummary(
|
|
518
|
+
id=trial.id,
|
|
519
|
+
sequence=trial.sequence,
|
|
520
|
+
candidate=candidate,
|
|
521
|
+
status=TrialStatus.FAILED,
|
|
522
|
+
metrics=(),
|
|
523
|
+
usage=tuple(accumulated_usage),
|
|
524
|
+
objective=None,
|
|
525
|
+
attempts=attempt.number + 1,
|
|
526
|
+
cached=cached_metrics is not None,
|
|
527
|
+
error=failure_message,
|
|
528
|
+
)
|
|
529
|
+
except Exception as error:
|
|
530
|
+
classified_failure = _TrialFailure(
|
|
531
|
+
"trial", type(error).__name__, (), self._retry_policy.is_retryable(error)
|
|
532
|
+
)
|
|
533
|
+
accumulated_usage.extend(classified_failure.usage)
|
|
534
|
+
failure_message = (
|
|
535
|
+
f"{classified_failure.stage} failed: {classified_failure.cause_type}"
|
|
536
|
+
)
|
|
537
|
+
if classified_failure.retryable and attempt.number < self._retry_policy.max_retries:
|
|
538
|
+
self._store.finish_attempt(
|
|
539
|
+
attempt.id, AttemptStatus.FAILED, error=failure_message
|
|
540
|
+
)
|
|
541
|
+
delay = self._retry_policy.delay_before_retry(attempt.number)
|
|
542
|
+
self._store.append_event(
|
|
543
|
+
trial.study_id,
|
|
544
|
+
"trial.retry_scheduled",
|
|
545
|
+
{
|
|
546
|
+
"trial_id": trial.id,
|
|
547
|
+
"attempt": attempt.number,
|
|
548
|
+
"delay_seconds": delay,
|
|
549
|
+
},
|
|
550
|
+
)
|
|
551
|
+
await self._dispatch_events(trial.study_id)
|
|
552
|
+
try:
|
|
553
|
+
await self._sleeper(delay)
|
|
554
|
+
except asyncio.CancelledError:
|
|
555
|
+
self._store.transition_trial(trial.id, TrialStatus.CANCELLED)
|
|
556
|
+
raise
|
|
557
|
+
cached_metrics = None
|
|
558
|
+
continue
|
|
559
|
+
self._store.commit_trial_result(
|
|
560
|
+
trial.id,
|
|
561
|
+
status=TrialStatus.FAILED,
|
|
562
|
+
attempt_id=attempt.id,
|
|
563
|
+
error=failure_message,
|
|
564
|
+
)
|
|
565
|
+
await self._dispatch_events(trial.study_id)
|
|
566
|
+
return TrialSummary(
|
|
567
|
+
id=trial.id,
|
|
568
|
+
sequence=trial.sequence,
|
|
569
|
+
candidate=candidate,
|
|
570
|
+
status=TrialStatus.FAILED,
|
|
571
|
+
metrics=(),
|
|
572
|
+
usage=tuple(accumulated_usage),
|
|
573
|
+
objective=None,
|
|
574
|
+
attempts=attempt.number + 1,
|
|
575
|
+
cached=False,
|
|
576
|
+
error=failure_message,
|
|
577
|
+
)
|
|
578
|
+
|
|
579
|
+
async def _execute_candidate(
|
|
580
|
+
self, dataset: Dataset, candidate: Candidate
|
|
581
|
+
) -> tuple[tuple[MetricValue, ...], tuple[UsageRecord, ...]]:
|
|
582
|
+
observations: dict[str, list[MetricValue]] = defaultdict(list)
|
|
583
|
+
usage: list[UsageRecord] = []
|
|
584
|
+
try:
|
|
585
|
+
examples = [example async for example in dataset]
|
|
586
|
+
except Exception as error:
|
|
587
|
+
raise _TrialFailure(
|
|
588
|
+
"dataset",
|
|
589
|
+
type(error).__name__,
|
|
590
|
+
(),
|
|
591
|
+
self._retry_policy.is_retryable(error),
|
|
592
|
+
) from error
|
|
593
|
+
if not examples:
|
|
594
|
+
raise _TrialFailure("dataset", "empty", tuple(usage))
|
|
595
|
+
|
|
596
|
+
semaphore = asyncio.Semaphore(self._settings.sample_concurrency)
|
|
597
|
+
|
|
598
|
+
async def bounded(example: Any) -> _SampleResult:
|
|
599
|
+
async with semaphore:
|
|
600
|
+
return await self._execute_sample(candidate, example)
|
|
601
|
+
|
|
602
|
+
outcomes = await asyncio.gather(
|
|
603
|
+
*(bounded(example) for example in examples), return_exceptions=True
|
|
604
|
+
)
|
|
605
|
+
first_failure: _TrialFailure | None = None
|
|
606
|
+
for outcome in outcomes:
|
|
607
|
+
if isinstance(outcome, _TrialFailure):
|
|
608
|
+
usage.extend(outcome.usage)
|
|
609
|
+
if first_failure is None:
|
|
610
|
+
first_failure = outcome
|
|
611
|
+
continue
|
|
612
|
+
if isinstance(outcome, BaseException):
|
|
613
|
+
raise outcome
|
|
614
|
+
usage.extend(outcome.usage)
|
|
615
|
+
for metric in outcome.metrics:
|
|
616
|
+
observations[metric.name].append(metric)
|
|
617
|
+
if first_failure is not None:
|
|
618
|
+
raise _TrialFailure(
|
|
619
|
+
first_failure.stage,
|
|
620
|
+
first_failure.cause_type,
|
|
621
|
+
tuple(usage),
|
|
622
|
+
first_failure.retryable,
|
|
623
|
+
) from first_failure
|
|
624
|
+
if not observations:
|
|
625
|
+
raise _TrialFailure("evaluation", "no_metrics", tuple(usage))
|
|
626
|
+
return _aggregate_metrics(observations, len(examples)), tuple(usage)
|
|
627
|
+
|
|
628
|
+
async def _execute_sample(self, candidate: Candidate, example: Any) -> _SampleResult:
|
|
629
|
+
usage: tuple[UsageRecord, ...] = ()
|
|
630
|
+
stage = "adapter"
|
|
631
|
+
try:
|
|
632
|
+
output, sample_usage = await _wait_for(
|
|
633
|
+
self._adapter.run(candidate, example),
|
|
634
|
+
self._retry_policy.sample_timeout_seconds,
|
|
635
|
+
)
|
|
636
|
+
usage = tuple(sample_usage)
|
|
637
|
+
if any(not isinstance(item, UsageRecord) for item in usage):
|
|
638
|
+
raise TypeError("adapter usage must contain UsageRecord values")
|
|
639
|
+
sample_metrics: list[MetricValue] = []
|
|
640
|
+
sample_names: set[str] = set()
|
|
641
|
+
for evaluator in self._evaluators:
|
|
642
|
+
stage = "evaluation"
|
|
643
|
+
values = tuple(
|
|
644
|
+
await _wait_for(
|
|
645
|
+
evaluator.evaluate(output, example),
|
|
646
|
+
self._retry_policy.sample_timeout_seconds,
|
|
647
|
+
)
|
|
648
|
+
)
|
|
649
|
+
for metric in values:
|
|
650
|
+
if not isinstance(metric, MetricValue):
|
|
651
|
+
raise TypeError("evaluator output must contain MetricValue values")
|
|
652
|
+
if metric.name == _OBJECTIVE_METRIC:
|
|
653
|
+
raise ValueError(f"metric name is reserved: {_OBJECTIVE_METRIC}")
|
|
654
|
+
if metric.name in sample_names:
|
|
655
|
+
raise ValueError(f"duplicate metric for sample: {metric.name}")
|
|
656
|
+
sample_names.add(metric.name)
|
|
657
|
+
sample_metrics.append(metric)
|
|
658
|
+
return _SampleResult(tuple(sample_metrics), usage)
|
|
659
|
+
except Exception as error:
|
|
660
|
+
raise _TrialFailure(
|
|
661
|
+
stage,
|
|
662
|
+
type(error).__name__,
|
|
663
|
+
usage,
|
|
664
|
+
self._retry_policy.is_retryable(error),
|
|
665
|
+
) from error
|
|
666
|
+
|
|
667
|
+
def _evaluate_objective(
|
|
668
|
+
self, metrics: Sequence[MetricValue], usage: tuple[UsageRecord, ...]
|
|
669
|
+
) -> ObjectiveResult:
|
|
670
|
+
try:
|
|
671
|
+
objective = self._objective.evaluate(metrics)
|
|
672
|
+
except Exception as error:
|
|
673
|
+
raise _TrialFailure(
|
|
674
|
+
"objective",
|
|
675
|
+
type(error).__name__,
|
|
676
|
+
usage,
|
|
677
|
+
self._retry_policy.is_retryable(error),
|
|
678
|
+
) from error
|
|
679
|
+
if objective.score is None:
|
|
680
|
+
raise _TrialFailure("objective", "unavailable", usage)
|
|
681
|
+
return objective
|
|
682
|
+
|
|
683
|
+
def _read_cache(self, trial: TrialRecord, key: CacheKey) -> tuple[MetricValue, ...] | None:
|
|
684
|
+
if self._cache is None or self._settings.cache_policy not in {
|
|
685
|
+
CachePolicy.READ_ONLY,
|
|
686
|
+
CachePolicy.READ_WRITE,
|
|
687
|
+
}:
|
|
688
|
+
return None
|
|
689
|
+
try:
|
|
690
|
+
lookup = self._cache.get(key)
|
|
691
|
+
except Exception as error:
|
|
692
|
+
self._store.append_event(
|
|
693
|
+
trial.study_id,
|
|
694
|
+
"trial.cache_read_failed",
|
|
695
|
+
{"trial_id": trial.id, "cause": type(error).__name__},
|
|
696
|
+
)
|
|
697
|
+
return None
|
|
698
|
+
if lookup.status is not CacheLookupStatus.HIT or lookup.payload is None:
|
|
699
|
+
return None
|
|
700
|
+
try:
|
|
701
|
+
return _decode_cached_metrics(lookup.payload)
|
|
702
|
+
except (TypeError, ValueError, json.JSONDecodeError) as error:
|
|
703
|
+
self._store.append_event(
|
|
704
|
+
trial.study_id,
|
|
705
|
+
"trial.cache_decode_failed",
|
|
706
|
+
{"trial_id": trial.id, "cause": type(error).__name__},
|
|
707
|
+
)
|
|
708
|
+
return None
|
|
709
|
+
|
|
710
|
+
def _write_cache(
|
|
711
|
+
self, trial: TrialRecord, key: CacheKey, metrics: Sequence[MetricValue]
|
|
712
|
+
) -> None:
|
|
713
|
+
if self._cache is None or self._settings.cache_policy not in {
|
|
714
|
+
CachePolicy.WRITE_ONLY,
|
|
715
|
+
CachePolicy.READ_WRITE,
|
|
716
|
+
}:
|
|
717
|
+
return
|
|
718
|
+
try:
|
|
719
|
+
self._cache.put(
|
|
720
|
+
key,
|
|
721
|
+
_encode_cached_metrics(metrics),
|
|
722
|
+
privacy=CachePrivacy.SENSITIVE,
|
|
723
|
+
)
|
|
724
|
+
except Exception as error:
|
|
725
|
+
self._store.append_event(
|
|
726
|
+
trial.study_id,
|
|
727
|
+
"trial.cache_write_failed",
|
|
728
|
+
{"trial_id": trial.id, "cause": type(error).__name__},
|
|
729
|
+
)
|
|
730
|
+
|
|
731
|
+
def _cache_key(self, dataset: Dataset, candidate: Candidate) -> CacheKey:
|
|
732
|
+
return CacheKey.from_parts(
|
|
733
|
+
"evaluation",
|
|
734
|
+
{
|
|
735
|
+
"payload_version": _CACHE_PAYLOAD_VERSION,
|
|
736
|
+
"dataset_version": dataset.version,
|
|
737
|
+
"candidate": candidate.parameters,
|
|
738
|
+
"adapter": _component_identity(self._adapter),
|
|
739
|
+
"evaluators": [_component_identity(item) for item in self._evaluators],
|
|
740
|
+
},
|
|
741
|
+
format_version=_CACHE_PAYLOAD_VERSION,
|
|
742
|
+
)
|
|
743
|
+
|
|
744
|
+
def _matching_stop_reason(self, summaries: Sequence[TrialSummary], started_at: float) -> str:
|
|
745
|
+
context = _stop_context(summaries, max(0.0, self._clock() - started_at))
|
|
746
|
+
for condition in self._stop_conditions:
|
|
747
|
+
if condition.should_stop(context):
|
|
748
|
+
return type(condition).__name__
|
|
749
|
+
return ""
|
|
750
|
+
|
|
751
|
+
def _study_configuration(self, dataset: Dataset) -> dict[str, Any]:
|
|
752
|
+
return {
|
|
753
|
+
"dataset_version": dataset.version,
|
|
754
|
+
"adapter": _component_identity(self._adapter),
|
|
755
|
+
"strategy": _component_identity(self._strategy),
|
|
756
|
+
"evaluators": [_component_identity(item) for item in self._evaluators],
|
|
757
|
+
"objective": [
|
|
758
|
+
{
|
|
759
|
+
"metric": term.metric,
|
|
760
|
+
"weight": term.weight,
|
|
761
|
+
"direction": term.direction.value,
|
|
762
|
+
"minimum": term.minimum,
|
|
763
|
+
"maximum": term.maximum,
|
|
764
|
+
}
|
|
765
|
+
for term in self._objective.terms
|
|
766
|
+
],
|
|
767
|
+
"max_trials": self._settings.max_trials,
|
|
768
|
+
"max_retries": self._retry_policy.max_retries,
|
|
769
|
+
"sample_timeout_seconds": self._retry_policy.sample_timeout_seconds,
|
|
770
|
+
"trial_timeout_seconds": self._retry_policy.trial_timeout_seconds,
|
|
771
|
+
"cache_policy": self._settings.cache_policy.value,
|
|
772
|
+
"stop_conditions": [
|
|
773
|
+
_component_identity(condition) for condition in self._stop_conditions
|
|
774
|
+
],
|
|
775
|
+
}
|
|
776
|
+
|
|
777
|
+
|
|
778
|
+
async def _wait_for(awaitable: Awaitable[_T], timeout: float | None) -> _T:
|
|
779
|
+
if timeout is None:
|
|
780
|
+
return await awaitable
|
|
781
|
+
return await asyncio.wait_for(awaitable, timeout)
|
|
782
|
+
|
|
783
|
+
|
|
784
|
+
def _encode_cached_metrics(metrics: Sequence[MetricValue]) -> bytes:
|
|
785
|
+
document = {
|
|
786
|
+
"version": _CACHE_PAYLOAD_VERSION,
|
|
787
|
+
"metrics": [
|
|
788
|
+
{
|
|
789
|
+
"name": metric.name,
|
|
790
|
+
"value": metric.value,
|
|
791
|
+
"status": metric.status.value,
|
|
792
|
+
"coverage": metric.coverage,
|
|
793
|
+
}
|
|
794
|
+
for metric in metrics
|
|
795
|
+
],
|
|
796
|
+
}
|
|
797
|
+
return json.dumps(document, ensure_ascii=True, separators=(",", ":"), sort_keys=True).encode()
|
|
798
|
+
|
|
799
|
+
|
|
800
|
+
def _decode_cached_metrics(payload: bytes) -> tuple[MetricValue, ...]:
|
|
801
|
+
document = json.loads(payload.decode("utf-8"))
|
|
802
|
+
if not isinstance(document, dict) or document.get("version") != _CACHE_PAYLOAD_VERSION:
|
|
803
|
+
raise ValueError("unsupported cached evaluation payload")
|
|
804
|
+
items = document.get("metrics")
|
|
805
|
+
if not isinstance(items, list) or not items:
|
|
806
|
+
raise ValueError("cached evaluation metrics must be a nonempty list")
|
|
807
|
+
metrics: list[MetricValue] = []
|
|
808
|
+
for item in items:
|
|
809
|
+
if not isinstance(item, Mapping):
|
|
810
|
+
raise TypeError("cached metric must be an object")
|
|
811
|
+
metrics.append(
|
|
812
|
+
MetricValue(
|
|
813
|
+
name=_required_string(item, "name"),
|
|
814
|
+
value=_optional_number(item.get("value")),
|
|
815
|
+
status=ValueStatus(_required_string(item, "status")),
|
|
816
|
+
coverage=_required_number(item, "coverage"),
|
|
817
|
+
)
|
|
818
|
+
)
|
|
819
|
+
if len({metric.name for metric in metrics}) != len(metrics):
|
|
820
|
+
raise ValueError("cached evaluation contains duplicate metric names")
|
|
821
|
+
return tuple(metrics)
|
|
822
|
+
|
|
823
|
+
|
|
824
|
+
def _required_string(item: Mapping[str, Any], key: str) -> str:
|
|
825
|
+
value = item.get(key)
|
|
826
|
+
if not isinstance(value, str):
|
|
827
|
+
raise TypeError(f"cached metric {key} must be a string")
|
|
828
|
+
return value
|
|
829
|
+
|
|
830
|
+
|
|
831
|
+
def _required_number(item: Mapping[str, Any], key: str) -> float:
|
|
832
|
+
value = item.get(key)
|
|
833
|
+
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
|
834
|
+
raise TypeError(f"cached metric {key} must be numeric")
|
|
835
|
+
return float(value)
|
|
836
|
+
|
|
837
|
+
|
|
838
|
+
def _optional_number(value: Any) -> float | None:
|
|
839
|
+
if value is None:
|
|
840
|
+
return None
|
|
841
|
+
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
|
842
|
+
raise TypeError("cached metric value must be numeric or null")
|
|
843
|
+
return float(value)
|
|
844
|
+
|
|
845
|
+
|
|
846
|
+
def _candidate_from_trial(trial: TrialRecord) -> Candidate:
|
|
847
|
+
parameters = json.loads(trial.candidate_json)
|
|
848
|
+
if not isinstance(parameters, Mapping):
|
|
849
|
+
raise ResumeError(
|
|
850
|
+
"persisted candidate is not an object",
|
|
851
|
+
stage="resume.replay",
|
|
852
|
+
details={"trial_id": trial.id},
|
|
853
|
+
)
|
|
854
|
+
try:
|
|
855
|
+
return Candidate(dict(parameters))
|
|
856
|
+
except (TypeError, ValueError) as error:
|
|
857
|
+
raise ResumeError(
|
|
858
|
+
"persisted candidate is invalid",
|
|
859
|
+
stage="resume.replay",
|
|
860
|
+
details={"trial_id": trial.id, "cause": type(error).__name__},
|
|
861
|
+
) from error
|
|
862
|
+
|
|
863
|
+
|
|
864
|
+
def _aggregate_metrics(
|
|
865
|
+
observations: dict[str, list[MetricValue]], sample_count: int
|
|
866
|
+
) -> tuple[MetricValue, ...]:
|
|
867
|
+
aggregated: list[MetricValue] = []
|
|
868
|
+
for name in sorted(observations):
|
|
869
|
+
values = observations[name]
|
|
870
|
+
available = [float(item.value) for item in values if item.value is not None]
|
|
871
|
+
coverage = sum(item.coverage for item in values) / sample_count
|
|
872
|
+
if not available:
|
|
873
|
+
aggregated.append(MetricValue(name, None, ValueStatus.UNAVAILABLE, coverage))
|
|
874
|
+
continue
|
|
875
|
+
status = (
|
|
876
|
+
ValueStatus.ESTIMATED
|
|
877
|
+
if any(item.status is not ValueStatus.EXACT for item in values)
|
|
878
|
+
else ValueStatus.EXACT
|
|
879
|
+
)
|
|
880
|
+
aggregated.append(MetricValue(name, fmean(available), status, coverage))
|
|
881
|
+
return tuple(aggregated)
|
|
882
|
+
|
|
883
|
+
|
|
884
|
+
def _stop_context(summaries: Sequence[TrialSummary], elapsed_seconds: float) -> StopContext:
|
|
885
|
+
exact_cost = 0.0
|
|
886
|
+
estimated_cost = 0.0
|
|
887
|
+
for summary in summaries:
|
|
888
|
+
for item in summary.usage:
|
|
889
|
+
if item.cost is None:
|
|
890
|
+
continue
|
|
891
|
+
if item.status is ValueStatus.EXACT:
|
|
892
|
+
exact_cost += item.cost
|
|
893
|
+
else:
|
|
894
|
+
estimated_cost += item.cost
|
|
895
|
+
successful = sum(summary.status is TrialStatus.SUCCEEDED for summary in summaries)
|
|
896
|
+
failed = sum(summary.status is TrialStatus.FAILED for summary in summaries)
|
|
897
|
+
scores = tuple(
|
|
898
|
+
summary.objective.score
|
|
899
|
+
if summary.status is TrialStatus.SUCCEEDED and summary.objective is not None
|
|
900
|
+
else None
|
|
901
|
+
for summary in summaries
|
|
902
|
+
)
|
|
903
|
+
return StopContext(
|
|
904
|
+
completed_trials=len(summaries),
|
|
905
|
+
exact_cost=exact_cost,
|
|
906
|
+
estimated_cost=estimated_cost,
|
|
907
|
+
elapsed_seconds=elapsed_seconds,
|
|
908
|
+
successful_trials=successful,
|
|
909
|
+
failed_trials=failed,
|
|
910
|
+
objective_scores=scores,
|
|
911
|
+
)
|
|
912
|
+
|
|
913
|
+
|
|
914
|
+
def _build_result(
|
|
915
|
+
study_id: str,
|
|
916
|
+
study_status: StudyStatus,
|
|
917
|
+
stop_reason: str,
|
|
918
|
+
summaries: Sequence[TrialSummary],
|
|
919
|
+
) -> OptimizationResult:
|
|
920
|
+
successful = [
|
|
921
|
+
summary
|
|
922
|
+
for summary in summaries
|
|
923
|
+
if summary.status is TrialStatus.SUCCEEDED
|
|
924
|
+
and summary.objective is not None
|
|
925
|
+
and summary.objective.score is not None
|
|
926
|
+
]
|
|
927
|
+
best = (
|
|
928
|
+
max(successful, key=lambda item: _objective_score(item.objective)) if successful else None
|
|
929
|
+
)
|
|
930
|
+
all_usage = [item for summary in summaries for item in summary.usage]
|
|
931
|
+
return OptimizationResult(
|
|
932
|
+
study_id=study_id,
|
|
933
|
+
study_status=study_status,
|
|
934
|
+
stop_reason=stop_reason,
|
|
935
|
+
trials=tuple(summaries),
|
|
936
|
+
best_trial=best,
|
|
937
|
+
usage=_usage_summary(all_usage),
|
|
938
|
+
)
|
|
939
|
+
|
|
940
|
+
|
|
941
|
+
def _objective_score(result: ObjectiveResult | None) -> float:
|
|
942
|
+
if result is None or result.score is None:
|
|
943
|
+
raise ValueError("successful trial objective is unavailable")
|
|
944
|
+
return result.score
|
|
945
|
+
|
|
946
|
+
|
|
947
|
+
def _usage_summary(usage: Sequence[UsageRecord]) -> UsageSummary:
|
|
948
|
+
input_tokens = sum(item.input_tokens or 0 for item in usage)
|
|
949
|
+
output_tokens = sum(item.output_tokens or 0 for item in usage)
|
|
950
|
+
exact_cost = sum(item.cost or 0.0 for item in usage if item.status is ValueStatus.EXACT)
|
|
951
|
+
estimated_cost = sum(item.cost or 0.0 for item in usage if item.status is not ValueStatus.EXACT)
|
|
952
|
+
latency = sum(item.latency_seconds or 0.0 for item in usage)
|
|
953
|
+
statuses = {item.status for item in usage}
|
|
954
|
+
if not usage or statuses == {ValueStatus.UNAVAILABLE}:
|
|
955
|
+
status = ValueStatus.UNAVAILABLE
|
|
956
|
+
elif statuses == {ValueStatus.EXACT}:
|
|
957
|
+
status = ValueStatus.EXACT
|
|
958
|
+
else:
|
|
959
|
+
status = ValueStatus.ESTIMATED
|
|
960
|
+
return UsageSummary(
|
|
961
|
+
input_tokens,
|
|
962
|
+
output_tokens,
|
|
963
|
+
exact_cost,
|
|
964
|
+
estimated_cost,
|
|
965
|
+
latency,
|
|
966
|
+
status,
|
|
967
|
+
)
|
|
968
|
+
|
|
969
|
+
|
|
970
|
+
def _type_name(value: object) -> str:
|
|
971
|
+
value_type = type(value)
|
|
972
|
+
return f"{value_type.__module__}.{value_type.__qualname__}"
|
|
973
|
+
|
|
974
|
+
|
|
975
|
+
def _component_identity(value: object) -> dict[str, Any]:
|
|
976
|
+
identity: dict[str, Any] = {"type": _type_name(value)}
|
|
977
|
+
if isinstance(value, Fingerprintable):
|
|
978
|
+
identity["fingerprint"] = value.fingerprint()
|
|
979
|
+
return identity
|