Skip to content

Commit 79a879f

Browse files
ItsMactocopybara-github
authored andcommitted
fix: register custom metrics from eval config in LocalEvalSampler
Merge #6424 Closes #6177 PiperOrigin-RevId: 951011003
1 parent 7c5008f commit 79a879f

6 files changed

Lines changed: 211 additions & 36 deletions

File tree

src/google/adk/cli/cli_eval.py

Lines changed: 0 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -36,9 +36,6 @@
3636
from ..evaluation.eval_case import get_all_tool_calls
3737
from ..evaluation.eval_case import IntermediateDataType
3838
from ..evaluation.eval_metrics import EvalMetric
39-
from ..evaluation.eval_metrics import Interval
40-
from ..evaluation.eval_metrics import MetricInfo
41-
from ..evaluation.eval_metrics import MetricValueInfo
4239
from ..evaluation.eval_result import EvalCaseResult
4340
from ..evaluation.eval_sets_manager import EvalSetsManager
4441
from ..utils.context_utils import Aclosing
@@ -77,19 +74,6 @@ def _get_agent_module(agent_module_file_path: str) -> ModuleType:
7774
return _import_from_path(module_name, file_path)
7875

7976

80-
def get_default_metric_info(
81-
metric_name: str, description: str = ""
82-
) -> MetricInfo:
83-
"""Returns a default MetricInfo for a metric."""
84-
return MetricInfo(
85-
metric_name=metric_name,
86-
description=description,
87-
metric_value_info=MetricValueInfo(
88-
interval=Interval(min_value=0.0, max_value=1.0)
89-
),
90-
)
91-
92-
9377
def get_root_agent(agent_module_file_path: str) -> Agent:
9478
"""Returns root agent given the agent module."""
9579
agent_module = _get_agent_module(agent_module_file_path)

src/google/adk/cli/cli_tools_click.py

