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/retry.py
ADDED
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
"""Deterministic retry and timeout policy."""
|
|
2
|
+
|
|
3
|
+
import asyncio
|
|
4
|
+
import math
|
|
5
|
+
from dataclasses import dataclass
|
|
6
|
+
|
|
7
|
+
from .errors import TunaRAGError
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
@dataclass(frozen=True, slots=True)
|
|
11
|
+
class RetryPolicy:
|
|
12
|
+
"""Bound retries, exponential backoff, and execution timeouts."""
|
|
13
|
+
|
|
14
|
+
max_retries: int = 2
|
|
15
|
+
sample_timeout_seconds: float | None = 60.0
|
|
16
|
+
trial_timeout_seconds: float | None = 300.0
|
|
17
|
+
base_delay_seconds: float = 0.5
|
|
18
|
+
max_delay_seconds: float = 30.0
|
|
19
|
+
|
|
20
|
+
def __post_init__(self) -> None:
|
|
21
|
+
if (
|
|
22
|
+
not isinstance(self.max_retries, int)
|
|
23
|
+
or isinstance(self.max_retries, bool)
|
|
24
|
+
or self.max_retries < 0
|
|
25
|
+
):
|
|
26
|
+
raise ValueError("max_retries must be a nonnegative integer")
|
|
27
|
+
_validate_optional_timeout("sample timeout", self.sample_timeout_seconds)
|
|
28
|
+
_validate_optional_timeout("trial timeout", self.trial_timeout_seconds)
|
|
29
|
+
_validate_nonnegative_finite("base retry delay", self.base_delay_seconds)
|
|
30
|
+
_validate_nonnegative_finite("maximum retry delay", self.max_delay_seconds)
|
|
31
|
+
if self.base_delay_seconds > self.max_delay_seconds:
|
|
32
|
+
raise ValueError("base retry delay must not exceed maximum retry delay")
|
|
33
|
+
|
|
34
|
+
@property
|
|
35
|
+
def max_attempts(self) -> int:
|
|
36
|
+
"""Return the initial attempt plus configured retries."""
|
|
37
|
+
|
|
38
|
+
return self.max_retries + 1
|
|
39
|
+
|
|
40
|
+
@property
|
|
41
|
+
def lease_seconds(self) -> float:
|
|
42
|
+
"""Return a recovery lease longer than the configured trial timeout."""
|
|
43
|
+
|
|
44
|
+
return 86_400.0 if self.trial_timeout_seconds is None else self.trial_timeout_seconds + 60.0
|
|
45
|
+
|
|
46
|
+
def delay_before_retry(self, failed_attempt_number: int) -> float:
|
|
47
|
+
"""Return deterministic capped exponential backoff."""
|
|
48
|
+
|
|
49
|
+
if failed_attempt_number < 0:
|
|
50
|
+
raise ValueError("failed attempt number must not be negative")
|
|
51
|
+
return float(
|
|
52
|
+
min(self.base_delay_seconds * (2**failed_attempt_number), self.max_delay_seconds)
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
def is_retryable(self, error: BaseException) -> bool:
|
|
56
|
+
"""Classify known transient failures without parsing messages."""
|
|
57
|
+
|
|
58
|
+
if isinstance(error, TunaRAGError):
|
|
59
|
+
return error.retryable
|
|
60
|
+
return isinstance(error, (TimeoutError, asyncio.TimeoutError, ConnectionError))
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def _validate_optional_timeout(name: str, value: float | None) -> None:
|
|
64
|
+
if value is not None and (not math.isfinite(value) or value <= 0):
|
|
65
|
+
raise ValueError(f"{name} must be a positive finite number or None")
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def _validate_nonnegative_finite(name: str, value: float) -> None:
|
|
69
|
+
if not math.isfinite(value) or value < 0:
|
|
70
|
+
raise ValueError(f"{name} must be a nonnegative finite number")
|
tunarag/search.py
ADDED
|
@@ -0,0 +1,286 @@
|
|
|
1
|
+
"""Typed search spaces and the deterministic random strategy."""
|
|
2
|
+
|
|
3
|
+
import math
|
|
4
|
+
import random
|
|
5
|
+
from collections.abc import Mapping, Sequence
|
|
6
|
+
from dataclasses import dataclass
|
|
7
|
+
from typing import Any, Protocol, cast
|
|
8
|
+
|
|
9
|
+
from .contracts import Fingerprintable
|
|
10
|
+
from .domain import Candidate, MetricValue, Secret, UsageRecord
|
|
11
|
+
from .errors import SearchSpaceError, redact
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class Parameter(Protocol):
|
|
15
|
+
"""Sample one typed parameter value from a random stream."""
|
|
16
|
+
|
|
17
|
+
def sample(self, generator: random.Random) -> Any: ...
|
|
18
|
+
|
|
19
|
+
def accepts(self, value: Any) -> bool: ...
|
|
20
|
+
|
|
21
|
+
def describe(self) -> Mapping[str, Any]: ...
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
@dataclass(frozen=True, slots=True)
|
|
25
|
+
class IntegerRange:
|
|
26
|
+
"""Inclusive integer range."""
|
|
27
|
+
|
|
28
|
+
minimum: int
|
|
29
|
+
maximum: int
|
|
30
|
+
|
|
31
|
+
def __post_init__(self) -> None:
|
|
32
|
+
if any(
|
|
33
|
+
isinstance(value, bool) or not isinstance(value, int)
|
|
34
|
+
for value in (self.minimum, self.maximum)
|
|
35
|
+
):
|
|
36
|
+
raise TypeError("integer range bounds must be integers")
|
|
37
|
+
if self.minimum > self.maximum:
|
|
38
|
+
raise ValueError("integer range minimum must not exceed maximum")
|
|
39
|
+
|
|
40
|
+
def sample(self, generator: random.Random) -> int:
|
|
41
|
+
return generator.randint(self.minimum, self.maximum)
|
|
42
|
+
|
|
43
|
+
def accepts(self, value: Any) -> bool:
|
|
44
|
+
return (
|
|
45
|
+
isinstance(value, int)
|
|
46
|
+
and not isinstance(value, bool)
|
|
47
|
+
and self.minimum <= value <= self.maximum
|
|
48
|
+
)
|
|
49
|
+
|
|
50
|
+
def describe(self) -> Mapping[str, Any]:
|
|
51
|
+
return {"type": "integer", "minimum": self.minimum, "maximum": self.maximum}
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
@dataclass(frozen=True, slots=True)
|
|
55
|
+
class FloatRange:
|
|
56
|
+
"""Continuous float range."""
|
|
57
|
+
|
|
58
|
+
minimum: float
|
|
59
|
+
maximum: float
|
|
60
|
+
|
|
61
|
+
def __post_init__(self) -> None:
|
|
62
|
+
if any(
|
|
63
|
+
isinstance(value, bool)
|
|
64
|
+
or not isinstance(value, (int, float))
|
|
65
|
+
or not math.isfinite(value)
|
|
66
|
+
for value in (self.minimum, self.maximum)
|
|
67
|
+
):
|
|
68
|
+
raise ValueError("float range bounds must be finite numbers")
|
|
69
|
+
if self.minimum > self.maximum:
|
|
70
|
+
raise ValueError("float range minimum must not exceed maximum")
|
|
71
|
+
|
|
72
|
+
def sample(self, generator: random.Random) -> float:
|
|
73
|
+
return generator.uniform(self.minimum, self.maximum)
|
|
74
|
+
|
|
75
|
+
def accepts(self, value: Any) -> bool:
|
|
76
|
+
return (
|
|
77
|
+
isinstance(value, (int, float))
|
|
78
|
+
and not isinstance(value, bool)
|
|
79
|
+
and self.minimum <= value <= self.maximum
|
|
80
|
+
)
|
|
81
|
+
|
|
82
|
+
def describe(self) -> Mapping[str, Any]:
|
|
83
|
+
return {"type": "number", "minimum": self.minimum, "maximum": self.maximum}
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
@dataclass(frozen=True, slots=True)
|
|
87
|
+
class Categorical:
|
|
88
|
+
"""A non-empty collection of candidate values."""
|
|
89
|
+
|
|
90
|
+
values: tuple[Any, ...]
|
|
91
|
+
|
|
92
|
+
def __post_init__(self) -> None:
|
|
93
|
+
if not self.values:
|
|
94
|
+
raise ValueError("categorical values must not be empty")
|
|
95
|
+
|
|
96
|
+
def sample(self, generator: random.Random) -> Any:
|
|
97
|
+
return generator.choice(self.values)
|
|
98
|
+
|
|
99
|
+
def accepts(self, value: Any) -> bool:
|
|
100
|
+
return value in self.values
|
|
101
|
+
|
|
102
|
+
def describe(self) -> Mapping[str, Any]:
|
|
103
|
+
return {
|
|
104
|
+
"type": "categorical",
|
|
105
|
+
"values": tuple(value for value in self.values if not isinstance(value, Secret)),
|
|
106
|
+
}
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
@dataclass(frozen=True, slots=True)
|
|
110
|
+
class SearchSpace:
|
|
111
|
+
"""Named typed parameters from which candidates can be sampled."""
|
|
112
|
+
|
|
113
|
+
parameters: dict[str, Parameter]
|
|
114
|
+
|
|
115
|
+
def __post_init__(self) -> None:
|
|
116
|
+
if not self.parameters:
|
|
117
|
+
raise ValueError("search space must contain at least one parameter")
|
|
118
|
+
if any(not name.strip() for name in self.parameters):
|
|
119
|
+
raise ValueError("search-space parameter names must not be empty")
|
|
120
|
+
object.__setattr__(self, "parameters", dict(self.parameters))
|
|
121
|
+
|
|
122
|
+
def sample(self, generator: random.Random) -> Candidate:
|
|
123
|
+
return Candidate(
|
|
124
|
+
{name: parameter.sample(generator) for name, parameter in self.parameters.items()}
|
|
125
|
+
)
|
|
126
|
+
|
|
127
|
+
def validate(self, parameters: Mapping[str, Any]) -> Candidate:
|
|
128
|
+
"""Validate an exact candidate mapping against every parameter."""
|
|
129
|
+
|
|
130
|
+
if set(parameters) != set(self.parameters):
|
|
131
|
+
raise ValueError("candidate keys must exactly match the search space")
|
|
132
|
+
for name, parameter in self.parameters.items():
|
|
133
|
+
if not parameter.accepts(parameters[name]):
|
|
134
|
+
raise ValueError(f"candidate value is outside the search space: {name}")
|
|
135
|
+
return Candidate(dict(parameters))
|
|
136
|
+
|
|
137
|
+
def describe(self) -> Mapping[str, Mapping[str, Any]]:
|
|
138
|
+
"""Return a provider-safe structured search-space description."""
|
|
139
|
+
|
|
140
|
+
return {name: parameter.describe() for name, parameter in self.parameters.items()}
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
class RandomSearch:
|
|
144
|
+
"""Seeded random strategy implementing the common search contract."""
|
|
145
|
+
|
|
146
|
+
def __init__(self, space: SearchSpace, *, seed: int | None = None) -> None:
|
|
147
|
+
self._space = space
|
|
148
|
+
self._seed = seed
|
|
149
|
+
self._generator = random.Random(seed)
|
|
150
|
+
|
|
151
|
+
async def suggest(self) -> Candidate:
|
|
152
|
+
return self._space.sample(self._generator)
|
|
153
|
+
|
|
154
|
+
async def observe(self, candidate: Candidate, metrics: Any, usage: Any) -> None:
|
|
155
|
+
del candidate, metrics, usage
|
|
156
|
+
|
|
157
|
+
async def replay(self, candidate: Candidate, metrics: Any, usage: Any) -> None:
|
|
158
|
+
"""Advance and verify deterministic generator state during resume."""
|
|
159
|
+
|
|
160
|
+
del metrics, usage
|
|
161
|
+
expected = self._space.sample(self._generator)
|
|
162
|
+
if expected != candidate:
|
|
163
|
+
raise ValueError("persisted candidate does not match seeded random search state")
|
|
164
|
+
|
|
165
|
+
def fingerprint(self) -> Any:
|
|
166
|
+
"""Return stable search-space and seed semantics."""
|
|
167
|
+
|
|
168
|
+
return {"space": self._space, "seed": self._seed}
|
|
169
|
+
|
|
170
|
+
|
|
171
|
+
@dataclass(frozen=True, slots=True)
|
|
172
|
+
class SearchObservation:
|
|
173
|
+
"""A prior reserved candidate and any committed metrics."""
|
|
174
|
+
|
|
175
|
+
candidate: Candidate
|
|
176
|
+
metrics: tuple[MetricValue, ...]
|
|
177
|
+
|
|
178
|
+
|
|
179
|
+
@dataclass(frozen=True, slots=True)
|
|
180
|
+
class LLMSuggestionRequest:
|
|
181
|
+
"""Provider-neutral structured input for one suggestion attempt."""
|
|
182
|
+
|
|
183
|
+
search_space: Mapping[str, Mapping[str, Any]]
|
|
184
|
+
observations: tuple[SearchObservation, ...]
|
|
185
|
+
suggestion_number: int
|
|
186
|
+
attempt_number: int
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
class LLMSearchProvider(Protocol):
|
|
190
|
+
"""Produce one structured candidate proposal without provider coupling."""
|
|
191
|
+
|
|
192
|
+
async def suggest(self, request: LLMSuggestionRequest) -> Mapping[str, Any]: ...
|
|
193
|
+
|
|
194
|
+
|
|
195
|
+
class LLMSearch:
|
|
196
|
+
"""Validate bounded provider suggestions with deterministic random fallback."""
|
|
197
|
+
|
|
198
|
+
def __init__(
|
|
199
|
+
self,
|
|
200
|
+
space: SearchSpace,
|
|
201
|
+
provider: LLMSearchProvider,
|
|
202
|
+
*,
|
|
203
|
+
max_suggestion_attempts: int = 3,
|
|
204
|
+
fallback_seed: int = 0,
|
|
205
|
+
fallback: bool = True,
|
|
206
|
+
) -> None:
|
|
207
|
+
if (
|
|
208
|
+
not isinstance(max_suggestion_attempts, int)
|
|
209
|
+
or isinstance(max_suggestion_attempts, bool)
|
|
210
|
+
or max_suggestion_attempts < 1
|
|
211
|
+
):
|
|
212
|
+
raise ValueError("maximum suggestion attempts must be positive")
|
|
213
|
+
self._space = space
|
|
214
|
+
self._provider = provider
|
|
215
|
+
self._max_suggestion_attempts = max_suggestion_attempts
|
|
216
|
+
self._fallback_seed = fallback_seed
|
|
217
|
+
self._fallback = fallback
|
|
218
|
+
self._suggestion_number = 0
|
|
219
|
+
self._observations: list[SearchObservation] = []
|
|
220
|
+
|
|
221
|
+
async def suggest(self) -> Candidate:
|
|
222
|
+
causes: list[str] = []
|
|
223
|
+
for attempt_number in range(self._max_suggestion_attempts):
|
|
224
|
+
request = LLMSuggestionRequest(
|
|
225
|
+
self._space.describe(),
|
|
226
|
+
tuple(_provider_observation(item) for item in self._observations),
|
|
227
|
+
self._suggestion_number,
|
|
228
|
+
attempt_number,
|
|
229
|
+
)
|
|
230
|
+
try:
|
|
231
|
+
proposal = await self._provider.suggest(request)
|
|
232
|
+
candidate = self._space.validate(proposal)
|
|
233
|
+
except Exception as error:
|
|
234
|
+
causes.append(type(error).__name__)
|
|
235
|
+
continue
|
|
236
|
+
self._suggestion_number += 1
|
|
237
|
+
return candidate
|
|
238
|
+
if self._fallback:
|
|
239
|
+
generator = random.Random(f"{self._fallback_seed}:{self._suggestion_number}")
|
|
240
|
+
candidate = self._space.sample(generator)
|
|
241
|
+
self._suggestion_number += 1
|
|
242
|
+
return candidate
|
|
243
|
+
raise SearchSpaceError(
|
|
244
|
+
"LLM search exhausted invalid suggestion attempts",
|
|
245
|
+
stage="strategy.suggest",
|
|
246
|
+
details={"causes": causes, "attempts": self._max_suggestion_attempts},
|
|
247
|
+
)
|
|
248
|
+
|
|
249
|
+
async def observe(
|
|
250
|
+
self,
|
|
251
|
+
candidate: Candidate,
|
|
252
|
+
metrics: Sequence[MetricValue],
|
|
253
|
+
usage: Sequence[UsageRecord],
|
|
254
|
+
) -> None:
|
|
255
|
+
del usage
|
|
256
|
+
self._observations.append(SearchObservation(candidate, tuple(metrics)))
|
|
257
|
+
|
|
258
|
+
async def replay(
|
|
259
|
+
self,
|
|
260
|
+
candidate: Candidate,
|
|
261
|
+
metrics: Sequence[MetricValue],
|
|
262
|
+
usage: Sequence[UsageRecord],
|
|
263
|
+
) -> None:
|
|
264
|
+
del usage
|
|
265
|
+
self._suggestion_number += 1
|
|
266
|
+
self._observations.append(SearchObservation(candidate, tuple(metrics)))
|
|
267
|
+
|
|
268
|
+
def fingerprint(self) -> Any:
|
|
269
|
+
provider_type = type(self._provider)
|
|
270
|
+
provider: dict[str, Any] = {
|
|
271
|
+
"type": f"{provider_type.__module__}.{provider_type.__qualname__}"
|
|
272
|
+
}
|
|
273
|
+
if isinstance(self._provider, Fingerprintable):
|
|
274
|
+
provider["fingerprint"] = self._provider.fingerprint()
|
|
275
|
+
return {
|
|
276
|
+
"space": self._space,
|
|
277
|
+
"provider": provider,
|
|
278
|
+
"max_suggestion_attempts": self._max_suggestion_attempts,
|
|
279
|
+
"fallback_seed": self._fallback_seed,
|
|
280
|
+
"fallback": self._fallback,
|
|
281
|
+
}
|
|
282
|
+
|
|
283
|
+
|
|
284
|
+
def _provider_observation(observation: SearchObservation) -> SearchObservation:
|
|
285
|
+
parameters = cast(dict[str, Any], redact(observation.candidate.parameters))
|
|
286
|
+
return SearchObservation(Candidate(parameters), observation.metrics)
|
tunarag/serialization.py
ADDED
|
@@ -0,0 +1,78 @@
|
|
|
1
|
+
"""Canonical, secret-safe serialization for identity and cache keys."""
|
|
2
|
+
|
|
3
|
+
import hashlib
|
|
4
|
+
import json
|
|
5
|
+
import math
|
|
6
|
+
from collections.abc import Mapping, Sequence
|
|
7
|
+
from dataclasses import asdict, is_dataclass
|
|
8
|
+
from typing import Any, cast
|
|
9
|
+
|
|
10
|
+
from .domain import Secret
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def canonicalize(value: Any) -> Any:
|
|
14
|
+
"""Convert supported values to deterministic JSON-compatible structures."""
|
|
15
|
+
|
|
16
|
+
if isinstance(value, Secret):
|
|
17
|
+
return {"__secret__": True}
|
|
18
|
+
if is_dataclass(value):
|
|
19
|
+
return canonicalize(asdict(cast(Any, value)))
|
|
20
|
+
if isinstance(value, Mapping):
|
|
21
|
+
return {
|
|
22
|
+
str(key): canonicalize(item)
|
|
23
|
+
for key, item in sorted(value.items(), key=lambda entry: str(entry[0]))
|
|
24
|
+
if not isinstance(item, Secret)
|
|
25
|
+
}
|
|
26
|
+
if isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
|
|
27
|
+
return [canonicalize(item) for item in value]
|
|
28
|
+
if isinstance(value, float):
|
|
29
|
+
if not math.isfinite(value):
|
|
30
|
+
raise ValueError("nonfinite numbers are not supported")
|
|
31
|
+
return format(value, ".15g")
|
|
32
|
+
if value is None or isinstance(value, (bool, int, str)):
|
|
33
|
+
return value
|
|
34
|
+
raise TypeError(f"unsupported canonical value: {type(value).__name__}")
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def canonical_json(value: Any) -> bytes:
|
|
38
|
+
"""Serialize a value using stable key ordering and schema-independent JSON."""
|
|
39
|
+
|
|
40
|
+
return json.dumps(
|
|
41
|
+
canonicalize(value), ensure_ascii=True, separators=(",", ":"), sort_keys=True
|
|
42
|
+
).encode("utf-8")
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def storage_json(value: Any) -> bytes:
|
|
46
|
+
"""Serialize supported values reversibly while omitting mapping secrets."""
|
|
47
|
+
|
|
48
|
+
return json.dumps(
|
|
49
|
+
_storage_value(value), ensure_ascii=True, separators=(",", ":"), sort_keys=True
|
|
50
|
+
).encode("utf-8")
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def content_hash(value: Any, *, namespace: str = "autorag-tuner:v1") -> str:
|
|
54
|
+
"""Return a stable SHA-256 digest for a canonical value."""
|
|
55
|
+
|
|
56
|
+
return hashlib.sha256(namespace.encode() + b"\0" + canonical_json(value)).hexdigest()
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def _storage_value(value: Any) -> Any:
|
|
60
|
+
if isinstance(value, Secret):
|
|
61
|
+
return {"__secret__": True}
|
|
62
|
+
if is_dataclass(value):
|
|
63
|
+
return _storage_value(asdict(cast(Any, value)))
|
|
64
|
+
if isinstance(value, Mapping):
|
|
65
|
+
return {
|
|
66
|
+
str(key): _storage_value(item)
|
|
67
|
+
for key, item in sorted(value.items(), key=lambda entry: str(entry[0]))
|
|
68
|
+
if not isinstance(item, Secret)
|
|
69
|
+
}
|
|
70
|
+
if isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
|
|
71
|
+
return [_storage_value(item) for item in value]
|
|
72
|
+
if isinstance(value, float):
|
|
73
|
+
if not math.isfinite(value):
|
|
74
|
+
raise ValueError("nonfinite numbers are not supported")
|
|
75
|
+
return value
|
|
76
|
+
if value is None or isinstance(value, (bool, int, str)):
|
|
77
|
+
return value
|
|
78
|
+
raise TypeError(f"unsupported storage value: {type(value).__name__}")
|
tunarag/stopping.py
ADDED
|
@@ -0,0 +1,193 @@
|
|
|
1
|
+
"""Composable stop conditions for study execution."""
|
|
2
|
+
|
|
3
|
+
import math
|
|
4
|
+
from dataclasses import dataclass
|
|
5
|
+
from typing import Protocol
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
@dataclass(frozen=True, slots=True)
|
|
9
|
+
class StopContext:
|
|
10
|
+
"""Current study totals supplied to stop conditions."""
|
|
11
|
+
|
|
12
|
+
completed_trials: int
|
|
13
|
+
exact_cost: float = 0.0
|
|
14
|
+
estimated_cost: float = 0.0
|
|
15
|
+
elapsed_seconds: float = 0.0
|
|
16
|
+
successful_trials: int = 0
|
|
17
|
+
failed_trials: int = 0
|
|
18
|
+
objective_scores: tuple[float | None, ...] = ()
|
|
19
|
+
|
|
20
|
+
def __post_init__(self) -> None:
|
|
21
|
+
if (
|
|
22
|
+
isinstance(self.completed_trials, bool)
|
|
23
|
+
or not isinstance(self.completed_trials, int)
|
|
24
|
+
or self.completed_trials < 0
|
|
25
|
+
):
|
|
26
|
+
raise ValueError("completed trials must not be negative")
|
|
27
|
+
if any(
|
|
28
|
+
isinstance(value, bool) or not isinstance(value, int) or value < 0
|
|
29
|
+
for value in (self.successful_trials, self.failed_trials)
|
|
30
|
+
):
|
|
31
|
+
raise ValueError("trial status counts must not be negative")
|
|
32
|
+
if self.successful_trials + self.failed_trials > self.completed_trials:
|
|
33
|
+
raise ValueError("trial status counts exceed completed trials")
|
|
34
|
+
if self.objective_scores and len(self.objective_scores) != self.completed_trials:
|
|
35
|
+
raise ValueError("objective score history must match completed trials")
|
|
36
|
+
for name in ("exact_cost", "estimated_cost", "elapsed_seconds"):
|
|
37
|
+
value = getattr(self, name)
|
|
38
|
+
if not math.isfinite(value) or value < 0:
|
|
39
|
+
raise ValueError(f"{name} must be finite and nonnegative")
|
|
40
|
+
if any(score is not None and not math.isfinite(score) for score in self.objective_scores):
|
|
41
|
+
raise ValueError("objective scores must be finite or None")
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
class StopCondition(Protocol):
|
|
45
|
+
"""Return whether execution should stop for the current context."""
|
|
46
|
+
|
|
47
|
+
def should_stop(self, context: StopContext) -> bool: ...
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
@dataclass(frozen=True, slots=True)
|
|
51
|
+
class MaxTrials:
|
|
52
|
+
"""Stop after a maximum number of completed trials."""
|
|
53
|
+
|
|
54
|
+
limit: int
|
|
55
|
+
|
|
56
|
+
def __post_init__(self) -> None:
|
|
57
|
+
if isinstance(self.limit, bool) or not isinstance(self.limit, int) or self.limit < 1:
|
|
58
|
+
raise ValueError("trial limit must be at least one")
|
|
59
|
+
|
|
60
|
+
def should_stop(self, context: StopContext) -> bool:
|
|
61
|
+
return context.completed_trials >= self.limit
|
|
62
|
+
|
|
63
|
+
def fingerprint(self) -> object:
|
|
64
|
+
return {"limit": self.limit}
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
@dataclass(frozen=True, slots=True)
|
|
68
|
+
class CostBudget:
|
|
69
|
+
"""Stop when exact cost, or exact plus estimated cost, reaches a budget."""
|
|
70
|
+
|
|
71
|
+
limit: float
|
|
72
|
+
count_estimated: bool = True
|
|
73
|
+
|
|
74
|
+
def __post_init__(self) -> None:
|
|
75
|
+
if isinstance(self.limit, bool) or not math.isfinite(self.limit) or self.limit < 0:
|
|
76
|
+
raise ValueError("cost budget must be finite and nonnegative")
|
|
77
|
+
if not isinstance(self.count_estimated, bool):
|
|
78
|
+
raise TypeError("count_estimated must be a boolean")
|
|
79
|
+
|
|
80
|
+
def should_stop(self, context: StopContext) -> bool:
|
|
81
|
+
total = context.exact_cost + (context.estimated_cost if self.count_estimated else 0.0)
|
|
82
|
+
return total >= self.limit
|
|
83
|
+
|
|
84
|
+
def fingerprint(self) -> object:
|
|
85
|
+
return {"limit": self.limit, "count_estimated": self.count_estimated}
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
@dataclass(frozen=True, slots=True)
|
|
89
|
+
class TimeBudget:
|
|
90
|
+
"""Stop when elapsed wall-clock execution reaches a limit."""
|
|
91
|
+
|
|
92
|
+
seconds: float
|
|
93
|
+
|
|
94
|
+
def __post_init__(self) -> None:
|
|
95
|
+
if isinstance(self.seconds, bool) or not math.isfinite(self.seconds) or self.seconds < 0:
|
|
96
|
+
raise ValueError("time budget must be finite and nonnegative")
|
|
97
|
+
|
|
98
|
+
def should_stop(self, context: StopContext) -> bool:
|
|
99
|
+
return context.elapsed_seconds >= self.seconds
|
|
100
|
+
|
|
101
|
+
def fingerprint(self) -> object:
|
|
102
|
+
return {"seconds": self.seconds}
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
@dataclass(frozen=True, slots=True)
|
|
106
|
+
class TargetScore:
|
|
107
|
+
"""Stop when any successful objective reaches a benefit-oriented target."""
|
|
108
|
+
|
|
109
|
+
target: float
|
|
110
|
+
|
|
111
|
+
def __post_init__(self) -> None:
|
|
112
|
+
if isinstance(self.target, bool) or not math.isfinite(self.target):
|
|
113
|
+
raise ValueError("target score must be finite")
|
|
114
|
+
|
|
115
|
+
def should_stop(self, context: StopContext) -> bool:
|
|
116
|
+
return any(score is not None and score >= self.target for score in context.objective_scores)
|
|
117
|
+
|
|
118
|
+
def fingerprint(self) -> object:
|
|
119
|
+
return {"target": self.target}
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
@dataclass(frozen=True, slots=True)
|
|
123
|
+
class MaxFailures:
|
|
124
|
+
"""Stop after a total or consecutive number of failed trials."""
|
|
125
|
+
|
|
126
|
+
limit: int
|
|
127
|
+
consecutive: bool = False
|
|
128
|
+
|
|
129
|
+
def __post_init__(self) -> None:
|
|
130
|
+
if isinstance(self.limit, bool) or not isinstance(self.limit, int) or self.limit < 1:
|
|
131
|
+
raise ValueError("failure limit must be a positive integer")
|
|
132
|
+
|
|
133
|
+
def should_stop(self, context: StopContext) -> bool:
|
|
134
|
+
if not self.consecutive:
|
|
135
|
+
return context.failed_trials >= self.limit
|
|
136
|
+
failures = 0
|
|
137
|
+
for score in reversed(context.objective_scores):
|
|
138
|
+
if score is not None:
|
|
139
|
+
break
|
|
140
|
+
failures += 1
|
|
141
|
+
return failures >= self.limit
|
|
142
|
+
|
|
143
|
+
def fingerprint(self) -> object:
|
|
144
|
+
return {"limit": self.limit, "consecutive": self.consecutive}
|
|
145
|
+
|
|
146
|
+
|
|
147
|
+
@dataclass(frozen=True, slots=True)
|
|
148
|
+
class NoImprovement:
|
|
149
|
+
"""Stop after completed trials fail to improve the best objective."""
|
|
150
|
+
|
|
151
|
+
patience: int
|
|
152
|
+
min_delta: float = 0.0
|
|
153
|
+
warmup_trials: int = 1
|
|
154
|
+
|
|
155
|
+
def __post_init__(self) -> None:
|
|
156
|
+
if (
|
|
157
|
+
isinstance(self.patience, bool)
|
|
158
|
+
or not isinstance(self.patience, int)
|
|
159
|
+
or self.patience < 1
|
|
160
|
+
):
|
|
161
|
+
raise ValueError("no-improvement patience must be a positive integer")
|
|
162
|
+
if (
|
|
163
|
+
isinstance(self.min_delta, bool)
|
|
164
|
+
or not math.isfinite(self.min_delta)
|
|
165
|
+
or self.min_delta < 0
|
|
166
|
+
):
|
|
167
|
+
raise ValueError("minimum improvement must be finite and nonnegative")
|
|
168
|
+
if (
|
|
169
|
+
isinstance(self.warmup_trials, bool)
|
|
170
|
+
or not isinstance(self.warmup_trials, int)
|
|
171
|
+
or self.warmup_trials < 1
|
|
172
|
+
):
|
|
173
|
+
raise ValueError("warmup trials must be a positive integer")
|
|
174
|
+
|
|
175
|
+
def should_stop(self, context: StopContext) -> bool:
|
|
176
|
+
if context.completed_trials < self.warmup_trials or not context.objective_scores:
|
|
177
|
+
return False
|
|
178
|
+
best: float | None = None
|
|
179
|
+
since_improvement = 0
|
|
180
|
+
for score in context.objective_scores:
|
|
181
|
+
if score is not None and (best is None or score > best + self.min_delta):
|
|
182
|
+
best = score
|
|
183
|
+
since_improvement = 0
|
|
184
|
+
else:
|
|
185
|
+
since_improvement += 1
|
|
186
|
+
return best is not None and since_improvement >= self.patience
|
|
187
|
+
|
|
188
|
+
def fingerprint(self) -> object:
|
|
189
|
+
return {
|
|
190
|
+
"patience": self.patience,
|
|
191
|
+
"min_delta": self.min_delta,
|
|
192
|
+
"warmup_trials": self.warmup_trials,
|
|
193
|
+
}
|