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/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