Skip to content

Commit 697c9b1

Browse files
committed
Consolidates correctness and groundeness evaluation requests
1 parent 1ccb231 commit 697c9b1

9 files changed

Lines changed: 259 additions & 60 deletions

File tree

src/elastic_evals/evaluators/__init__.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66

77
from .correctness import (
88
create_correctness_analysis_evaluator,
9+
create_correctness_evaluators,
910
create_quantitative_correctness_evaluators,
1011
)
1112
from .criteria import (
@@ -15,6 +16,7 @@
1516
)
1617
from .groundedness import (
1718
create_groundedness_analysis_evaluator,
19+
create_groundedness_evaluators,
1820
create_quantitative_groundedness_evaluator,
1921
)
2022
from .input_tokens import create_input_tokens_evaluator
@@ -33,8 +35,10 @@
3335
"KibanaEvaluatorConfig",
3436
"KibanaSubScore",
3537
"create_correctness_analysis_evaluator",
38+
"create_correctness_evaluators",
3639
"create_criteria_evaluator",
3740
"create_groundedness_analysis_evaluator",
41+
"create_groundedness_evaluators",
3842
"create_input_tokens_evaluator",
3943
"create_latency_evaluator",
4044
"create_output_tokens_evaluator",

src/elastic_evals/evaluators/correctness/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,10 +6,12 @@
66

77
from .evaluator import (
88
create_correctness_analysis_evaluator,
9+
create_correctness_evaluators,
910
create_quantitative_correctness_evaluators,
1011
)
1112

1213
__all__ = [
1314
"create_correctness_analysis_evaluator",
15+
"create_correctness_evaluators",
1416
"create_quantitative_correctness_evaluators",
1517
]

src/elastic_evals/evaluators/correctness/evaluator.py

