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/errors.py
ADDED
|
@@ -0,0 +1,179 @@
|
|
|
1
|
+
"""Structured public errors with secret-safe diagnostic details."""
|
|
2
|
+
|
|
3
|
+
from collections.abc import Mapping, Sequence
|
|
4
|
+
from dataclasses import fields, is_dataclass
|
|
5
|
+
from enum import Enum
|
|
6
|
+
from types import MappingProxyType
|
|
7
|
+
from typing import Any, cast
|
|
8
|
+
|
|
9
|
+
from .domain import Secret
|
|
10
|
+
|
|
11
|
+
_SENSITIVE_KEYS = {
|
|
12
|
+
"access_token",
|
|
13
|
+
"api_key",
|
|
14
|
+
"authorization",
|
|
15
|
+
"credential",
|
|
16
|
+
"credentials",
|
|
17
|
+
"password",
|
|
18
|
+
"refresh_token",
|
|
19
|
+
"secret",
|
|
20
|
+
"token",
|
|
21
|
+
}
|
|
22
|
+
_REDACTED = "***redacted***"
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class ErrorCode(str, Enum):
|
|
26
|
+
"""Stable machine-readable error identifiers."""
|
|
27
|
+
|
|
28
|
+
CONFIGURATION = "configuration.invalid"
|
|
29
|
+
SEARCH_SPACE = "search_space.invalid"
|
|
30
|
+
DATASET = "dataset.invalid"
|
|
31
|
+
ADAPTER = "adapter.failed"
|
|
32
|
+
EVALUATION = "evaluation.failed"
|
|
33
|
+
STORE = "store.failed"
|
|
34
|
+
RESUME = "resume.incompatible"
|
|
35
|
+
BUDGET = "budget.exhausted"
|
|
36
|
+
OPTIONAL_DEPENDENCY = "dependency.optional_missing"
|
|
37
|
+
SYNC_IN_ASYNC = "runtime.sync_in_async"
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class TunaRAGError(Exception):
|
|
41
|
+
"""Base class for public package errors."""
|
|
42
|
+
|
|
43
|
+
code = ErrorCode.CONFIGURATION
|
|
44
|
+
|
|
45
|
+
def __init__(
|
|
46
|
+
self,
|
|
47
|
+
message: str,
|
|
48
|
+
*,
|
|
49
|
+
stage: str | None = None,
|
|
50
|
+
retryable: bool = False,
|
|
51
|
+
details: Mapping[str, Any] | None = None,
|
|
52
|
+
) -> None:
|
|
53
|
+
if not message.strip():
|
|
54
|
+
raise ValueError("error message must not be empty")
|
|
55
|
+
if stage is not None and not stage.strip():
|
|
56
|
+
raise ValueError("error stage must not be empty")
|
|
57
|
+
self.message = message
|
|
58
|
+
self.stage = stage
|
|
59
|
+
self.retryable = retryable
|
|
60
|
+
safe_details = cast(dict[str, Any], redact(details or {}))
|
|
61
|
+
self.details: Mapping[str, Any] = MappingProxyType(safe_details)
|
|
62
|
+
super().__init__(message)
|
|
63
|
+
|
|
64
|
+
def as_dict(self) -> dict[str, Any]:
|
|
65
|
+
"""Return a JSON-compatible diagnostic representation."""
|
|
66
|
+
|
|
67
|
+
return {
|
|
68
|
+
"code": self.code.value,
|
|
69
|
+
"message": self.message,
|
|
70
|
+
"stage": self.stage,
|
|
71
|
+
"retryable": self.retryable,
|
|
72
|
+
"details": dict(self.details),
|
|
73
|
+
}
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
class ConfigurationError(TunaRAGError):
|
|
77
|
+
"""Configuration could not be parsed or validated."""
|
|
78
|
+
|
|
79
|
+
code = ErrorCode.CONFIGURATION
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
class SearchSpaceError(TunaRAGError):
|
|
83
|
+
"""A search-space definition or candidate is invalid."""
|
|
84
|
+
|
|
85
|
+
code = ErrorCode.SEARCH_SPACE
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
class DatasetError(TunaRAGError):
|
|
89
|
+
"""An evaluation dataset could not be loaded or validated."""
|
|
90
|
+
|
|
91
|
+
code = ErrorCode.DATASET
|
|
92
|
+
|
|
93
|
+
def __init__(self, message: str, *, details: Mapping[str, Any] | None = None) -> None:
|
|
94
|
+
super().__init__(message, stage="dataset", details=details)
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
class AdapterError(TunaRAGError):
|
|
98
|
+
"""The user-owned RAG adapter failed."""
|
|
99
|
+
|
|
100
|
+
code = ErrorCode.ADAPTER
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
class EvaluationError(TunaRAGError):
|
|
104
|
+
"""An evaluator failed or returned invalid observations."""
|
|
105
|
+
|
|
106
|
+
code = ErrorCode.EVALUATION
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
class StoreError(TunaRAGError):
|
|
110
|
+
"""Durable state could not be read or written."""
|
|
111
|
+
|
|
112
|
+
code = ErrorCode.STORE
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
class ResumeError(TunaRAGError):
|
|
116
|
+
"""A study cannot resume under the supplied semantics."""
|
|
117
|
+
|
|
118
|
+
code = ErrorCode.RESUME
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
class BudgetError(TunaRAGError):
|
|
122
|
+
"""A configured scheduling budget has been exhausted."""
|
|
123
|
+
|
|
124
|
+
code = ErrorCode.BUDGET
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
class MissingOptionalDependencyError(TunaRAGError):
|
|
128
|
+
"""A requested integration extra is not installed."""
|
|
129
|
+
|
|
130
|
+
code = ErrorCode.OPTIONAL_DEPENDENCY
|
|
131
|
+
|
|
132
|
+
def __init__(self, feature: str, extra: str) -> None:
|
|
133
|
+
if not feature.strip() or not extra.strip():
|
|
134
|
+
raise ValueError("feature and extra must not be empty")
|
|
135
|
+
super().__init__(
|
|
136
|
+
f"{feature} requires the optional '{extra}' extra",
|
|
137
|
+
stage="dependency",
|
|
138
|
+
details={"feature": feature, "extra": extra},
|
|
139
|
+
)
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
class SyncInAsyncContextError(TunaRAGError):
|
|
143
|
+
"""A synchronous entry point was called from an active event loop."""
|
|
144
|
+
|
|
145
|
+
code = ErrorCode.SYNC_IN_ASYNC
|
|
146
|
+
|
|
147
|
+
def __init__(self) -> None:
|
|
148
|
+
super().__init__(
|
|
149
|
+
"the synchronous API cannot run inside an active event loop; use the async API",
|
|
150
|
+
stage="runtime",
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def redact(value: Any, *, key: str | None = None) -> Any:
|
|
155
|
+
"""Return a recursively redacted, JSON-compatible diagnostic value."""
|
|
156
|
+
|
|
157
|
+
if key is not None and _is_sensitive_key(key):
|
|
158
|
+
return _REDACTED
|
|
159
|
+
if isinstance(value, Secret):
|
|
160
|
+
return _REDACTED
|
|
161
|
+
if is_dataclass(value):
|
|
162
|
+
return {
|
|
163
|
+
field.name: redact(getattr(value, field.name), key=field.name)
|
|
164
|
+
for field in fields(cast(Any, value))
|
|
165
|
+
}
|
|
166
|
+
if isinstance(value, Mapping):
|
|
167
|
+
return {str(item_key): redact(item, key=str(item_key)) for item_key, item in value.items()}
|
|
168
|
+
if isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
|
|
169
|
+
return [redact(item) for item in value]
|
|
170
|
+
if value is None or isinstance(value, (bool, int, float, str)):
|
|
171
|
+
return value
|
|
172
|
+
return f"<{type(value).__name__}>"
|
|
173
|
+
|
|
174
|
+
|
|
175
|
+
def _is_sensitive_key(key: str) -> bool:
|
|
176
|
+
normalized = key.lower().replace("-", "_")
|
|
177
|
+
return normalized in _SENSITIVE_KEYS or normalized.endswith(
|
|
178
|
+
("_api_key", "_credential", "_password", "_secret", "_token")
|
|
179
|
+
)
|
tunarag/evaluators.py
ADDED
|
@@ -0,0 +1,186 @@
|
|
|
1
|
+
"""Built-in evaluator helpers and optional third-party integrations."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import inspect
|
|
6
|
+
import math
|
|
7
|
+
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
|
8
|
+
from importlib import import_module
|
|
9
|
+
from numbers import Real
|
|
10
|
+
from typing import Any
|
|
11
|
+
|
|
12
|
+
from .contracts import Fingerprintable
|
|
13
|
+
from .domain import MetricValue, ValueStatus
|
|
14
|
+
from .errors import EvaluationError, MissingOptionalDependencyError
|
|
15
|
+
|
|
16
|
+
MetricScorer = Callable[[Any, Any], float | None | Awaitable[float | None]]
|
|
17
|
+
RagasSampleMapper = Callable[[Any, Any], Mapping[str, Any]]
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class MetricEvaluator:
|
|
21
|
+
"""Adapt one synchronous or asynchronous numeric scorer to an evaluator."""
|
|
22
|
+
|
|
23
|
+
def __init__(
|
|
24
|
+
self,
|
|
25
|
+
name: str,
|
|
26
|
+
scorer: MetricScorer,
|
|
27
|
+
*,
|
|
28
|
+
status: ValueStatus = ValueStatus.EXACT,
|
|
29
|
+
) -> None:
|
|
30
|
+
if not name.strip():
|
|
31
|
+
raise ValueError("metric name must not be empty")
|
|
32
|
+
if name == "tunarag.objective":
|
|
33
|
+
raise ValueError("metric name is reserved: tunarag.objective")
|
|
34
|
+
if not callable(scorer):
|
|
35
|
+
raise TypeError("metric scorer must be callable")
|
|
36
|
+
self._name = name
|
|
37
|
+
self._scorer = scorer
|
|
38
|
+
self._status = status
|
|
39
|
+
|
|
40
|
+
async def evaluate(self, output: Any, example: Any) -> Sequence[MetricValue]:
|
|
41
|
+
"""Score one adapter output and evaluation example."""
|
|
42
|
+
|
|
43
|
+
result = self._scorer(output, example)
|
|
44
|
+
value = await result if inspect.isawaitable(result) else result
|
|
45
|
+
normalized = _numeric_score(value, metric=self._name)
|
|
46
|
+
if normalized is None:
|
|
47
|
+
return (MetricValue(self._name, None, ValueStatus.UNAVAILABLE, 0.0),)
|
|
48
|
+
return (MetricValue(self._name, normalized, self._status),)
|
|
49
|
+
|
|
50
|
+
def fingerprint(self) -> Any:
|
|
51
|
+
"""Return stable configuration used by study and cache identity."""
|
|
52
|
+
|
|
53
|
+
return {
|
|
54
|
+
"name": self._name,
|
|
55
|
+
"status": self._status.value,
|
|
56
|
+
"scorer": _callable_identity(self._scorer),
|
|
57
|
+
}
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
class RagasEvaluator:
|
|
61
|
+
"""Evaluate one sample with optional RAGAS metrics without hard coupling."""
|
|
62
|
+
|
|
63
|
+
def __init__(
|
|
64
|
+
self,
|
|
65
|
+
metrics: Sequence[Any],
|
|
66
|
+
sample_mapper: RagasSampleMapper,
|
|
67
|
+
*,
|
|
68
|
+
metric_prefix: str = "ragas.",
|
|
69
|
+
llm: Any = None,
|
|
70
|
+
embeddings: Any = None,
|
|
71
|
+
raise_exceptions: bool = False,
|
|
72
|
+
) -> None:
|
|
73
|
+
if not metrics:
|
|
74
|
+
raise ValueError("RAGAS evaluator requires at least one metric")
|
|
75
|
+
if not callable(sample_mapper):
|
|
76
|
+
raise TypeError("RAGAS sample mapper must be callable")
|
|
77
|
+
if metric_prefix == "tunarag.objective" or metric_prefix.startswith("tunarag.objective."):
|
|
78
|
+
raise ValueError("RAGAS metric prefix uses a reserved name")
|
|
79
|
+
self._metrics = tuple(metrics)
|
|
80
|
+
self._sample_mapper = sample_mapper
|
|
81
|
+
self._metric_prefix = metric_prefix
|
|
82
|
+
self._llm = llm
|
|
83
|
+
self._embeddings = embeddings
|
|
84
|
+
self._raise_exceptions = raise_exceptions
|
|
85
|
+
|
|
86
|
+
async def evaluate(self, output: Any, example: Any) -> Sequence[MetricValue]:
|
|
87
|
+
"""Map and evaluate exactly one sample through RAGAS's async API."""
|
|
88
|
+
|
|
89
|
+
try:
|
|
90
|
+
sample_data = self._sample_mapper(output, example)
|
|
91
|
+
except Exception as error:
|
|
92
|
+
raise EvaluationError(
|
|
93
|
+
"RAGAS sample mapping failed",
|
|
94
|
+
stage="evaluation",
|
|
95
|
+
details={"cause": type(error).__name__},
|
|
96
|
+
) from error
|
|
97
|
+
if not isinstance(sample_data, Mapping):
|
|
98
|
+
raise EvaluationError("RAGAS sample mapper must return a mapping", stage="evaluation")
|
|
99
|
+
ragas, dataset_schema = _load_ragas()
|
|
100
|
+
try:
|
|
101
|
+
sample = dataset_schema.SingleTurnSample(**dict(sample_data))
|
|
102
|
+
dataset = ragas.EvaluationDataset(samples=[sample])
|
|
103
|
+
result = await ragas.aevaluate(
|
|
104
|
+
dataset=dataset,
|
|
105
|
+
metrics=list(self._metrics),
|
|
106
|
+
llm=self._llm,
|
|
107
|
+
embeddings=self._embeddings,
|
|
108
|
+
raise_exceptions=self._raise_exceptions,
|
|
109
|
+
show_progress=False,
|
|
110
|
+
)
|
|
111
|
+
scores = result.scores
|
|
112
|
+
except Exception as error:
|
|
113
|
+
raise EvaluationError(
|
|
114
|
+
"RAGAS evaluation failed",
|
|
115
|
+
stage="evaluation",
|
|
116
|
+
details={"cause": type(error).__name__},
|
|
117
|
+
) from error
|
|
118
|
+
if not isinstance(scores, Sequence) or len(scores) != 1:
|
|
119
|
+
raise EvaluationError("RAGAS returned an invalid per-sample result", stage="evaluation")
|
|
120
|
+
row = scores[0]
|
|
121
|
+
if not isinstance(row, Mapping) or not row:
|
|
122
|
+
raise EvaluationError("RAGAS returned no metric scores", stage="evaluation")
|
|
123
|
+
observations: list[MetricValue] = []
|
|
124
|
+
names: set[str] = set()
|
|
125
|
+
for raw_name, raw_value in row.items():
|
|
126
|
+
name = f"{self._metric_prefix}{raw_name}"
|
|
127
|
+
if not name.strip() or name == "tunarag.objective":
|
|
128
|
+
raise EvaluationError("RAGAS returned an invalid metric name", stage="evaluation")
|
|
129
|
+
if name in names:
|
|
130
|
+
raise EvaluationError("RAGAS returned duplicate metric names", stage="evaluation")
|
|
131
|
+
names.add(name)
|
|
132
|
+
value = _numeric_score(raw_value, metric=name)
|
|
133
|
+
if value is None:
|
|
134
|
+
observations.append(MetricValue(name, None, ValueStatus.UNAVAILABLE, 0.0))
|
|
135
|
+
else:
|
|
136
|
+
observations.append(MetricValue(name, value))
|
|
137
|
+
return tuple(observations)
|
|
138
|
+
|
|
139
|
+
def fingerprint(self) -> Any:
|
|
140
|
+
"""Return secret-safe semantic integration configuration."""
|
|
141
|
+
|
|
142
|
+
return {
|
|
143
|
+
"metrics": [_component_identity(metric) for metric in self._metrics],
|
|
144
|
+
"sample_mapper": _callable_identity(self._sample_mapper),
|
|
145
|
+
"metric_prefix": self._metric_prefix,
|
|
146
|
+
"llm": _component_identity(self._llm),
|
|
147
|
+
"embeddings": _component_identity(self._embeddings),
|
|
148
|
+
"raise_exceptions": self._raise_exceptions,
|
|
149
|
+
}
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
def _load_ragas() -> tuple[Any, Any]:
|
|
153
|
+
try:
|
|
154
|
+
return import_module("ragas"), import_module("ragas.dataset_schema")
|
|
155
|
+
except ImportError as error:
|
|
156
|
+
raise MissingOptionalDependencyError("RAGAS evaluation", "ragas") from error
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
def _numeric_score(value: Any, *, metric: str) -> float | None:
|
|
160
|
+
if value is None:
|
|
161
|
+
return None
|
|
162
|
+
if isinstance(value, bool) or not isinstance(value, Real):
|
|
163
|
+
raise EvaluationError(
|
|
164
|
+
"evaluator score must be numeric or None",
|
|
165
|
+
stage="evaluation",
|
|
166
|
+
details={"metric": metric, "value_type": type(value).__name__},
|
|
167
|
+
)
|
|
168
|
+
normalized = float(value)
|
|
169
|
+
return normalized if math.isfinite(normalized) else None
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
def _callable_identity(function: Callable[..., Any]) -> str:
|
|
173
|
+
module = getattr(function, "__module__", type(function).__module__)
|
|
174
|
+
name = getattr(function, "__qualname__", type(function).__qualname__)
|
|
175
|
+
return f"{module}.{name}"
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
def _component_identity(component: Any) -> Any:
|
|
179
|
+
if component is None:
|
|
180
|
+
return None
|
|
181
|
+
identity: dict[str, Any] = {
|
|
182
|
+
"type": f"{type(component).__module__}.{type(component).__qualname__}"
|
|
183
|
+
}
|
|
184
|
+
if isinstance(component, Fingerprintable):
|
|
185
|
+
identity["fingerprint"] = component.fingerprint()
|
|
186
|
+
return identity
|
|
@@ -0,0 +1,25 @@
|
|
|
1
|
+
"""Optional integration adapters."""
|
|
2
|
+
|
|
3
|
+
from .mlflow import MLflowCallback
|
|
4
|
+
from .runnables import (
|
|
5
|
+
AsyncRunnable,
|
|
6
|
+
LangChainAdapter,
|
|
7
|
+
LangGraphAdapter,
|
|
8
|
+
RunnableConfigMapper,
|
|
9
|
+
RunnableFactory,
|
|
10
|
+
RunnableInputMapper,
|
|
11
|
+
RunnableOutputMapper,
|
|
12
|
+
RunnableUsageMapper,
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
__all__ = [
|
|
16
|
+
"AsyncRunnable",
|
|
17
|
+
"LangChainAdapter",
|
|
18
|
+
"LangGraphAdapter",
|
|
19
|
+
"MLflowCallback",
|
|
20
|
+
"RunnableConfigMapper",
|
|
21
|
+
"RunnableFactory",
|
|
22
|
+
"RunnableInputMapper",
|
|
23
|
+
"RunnableOutputMapper",
|
|
24
|
+
"RunnableUsageMapper",
|
|
25
|
+
]
|
|
@@ -0,0 +1,155 @@
|
|
|
1
|
+
"""Optional post-commit MLflow tracking mirror."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import hashlib
|
|
7
|
+
import json
|
|
8
|
+
import re
|
|
9
|
+
from collections.abc import Mapping
|
|
10
|
+
from dataclasses import dataclass
|
|
11
|
+
from importlib import import_module
|
|
12
|
+
from typing import Any, Protocol
|
|
13
|
+
|
|
14
|
+
from ..errors import MissingOptionalDependencyError, redact
|
|
15
|
+
from ..store import SQLiteStore, StudyEvent, TrialStatus
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class MLflowClient(Protocol):
|
|
19
|
+
"""Subset of the low-level MLflow client used by the mirror."""
|
|
20
|
+
|
|
21
|
+
def create_run(self, experiment_id: str, *, tags: Mapping[str, str]) -> Any: ...
|
|
22
|
+
|
|
23
|
+
def log_metric(
|
|
24
|
+
self, run_id: str, key: str, value: float, *, step: int | None = None
|
|
25
|
+
) -> Any: ...
|
|
26
|
+
|
|
27
|
+
def log_param(self, run_id: str, key: str, value: Any) -> Any: ...
|
|
28
|
+
|
|
29
|
+
def set_tag(self, run_id: str, key: str, value: Any) -> None: ...
|
|
30
|
+
|
|
31
|
+
def set_terminated(self, run_id: str, *, status: str | None = None) -> None: ...
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
@dataclass(frozen=True, slots=True)
|
|
35
|
+
class _TrialMirrorData:
|
|
36
|
+
sequence: int
|
|
37
|
+
status: TrialStatus
|
|
38
|
+
candidate: Mapping[str, Any]
|
|
39
|
+
metrics: Mapping[str, float]
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class MLflowCallback:
|
|
43
|
+
"""Mirror committed study events into one low-level MLflow run."""
|
|
44
|
+
|
|
45
|
+
def __init__(
|
|
46
|
+
self,
|
|
47
|
+
store: SQLiteStore,
|
|
48
|
+
*,
|
|
49
|
+
experiment_id: str = "0",
|
|
50
|
+
tracking_uri: str | None = None,
|
|
51
|
+
tags: Mapping[str, str] | None = None,
|
|
52
|
+
client: MLflowClient | None = None,
|
|
53
|
+
) -> None:
|
|
54
|
+
if not experiment_id.strip():
|
|
55
|
+
raise ValueError("MLflow experiment id must not be empty")
|
|
56
|
+
self._store = store
|
|
57
|
+
self._experiment_id = experiment_id
|
|
58
|
+
self._tracking_uri = tracking_uri
|
|
59
|
+
safe_tags = redact(tags or {})
|
|
60
|
+
if not isinstance(safe_tags, Mapping):
|
|
61
|
+
raise TypeError("MLflow tags must be a mapping")
|
|
62
|
+
self._tags = {str(key): str(value) for key, value in safe_tags.items()}
|
|
63
|
+
self._client = client
|
|
64
|
+
self._run_ids: dict[str, str] = {}
|
|
65
|
+
|
|
66
|
+
async def on_event(self, event: StudyEvent) -> None:
|
|
67
|
+
"""Mirror one durable event after gathering committed local state."""
|
|
68
|
+
|
|
69
|
+
trial_data = self._trial_data(event)
|
|
70
|
+
await asyncio.to_thread(self._mirror, event, trial_data)
|
|
71
|
+
|
|
72
|
+
def _trial_data(self, event: StudyEvent) -> _TrialMirrorData | None:
|
|
73
|
+
if event.type not in {"trial.succeeded", "trial.failed", "trial.cancelled"}:
|
|
74
|
+
return None
|
|
75
|
+
trial_id = event.payload.get("trial_id")
|
|
76
|
+
if not isinstance(trial_id, str):
|
|
77
|
+
return None
|
|
78
|
+
trial = self._store.get_trial(trial_id)
|
|
79
|
+
candidate = json.loads(trial.candidate_json)
|
|
80
|
+
if not isinstance(candidate, Mapping):
|
|
81
|
+
candidate = {}
|
|
82
|
+
safe_candidate = redact(candidate)
|
|
83
|
+
if not isinstance(safe_candidate, Mapping):
|
|
84
|
+
safe_candidate = {}
|
|
85
|
+
metrics = {
|
|
86
|
+
metric.name: float(metric.value)
|
|
87
|
+
for metric in self._store.list_metrics(trial_id)
|
|
88
|
+
if metric.value is not None
|
|
89
|
+
}
|
|
90
|
+
return _TrialMirrorData(trial.sequence, trial.status, safe_candidate, metrics)
|
|
91
|
+
|
|
92
|
+
def _mirror(self, event: StudyEvent, trial: _TrialMirrorData | None) -> None:
|
|
93
|
+
client = self._client_or_import()
|
|
94
|
+
run_id = self._run_ids.get(event.study_id)
|
|
95
|
+
if run_id is None:
|
|
96
|
+
tags = {
|
|
97
|
+
**self._tags,
|
|
98
|
+
"tunarag.study_id": event.study_id,
|
|
99
|
+
"tunarag.mirror": "event-callback",
|
|
100
|
+
}
|
|
101
|
+
run = client.create_run(self._experiment_id, tags=tags)
|
|
102
|
+
run_id = str(run.info.run_id)
|
|
103
|
+
self._run_ids[event.study_id] = run_id
|
|
104
|
+
client.set_tag(run_id, "tunarag.last_event", event.type)
|
|
105
|
+
client.set_tag(run_id, "tunarag.last_event_sequence", event.sequence)
|
|
106
|
+
if trial is not None:
|
|
107
|
+
self._mirror_trial(client, run_id, trial)
|
|
108
|
+
terminal_status = {
|
|
109
|
+
"study.completed": "FINISHED",
|
|
110
|
+
"study.failed": "FAILED",
|
|
111
|
+
"study.cancelled": "KILLED",
|
|
112
|
+
}.get(event.type)
|
|
113
|
+
if terminal_status is not None:
|
|
114
|
+
client.set_terminated(run_id, status=terminal_status)
|
|
115
|
+
|
|
116
|
+
def _mirror_trial(self, client: MLflowClient, run_id: str, trial: _TrialMirrorData) -> None:
|
|
117
|
+
prefix = f"trial.{trial.sequence}"
|
|
118
|
+
client.set_tag(run_id, f"{prefix}.status", trial.status.value)
|
|
119
|
+
for key, value in sorted(trial.candidate.items()):
|
|
120
|
+
client.log_param(
|
|
121
|
+
run_id, _tracking_key(f"{prefix}.candidate.{key}"), _parameter_value(value)
|
|
122
|
+
)
|
|
123
|
+
for name, value in sorted(trial.metrics.items()):
|
|
124
|
+
client.log_metric(run_id, _tracking_key(name), value, step=trial.sequence)
|
|
125
|
+
client.log_metric(
|
|
126
|
+
run_id,
|
|
127
|
+
"tunarag.trial_failed",
|
|
128
|
+
1.0 if trial.status is TrialStatus.FAILED else 0.0,
|
|
129
|
+
step=trial.sequence,
|
|
130
|
+
)
|
|
131
|
+
|
|
132
|
+
def _client_or_import(self) -> MLflowClient:
|
|
133
|
+
if self._client is not None:
|
|
134
|
+
return self._client
|
|
135
|
+
try:
|
|
136
|
+
module = import_module("mlflow")
|
|
137
|
+
except ImportError as error:
|
|
138
|
+
raise MissingOptionalDependencyError("MLflow tracking", "mlflow") from error
|
|
139
|
+
self._client = module.MlflowClient(tracking_uri=self._tracking_uri)
|
|
140
|
+
return self._client
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
def _parameter_value(value: Any) -> str:
|
|
144
|
+
safe = redact(value)
|
|
145
|
+
if safe is None or isinstance(safe, (bool, int, float, str)):
|
|
146
|
+
return str(safe)
|
|
147
|
+
return json.dumps(safe, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
def _tracking_key(value: str) -> str:
|
|
151
|
+
normalized = re.sub(r"[^A-Za-z0-9_.\-/ ]", "_", value)
|
|
152
|
+
if normalized == value and len(normalized) <= 250:
|
|
153
|
+
return normalized
|
|
154
|
+
digest = hashlib.sha256(value.encode("utf-8")).hexdigest()[:12]
|
|
155
|
+
return f"{normalized[:237]}.{digest}"
|