Lines changed: 2 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -1005,7 +1005,6 @@ def cli_eval(
10051005

10061006
from ..evaluation.base_eval_service import InferenceConfig
10071007
from ..evaluation.base_eval_service import InferenceRequest
1008-
from ..evaluation.custom_metric_evaluator import _CustomMetricEvaluator
10091008
from ..evaluation.eval_config import get_eval_metrics_from_config
10101009
from ..evaluation.eval_config import get_evaluation_criteria_or_default
10111010
from ..evaluation.evaluator import EvalStatus
@@ -1014,11 +1013,10 @@ def cli_eval(
10141013
from ..evaluation.local_eval_set_results_manager import LocalEvalSetResultsManager
10151014
from ..evaluation.local_eval_sets_manager import load_eval_set_from_file
10161015
from ..evaluation.local_eval_sets_manager import LocalEvalSetsManager
1017-
from ..evaluation.metric_evaluator_registry import DEFAULT_METRIC_EVALUATOR_REGISTRY
1016+
from ..evaluation.metric_evaluator_registry import register_custom_metrics_from_config
10181017
from ..evaluation.simulation.user_simulator_provider import UserSimulatorProvider
10191018
from .cli_eval import _collect_eval_results
10201019
from .cli_eval import _collect_inferences
1021-
from .cli_eval import get_default_metric_info
10221020
from .cli_eval import get_root_agent
10231021
from .cli_eval import parse_and_get_evals_to_run
10241022
from .cli_eval import pretty_print_eval_result
@@ -1113,23 +1111,7 @@ def cli_eval(
11131111
)
11141112

11151113
try:
1116-
metric_evaluator_registry = DEFAULT_METRIC_EVALUATOR_REGISTRY
1117-
if eval_config.custom_metrics:
1118-
for (
1119-
metric_name,
1120-
config,
1121-
) in eval_config.custom_metrics.items():
1122-
if config.metric_info:
1123-
metric_info = config.metric_info.model_copy()
1124-
metric_info.metric_name = metric_name
1125-
else:
1126-
metric_info = get_default_metric_info(
1127-
metric_name=metric_name, description=config.description
1128-
)
1129-
1130-
metric_evaluator_registry.register_evaluator(
1131-
metric_info, _CustomMetricEvaluator
1132-
)
1114+
metric_evaluator_registry = register_custom_metrics_from_config(eval_config)
11331115

11341116
eval_service = LocalEvalService(
11351117
root_agent=root_agent,

src/google/adk/evaluation/metric_evaluator_registry.py

Lines changed: 52 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,12 +15,16 @@
1515
from __future__ import annotations
1616

1717
import logging
18+
from typing import Optional
1819

1920
from ..errors.not_found_error import NotFoundError
2021
from ..utils.feature_decorator import experimental
2122
from .custom_metric_evaluator import _CustomMetricEvaluator
23+
from .eval_config import EvalConfig
2224
from .eval_metrics import EvalMetric
25+
from .eval_metrics import Interval
2326
from .eval_metrics import MetricInfo
27+
from .eval_metrics import MetricValueInfo
2428
from .eval_metrics import PrebuiltMetrics
2529
from .evaluator import Evaluator
2630
from .final_response_match_v2 import FinalResponseMatchV2Evaluator
@@ -175,3 +179,51 @@ def _get_default_metric_evaluator_registry() -> MetricEvaluatorRegistry:
175179

176180

177181
DEFAULT_METRIC_EVALUATOR_REGISTRY = _get_default_metric_evaluator_registry()
182+
183+
184+
def _get_default_metric_info(
185+
metric_name: str, description: str = ""
186+
) -> MetricInfo:
187+
"""Returns a default MetricInfo for a metric."""
188+
return MetricInfo(
189+
metric_name=metric_name,
190+
description=description,
191+
metric_value_info=MetricValueInfo(
192+
interval=Interval(min_value=0.0, max_value=1.0)
193+
),
194+
)
195+
196+
197+
def register_custom_metrics_from_config(
198+
eval_config: EvalConfig,
199+
metric_evaluator_registry: Optional[MetricEvaluatorRegistry] = None,
200+
) -> MetricEvaluatorRegistry:
201+
"""Registers custom metrics declared in the given eval config.
202+
203+
Args:
204+
eval_config: The eval config whose custom_metrics entries should be
205+
registered. Entries without a metric_info get a default one with a
206+
[0.0, 1.0] value interval.
207+
metric_evaluator_registry: The registry to register the metrics in.
208+
Defaults to DEFAULT_METRIC_EVALUATOR_REGISTRY.
209+
210+
Returns:
211+
The registry the metrics were registered in.
212+
"""
213+
if metric_evaluator_registry is None:
214+
metric_evaluator_registry = DEFAULT_METRIC_EVALUATOR_REGISTRY
215+
if not eval_config.custom_metrics:
216+
return metric_evaluator_registry
217+
218+
for metric_name, config in eval_config.custom_metrics.items():
219+
if config.metric_info:
220+
metric_info = config.metric_info.model_copy()
221+
metric_info.metric_name = metric_name
222+
else:
223+
metric_info = _get_default_metric_info(
224+
metric_name=metric_name, description=config.description
225+
)
226+
metric_evaluator_registry.register_evaluator(
227+
metric_info, _CustomMetricEvaluator
228+
)
229+
return metric_evaluator_registry

src/google/adk/optimization/local_eval_sampler.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@
3838
from ..evaluation.eval_result import EvalCaseResult
3939
from ..evaluation.eval_sets_manager import EvalSetsManager
4040
from ..evaluation.local_eval_service import LocalEvalService
41+
from ..evaluation.metric_evaluator_registry import register_custom_metrics_from_config
4142
from ..evaluation.simulation.user_simulator_provider import UserSimulatorProvider
4243
from ..utils.context_utils import Aclosing
4344
from .data_types import UnstructuredSamplingResult
@@ -154,6 +155,9 @@ def __init__(
154155
):
155156
self._config = config
156157
self._eval_sets_manager = eval_sets_manager
158+
self._metric_evaluator_registry = register_custom_metrics_from_config(
159+
self._config.eval_config
160+
)
157161

158162
self._train_eval_set = self._config.train_eval_set
159163
self._train_eval_case_ids = (
@@ -236,6 +240,7 @@ async def _evaluate_agent(
236240
eval_service = LocalEvalService(
237241
root_agent=agent,
238242
eval_sets_manager=self._eval_sets_manager,
243+
metric_evaluator_registry=self._metric_evaluator_registry,
239244
user_simulator_provider=user_simulator_provider,
240245
)
241246

tests/unittests/evaluation/test_metric_evaluator_registry.py

Lines changed: 112 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,17 +14,23 @@
1414

1515
from __future__ import annotations
1616

17+
from google.adk.agents.common_configs import CodeConfig
1718
from google.adk.errors.not_found_error import NotFoundError
19+
from google.adk.evaluation.custom_metric_evaluator import _CustomMetricEvaluator
20+
from google.adk.evaluation.eval_config import CustomMetricConfig
21+
from google.adk.evaluation.eval_config import EvalConfig
1822
from google.adk.evaluation.eval_metrics import EvalMetric
1923
from google.adk.evaluation.eval_metrics import Interval
2024
from google.adk.evaluation.eval_metrics import MetricInfo
2125
from google.adk.evaluation.eval_metrics import MetricValueInfo
2226
from google.adk.evaluation.eval_metrics import PrebuiltMetrics
2327
from google.adk.evaluation.evaluator import Evaluator
28+
from google.adk.evaluation.metric_evaluator_registry import DEFAULT_METRIC_EVALUATOR_REGISTRY
2429
from google.adk.evaluation.metric_evaluator_registry import FinalResponseMatchV2EvaluatorMetricInfoProvider
2530
from google.adk.evaluation.metric_evaluator_registry import HallucinationsV1EvaluatorMetricInfoProvider
2631
from google.adk.evaluation.metric_evaluator_registry import MetricEvaluatorRegistry
2732
from google.adk.evaluation.metric_evaluator_registry import PerTurnUserSimulatorQualityV1MetricInfoProvider
33+
from google.adk.evaluation.metric_evaluator_registry import register_custom_metrics_from_config
2834
from google.adk.evaluation.metric_evaluator_registry import ResponseEvaluatorMetricInfoProvider
2935
from google.adk.evaluation.metric_evaluator_registry import RubricBasedFinalResponseQualityV1EvaluatorMetricInfoProvider
3036
from google.adk.evaluation.metric_evaluator_registry import RubricBasedMultiTurnTrajectoryMetricInfoProvider
@@ -120,6 +126,112 @@ def test_get_evaluator_not_found(self, registry):
120126
registry.get_evaluator(eval_metric)
121127

122128

129+
class TestRegisterCustomMetricsFromConfig:
130+
"""Test cases for register_custom_metrics_from_config."""
131+
132+
_CUSTOM_METRIC_NAME = "custom_metric_for_registry_test"
133+
134+
@pytest.fixture
135+
def registry(self):
136+
registry = MetricEvaluatorRegistry()
137+
yield registry
138+
# The registry dict is shared class-level state; remove what we added.
139+
registry._registry.pop(self._CUSTOM_METRIC_NAME, None)
140+
141+
def _registered_metric_info(self, registry, metric_name):
142+
return next(
143+
metric_info
144+
for metric_info in registry.get_registered_metrics()
145+
if metric_info.metric_name == metric_name
146+
)
147+
148+
def test_registers_custom_metric_with_provided_metric_info(self, registry):
149+
metric_info = MetricInfo(
150+
metric_name="name_to_be_overridden",
151+
description="Custom metric description",
152+
metric_value_info=MetricValueInfo(
153+
interval=Interval(min_value=0.0, max_value=5.0)
154+
),
155+
)
156+
eval_config = EvalConfig(
157+
custom_metrics={
158+
self._CUSTOM_METRIC_NAME: CustomMetricConfig(
159+
code_config=CodeConfig(name="math.sqrt"),
160+
metric_info=metric_info,
161+
)
162+
}
163+
)
164+
165+
result = register_custom_metrics_from_config(eval_config, registry)
166+
167+
assert result is registry
168+
registered_info = self._registered_metric_info(
169+
registry, self._CUSTOM_METRIC_NAME
170+
)
171+
assert registered_info.metric_value_info.interval.max_value == 5.0
172+
assert all(
173+
metric_info.metric_name != "name_to_be_overridden"
174+
for metric_info in registry.get_registered_metrics()
175+
)
176+
evaluator = registry.get_evaluator(
177+
EvalMetric(
178+
metric_name=self._CUSTOM_METRIC_NAME,
179+
threshold=0.5,
180+
custom_function_path="math.sqrt",
181+
)
182+
)
183+
assert isinstance(evaluator, _CustomMetricEvaluator)
184+
185+
def test_registers_custom_metric_with_default_metric_info(self, registry):
186+
eval_config = EvalConfig(
187+
custom_metrics={
188+
self._CUSTOM_METRIC_NAME: CustomMetricConfig(
189+
code_config=CodeConfig(name="math.sqrt"),
190+
description="A custom metric",
191+
)
192+
}
193+
)
194+
195+
register_custom_metrics_from_config(eval_config, registry)
196+
197+
registered_info = self._registered_metric_info(
198+
registry, self._CUSTOM_METRIC_NAME
199+
)
200+
assert registered_info.description == "A custom metric"
201+
assert registered_info.metric_value_info.interval.min_value == 0.0
202+
assert registered_info.metric_value_info.interval.max_value == 1.0
203+
204+
def test_no_custom_metrics_is_a_no_op(self, registry):
205+
registered_before = registry.get_registered_metrics()
206+
207+
result = register_custom_metrics_from_config(EvalConfig(), registry)
208+
209+
assert result is registry
210+
assert registry.get_registered_metrics() == registered_before
211+
212+
def test_defaults_to_the_default_registry(self):
213+
eval_config = EvalConfig(
214+
custom_metrics={
215+
self._CUSTOM_METRIC_NAME: CustomMetricConfig(
216+
code_config=CodeConfig(name="math.sqrt"),
217+
)
218+
}
219+
)
220+
221+
try:
222+
result = register_custom_metrics_from_config(eval_config)
223+
224+
assert result is DEFAULT_METRIC_EVALUATOR_REGISTRY
225+
registered_info = self._registered_metric_info(
226+
DEFAULT_METRIC_EVALUATOR_REGISTRY, self._CUSTOM_METRIC_NAME
227+
)
228+
assert registered_info.metric_name == self._CUSTOM_METRIC_NAME
229+
finally:
230+
DEFAULT_METRIC_EVALUATOR_REGISTRY._registry.pop(
231+
self._CUSTOM_METRIC_NAME, None
232+
)
233+
234+
123235
class TestMetricInfoProviders:
124236
"""Test cases for MetricInfoProviders."""
125237

tests/unittests/optimization/local_eval_sampler_test.py

Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,22 +14,26 @@
1414

1515
from __future__ import annotations
1616

17+
from google.adk.agents.common_configs import CodeConfig
1718
from google.adk.agents.llm_agent import Agent
1819
from google.adk.evaluation.base_eval_service import EvaluateConfig
1920
from google.adk.evaluation.base_eval_service import EvaluateRequest
2021
from google.adk.evaluation.base_eval_service import InferenceConfig
2122
from google.adk.evaluation.base_eval_service import InferenceRequest
2223
from google.adk.evaluation.base_eval_service import InferenceResult
24+
from google.adk.evaluation.custom_metric_evaluator import _CustomMetricEvaluator
2325
from google.adk.evaluation.eval_case import Invocation
2426
from google.adk.evaluation.eval_case import InvocationEvent
2527
from google.adk.evaluation.eval_case import InvocationEvents
28+
from google.adk.evaluation.eval_config import CustomMetricConfig
2629
from google.adk.evaluation.eval_config import EvalConfig
2730
from google.adk.evaluation.eval_config import EvalMetric
2831
from google.adk.evaluation.eval_metrics import EvalMetricResult
2932
from google.adk.evaluation.eval_metrics import EvalMetricResultPerInvocation
3033
from google.adk.evaluation.eval_metrics import EvalStatus
3134
from google.adk.evaluation.eval_result import EvalCaseResult
3235
from google.adk.evaluation.eval_sets_manager import EvalSetsManager
36+
from google.adk.evaluation.metric_evaluator_registry import DEFAULT_METRIC_EVALUATOR_REGISTRY
3337
from google.adk.optimization.local_eval_sampler import _log_eval_summary
3438
from google.adk.optimization.local_eval_sampler import extract_single_invocation_info
3539
from google.adk.optimization.local_eval_sampler import extract_tool_call_data
@@ -218,6 +222,42 @@ def mock_get_eval_case_ids(self, eval_set_id):
218222
assert getattr(interface, attr) == expected_value
219223

220224

225+
def test_init_registers_custom_metrics(mocker):
226+
mocker.patch.object(
227+
LocalEvalSampler,
228+
"_get_eval_case_ids",
229+
autospec=True,
230+
return_value=["t1"],
231+
)
232+
custom_metric_name = "custom_metric_for_sampler_test"
233+
config = LocalEvalSamplerConfig(
234+
eval_config=EvalConfig(
235+
custom_metrics={
236+
custom_metric_name: CustomMetricConfig(
237+
code_config=CodeConfig(name="math.sqrt")
238+
)
239+
}
240+
),
241+
app_name="test_app",
242+
train_eval_set="train_set",
243+
)
244+
245+
try:
246+
LocalEvalSampler(config, mocker.MagicMock(spec=EvalSetsManager))
247+
248+
evaluator = DEFAULT_METRIC_EVALUATOR_REGISTRY.get_evaluator(
249+
EvalMetric(
250+
metric_name=custom_metric_name,
251+
threshold=0.5,
252+
custom_function_path="math.sqrt",
253+
)
254+
)
255+
assert isinstance(evaluator, _CustomMetricEvaluator)
256+
finally:
257+
# The registry dict is shared class-level state; remove what we added.
258+
DEFAULT_METRIC_EVALUATOR_REGISTRY._registry.pop(custom_metric_name, None)
259+
260+
221261
@pytest.mark.asyncio
222262
async def test_evaluate_agent(mocker):
223263
# Mocking LocalEvalService and its methods

0 commit comments

Comments
 (0)