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/synthetic.py
ADDED
|
@@ -0,0 +1,257 @@
|
|
|
1
|
+
"""Provider-neutral synthetic QA dataset generation."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import math
|
|
7
|
+
import unicodedata
|
|
8
|
+
from collections.abc import Mapping, Sequence
|
|
9
|
+
from dataclasses import dataclass, field
|
|
10
|
+
from typing import Any, Protocol
|
|
11
|
+
|
|
12
|
+
from .contracts import Fingerprintable
|
|
13
|
+
from .dataset import EvaluationDataset, EvaluationExample
|
|
14
|
+
from .domain import UsageRecord
|
|
15
|
+
from .errors import DatasetError
|
|
16
|
+
from .serialization import content_hash
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
@dataclass(frozen=True, slots=True)
|
|
20
|
+
class SourceDocument:
|
|
21
|
+
"""One source document approved for synthetic generation."""
|
|
22
|
+
|
|
23
|
+
id: str
|
|
24
|
+
text: str
|
|
25
|
+
metadata: Mapping[str, Any] = field(default_factory=dict)
|
|
26
|
+
|
|
27
|
+
def __post_init__(self) -> None:
|
|
28
|
+
object.__setattr__(self, "id", _required_text(self.id, "source document id"))
|
|
29
|
+
object.__setattr__(self, "text", _required_text(self.text, "source document text"))
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
@dataclass(frozen=True, slots=True)
|
|
33
|
+
class SourceSpan:
|
|
34
|
+
"""Half-open source character range supporting a generated QA pair."""
|
|
35
|
+
|
|
36
|
+
start: int
|
|
37
|
+
end: int
|
|
38
|
+
|
|
39
|
+
def __post_init__(self) -> None:
|
|
40
|
+
if any(
|
|
41
|
+
isinstance(value, bool) or not isinstance(value, int)
|
|
42
|
+
for value in (self.start, self.end)
|
|
43
|
+
):
|
|
44
|
+
raise TypeError("source span offsets must be integers")
|
|
45
|
+
if self.start < 0 or self.end <= self.start:
|
|
46
|
+
raise ValueError("source span must be a nonempty half-open range")
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
@dataclass(frozen=True, slots=True)
|
|
50
|
+
class SyntheticQA:
|
|
51
|
+
"""One provider-produced question, answer, and supporting provenance."""
|
|
52
|
+
|
|
53
|
+
query: str
|
|
54
|
+
answer: str
|
|
55
|
+
source_spans: tuple[SourceSpan, ...] = ()
|
|
56
|
+
confidence: float | None = None
|
|
57
|
+
|
|
58
|
+
def __post_init__(self) -> None:
|
|
59
|
+
object.__setattr__(self, "query", _required_text(self.query, "synthetic query"))
|
|
60
|
+
object.__setattr__(self, "answer", _required_text(self.answer, "synthetic answer"))
|
|
61
|
+
if self.confidence is not None:
|
|
62
|
+
if (
|
|
63
|
+
isinstance(self.confidence, bool)
|
|
64
|
+
or not math.isfinite(self.confidence)
|
|
65
|
+
or not 0.0 <= self.confidence <= 1.0
|
|
66
|
+
):
|
|
67
|
+
raise ValueError("synthetic confidence must be between 0 and 1")
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
@dataclass(frozen=True, slots=True)
|
|
71
|
+
class SyntheticGenerationRequest:
|
|
72
|
+
"""Structured request sent to a synthetic QA provider."""
|
|
73
|
+
|
|
74
|
+
document: SourceDocument
|
|
75
|
+
count: int
|
|
76
|
+
seed: int
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
@dataclass(frozen=True, slots=True)
|
|
80
|
+
class SyntheticGenerationResponse:
|
|
81
|
+
"""Provider output with exact or estimated usage records."""
|
|
82
|
+
|
|
83
|
+
items: tuple[SyntheticQA, ...]
|
|
84
|
+
usage: tuple[UsageRecord, ...] = ()
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
class SyntheticQAProvider(Protocol):
|
|
88
|
+
"""Generate typed QA pairs for one source document."""
|
|
89
|
+
|
|
90
|
+
async def generate(
|
|
91
|
+
self, request: SyntheticGenerationRequest
|
|
92
|
+
) -> SyntheticGenerationResponse: ...
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
@dataclass(frozen=True, slots=True)
|
|
96
|
+
class SyntheticDatasetResult:
|
|
97
|
+
"""Generated evaluation dataset and provider usage."""
|
|
98
|
+
|
|
99
|
+
dataset: EvaluationDataset
|
|
100
|
+
usage: tuple[UsageRecord, ...]
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
class SyntheticDatasetBuilder:
|
|
104
|
+
"""Generate a validated dataset with bounded document concurrency."""
|
|
105
|
+
|
|
106
|
+
def __init__(
|
|
107
|
+
self,
|
|
108
|
+
provider: SyntheticQAProvider,
|
|
109
|
+
*,
|
|
110
|
+
samples_per_document: int = 3,
|
|
111
|
+
concurrency: int = 4,
|
|
112
|
+
seed: int = 0,
|
|
113
|
+
generator_version: str = "unspecified",
|
|
114
|
+
) -> None:
|
|
115
|
+
for name, value in (
|
|
116
|
+
("samples_per_document", samples_per_document),
|
|
117
|
+
("concurrency", concurrency),
|
|
118
|
+
):
|
|
119
|
+
if isinstance(value, bool) or not isinstance(value, int) or value < 1:
|
|
120
|
+
raise ValueError(f"{name} must be a positive integer")
|
|
121
|
+
if isinstance(seed, bool) or not isinstance(seed, int):
|
|
122
|
+
raise TypeError("synthetic generation seed must be an integer")
|
|
123
|
+
self._provider = provider
|
|
124
|
+
self._samples_per_document = samples_per_document
|
|
125
|
+
self._concurrency = concurrency
|
|
126
|
+
self._seed = seed
|
|
127
|
+
self._generator_version = _required_text(generator_version, "generator version")
|
|
128
|
+
|
|
129
|
+
async def generate(self, documents: Sequence[SourceDocument]) -> SyntheticDatasetResult:
|
|
130
|
+
"""Generate and validate examples in source-document order."""
|
|
131
|
+
|
|
132
|
+
if not documents:
|
|
133
|
+
raise DatasetError("synthetic generation requires source documents")
|
|
134
|
+
if len({document.id for document in documents}) != len(documents):
|
|
135
|
+
raise DatasetError("source document ids must be unique")
|
|
136
|
+
semaphore = asyncio.Semaphore(self._concurrency)
|
|
137
|
+
|
|
138
|
+
async def bounded(document: SourceDocument) -> SyntheticGenerationResponse:
|
|
139
|
+
async with semaphore:
|
|
140
|
+
request = SyntheticGenerationRequest(
|
|
141
|
+
document,
|
|
142
|
+
self._samples_per_document,
|
|
143
|
+
_document_seed(self._seed, document.id),
|
|
144
|
+
)
|
|
145
|
+
try:
|
|
146
|
+
response = await self._provider.generate(request)
|
|
147
|
+
except Exception as error:
|
|
148
|
+
raise DatasetError(
|
|
149
|
+
"synthetic QA provider failed",
|
|
150
|
+
details={"document_id": document.id, "cause": type(error).__name__},
|
|
151
|
+
) from error
|
|
152
|
+
return _validate_response(response, request)
|
|
153
|
+
|
|
154
|
+
responses = await asyncio.gather(*(bounded(document) for document in documents))
|
|
155
|
+
examples: list[EvaluationExample] = []
|
|
156
|
+
usage: list[UsageRecord] = []
|
|
157
|
+
provider_identity = _provider_identity(self._provider)
|
|
158
|
+
for document, response in zip(documents, responses, strict=True):
|
|
159
|
+
request_seed = _document_seed(self._seed, document.id)
|
|
160
|
+
usage.extend(response.usage)
|
|
161
|
+
for index, item in enumerate(response.items):
|
|
162
|
+
contexts = tuple(document.text[span.start : span.end] for span in item.source_spans)
|
|
163
|
+
if not contexts:
|
|
164
|
+
contexts = (document.text,)
|
|
165
|
+
metadata = {
|
|
166
|
+
"source_document_id": document.id,
|
|
167
|
+
"source_document_metadata": document.metadata,
|
|
168
|
+
"source_spans": [
|
|
169
|
+
{"start": span.start, "end": span.end} for span in item.source_spans
|
|
170
|
+
],
|
|
171
|
+
"generator": provider_identity,
|
|
172
|
+
"generator_version": self._generator_version,
|
|
173
|
+
"seed": request_seed,
|
|
174
|
+
"confidence": item.confidence,
|
|
175
|
+
"review_status": "unreviewed",
|
|
176
|
+
}
|
|
177
|
+
example_id = content_hash(
|
|
178
|
+
{
|
|
179
|
+
"document_id": document.id,
|
|
180
|
+
"index": index,
|
|
181
|
+
"qa": metadata,
|
|
182
|
+
"query": item.query,
|
|
183
|
+
"answer": item.answer,
|
|
184
|
+
},
|
|
185
|
+
namespace="tunarag:synthetic-example:v1",
|
|
186
|
+
)
|
|
187
|
+
examples.append(
|
|
188
|
+
EvaluationExample(
|
|
189
|
+
id=example_id,
|
|
190
|
+
query=item.query,
|
|
191
|
+
reference_answer=item.answer,
|
|
192
|
+
reference_contexts=contexts,
|
|
193
|
+
relevant_document_ids=(document.id,),
|
|
194
|
+
tags=("synthetic",),
|
|
195
|
+
metadata=metadata,
|
|
196
|
+
synthetic=True,
|
|
197
|
+
)
|
|
198
|
+
)
|
|
199
|
+
return SyntheticDatasetResult(EvaluationDataset(tuple(examples)), tuple(usage))
|
|
200
|
+
|
|
201
|
+
|
|
202
|
+
def _validate_response(
|
|
203
|
+
response: Any, request: SyntheticGenerationRequest
|
|
204
|
+
) -> SyntheticGenerationResponse:
|
|
205
|
+
if not isinstance(response, SyntheticGenerationResponse):
|
|
206
|
+
raise DatasetError(
|
|
207
|
+
"synthetic provider returned an invalid response",
|
|
208
|
+
details={"document_id": request.document.id},
|
|
209
|
+
)
|
|
210
|
+
if len(response.items) != request.count:
|
|
211
|
+
raise DatasetError(
|
|
212
|
+
"synthetic provider returned an unexpected item count",
|
|
213
|
+
details={
|
|
214
|
+
"document_id": request.document.id,
|
|
215
|
+
"expected": request.count,
|
|
216
|
+
"actual": len(response.items),
|
|
217
|
+
},
|
|
218
|
+
)
|
|
219
|
+
if any(not isinstance(item, SyntheticQA) for item in response.items):
|
|
220
|
+
raise DatasetError("synthetic provider returned invalid QA items")
|
|
221
|
+
if any(not isinstance(item, UsageRecord) for item in response.usage):
|
|
222
|
+
raise DatasetError("synthetic provider returned invalid usage records")
|
|
223
|
+
for item in response.items:
|
|
224
|
+
if any(span.end > len(request.document.text) for span in item.source_spans):
|
|
225
|
+
raise DatasetError(
|
|
226
|
+
"synthetic source span exceeds its document",
|
|
227
|
+
details={"document_id": request.document.id},
|
|
228
|
+
)
|
|
229
|
+
return response
|
|
230
|
+
|
|
231
|
+
|
|
232
|
+
def _document_seed(seed: int, document_id: str) -> int:
|
|
233
|
+
digest = content_hash(
|
|
234
|
+
{"seed": seed, "document_id": document_id}, namespace="tunarag:synthetic-seed:v1"
|
|
235
|
+
)
|
|
236
|
+
return int(digest[:16], 16)
|
|
237
|
+
|
|
238
|
+
|
|
239
|
+
def _provider_identity(provider: SyntheticQAProvider) -> Mapping[str, Any]:
|
|
240
|
+
identity: dict[str, Any] = {
|
|
241
|
+
"type": f"{type(provider).__module__}.{type(provider).__qualname__}"
|
|
242
|
+
}
|
|
243
|
+
if isinstance(provider, Fingerprintable):
|
|
244
|
+
identity["fingerprint"] = provider.fingerprint()
|
|
245
|
+
return identity
|
|
246
|
+
|
|
247
|
+
|
|
248
|
+
def _required_text(value: str, label: str) -> str:
|
|
249
|
+
if not isinstance(value, str):
|
|
250
|
+
raise TypeError(f"{label} must be a string")
|
|
251
|
+
normalized = unicodedata.normalize("NFC", value.lstrip("\ufeff"))
|
|
252
|
+
normalized = normalized.replace("\r\n", "\n").replace("\r", "\n").strip()
|
|
253
|
+
if not normalized:
|
|
254
|
+
raise ValueError(f"{label} must not be empty")
|
|
255
|
+
if any(unicodedata.category(char) == "Cc" and char not in {"\n", "\t"} for char in normalized):
|
|
256
|
+
raise ValueError(f"{label} contains a disallowed control character")
|
|
257
|
+
return normalized
|