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/dataset.py
ADDED
|
@@ -0,0 +1,481 @@
|
|
|
1
|
+
"""Repeatable datasets and normalized evaluation-example ingestion."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import csv
|
|
6
|
+
import json
|
|
7
|
+
import math
|
|
8
|
+
import unicodedata
|
|
9
|
+
from collections.abc import AsyncIterator, Iterable, Mapping, Sequence
|
|
10
|
+
from dataclasses import dataclass, field, replace
|
|
11
|
+
from enum import Enum
|
|
12
|
+
from pathlib import Path
|
|
13
|
+
from types import MappingProxyType
|
|
14
|
+
from typing import Any
|
|
15
|
+
|
|
16
|
+
from .errors import DatasetError
|
|
17
|
+
from .serialization import content_hash
|
|
18
|
+
|
|
19
|
+
_SCHEMA_VERSION = 1
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class DatasetSplit(str, Enum):
|
|
23
|
+
"""Canonical purpose assigned to an evaluation example."""
|
|
24
|
+
|
|
25
|
+
OPTIMIZE = "optimize"
|
|
26
|
+
VALIDATION = "validation"
|
|
27
|
+
TEST = "test"
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
@dataclass(frozen=True, slots=True)
|
|
31
|
+
class SplitRatios:
|
|
32
|
+
"""Deterministic split thresholds."""
|
|
33
|
+
|
|
34
|
+
optimize: float = 0.8
|
|
35
|
+
validation: float = 0.1
|
|
36
|
+
test: float = 0.1
|
|
37
|
+
|
|
38
|
+
def __post_init__(self) -> None:
|
|
39
|
+
values = (self.optimize, self.validation, self.test)
|
|
40
|
+
if any(
|
|
41
|
+
isinstance(value, bool) or not math.isfinite(value) or value < 0 for value in values
|
|
42
|
+
):
|
|
43
|
+
raise ValueError("split ratios must be finite nonnegative numbers")
|
|
44
|
+
if not math.isclose(sum(values), 1.0, rel_tol=0.0, abs_tol=1e-12):
|
|
45
|
+
raise ValueError("split ratios must sum to 1")
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
@dataclass(frozen=True, slots=True)
|
|
49
|
+
class DatasetFieldMap:
|
|
50
|
+
"""Map source record fields to the canonical evaluation schema."""
|
|
51
|
+
|
|
52
|
+
id: str = "id"
|
|
53
|
+
query: str = "query"
|
|
54
|
+
reference_answer: str = "reference_answer"
|
|
55
|
+
reference_contexts: str = "reference_contexts"
|
|
56
|
+
relevant_document_ids: str = "relevant_document_ids"
|
|
57
|
+
tags: str = "tags"
|
|
58
|
+
metadata: str = "metadata"
|
|
59
|
+
split: str = "split"
|
|
60
|
+
|
|
61
|
+
def __post_init__(self) -> None:
|
|
62
|
+
values = tuple(getattr(self, name) for name in self.__dataclass_fields__)
|
|
63
|
+
if any(not value.strip() for value in values):
|
|
64
|
+
raise ValueError("dataset field names must not be empty")
|
|
65
|
+
if len(set(values)) != len(values):
|
|
66
|
+
raise ValueError("dataset field names must be unique")
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
@dataclass(frozen=True, slots=True)
|
|
70
|
+
class EvaluationExample:
|
|
71
|
+
"""One normalized, provenance-bearing evaluation example."""
|
|
72
|
+
|
|
73
|
+
id: str
|
|
74
|
+
query: str
|
|
75
|
+
reference_answer: str | None = None
|
|
76
|
+
reference_contexts: tuple[str, ...] = ()
|
|
77
|
+
relevant_document_ids: tuple[str, ...] = ()
|
|
78
|
+
tags: tuple[str, ...] = ()
|
|
79
|
+
metadata: Mapping[str, Any] = field(default_factory=dict)
|
|
80
|
+
split: DatasetSplit = DatasetSplit.OPTIMIZE
|
|
81
|
+
synthetic: bool = False
|
|
82
|
+
|
|
83
|
+
def __post_init__(self) -> None:
|
|
84
|
+
object.__setattr__(self, "id", _normalize_required(self.id, "example id"))
|
|
85
|
+
object.__setattr__(self, "query", _normalize_required(self.query, "query"))
|
|
86
|
+
if self.reference_answer is not None:
|
|
87
|
+
normalized_answer = _normalize_text(self.reference_answer)
|
|
88
|
+
object.__setattr__(self, "reference_answer", normalized_answer or None)
|
|
89
|
+
object.__setattr__(
|
|
90
|
+
self,
|
|
91
|
+
"reference_contexts",
|
|
92
|
+
_normalize_unique(self.reference_contexts, "reference context", unique=False),
|
|
93
|
+
)
|
|
94
|
+
object.__setattr__(
|
|
95
|
+
self,
|
|
96
|
+
"relevant_document_ids",
|
|
97
|
+
_normalize_unique(self.relevant_document_ids, "relevant document id"),
|
|
98
|
+
)
|
|
99
|
+
object.__setattr__(self, "tags", _normalize_unique(self.tags, "tag"))
|
|
100
|
+
normalized_metadata = _json_value(dict(self.metadata), path="metadata")
|
|
101
|
+
object.__setattr__(self, "metadata", MappingProxyType(normalized_metadata))
|
|
102
|
+
if not isinstance(self.split, DatasetSplit):
|
|
103
|
+
try:
|
|
104
|
+
object.__setattr__(self, "split", DatasetSplit(self.split))
|
|
105
|
+
except (TypeError, ValueError) as error:
|
|
106
|
+
raise ValueError("invalid dataset split") from error
|
|
107
|
+
if not isinstance(self.synthetic, bool):
|
|
108
|
+
raise TypeError("synthetic must be a boolean")
|
|
109
|
+
|
|
110
|
+
def semantic_data(self) -> dict[str, Any]:
|
|
111
|
+
"""Return normalized content used for version hashing."""
|
|
112
|
+
|
|
113
|
+
return {
|
|
114
|
+
"id": self.id,
|
|
115
|
+
"query": self.query,
|
|
116
|
+
"reference_answer": self.reference_answer,
|
|
117
|
+
"reference_contexts": self.reference_contexts,
|
|
118
|
+
"relevant_document_ids": self.relevant_document_ids,
|
|
119
|
+
"tags": self.tags,
|
|
120
|
+
"metadata": self.metadata,
|
|
121
|
+
"split": self.split.value,
|
|
122
|
+
"synthetic": self.synthetic,
|
|
123
|
+
}
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
@dataclass(frozen=True, slots=True)
|
|
127
|
+
class EvaluationDataset:
|
|
128
|
+
"""A normalized, immutable dataset with a content-derived version."""
|
|
129
|
+
|
|
130
|
+
examples: tuple[EvaluationExample, ...]
|
|
131
|
+
version: str = field(init=False)
|
|
132
|
+
|
|
133
|
+
def __post_init__(self) -> None:
|
|
134
|
+
if not self.examples:
|
|
135
|
+
raise ValueError("dataset must contain at least one example")
|
|
136
|
+
ids = [example.id for example in self.examples]
|
|
137
|
+
if len(ids) != len(set(ids)):
|
|
138
|
+
raise ValueError("dataset example ids must be unique")
|
|
139
|
+
version = content_hash(
|
|
140
|
+
{
|
|
141
|
+
"schema_version": _SCHEMA_VERSION,
|
|
142
|
+
"examples": [example.semantic_data() for example in self.examples],
|
|
143
|
+
},
|
|
144
|
+
namespace="tunarag:dataset:v1",
|
|
145
|
+
)
|
|
146
|
+
object.__setattr__(self, "version", version)
|
|
147
|
+
|
|
148
|
+
@classmethod
|
|
149
|
+
def from_records(
|
|
150
|
+
cls,
|
|
151
|
+
records: Iterable[Mapping[str, Any]],
|
|
152
|
+
*,
|
|
153
|
+
field_map: DatasetFieldMap | None = None,
|
|
154
|
+
source: str = "records",
|
|
155
|
+
) -> EvaluationDataset:
|
|
156
|
+
"""Normalize mapping records into an immutable dataset."""
|
|
157
|
+
|
|
158
|
+
mapping = field_map or DatasetFieldMap()
|
|
159
|
+
examples: list[EvaluationExample] = []
|
|
160
|
+
for index, record in enumerate(records, start=1):
|
|
161
|
+
if not isinstance(record, Mapping):
|
|
162
|
+
raise DatasetError(
|
|
163
|
+
"dataset record must be an object",
|
|
164
|
+
details={"source": source, "record": index},
|
|
165
|
+
)
|
|
166
|
+
try:
|
|
167
|
+
examples.append(_example_from_record(record, mapping))
|
|
168
|
+
except (TypeError, ValueError, json.JSONDecodeError) as error:
|
|
169
|
+
raise DatasetError(
|
|
170
|
+
"dataset record is invalid",
|
|
171
|
+
details={
|
|
172
|
+
"source": source,
|
|
173
|
+
"record": index,
|
|
174
|
+
"cause": type(error).__name__,
|
|
175
|
+
},
|
|
176
|
+
) from error
|
|
177
|
+
try:
|
|
178
|
+
return cls(tuple(examples))
|
|
179
|
+
except ValueError as error:
|
|
180
|
+
raise DatasetError(
|
|
181
|
+
"dataset validation failed",
|
|
182
|
+
details={"source": source, "cause": type(error).__name__},
|
|
183
|
+
) from error
|
|
184
|
+
|
|
185
|
+
@classmethod
|
|
186
|
+
def from_jsonl(
|
|
187
|
+
cls, path: str | Path, *, field_map: DatasetFieldMap | None = None
|
|
188
|
+
) -> EvaluationDataset:
|
|
189
|
+
"""Load UTF-8 JSON Lines, ignoring blank lines."""
|
|
190
|
+
|
|
191
|
+
source = str(path)
|
|
192
|
+
records: list[Mapping[str, Any]] = []
|
|
193
|
+
try:
|
|
194
|
+
with Path(path).open("r", encoding="utf-8-sig") as stream:
|
|
195
|
+
for line_number, line in enumerate(stream, start=1):
|
|
196
|
+
if not line.strip():
|
|
197
|
+
continue
|
|
198
|
+
try:
|
|
199
|
+
value = json.loads(line)
|
|
200
|
+
except json.JSONDecodeError as error:
|
|
201
|
+
raise DatasetError(
|
|
202
|
+
"could not load JSONL dataset",
|
|
203
|
+
details={
|
|
204
|
+
"source": source,
|
|
205
|
+
"record": line_number,
|
|
206
|
+
"cause": type(error).__name__,
|
|
207
|
+
},
|
|
208
|
+
) from error
|
|
209
|
+
if not isinstance(value, Mapping):
|
|
210
|
+
raise DatasetError(
|
|
211
|
+
"JSONL row must be an object",
|
|
212
|
+
details={"source": source, "record": line_number},
|
|
213
|
+
)
|
|
214
|
+
records.append(value)
|
|
215
|
+
except DatasetError:
|
|
216
|
+
raise
|
|
217
|
+
except (OSError, UnicodeError) as error:
|
|
218
|
+
raise DatasetError(
|
|
219
|
+
"could not load JSONL dataset",
|
|
220
|
+
details={"source": source, "cause": type(error).__name__},
|
|
221
|
+
) from error
|
|
222
|
+
return cls.from_records(records, field_map=field_map, source=source)
|
|
223
|
+
|
|
224
|
+
@classmethod
|
|
225
|
+
def from_json(
|
|
226
|
+
cls, path: str | Path, *, field_map: DatasetFieldMap | None = None
|
|
227
|
+
) -> EvaluationDataset:
|
|
228
|
+
"""Load a UTF-8 JSON array of objects."""
|
|
229
|
+
|
|
230
|
+
source = str(path)
|
|
231
|
+
try:
|
|
232
|
+
with Path(path).open("r", encoding="utf-8-sig") as stream:
|
|
233
|
+
value = json.load(stream)
|
|
234
|
+
except (OSError, UnicodeError, json.JSONDecodeError) as error:
|
|
235
|
+
raise DatasetError(
|
|
236
|
+
"could not load JSON dataset",
|
|
237
|
+
details={"source": source, "cause": type(error).__name__},
|
|
238
|
+
) from error
|
|
239
|
+
if not isinstance(value, list):
|
|
240
|
+
raise DatasetError("JSON dataset must be an array", details={"source": source})
|
|
241
|
+
return cls.from_records(value, field_map=field_map, source=source)
|
|
242
|
+
|
|
243
|
+
@classmethod
|
|
244
|
+
def from_csv(
|
|
245
|
+
cls, path: str | Path, *, field_map: DatasetFieldMap | None = None
|
|
246
|
+
) -> EvaluationDataset:
|
|
247
|
+
"""Load a UTF-8 CSV with JSON-encoded collection fields."""
|
|
248
|
+
|
|
249
|
+
source = str(path)
|
|
250
|
+
try:
|
|
251
|
+
with Path(path).open("r", encoding="utf-8-sig", newline="") as stream:
|
|
252
|
+
records = list(csv.DictReader(stream))
|
|
253
|
+
except (OSError, UnicodeError, csv.Error) as error:
|
|
254
|
+
raise DatasetError(
|
|
255
|
+
"could not load CSV dataset",
|
|
256
|
+
details={"source": source, "cause": type(error).__name__},
|
|
257
|
+
) from error
|
|
258
|
+
return cls.from_records(records, field_map=field_map, source=source)
|
|
259
|
+
|
|
260
|
+
def with_splits(
|
|
261
|
+
self,
|
|
262
|
+
*,
|
|
263
|
+
seed: int = 0,
|
|
264
|
+
ratios: SplitRatios | None = None,
|
|
265
|
+
group_metadata_key: str | None = None,
|
|
266
|
+
) -> EvaluationDataset:
|
|
267
|
+
"""Return a deterministically split copy, optionally keeping groups together."""
|
|
268
|
+
|
|
269
|
+
if isinstance(seed, bool) or not isinstance(seed, int):
|
|
270
|
+
raise TypeError("split seed must be an integer")
|
|
271
|
+
if group_metadata_key is not None and not group_metadata_key.strip():
|
|
272
|
+
raise ValueError("group metadata key must not be empty")
|
|
273
|
+
thresholds = ratios or SplitRatios()
|
|
274
|
+
examples = tuple(
|
|
275
|
+
replace(
|
|
276
|
+
example,
|
|
277
|
+
split=_select_split(
|
|
278
|
+
seed,
|
|
279
|
+
_group_identity(example, group_metadata_key),
|
|
280
|
+
thresholds,
|
|
281
|
+
),
|
|
282
|
+
)
|
|
283
|
+
for example in self.examples
|
|
284
|
+
)
|
|
285
|
+
return EvaluationDataset(examples)
|
|
286
|
+
|
|
287
|
+
def select(self, split: DatasetSplit) -> EvaluationDataset:
|
|
288
|
+
"""Return one nonempty canonical split."""
|
|
289
|
+
|
|
290
|
+
selected = tuple(example for example in self.examples if example.split is split)
|
|
291
|
+
if not selected:
|
|
292
|
+
raise DatasetError("selected dataset split is empty", details={"split": split.value})
|
|
293
|
+
return EvaluationDataset(selected)
|
|
294
|
+
|
|
295
|
+
def __aiter__(self) -> AsyncIterator[EvaluationExample]:
|
|
296
|
+
return self._iterate()
|
|
297
|
+
|
|
298
|
+
async def _iterate(self) -> AsyncIterator[EvaluationExample]:
|
|
299
|
+
for example in self.examples:
|
|
300
|
+
yield example
|
|
301
|
+
|
|
302
|
+
|
|
303
|
+
@dataclass(frozen=True, slots=True)
|
|
304
|
+
class InMemoryDataset:
|
|
305
|
+
"""A versioned, repeatable dataset backed by an immutable tuple."""
|
|
306
|
+
|
|
307
|
+
examples: tuple[Any, ...]
|
|
308
|
+
version: str
|
|
309
|
+
|
|
310
|
+
def __post_init__(self) -> None:
|
|
311
|
+
if not self.examples:
|
|
312
|
+
raise ValueError("dataset must contain at least one example")
|
|
313
|
+
if not self.version.strip():
|
|
314
|
+
raise ValueError("dataset version must not be empty")
|
|
315
|
+
|
|
316
|
+
@classmethod
|
|
317
|
+
def from_iterable(cls, examples: Iterable[Any], *, version: str) -> InMemoryDataset:
|
|
318
|
+
"""Materialize a repeatable dataset from an iterable."""
|
|
319
|
+
|
|
320
|
+
return cls(tuple(examples), version)
|
|
321
|
+
|
|
322
|
+
def __aiter__(self) -> AsyncIterator[Any]:
|
|
323
|
+
return self._iterate()
|
|
324
|
+
|
|
325
|
+
async def _iterate(self) -> AsyncIterator[Any]:
|
|
326
|
+
for example in self.examples:
|
|
327
|
+
yield example
|
|
328
|
+
|
|
329
|
+
|
|
330
|
+
def _example_from_record(record: Mapping[str, Any], fields: DatasetFieldMap) -> EvaluationExample:
|
|
331
|
+
query = _required_record_value(record, fields.query)
|
|
332
|
+
example_id = record.get(fields.id)
|
|
333
|
+
reference_answer = _optional_string(record.get(fields.reference_answer))
|
|
334
|
+
reference_contexts = _string_sequence(record.get(fields.reference_contexts))
|
|
335
|
+
relevant_document_ids = _string_sequence(record.get(fields.relevant_document_ids))
|
|
336
|
+
tags = _string_sequence(record.get(fields.tags))
|
|
337
|
+
metadata = _metadata(record.get(fields.metadata))
|
|
338
|
+
split = _split(record.get(fields.split))
|
|
339
|
+
if example_id is None or (isinstance(example_id, str) and not example_id.strip()):
|
|
340
|
+
example_id = content_hash(
|
|
341
|
+
{
|
|
342
|
+
"query": query,
|
|
343
|
+
"reference_answer": reference_answer,
|
|
344
|
+
"reference_contexts": reference_contexts,
|
|
345
|
+
"relevant_document_ids": relevant_document_ids,
|
|
346
|
+
"tags": tags,
|
|
347
|
+
"metadata": metadata,
|
|
348
|
+
"split": split.value,
|
|
349
|
+
"synthetic": False,
|
|
350
|
+
},
|
|
351
|
+
namespace="tunarag:example-id:v1",
|
|
352
|
+
)
|
|
353
|
+
if not isinstance(example_id, str):
|
|
354
|
+
raise TypeError("example id must be a string")
|
|
355
|
+
return EvaluationExample(
|
|
356
|
+
id=example_id,
|
|
357
|
+
query=query,
|
|
358
|
+
reference_answer=reference_answer,
|
|
359
|
+
reference_contexts=reference_contexts,
|
|
360
|
+
relevant_document_ids=relevant_document_ids,
|
|
361
|
+
tags=tags,
|
|
362
|
+
metadata=metadata,
|
|
363
|
+
split=split,
|
|
364
|
+
synthetic=False,
|
|
365
|
+
)
|
|
366
|
+
|
|
367
|
+
|
|
368
|
+
def _required_record_value(record: Mapping[str, Any], key: str) -> str:
|
|
369
|
+
value = record.get(key)
|
|
370
|
+
if not isinstance(value, str):
|
|
371
|
+
raise TypeError(f"{key} must be a string")
|
|
372
|
+
return value
|
|
373
|
+
|
|
374
|
+
|
|
375
|
+
def _optional_string(value: Any) -> str | None:
|
|
376
|
+
if value is None or value == "":
|
|
377
|
+
return None
|
|
378
|
+
if not isinstance(value, str):
|
|
379
|
+
raise TypeError("optional text field must be a string")
|
|
380
|
+
return value
|
|
381
|
+
|
|
382
|
+
|
|
383
|
+
def _string_sequence(value: Any) -> tuple[str, ...]:
|
|
384
|
+
if value is None or value == "":
|
|
385
|
+
return ()
|
|
386
|
+
if isinstance(value, str):
|
|
387
|
+
value = json.loads(value)
|
|
388
|
+
if not isinstance(value, Sequence) or isinstance(value, (str, bytes, bytearray)):
|
|
389
|
+
raise TypeError("collection field must be an array of strings")
|
|
390
|
+
if any(not isinstance(item, str) for item in value):
|
|
391
|
+
raise TypeError("collection field must contain only strings")
|
|
392
|
+
return tuple(value)
|
|
393
|
+
|
|
394
|
+
|
|
395
|
+
def _metadata(value: Any) -> Mapping[str, Any]:
|
|
396
|
+
if value is None or value == "":
|
|
397
|
+
return {}
|
|
398
|
+
if isinstance(value, str):
|
|
399
|
+
value = json.loads(value)
|
|
400
|
+
if not isinstance(value, Mapping):
|
|
401
|
+
raise TypeError("metadata must be an object")
|
|
402
|
+
return value
|
|
403
|
+
|
|
404
|
+
|
|
405
|
+
def _split(value: Any) -> DatasetSplit:
|
|
406
|
+
if value is None or value == "":
|
|
407
|
+
return DatasetSplit.OPTIMIZE
|
|
408
|
+
if not isinstance(value, str):
|
|
409
|
+
raise TypeError("split must be a string")
|
|
410
|
+
return DatasetSplit(value.strip().lower())
|
|
411
|
+
|
|
412
|
+
|
|
413
|
+
def _normalize_required(value: str, label: str) -> str:
|
|
414
|
+
normalized = _normalize_text(value)
|
|
415
|
+
if not normalized:
|
|
416
|
+
raise ValueError(f"{label} must not be empty")
|
|
417
|
+
return normalized
|
|
418
|
+
|
|
419
|
+
|
|
420
|
+
def _normalize_text(value: str) -> str:
|
|
421
|
+
if not isinstance(value, str):
|
|
422
|
+
raise TypeError("text value must be a string")
|
|
423
|
+
normalized = unicodedata.normalize("NFC", value.lstrip("\ufeff"))
|
|
424
|
+
normalized = normalized.replace("\r\n", "\n").replace("\r", "\n").strip()
|
|
425
|
+
if any(unicodedata.category(char) == "Cc" and char not in {"\n", "\t"} for char in normalized):
|
|
426
|
+
raise ValueError("text contains a disallowed control character")
|
|
427
|
+
return normalized
|
|
428
|
+
|
|
429
|
+
|
|
430
|
+
def _normalize_unique(values: Iterable[str], label: str, *, unique: bool = True) -> tuple[str, ...]:
|
|
431
|
+
normalized = tuple(_normalize_required(value, label) for value in values)
|
|
432
|
+
if unique and len(normalized) != len(set(normalized)):
|
|
433
|
+
raise ValueError(f"{label} values must be unique")
|
|
434
|
+
return normalized
|
|
435
|
+
|
|
436
|
+
|
|
437
|
+
def _json_value(value: Any, *, path: str) -> Any:
|
|
438
|
+
if value is None or isinstance(value, (bool, int, str)):
|
|
439
|
+
return value
|
|
440
|
+
if isinstance(value, float):
|
|
441
|
+
if not math.isfinite(value):
|
|
442
|
+
raise ValueError(f"{path} contains a nonfinite number")
|
|
443
|
+
return value
|
|
444
|
+
if isinstance(value, Mapping):
|
|
445
|
+
if any(not isinstance(key, str) for key in value):
|
|
446
|
+
raise TypeError(f"{path} keys must be strings")
|
|
447
|
+
return MappingProxyType(
|
|
448
|
+
{key: _json_value(item, path=f"{path}.{key}") for key, item in value.items()}
|
|
449
|
+
)
|
|
450
|
+
if isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
|
|
451
|
+
return tuple(_json_value(item, path=f"{path}[]") for item in value)
|
|
452
|
+
raise TypeError(f"{path} contains unsupported value type {type(value).__name__}")
|
|
453
|
+
|
|
454
|
+
|
|
455
|
+
def _group_identity(example: EvaluationExample, metadata_key: str | None) -> str:
|
|
456
|
+
if metadata_key is None:
|
|
457
|
+
return example.id
|
|
458
|
+
if metadata_key not in example.metadata:
|
|
459
|
+
raise DatasetError(
|
|
460
|
+
"group metadata is missing",
|
|
461
|
+
details={"example_id": example.id, "metadata_key": metadata_key},
|
|
462
|
+
)
|
|
463
|
+
value = example.metadata[metadata_key]
|
|
464
|
+
if not isinstance(value, (str, int)) or isinstance(value, bool):
|
|
465
|
+
raise DatasetError(
|
|
466
|
+
"group metadata must be a string or integer",
|
|
467
|
+
details={"example_id": example.id, "metadata_key": metadata_key},
|
|
468
|
+
)
|
|
469
|
+
return str(value)
|
|
470
|
+
|
|
471
|
+
|
|
472
|
+
def _select_split(seed: int, identity: str, ratios: SplitRatios) -> DatasetSplit:
|
|
473
|
+
digest = content_hash(
|
|
474
|
+
{"seed": seed, "identity": identity}, namespace="tunarag:dataset-split:v1"
|
|
475
|
+
)
|
|
476
|
+
fraction = int(digest[:16], 16) / float(2**64)
|
|
477
|
+
if fraction < ratios.optimize:
|
|
478
|
+
return DatasetSplit.OPTIMIZE
|
|
479
|
+
if fraction < ratios.optimize + ratios.validation:
|
|
480
|
+
return DatasetSplit.VALIDATION
|
|
481
|
+
return DatasetSplit.TEST
|
tunarag/domain.py
ADDED
|
@@ -0,0 +1,83 @@
|
|
|
1
|
+
"""Immutable domain values shared by the package contracts."""
|
|
2
|
+
|
|
3
|
+
from collections.abc import Mapping
|
|
4
|
+
from dataclasses import dataclass
|
|
5
|
+
from enum import Enum
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class ValueStatus(str, Enum):
|
|
10
|
+
"""Confidence state for an observed or derived value."""
|
|
11
|
+
|
|
12
|
+
EXACT = "exact"
|
|
13
|
+
ESTIMATED = "estimated"
|
|
14
|
+
UNAVAILABLE = "unavailable"
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
@dataclass(frozen=True, slots=True)
|
|
18
|
+
class Secret:
|
|
19
|
+
"""A value whose contents are never exposed by representation or hashing."""
|
|
20
|
+
|
|
21
|
+
value: str
|
|
22
|
+
|
|
23
|
+
def __repr__(self) -> str:
|
|
24
|
+
return "Secret(***redacted***)"
|
|
25
|
+
|
|
26
|
+
def __str__(self) -> str:
|
|
27
|
+
return "***redacted***"
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
@dataclass(frozen=True, slots=True)
|
|
31
|
+
class MetricValue:
|
|
32
|
+
"""A metric observation with its coverage and confidence state."""
|
|
33
|
+
|
|
34
|
+
name: str
|
|
35
|
+
value: float | None
|
|
36
|
+
status: ValueStatus = ValueStatus.EXACT
|
|
37
|
+
coverage: float = 1.0
|
|
38
|
+
|
|
39
|
+
def __post_init__(self) -> None:
|
|
40
|
+
if not self.name.strip():
|
|
41
|
+
raise ValueError("metric name must not be empty")
|
|
42
|
+
if not 0.0 <= self.coverage <= 1.0:
|
|
43
|
+
raise ValueError("metric coverage must be between 0 and 1")
|
|
44
|
+
if self.value is not None and not isinstance(self.value, (int, float)):
|
|
45
|
+
raise TypeError("metric value must be numeric or None")
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
@dataclass(frozen=True, slots=True)
|
|
49
|
+
class UsageRecord:
|
|
50
|
+
"""Usage observed for one component of a trial."""
|
|
51
|
+
|
|
52
|
+
component: str
|
|
53
|
+
input_tokens: int | None = None
|
|
54
|
+
output_tokens: int | None = None
|
|
55
|
+
cost: float | None = None
|
|
56
|
+
latency_seconds: float | None = None
|
|
57
|
+
status: ValueStatus = ValueStatus.EXACT
|
|
58
|
+
pricing_version: str | None = None
|
|
59
|
+
|
|
60
|
+
def __post_init__(self) -> None:
|
|
61
|
+
if not self.component.strip():
|
|
62
|
+
raise ValueError("usage component must not be empty")
|
|
63
|
+
for name in ("input_tokens", "output_tokens"):
|
|
64
|
+
token_value: Any = getattr(self, name)
|
|
65
|
+
if token_value is not None and token_value < 0:
|
|
66
|
+
raise ValueError(f"{name} must not be negative")
|
|
67
|
+
for name in ("cost", "latency_seconds"):
|
|
68
|
+
usage_value: Any = getattr(self, name)
|
|
69
|
+
if usage_value is not None and usage_value < 0:
|
|
70
|
+
raise ValueError(f"{name} must not be negative")
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
@dataclass(frozen=True, slots=True)
|
|
74
|
+
class Candidate:
|
|
75
|
+
"""A typed candidate configuration proposed for a trial."""
|
|
76
|
+
|
|
77
|
+
parameters: Mapping[str, Any]
|
|
78
|
+
|
|
79
|
+
def __post_init__(self) -> None:
|
|
80
|
+
if not self.parameters:
|
|
81
|
+
raise ValueError("candidate parameters must not be empty")
|
|
82
|
+
if any(not name.strip() for name in self.parameters):
|
|
83
|
+
raise ValueError("candidate parameter names must not be empty")
|