Lines changed: 43 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -47,39 +47,61 @@ def create_quantitative_correctness_evaluators(
4747
instrumentation_profile: InstrumentationProfile = "elastic-inference",
4848
log: logging.Logger | None = None,
4949
) -> list[Evaluator]:
50-
return kibana_evaluators(
50+
return create_correctness_evaluators(
51+
client=client,
52+
connector_id=connector_id,
53+
instrumentation_profile=instrumentation_profile,
54+
log=log,
55+
)[1:]
56+
57+
58+
def create_correctness_evaluators(
59+
*,
60+
client: KibanaEvaluatorsClient,
61+
connector_id: str,
62+
instrumentation_profile: InstrumentationProfile = "elastic-inference",
63+
log: logging.Logger | None = None,
64+
) -> list[Evaluator]:
65+
quantitative = kibana_evaluators(
5166
[_config(connector_id)],
5267
client=client,
5368
instrumentation_profile=instrumentation_profile,
5469
log=log,
5570
)
71+
return [_analysis_evaluator(quantitative[0]), *quantitative]
5672

5773

5874
def create_correctness_analysis_evaluator(
5975
*,
60-
inference_client: KibanaInferenceClient,
76+
inference_client: KibanaInferenceClient | None = None,
77+
client: KibanaEvaluatorsClient | None = None,
78+
connector_id: str | None = None,
6179
log: logging.Logger,
6280
instrumentation_profile: InstrumentationProfile = "elastic-inference",
6381
) -> Evaluator:
64-
client = KibanaEvaluatorsClient(
65-
kibana_url=inference_client.kibana_url,
66-
api_key=inference_client.api_key,
67-
timeout=inference_client.timeout,
68-
)
69-
factuality = kibana_evaluators(
70-
[
71-
KibanaEvaluatorConfig(
72-
name="correctness",
73-
kind="LLM",
74-
connector_id=inference_client.connector_id,
75-
sub_scores=(KibanaSubScore(key="factuality", evaluator_name="correctness"),),
76-
)
77-
],
82+
if client is None:
83+
if inference_client is None:
84+
raise ValueError("client and connector_id are required")
85+
client = KibanaEvaluatorsClient(
86+
kibana_url=inference_client.kibana_url,
87+
api_key=inference_client.api_key,
88+
timeout=inference_client.timeout,
89+
)
90+
connector_id = inference_client.connector_id
91+
elif inference_client is not None:
92+
raise ValueError("Pass either client or inference_client, not both")
93+
elif not connector_id:
94+
raise ValueError("connector_id is required")
95+
96+
return create_correctness_evaluators(
7897
client=client,
98+
connector_id=connector_id,
7999
instrumentation_profile=instrumentation_profile,
80100
log=log,
81101
)[0]
82102

103+
104+
def _analysis_evaluator(factuality: Evaluator) -> Evaluator:
83105
async def evaluate(params: EvaluatorParams) -> EvaluationResult:
84106
result = await factuality.evaluate(params)
85107
if result.label in {"error", "unavailable"}:
@@ -94,7 +116,11 @@ async def evaluate(params: EvaluatorParams) -> EvaluationResult:
94116
metadata=result.metadata,
95117
)
96118

97-
return SimpleEvaluator(name=QUALITATIVE_EVALUATOR_NAME, kind="LLM", evaluate=evaluate)
119+
return SimpleEvaluator(
120+
name=QUALITATIVE_EVALUATOR_NAME,
121+
kind="LLM",
122+
evaluate=evaluate,
123+
)
98124

99125

100126
def _analysis_explanation(summary: Any, fallback: str | None) -> str | None:

src/elastic_evals/evaluators/groundedness/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,10 +6,12 @@
66

77
from .evaluator import (
88
create_groundedness_analysis_evaluator,
9+
create_groundedness_evaluators,
910
create_quantitative_groundedness_evaluator,
1011
)
1112

1213
__all__ = [
1314
"create_groundedness_analysis_evaluator",
15+
"create_groundedness_evaluators",
1416
"create_quantitative_groundedness_evaluator",
1517
]

src/elastic_evals/evaluators/groundedness/evaluator.py

Lines changed: 39 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,22 @@ def create_quantitative_groundedness_evaluator(
2929
instrumentation_profile: InstrumentationProfile = "elastic-inference",
3030
log: logging.Logger | None = None,
3131
) -> Evaluator:
32-
return kibana_evaluators(
32+
return create_groundedness_evaluators(
33+
client=client,
34+
connector_id=connector_id,
35+
instrumentation_profile=instrumentation_profile,
36+
log=log,
37+
)[1]
38+
39+
40+
def create_groundedness_evaluators(
41+
*,
42+
client: KibanaEvaluatorsClient,
43+
connector_id: str,
44+
instrumentation_profile: InstrumentationProfile = "elastic-inference",
45+
log: logging.Logger | None = None,
46+
) -> list[Evaluator]:
47+
quantitative = kibana_evaluators(
3348
[
3449
KibanaEvaluatorConfig(
3550
name="groundedness",
@@ -47,26 +62,40 @@ def create_quantitative_groundedness_evaluator(
4762
instrumentation_profile=instrumentation_profile,
4863
log=log,
4964
)[0]
65+
return [_analysis_evaluator(quantitative), quantitative]
5066

5167

5268
def create_groundedness_analysis_evaluator(
5369
*,
54-
inference_client: KibanaInferenceClient,
70+
inference_client: KibanaInferenceClient | None = None,
71+
client: KibanaEvaluatorsClient | None = None,
72+
connector_id: str | None = None,
5573
log: logging.Logger,
5674
instrumentation_profile: InstrumentationProfile = "elastic-inference",
5775
) -> Evaluator:
58-
client = KibanaEvaluatorsClient(
59-
kibana_url=inference_client.kibana_url,
60-
api_key=inference_client.api_key,
61-
timeout=inference_client.timeout,
62-
)
63-
groundedness = create_quantitative_groundedness_evaluator(
76+
if client is None:
77+
if inference_client is None:
78+
raise ValueError("client and connector_id are required")
79+
client = KibanaEvaluatorsClient(
80+
kibana_url=inference_client.kibana_url,
81+
api_key=inference_client.api_key,
82+
timeout=inference_client.timeout,
83+
)
84+
connector_id = inference_client.connector_id
85+
elif inference_client is not None:
86+
raise ValueError("Pass either client or inference_client, not both")
87+
elif not connector_id:
88+
raise ValueError("connector_id is required")
89+
90+
return create_groundedness_evaluators(
6491
client=client,
65-
connector_id=inference_client.connector_id,
92+
connector_id=connector_id,
6693
instrumentation_profile=instrumentation_profile,
6794
log=log,
68-
)
95+
)[0]
96+
6997

98+
def _analysis_evaluator(groundedness: Evaluator) -> Evaluator:
7099
async def evaluate(params: EvaluatorParams) -> EvaluationResult:
71100
result = await groundedness.evaluate(params)
72101
if result.label in {"error", "unavailable"}:

src/elastic_evals/evaluators/kibana.py

Lines changed: 28 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88

99
import asyncio
1010
import logging
11+
import weakref
1112
from dataclasses import dataclass
1213
from typing import Any, Literal, Sequence
1314

@@ -49,6 +50,12 @@ class _ScoreSelector:
4950
evaluator_name: str
5051

5152

53+
@dataclass(frozen=True)
54+
class _CachedEvaluation:
55+
request_key: str
56+
task: asyncio.Task[EvaluateResponse]
57+
58+
5259
class _EvaluationBatch:
5360
def __init__(
5461
self,
@@ -60,21 +67,35 @@ def __init__(
6067
self._client = client
6168
self._configs = tuple(configs)
6269
self._instrumentation_profile = instrumentation_profile
63-
self._evaluations: dict[str, asyncio.Task[EvaluateResponse]] = {}
70+
self._evaluations: weakref.WeakKeyDictionary[object, _CachedEvaluation] = weakref.WeakKeyDictionary()
6471
self._lock = asyncio.Lock()
6572

6673
async def evaluate(self, params: EvaluatorParams) -> EvaluateResponse:
6774
trace_id = _trace_id(params)
6875
if not trace_id:
6976
raise ValueError("A trace ID is required for Kibana evaluators")
7077

78+
request = self._build_request(params, trace_id)
79+
request_key = request.model_dump_json(exclude_none=True)
80+
scope = params._evaluation_scope
81+
7182
async with self._lock:
72-
evaluation = self._evaluations.get(trace_id)
73-
if evaluation is None:
74-
evaluation = asyncio.create_task(self._client.evaluate(self._build_request(params, trace_id)))
75-
self._evaluations[trace_id] = evaluation
83+
cached = self._evaluations.get(scope)
84+
if cached is None or cached.request_key != request_key:
85+
cached = _CachedEvaluation(
86+
request_key=request_key,
87+
task=asyncio.create_task(self._client.evaluate(request)),
88+
)
89+
self._evaluations[scope] = cached
7690

77-
return await asyncio.shield(evaluation)
91+
try:
92+
return await asyncio.shield(cached.task)
93+
except Exception:
94+
async with self._lock:
95+
current = self._evaluations.get(scope)
96+
if current is cached:
97+
del self._evaluations[scope]
98+
raise
7899

79100
def _build_request(self, params: EvaluatorParams, trace_id: str) -> EvaluateRequest:
80101
reference_data: dict[str, Any] | None
@@ -189,7 +210,7 @@ def kibana_evaluators(
189210
instrumentation_profile: InstrumentationProfile = "elastic-inference",
190211
log: logging.Logger | None = None,
191212
) -> list[Evaluator]:
192-
"""Create Python evaluators backed by one Kibana request per trace."""
213+
"""Create Python evaluators backed by one Kibana request per evaluation."""
193214
_validate_configs(configs)
194215
immutable_configs = tuple(configs)
195216
selectors = _build_selectors(immutable_configs)

src/elastic_evals/executor/client.py

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -163,15 +163,16 @@ async def task_runner() -> TaskOutput:
163163

164164
log_evaluation_start(example_index, repetition, len(evaluators))
165165

166+
params = EvaluatorParams(
167+
input=example.input,
168+
output=task_output,
169+
expected=example.output,
170+
metadata=example.metadata,
171+
trace_id=task_trace_id,
172+
)
173+
166174
for evaluator in evaluators:
167175
log_evaluator_start(evaluator.name, example_index, repetition)
168-
params = EvaluatorParams(
169-
input=example.input,
170-
output=task_output,
171-
expected=example.output,
172-
metadata=example.metadata,
173-
trace_id=task_trace_id,
174-
)
175176

176177
async def evaluator_runner() -> Any:
177178
return await evaluator.evaluate(params)

src/elastic_evals/types.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@
66

77
from __future__ import annotations
88

9-
from dataclasses import dataclass
9+
from dataclasses import dataclass, field
1010
from typing import Any, Generic, Literal, Protocol, TypeAlias, TypeVar
1111

1212
from pydantic import BaseModel
@@ -20,6 +20,10 @@
2020
TaskOutput: TypeAlias = Any
2121

2222

23+
class _EvaluationScope:
24+
pass
25+
26+
2327
class Example(BaseModel, Generic[TInput, TExpected, TMetadata]):
2428
input: TInput
2529
output: TExpected | None = None
@@ -64,6 +68,7 @@ class EvaluatorParams(Generic[TInput, TExpected, TMetadata, TTaskOutput]):
6468
expected: TExpected | None
6569
metadata: TMetadata
6670
trace_id: str | None = None
71+
_evaluation_scope: object = field(default_factory=_EvaluationScope, compare=False, repr=False)
6772

6873

6974
class Evaluator(Protocol, Generic[TInput, TExpected, TMetadata, TTaskOutput]):

0 commit comments

Comments
 (0)