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
|
@@ -0,0 +1,270 @@
|
|
|
1
|
+
"""Adapters for LangChain runnables and compiled LangGraph graphs."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import inspect
|
|
6
|
+
import time
|
|
7
|
+
from collections.abc import Callable, Mapping, Sequence
|
|
8
|
+
from typing import Any, Protocol
|
|
9
|
+
|
|
10
|
+
from ..dataset import EvaluationExample
|
|
11
|
+
from ..domain import Candidate, UsageRecord
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class AsyncRunnable(Protocol):
|
|
15
|
+
"""Shared asynchronous invocation surface implemented by both frameworks."""
|
|
16
|
+
|
|
17
|
+
async def ainvoke(
|
|
18
|
+
self,
|
|
19
|
+
input: Any,
|
|
20
|
+
config: Any = None,
|
|
21
|
+
**kwargs: Any,
|
|
22
|
+
) -> Any: ...
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
RunnableFactory = Callable[[Candidate], Any]
|
|
26
|
+
RunnableInputMapper = Callable[[Any, Candidate], Any]
|
|
27
|
+
RunnableConfigMapper = Callable[[Candidate], Mapping[str, Any] | None]
|
|
28
|
+
RunnableOutputMapper = Callable[[Any], Any]
|
|
29
|
+
RunnableUsageMapper = Callable[[Any], Sequence[UsageRecord]]
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class _RunnableAdapter:
|
|
33
|
+
def __init__(
|
|
34
|
+
self,
|
|
35
|
+
runnable: Any,
|
|
36
|
+
*,
|
|
37
|
+
factory: RunnableFactory | None,
|
|
38
|
+
input_key: str,
|
|
39
|
+
input_mapper: RunnableInputMapper | None,
|
|
40
|
+
config_mapper: RunnableConfigMapper | None,
|
|
41
|
+
output_mapper: RunnableOutputMapper | None,
|
|
42
|
+
usage_mapper: RunnableUsageMapper | None,
|
|
43
|
+
capture_message_usage: bool,
|
|
44
|
+
identity: Any,
|
|
45
|
+
framework: str,
|
|
46
|
+
) -> None:
|
|
47
|
+
if (runnable is None) == (factory is None):
|
|
48
|
+
raise ValueError("provide exactly one of runnable or factory")
|
|
49
|
+
if runnable is not None:
|
|
50
|
+
_require_runnable(runnable)
|
|
51
|
+
if factory is not None and not callable(factory):
|
|
52
|
+
raise TypeError("runnable factory must be callable")
|
|
53
|
+
if not input_key.strip():
|
|
54
|
+
raise ValueError("runnable input key must not be empty")
|
|
55
|
+
for name, mapper in (
|
|
56
|
+
("input_mapper", input_mapper),
|
|
57
|
+
("config_mapper", config_mapper),
|
|
58
|
+
("output_mapper", output_mapper),
|
|
59
|
+
("usage_mapper", usage_mapper),
|
|
60
|
+
):
|
|
61
|
+
if mapper is not None and not callable(mapper):
|
|
62
|
+
raise TypeError(f"{name} must be callable")
|
|
63
|
+
if not isinstance(capture_message_usage, bool):
|
|
64
|
+
raise TypeError("capture_message_usage must be a boolean")
|
|
65
|
+
self._runnable = runnable
|
|
66
|
+
self._factory = factory
|
|
67
|
+
self._input_key = input_key
|
|
68
|
+
self._input_mapper = input_mapper
|
|
69
|
+
self._config_mapper = config_mapper
|
|
70
|
+
self._output_mapper = output_mapper
|
|
71
|
+
self._usage_mapper = usage_mapper
|
|
72
|
+
self._capture_message_usage = capture_message_usage
|
|
73
|
+
self._identity = identity
|
|
74
|
+
self._framework = framework
|
|
75
|
+
|
|
76
|
+
async def run(self, candidate: Candidate, example: Any) -> tuple[Any, tuple[UsageRecord, ...]]:
|
|
77
|
+
"""Invoke one framework runnable for a TunaRAG candidate and example."""
|
|
78
|
+
|
|
79
|
+
runnable = self._resolve_runnable(candidate)
|
|
80
|
+
runnable_input = self._map_input(example, candidate)
|
|
81
|
+
config = self._map_config(candidate)
|
|
82
|
+
started_at = time.perf_counter()
|
|
83
|
+
raw_output = await runnable.ainvoke(runnable_input, config=config)
|
|
84
|
+
latency = time.perf_counter() - started_at
|
|
85
|
+
output = raw_output if self._output_mapper is None else self._output_mapper(raw_output)
|
|
86
|
+
usage = list(self._map_usage(raw_output))
|
|
87
|
+
usage.append(
|
|
88
|
+
UsageRecord(
|
|
89
|
+
component=f"{self._framework}.runnable",
|
|
90
|
+
latency_seconds=latency,
|
|
91
|
+
)
|
|
92
|
+
)
|
|
93
|
+
return output, tuple(usage)
|
|
94
|
+
|
|
95
|
+
def fingerprint(self) -> dict[str, Any]:
|
|
96
|
+
"""Return stable, provider-secret-free adapter semantics."""
|
|
97
|
+
|
|
98
|
+
runnable_identity = _type_name(self._runnable) if self._runnable is not None else None
|
|
99
|
+
return {
|
|
100
|
+
"integration": self._framework,
|
|
101
|
+
"schema_version": 1,
|
|
102
|
+
"runnable": runnable_identity,
|
|
103
|
+
"factory": _callable_identity(self._factory),
|
|
104
|
+
"input_key": self._input_key,
|
|
105
|
+
"input_mapper": _callable_identity(self._input_mapper),
|
|
106
|
+
"config_mapper": _callable_identity(self._config_mapper),
|
|
107
|
+
"output_mapper": _callable_identity(self._output_mapper),
|
|
108
|
+
"usage_mapper": _callable_identity(self._usage_mapper),
|
|
109
|
+
"capture_message_usage": self._capture_message_usage,
|
|
110
|
+
"identity": self._identity,
|
|
111
|
+
}
|
|
112
|
+
|
|
113
|
+
def _resolve_runnable(self, candidate: Candidate) -> Any:
|
|
114
|
+
if self._runnable is not None:
|
|
115
|
+
return self._runnable
|
|
116
|
+
if self._factory is None:
|
|
117
|
+
raise RuntimeError("runnable adapter is not configured")
|
|
118
|
+
runnable = self._factory(candidate)
|
|
119
|
+
_require_runnable(runnable)
|
|
120
|
+
return runnable
|
|
121
|
+
|
|
122
|
+
def _map_input(self, example: Any, candidate: Candidate) -> Any:
|
|
123
|
+
if self._input_mapper is not None:
|
|
124
|
+
return self._input_mapper(example, candidate)
|
|
125
|
+
if isinstance(example, EvaluationExample):
|
|
126
|
+
return {self._input_key: example.query}
|
|
127
|
+
return example
|
|
128
|
+
|
|
129
|
+
def _map_config(self, candidate: Candidate) -> Mapping[str, Any] | None:
|
|
130
|
+
if self._config_mapper is None:
|
|
131
|
+
return {"configurable": dict(candidate.parameters)}
|
|
132
|
+
config = self._config_mapper(candidate)
|
|
133
|
+
if config is not None and not isinstance(config, Mapping):
|
|
134
|
+
raise TypeError("config_mapper must return a mapping or None")
|
|
135
|
+
return config
|
|
136
|
+
|
|
137
|
+
def _map_usage(self, raw_output: Any) -> tuple[UsageRecord, ...]:
|
|
138
|
+
if self._usage_mapper is not None:
|
|
139
|
+
usage = tuple(self._usage_mapper(raw_output))
|
|
140
|
+
if any(not isinstance(item, UsageRecord) for item in usage):
|
|
141
|
+
raise TypeError("usage_mapper must return only UsageRecord values")
|
|
142
|
+
return usage
|
|
143
|
+
if not self._capture_message_usage:
|
|
144
|
+
return ()
|
|
145
|
+
return _message_usage(raw_output, framework=self._framework)
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
class LangChainAdapter(_RunnableAdapter):
|
|
149
|
+
"""Expose a LangChain runnable through TunaRAG's standard adapter contract."""
|
|
150
|
+
|
|
151
|
+
def __init__(
|
|
152
|
+
self,
|
|
153
|
+
runnable: Any = None,
|
|
154
|
+
*,
|
|
155
|
+
factory: RunnableFactory | None = None,
|
|
156
|
+
input_key: str = "question",
|
|
157
|
+
input_mapper: RunnableInputMapper | None = None,
|
|
158
|
+
config_mapper: RunnableConfigMapper | None = None,
|
|
159
|
+
output_mapper: RunnableOutputMapper | None = None,
|
|
160
|
+
usage_mapper: RunnableUsageMapper | None = None,
|
|
161
|
+
capture_message_usage: bool = True,
|
|
162
|
+
identity: Any = None,
|
|
163
|
+
) -> None:
|
|
164
|
+
super().__init__(
|
|
165
|
+
runnable,
|
|
166
|
+
factory=factory,
|
|
167
|
+
input_key=input_key,
|
|
168
|
+
input_mapper=input_mapper,
|
|
169
|
+
config_mapper=config_mapper,
|
|
170
|
+
output_mapper=output_mapper,
|
|
171
|
+
usage_mapper=usage_mapper,
|
|
172
|
+
capture_message_usage=capture_message_usage,
|
|
173
|
+
identity=identity,
|
|
174
|
+
framework="langchain",
|
|
175
|
+
)
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
class LangGraphAdapter(_RunnableAdapter):
|
|
179
|
+
"""Expose a compiled LangGraph graph through TunaRAG's adapter contract."""
|
|
180
|
+
|
|
181
|
+
def __init__(
|
|
182
|
+
self,
|
|
183
|
+
graph: Any = None,
|
|
184
|
+
*,
|
|
185
|
+
factory: RunnableFactory | None = None,
|
|
186
|
+
input_key: str = "question",
|
|
187
|
+
input_mapper: RunnableInputMapper | None = None,
|
|
188
|
+
config_mapper: RunnableConfigMapper | None = None,
|
|
189
|
+
output_mapper: RunnableOutputMapper | None = None,
|
|
190
|
+
usage_mapper: RunnableUsageMapper | None = None,
|
|
191
|
+
capture_message_usage: bool = True,
|
|
192
|
+
identity: Any = None,
|
|
193
|
+
) -> None:
|
|
194
|
+
super().__init__(
|
|
195
|
+
graph,
|
|
196
|
+
factory=factory,
|
|
197
|
+
input_key=input_key,
|
|
198
|
+
input_mapper=input_mapper,
|
|
199
|
+
config_mapper=config_mapper,
|
|
200
|
+
output_mapper=output_mapper,
|
|
201
|
+
usage_mapper=usage_mapper,
|
|
202
|
+
capture_message_usage=capture_message_usage,
|
|
203
|
+
identity=identity,
|
|
204
|
+
framework="langgraph",
|
|
205
|
+
)
|
|
206
|
+
|
|
207
|
+
|
|
208
|
+
def _message_usage(output: Any, *, framework: str) -> tuple[UsageRecord, ...]:
|
|
209
|
+
records: list[UsageRecord] = []
|
|
210
|
+
visited: set[int] = set()
|
|
211
|
+
|
|
212
|
+
def collect(value: Any) -> None:
|
|
213
|
+
if value is None or isinstance(value, (str, bytes, bytearray, bool, int, float)):
|
|
214
|
+
return
|
|
215
|
+
value_id = id(value)
|
|
216
|
+
if value_id in visited:
|
|
217
|
+
return
|
|
218
|
+
visited.add(value_id)
|
|
219
|
+
metadata = (
|
|
220
|
+
value.get("usage_metadata")
|
|
221
|
+
if isinstance(value, Mapping)
|
|
222
|
+
else getattr(value, "usage_metadata", None)
|
|
223
|
+
)
|
|
224
|
+
if isinstance(metadata, Mapping):
|
|
225
|
+
input_tokens = _token_count(metadata.get("input_tokens"))
|
|
226
|
+
output_tokens = _token_count(metadata.get("output_tokens"))
|
|
227
|
+
if input_tokens is not None or output_tokens is not None:
|
|
228
|
+
records.append(
|
|
229
|
+
UsageRecord(
|
|
230
|
+
component=f"{framework}.message",
|
|
231
|
+
input_tokens=input_tokens,
|
|
232
|
+
output_tokens=output_tokens,
|
|
233
|
+
)
|
|
234
|
+
)
|
|
235
|
+
if isinstance(value, Mapping):
|
|
236
|
+
for item in value.values():
|
|
237
|
+
collect(item)
|
|
238
|
+
elif isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
|
|
239
|
+
for item in value:
|
|
240
|
+
collect(item)
|
|
241
|
+
|
|
242
|
+
collect(output)
|
|
243
|
+
return tuple(records)
|
|
244
|
+
|
|
245
|
+
|
|
246
|
+
def _token_count(value: Any) -> int | None:
|
|
247
|
+
if value is None:
|
|
248
|
+
return None
|
|
249
|
+
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
|
|
250
|
+
return None
|
|
251
|
+
return value
|
|
252
|
+
|
|
253
|
+
|
|
254
|
+
def _require_runnable(value: Any) -> None:
|
|
255
|
+
method = getattr(value, "ainvoke", None)
|
|
256
|
+
if method is None or not callable(method):
|
|
257
|
+
raise TypeError("framework object must provide an async ainvoke method")
|
|
258
|
+
|
|
259
|
+
|
|
260
|
+
def _callable_identity(value: Callable[..., Any] | None) -> str | None:
|
|
261
|
+
if value is None:
|
|
262
|
+
return None
|
|
263
|
+
if inspect.ismethod(value) or inspect.isfunction(value):
|
|
264
|
+
return f"{value.__module__}.{value.__qualname__}"
|
|
265
|
+
return _type_name(value)
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
def _type_name(value: Any) -> str:
|
|
269
|
+
value_type = type(value)
|
|
270
|
+
return f"{value_type.__module__}.{value_type.__qualname__}"
|
tunarag/objective.py
ADDED
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
"""Objective aggregation for quality, cost, and latency metrics."""
|
|
2
|
+
|
|
3
|
+
from collections.abc import Sequence
|
|
4
|
+
from dataclasses import dataclass
|
|
5
|
+
from enum import Enum
|
|
6
|
+
|
|
7
|
+
from .domain import MetricValue, ValueStatus
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class ObjectiveDirection(str, Enum):
|
|
11
|
+
"""Whether larger or smaller metric values are beneficial."""
|
|
12
|
+
|
|
13
|
+
MAXIMIZE = "maximize"
|
|
14
|
+
MINIMIZE = "minimize"
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
@dataclass(frozen=True, slots=True)
|
|
18
|
+
class ObjectiveTerm:
|
|
19
|
+
"""One weighted metric contribution to an objective score."""
|
|
20
|
+
|
|
21
|
+
metric: str
|
|
22
|
+
weight: float = 1.0
|
|
23
|
+
direction: ObjectiveDirection = ObjectiveDirection.MAXIMIZE
|
|
24
|
+
minimum: float | None = None
|
|
25
|
+
maximum: float | None = None
|
|
26
|
+
|
|
27
|
+
def __post_init__(self) -> None:
|
|
28
|
+
if not self.metric.strip():
|
|
29
|
+
raise ValueError("objective metric must not be empty")
|
|
30
|
+
if self.weight < 0:
|
|
31
|
+
raise ValueError("objective weight must not be negative")
|
|
32
|
+
if self.minimum is not None and self.maximum is not None and self.minimum >= self.maximum:
|
|
33
|
+
raise ValueError("objective minimum must be less than maximum")
|
|
34
|
+
if (self.minimum is None) != (self.maximum is None):
|
|
35
|
+
raise ValueError("objective minimum and maximum must be provided together")
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
@dataclass(frozen=True, slots=True)
|
|
39
|
+
class ObjectiveResult:
|
|
40
|
+
"""Aggregated score and confidence state."""
|
|
41
|
+
|
|
42
|
+
score: float | None
|
|
43
|
+
status: ValueStatus
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
@dataclass(frozen=True, slots=True)
|
|
47
|
+
class Objective:
|
|
48
|
+
"""Combine metric observations into a benefit-oriented weighted score."""
|
|
49
|
+
|
|
50
|
+
terms: tuple[ObjectiveTerm, ...]
|
|
51
|
+
|
|
52
|
+
def __post_init__(self) -> None:
|
|
53
|
+
if not self.terms:
|
|
54
|
+
raise ValueError("objective must contain at least one term")
|
|
55
|
+
if sum(term.weight for term in self.terms) == 0:
|
|
56
|
+
raise ValueError("objective must contain a positive weight")
|
|
57
|
+
|
|
58
|
+
def evaluate(self, metrics: Sequence[MetricValue]) -> ObjectiveResult:
|
|
59
|
+
by_name = {metric.name: metric for metric in metrics}
|
|
60
|
+
score = 0.0
|
|
61
|
+
status = ValueStatus.EXACT
|
|
62
|
+
for term in self.terms:
|
|
63
|
+
metric = by_name.get(term.metric)
|
|
64
|
+
if metric is None or metric.value is None:
|
|
65
|
+
return ObjectiveResult(None, ValueStatus.UNAVAILABLE)
|
|
66
|
+
value = float(metric.value)
|
|
67
|
+
if term.minimum is not None and term.maximum is not None:
|
|
68
|
+
value = (value - term.minimum) / (term.maximum - term.minimum)
|
|
69
|
+
value = max(0.0, min(1.0, value))
|
|
70
|
+
if term.direction is ObjectiveDirection.MINIMIZE:
|
|
71
|
+
value = 1.0 - value if term.minimum is not None else -value
|
|
72
|
+
score += term.weight * value
|
|
73
|
+
if metric.status is ValueStatus.ESTIMATED:
|
|
74
|
+
status = ValueStatus.ESTIMATED
|
|
75
|
+
return ObjectiveResult(score, status)
|
tunarag/py.typed
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
|
tunarag/result.py
ADDED
|
@@ -0,0 +1,309 @@
|
|
|
1
|
+
"""Inspectable optimization result models."""
|
|
2
|
+
|
|
3
|
+
import csv
|
|
4
|
+
import json
|
|
5
|
+
from collections.abc import Mapping
|
|
6
|
+
from dataclasses import dataclass
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
from typing import Any, cast
|
|
9
|
+
|
|
10
|
+
from .domain import Candidate, MetricValue, UsageRecord, ValueStatus
|
|
11
|
+
from .errors import redact
|
|
12
|
+
from .objective import ObjectiveDirection, ObjectiveResult
|
|
13
|
+
from .store import StudyStatus, TrialStatus
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
@dataclass(frozen=True, slots=True)
|
|
17
|
+
class UsageSummary:
|
|
18
|
+
"""Aggregate usage with estimate provenance retained."""
|
|
19
|
+
|
|
20
|
+
input_tokens: int
|
|
21
|
+
output_tokens: int
|
|
22
|
+
exact_cost: float
|
|
23
|
+
estimated_cost: float
|
|
24
|
+
latency_seconds: float
|
|
25
|
+
status: ValueStatus
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
@dataclass(frozen=True, slots=True)
|
|
29
|
+
class TrialSummary:
|
|
30
|
+
"""Actionable terminal summary for one candidate trial."""
|
|
31
|
+
|
|
32
|
+
id: str
|
|
33
|
+
sequence: int
|
|
34
|
+
candidate: Candidate
|
|
35
|
+
status: TrialStatus
|
|
36
|
+
metrics: tuple[MetricValue, ...]
|
|
37
|
+
usage: tuple[UsageRecord, ...]
|
|
38
|
+
objective: ObjectiveResult | None
|
|
39
|
+
attempts: int
|
|
40
|
+
cached: bool = False
|
|
41
|
+
error: str | None = None
|
|
42
|
+
|
|
43
|
+
def metric(self, name: str) -> MetricValue | None:
|
|
44
|
+
"""Return one named metric, or None when it was not observed."""
|
|
45
|
+
|
|
46
|
+
return next((metric for metric in self.metrics if metric.name == name), None)
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
@dataclass(frozen=True, slots=True)
|
|
50
|
+
class OptimizationResult:
|
|
51
|
+
"""Completed study result, including failures and no-success outcomes."""
|
|
52
|
+
|
|
53
|
+
study_id: str
|
|
54
|
+
study_status: StudyStatus
|
|
55
|
+
stop_reason: str
|
|
56
|
+
trials: tuple[TrialSummary, ...]
|
|
57
|
+
best_trial: TrialSummary | None
|
|
58
|
+
usage: UsageSummary
|
|
59
|
+
|
|
60
|
+
@property
|
|
61
|
+
def best_candidate(self) -> Candidate | None:
|
|
62
|
+
"""Return the highest-scoring candidate when one succeeded."""
|
|
63
|
+
|
|
64
|
+
return None if self.best_trial is None else self.best_trial.candidate
|
|
65
|
+
|
|
66
|
+
@property
|
|
67
|
+
def successful_trials(self) -> tuple[TrialSummary, ...]:
|
|
68
|
+
"""Return successful trials in execution order."""
|
|
69
|
+
|
|
70
|
+
return tuple(trial for trial in self.trials if trial.status is TrialStatus.SUCCEEDED)
|
|
71
|
+
|
|
72
|
+
@property
|
|
73
|
+
def failed_trials(self) -> tuple[TrialSummary, ...]:
|
|
74
|
+
"""Return failed trials in execution order."""
|
|
75
|
+
|
|
76
|
+
return tuple(trial for trial in self.trials if trial.status is TrialStatus.FAILED)
|
|
77
|
+
|
|
78
|
+
@property
|
|
79
|
+
def ranked_trials(self) -> tuple[TrialSummary, ...]:
|
|
80
|
+
"""Return successful trials by descending objective, then sequence."""
|
|
81
|
+
|
|
82
|
+
return tuple(
|
|
83
|
+
sorted(
|
|
84
|
+
(
|
|
85
|
+
trial
|
|
86
|
+
for trial in self.successful_trials
|
|
87
|
+
if trial.objective is not None and trial.objective.score is not None
|
|
88
|
+
),
|
|
89
|
+
key=lambda trial: (-_objective_score(trial), trial.sequence),
|
|
90
|
+
)
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
@property
|
|
94
|
+
def metric_names(self) -> tuple[str, ...]:
|
|
95
|
+
"""Return all observed metric names in lexical order."""
|
|
96
|
+
|
|
97
|
+
return tuple(sorted({metric.name for trial in self.trials for metric in trial.metrics}))
|
|
98
|
+
|
|
99
|
+
def metric_leaders(
|
|
100
|
+
self,
|
|
101
|
+
name: str,
|
|
102
|
+
*,
|
|
103
|
+
direction: ObjectiveDirection = ObjectiveDirection.MAXIMIZE,
|
|
104
|
+
) -> tuple[TrialSummary, ...]:
|
|
105
|
+
"""Rank successful trials that have an available named metric."""
|
|
106
|
+
|
|
107
|
+
if not name.strip():
|
|
108
|
+
raise ValueError("metric name must not be empty")
|
|
109
|
+
trials = [
|
|
110
|
+
trial
|
|
111
|
+
for trial in self.successful_trials
|
|
112
|
+
if (metric := trial.metric(name)) is not None and metric.value is not None
|
|
113
|
+
]
|
|
114
|
+
multiplier = -1.0 if direction is ObjectiveDirection.MAXIMIZE else 1.0
|
|
115
|
+
return tuple(
|
|
116
|
+
sorted(
|
|
117
|
+
trials,
|
|
118
|
+
key=lambda trial: (
|
|
119
|
+
multiplier * cast(float, cast(MetricValue, trial.metric(name)).value),
|
|
120
|
+
trial.sequence,
|
|
121
|
+
),
|
|
122
|
+
)
|
|
123
|
+
)
|
|
124
|
+
|
|
125
|
+
def as_dict(self) -> dict[str, Any]:
|
|
126
|
+
"""Return a redacted, JSON-compatible report."""
|
|
127
|
+
|
|
128
|
+
return {
|
|
129
|
+
"schema_version": 1,
|
|
130
|
+
"study_id": self.study_id,
|
|
131
|
+
"study_status": self.study_status.value,
|
|
132
|
+
"stop_reason": self.stop_reason,
|
|
133
|
+
"best_trial_id": None if self.best_trial is None else self.best_trial.id,
|
|
134
|
+
"counts": {
|
|
135
|
+
"trials": len(self.trials),
|
|
136
|
+
"successful": len(self.successful_trials),
|
|
137
|
+
"failed": len(self.failed_trials),
|
|
138
|
+
},
|
|
139
|
+
"usage": _usage_dict(self.usage),
|
|
140
|
+
"trials": [_trial_dict(trial) for trial in self.trials],
|
|
141
|
+
}
|
|
142
|
+
|
|
143
|
+
def to_json(self, *, indent: int | None = 2) -> str:
|
|
144
|
+
"""Serialize the redacted report to deterministic JSON."""
|
|
145
|
+
|
|
146
|
+
return json.dumps(
|
|
147
|
+
self.as_dict(),
|
|
148
|
+
ensure_ascii=False,
|
|
149
|
+
allow_nan=False,
|
|
150
|
+
indent=indent,
|
|
151
|
+
sort_keys=True,
|
|
152
|
+
)
|
|
153
|
+
|
|
154
|
+
def write_json(self, path: str | Path, *, indent: int | None = 2) -> Path:
|
|
155
|
+
"""Write a redacted JSON report and return its path."""
|
|
156
|
+
|
|
157
|
+
target = Path(path)
|
|
158
|
+
target.write_text(self.to_json(indent=indent) + "\n", encoding="utf-8")
|
|
159
|
+
return target
|
|
160
|
+
|
|
161
|
+
def write_csv(self, path: str | Path) -> Path:
|
|
162
|
+
"""Write one redacted, flattened row per trial and return its path."""
|
|
163
|
+
|
|
164
|
+
target = Path(path)
|
|
165
|
+
metric_columns = [
|
|
166
|
+
column
|
|
167
|
+
for name in self.metric_names
|
|
168
|
+
for column in (
|
|
169
|
+
f"metric.{name}.value",
|
|
170
|
+
f"metric.{name}.status",
|
|
171
|
+
f"metric.{name}.coverage",
|
|
172
|
+
)
|
|
173
|
+
]
|
|
174
|
+
fieldnames = [
|
|
175
|
+
"study_id",
|
|
176
|
+
"trial_id",
|
|
177
|
+
"sequence",
|
|
178
|
+
"status",
|
|
179
|
+
"objective_score",
|
|
180
|
+
"objective_status",
|
|
181
|
+
"attempts",
|
|
182
|
+
"cached",
|
|
183
|
+
"error",
|
|
184
|
+
"candidate",
|
|
185
|
+
"input_tokens",
|
|
186
|
+
"output_tokens",
|
|
187
|
+
"exact_cost",
|
|
188
|
+
"estimated_cost",
|
|
189
|
+
"latency_seconds",
|
|
190
|
+
"usage_status",
|
|
191
|
+
*metric_columns,
|
|
192
|
+
]
|
|
193
|
+
with target.open("w", encoding="utf-8", newline="") as stream:
|
|
194
|
+
writer = csv.DictWriter(stream, fieldnames=fieldnames, extrasaction="raise")
|
|
195
|
+
writer.writeheader()
|
|
196
|
+
for trial in self.trials:
|
|
197
|
+
writer.writerow(_trial_csv_row(self.study_id, trial, self.metric_names))
|
|
198
|
+
return target
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
def _usage_dict(usage: UsageSummary) -> dict[str, Any]:
|
|
202
|
+
return {
|
|
203
|
+
"input_tokens": usage.input_tokens,
|
|
204
|
+
"output_tokens": usage.output_tokens,
|
|
205
|
+
"exact_cost": usage.exact_cost,
|
|
206
|
+
"estimated_cost": usage.estimated_cost,
|
|
207
|
+
"latency_seconds": usage.latency_seconds,
|
|
208
|
+
"status": usage.status.value,
|
|
209
|
+
}
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
def _trial_dict(trial: TrialSummary) -> dict[str, Any]:
|
|
213
|
+
candidate = cast(Mapping[str, Any], redact(trial.candidate.parameters))
|
|
214
|
+
return {
|
|
215
|
+
"id": trial.id,
|
|
216
|
+
"sequence": trial.sequence,
|
|
217
|
+
"status": trial.status.value,
|
|
218
|
+
"candidate": dict(candidate),
|
|
219
|
+
"metrics": [
|
|
220
|
+
{
|
|
221
|
+
"name": metric.name,
|
|
222
|
+
"value": metric.value,
|
|
223
|
+
"status": metric.status.value,
|
|
224
|
+
"coverage": metric.coverage,
|
|
225
|
+
}
|
|
226
|
+
for metric in trial.metrics
|
|
227
|
+
],
|
|
228
|
+
"usage": [
|
|
229
|
+
{
|
|
230
|
+
"component": item.component,
|
|
231
|
+
"input_tokens": item.input_tokens,
|
|
232
|
+
"output_tokens": item.output_tokens,
|
|
233
|
+
"cost": item.cost,
|
|
234
|
+
"latency_seconds": item.latency_seconds,
|
|
235
|
+
"status": item.status.value,
|
|
236
|
+
"pricing_version": item.pricing_version,
|
|
237
|
+
}
|
|
238
|
+
for item in trial.usage
|
|
239
|
+
],
|
|
240
|
+
"objective": (
|
|
241
|
+
None
|
|
242
|
+
if trial.objective is None
|
|
243
|
+
else {"score": trial.objective.score, "status": trial.objective.status.value}
|
|
244
|
+
),
|
|
245
|
+
"attempts": trial.attempts,
|
|
246
|
+
"cached": trial.cached,
|
|
247
|
+
"error": trial.error,
|
|
248
|
+
}
|
|
249
|
+
|
|
250
|
+
|
|
251
|
+
def _trial_csv_row(
|
|
252
|
+
study_id: str, trial: TrialSummary, metric_names: tuple[str, ...]
|
|
253
|
+
) -> dict[str, Any]:
|
|
254
|
+
usage = _summarize_usage(trial.usage)
|
|
255
|
+
candidate = redact(trial.candidate.parameters)
|
|
256
|
+
row: dict[str, Any] = {
|
|
257
|
+
"study_id": study_id,
|
|
258
|
+
"trial_id": trial.id,
|
|
259
|
+
"sequence": trial.sequence,
|
|
260
|
+
"status": trial.status.value,
|
|
261
|
+
"objective_score": None if trial.objective is None else trial.objective.score,
|
|
262
|
+
"objective_status": None if trial.objective is None else trial.objective.status.value,
|
|
263
|
+
"attempts": trial.attempts,
|
|
264
|
+
"cached": trial.cached,
|
|
265
|
+
"error": trial.error,
|
|
266
|
+
"candidate": json.dumps(
|
|
267
|
+
candidate, ensure_ascii=False, sort_keys=True, separators=(",", ":")
|
|
268
|
+
),
|
|
269
|
+
"input_tokens": usage.input_tokens,
|
|
270
|
+
"output_tokens": usage.output_tokens,
|
|
271
|
+
"exact_cost": usage.exact_cost,
|
|
272
|
+
"estimated_cost": usage.estimated_cost,
|
|
273
|
+
"latency_seconds": usage.latency_seconds,
|
|
274
|
+
"usage_status": usage.status.value,
|
|
275
|
+
}
|
|
276
|
+
for name in metric_names:
|
|
277
|
+
metric = trial.metric(name)
|
|
278
|
+
row[f"metric.{name}.value"] = None if metric is None else metric.value
|
|
279
|
+
row[f"metric.{name}.status"] = None if metric is None else metric.status.value
|
|
280
|
+
row[f"metric.{name}.coverage"] = None if metric is None else metric.coverage
|
|
281
|
+
return row
|
|
282
|
+
|
|
283
|
+
|
|
284
|
+
def _objective_score(trial: TrialSummary) -> float:
|
|
285
|
+
if trial.objective is None or trial.objective.score is None:
|
|
286
|
+
raise ValueError("ranked trial has no objective score")
|
|
287
|
+
return trial.objective.score
|
|
288
|
+
|
|
289
|
+
|
|
290
|
+
def _summarize_usage(records: tuple[UsageRecord, ...]) -> UsageSummary:
|
|
291
|
+
statuses = {record.status for record in records}
|
|
292
|
+
if not records or statuses == {ValueStatus.UNAVAILABLE}:
|
|
293
|
+
status = ValueStatus.UNAVAILABLE
|
|
294
|
+
elif statuses == {ValueStatus.EXACT}:
|
|
295
|
+
status = ValueStatus.EXACT
|
|
296
|
+
else:
|
|
297
|
+
status = ValueStatus.ESTIMATED
|
|
298
|
+
return UsageSummary(
|
|
299
|
+
input_tokens=sum(record.input_tokens or 0 for record in records),
|
|
300
|
+
output_tokens=sum(record.output_tokens or 0 for record in records),
|
|
301
|
+
exact_cost=sum(
|
|
302
|
+
record.cost or 0.0 for record in records if record.status is ValueStatus.EXACT
|
|
303
|
+
),
|
|
304
|
+
estimated_cost=sum(
|
|
305
|
+
record.cost or 0.0 for record in records if record.status is not ValueStatus.EXACT
|
|
306
|
+
),
|
|
307
|
+
latency_seconds=sum(record.latency_seconds or 0.0 for record in records),
|
|
308
|
+
status=status,
|
|
309
|
+
)
|