Skip to content

Commit 5fa1455

Browse files
vertex-sdk-botcopybara-github
authored andcommitted
feat: Add per-step judges and cross-region routing to evals
PiperOrigin-RevId: 991795392
1 parent 3647753 commit 5fa1455

8 files changed

Lines changed: 440 additions & 30 deletions

File tree

‎agentplatform/_genai/_evals_common.py‎

Lines changed: 21 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -2886,6 +2886,21 @@ def _resolve_dataset_inputs(
28862886
return processed_eval_dataset, num_response_candidates
28872887

28882888

2889+
def _prebuilt_evaluation_run_metric(
2890+
resolved_metric: types.Metric,
2891+
) -> types.EvaluationRunMetric:
2892+
"""Builds the evaluation run metric for a resolved RubricMetric."""
2893+
if resolved_metric.name in _evals_constant.SUPPORTED_PREDEFINED_METRICS:
2894+
metric_config = t.t_metrics([resolved_metric])[0]
2895+
else:
2896+
metric_config = {
2897+
"predefined_metric_spec": {"metric_spec_name": resolved_metric.name}
2898+
}
2899+
return types.EvaluationRunMetric(
2900+
metric=resolved_metric.name, metric_config=metric_config
2901+
)
2902+
2903+
28892904
def _resolve_evaluation_run_metrics(
28902905
metrics: Union[list[types.EvaluationRunMetric], list[types.Metric]], api_client: Any
28912906
) -> list[types.EvaluationRunMetric]:
@@ -2903,14 +2918,7 @@ def _resolve_evaluation_run_metrics(
29032918
resolved_metric = metric_instance.resolve(api_client=api_client)
29042919
if resolved_metric.name:
29052920
resolved_metrics_list.append(
2906-
types.EvaluationRunMetric(
2907-
metric=resolved_metric.name,
2908-
metric_config=types.UnifiedMetric(
2909-
predefined_metric_spec=genai_types.PredefinedMetricSpec(
2910-
metric_spec_name=resolved_metric.name,
2911-
)
2912-
),
2913-
)
2921+
_prebuilt_evaluation_run_metric(resolved_metric)
29142922
)
29152923
except Exception as e:
29162924
logger.error(
@@ -2944,14 +2952,7 @@ def _resolve_evaluation_run_metrics(
29442952
)
29452953
if resolved_metric.name:
29462954
resolved_metrics_list.append(
2947-
types.EvaluationRunMetric(
2948-
metric=resolved_metric.name,
2949-
metric_config=types.UnifiedMetric(
2950-
predefined_metric_spec=genai_types.PredefinedMetricSpec(
2951-
metric_spec_name=resolved_metric.name,
2952-
)
2953-
),
2954-
)
2955+
_prebuilt_evaluation_run_metric(resolved_metric)
29552956
)
29562957
else:
29572958
raise TypeError(
@@ -3020,6 +3021,7 @@ def _execute_evaluation( # type: ignore[no-untyped-def]
30203021
dest: Optional[str] = None,
30213022
location: Optional[str] = None,
30223023
evaluation_service_qps: Optional[float] = None,
3024+
allow_cross_region_model: Optional[bool] = None,
30233025
**kwargs,
30243026
) -> types.EvaluationResult:
30253027
"""Evaluates a dataset using the provided metrics.
@@ -3035,6 +3037,8 @@ def _execute_evaluation( # type: ignore[no-untyped-def]
30353037
evaluation_service_qps: The rate limit (queries per second) for calls
30363038
to the evaluation service. Defaults to 10. Increase this value if
30373039
your project has a higher EvaluateInstances API quota.
3040+
allow_cross_region_model: Opt-in flag to authorize cross-region
3041+
routing for judge models.
30383042
**kwargs: Extra arguments to pass to evaluation, such as `agent_info`.
30393043
30403044
Returns:
@@ -3117,6 +3121,7 @@ def _execute_evaluation( # type: ignore[no-untyped-def]
31173121
evaluation_result = _evals_metric_handlers.compute_metrics_and_aggregate(
31183122
evaluation_run_config,
31193123
evaluation_service_qps=evaluation_service_qps,
3124+
allow_cross_region_model=allow_cross_region_model,
31203125
)
31213126
t2 = time.perf_counter()
31223127
logger.info("Evaluation took: %f seconds", t2 - t1)

‎agentplatform/_genai/_evals_metric_handlers.py‎

Lines changed: 21 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -288,6 +288,7 @@ class MetricHandler(abc.ABC, Generic[T]):
288288
def __init__(self, module: "evals.Evals", metric: T):
289289
self.module = module
290290
self.metric: T = metric
291+
self.allow_cross_region_model: Optional[bool] = None
291292

292293
@property
293294
@abc.abstractmethod
@@ -761,6 +762,7 @@ def get_metric_result(
761762
lambda: self.module._evaluate_instances(
762763
metrics=[self.metric],
763764
instance=instance,
765+
allow_cross_region_model=self.allow_cross_region_model,
764766
),
765767
self.metric_name,
766768
)
@@ -981,15 +983,16 @@ def __init__(self, module: "evals.Evals", metric: types.Metric):
981983
raise ValueError(
982984
f"Metric '{self.metric.name}' is not a supported predefined metric."
983985
)
984-
if (
986+
if self.metric.name.startswith("multi_turn") and (
985987
self.metric.judge_model
986988
or self.metric.judge_model_generation_config
987989
or self.metric.judge_model_sampling_count
988990
):
989991
logger.warning(
990992
"Autorater config settings (judge_model, "
991993
"judge_model_generation_config, judge_model_sampling_count) "
992-
"are ignored for predefined metric '%s'.",
994+
"are ignored for multi-turn metric '%s'. Use "
995+
"judge_model_step_configs to set its judges.",
993996
self.metric.name,
994997
)
995998

@@ -1071,6 +1074,7 @@ def get_metric_result(
10711074
metrics=[self.metric],
10721075
instance=payload.get("instance"),
10731076
autorater_config=payload.get("autorater_config"),
1077+
allow_cross_region_model=self.allow_cross_region_model,
10741078
),
10751079
metric_name,
10761080
)
@@ -1318,6 +1322,7 @@ def get_metric_result(
13181322
metric_sources=[metric_source],
13191323
instance=payload.get("instance"),
13201324
autorater_config=payload.get("autorater_config"),
1325+
allow_cross_region_model=self.allow_cross_region_model,
13211326
),
13221327
metric_name,
13231328
)
@@ -1404,12 +1409,16 @@ def aggregate(
14041409

14051410

14061411
def get_handler_for_metric(
1407-
module: "evals.Evals", metric: types.Metric
1412+
module: "evals.Evals",
1413+
metric: types.Metric,
1414+
allow_cross_region_model: Optional[bool] = None,
14081415
) -> Union[MetricHandlerType, Any]:
14091416
"""Returns a metric handler for the given metric."""
14101417
for condition, handler_class in _METRIC_HANDLER_MAPPING:
14111418
if condition(metric): # type: ignore[no-untyped-call]
1412-
return handler_class(module=module, metric=metric)
1419+
handler = handler_class(module=module, metric=metric)
1420+
handler.allow_cross_region_model = allow_cross_region_model
1421+
return handler
14131422
raise ValueError(f"Unsupported metric: {metric.name}")
14141423

14151424

@@ -1548,6 +1557,7 @@ def _rate_limited_get_metric_result(
15481557
def compute_metrics_and_aggregate(
15491558
evaluation_run_config: EvaluationRunConfig,
15501559
evaluation_service_qps: Optional[float] = None,
1560+
allow_cross_region_model: Optional[bool] = None,
15511561
) -> types.EvaluationResult:
15521562
"""Computes metrics and aggregates them for a given evaluation run config.
15531563
@@ -1556,6 +1566,8 @@ def compute_metrics_and_aggregate(
15561566
evaluation_service_qps: Optional QPS limit for the evaluation service.
15571567
Defaults to _DEFAULT_EVAL_SERVICE_QPS (10). Users with higher
15581568
quotas can increase this value.
1569+
allow_cross_region_model: Opt-in flag to authorize cross-region
1570+
routing for judge models.
15591571
"""
15601572
metric_handlers = []
15611573
all_futures = []
@@ -1574,7 +1586,11 @@ def compute_metrics_and_aggregate(
15741586

15751587
for eval_metric in evaluation_run_config.metrics:
15761588
metric_handlers.append(
1577-
get_handler_for_metric(evaluation_run_config.evals_module, eval_metric)
1589+
get_handler_for_metric(
1590+
evaluation_run_config.evals_module,
1591+
eval_metric,
1592+
allow_cross_region_model=allow_cross_region_model,
1593+
)
15781594
)
15791595

15801596
eval_case_count = len(evaluation_run_config.dataset.eval_cases)

‎agentplatform/_genai/_evals_metric_loaders.py‎

Lines changed: 9 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -215,8 +215,11 @@ def resolve(self, api_client: Any) -> "types.Metric":
215215
if self._resolved_metric:
216216
return self._resolved_metric
217217

218+
# The shared cache is keyed by name and version only, so a metric with
219+
# overrides (such as judge_model_step_configs) must bypass it.
220+
use_cache = not self.metric_kwargs
218221
cache_key = f"{self.name}@{self.version or 'default'}"
219-
if cache_key in LazyLoadedPrebuiltMetric._cache:
222+
if use_cache and cache_key in LazyLoadedPrebuiltMetric._cache:
220223
self._resolved_metric = LazyLoadedPrebuiltMetric._cache[cache_key]
221224
logger.debug("Metric '%s' found in cache.", cache_key)
222225
return self._resolved_metric
@@ -225,7 +228,8 @@ def resolve(self, api_client: Any) -> "types.Metric":
225228
api_metric = self._resolve_api_predefined()
226229
if api_metric:
227230
self._resolved_metric = api_metric
228-
LazyLoadedPrebuiltMetric._cache[cache_key] = self._resolved_metric
231+
if use_cache:
232+
LazyLoadedPrebuiltMetric._cache[cache_key] = self._resolved_metric
229233
return self._resolved_metric
230234

231235
# Fallback to GCS loading for custom LLM-based Prebuilt Metrics
@@ -234,8 +238,9 @@ def resolve(self, api_client: Any) -> "types.Metric":
234238
)
235239
try:
236240
gcs_metric = self._fetch_and_parse(api_client)
237-
final_cache_key = f"{self.name}@{self.version}"
238-
LazyLoadedPrebuiltMetric._cache[final_cache_key] = gcs_metric
241+
if use_cache:
242+
final_cache_key = f"{self.name}@{self.version}"
243+
LazyLoadedPrebuiltMetric._cache[final_cache_key] = gcs_metric
239244
self._resolved_metric = gcs_metric
240245
return self._resolved_metric
241246
except Exception as e:

‎agentplatform/_genai/_transformers.py‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -64,10 +64,14 @@ def t_metrics(
6464
elif (
6565
metric_name and metric_name in _evals_constant.SUPPORTED_PREDEFINED_METRICS
6666
):
67-
metric_payload_item["predefined_metric_spec"] = {
67+
predefined_spec: dict[str, Any] = {
6868
"metric_spec_name": metric_name,
6969
"metric_spec_parameters": metric.metric_spec_parameters,
7070
}
71+
step_configs = getv(metric, ["judge_model_step_configs"])
72+
if step_configs:
73+
predefined_spec["step_autorater_configs"] = step_configs
74+
metric_payload_item["predefined_metric_spec"] = predefined_spec
7175
# Custom Code Execution Metric
7276
elif (
7377
hasattr(metric, "remote_custom_function") and metric.remote_custom_function

‎agentplatform/_genai/evals.py‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -392,6 +392,13 @@ def _EvaluateInstancesRequestParameters_to_vertex(
392392
],
393393
)
394394

395+
if getv(from_object, ["allow_cross_region_model"]) is not None:
396+
setv(
397+
to_object,
398+
["allowCrossRegionModel"],
399+
getv(from_object, ["allow_cross_region_model"]),
400+
)
401+
395402
if getv(from_object, ["config"]) is not None:
396403
setv(to_object, ["config"], getv(from_object, ["config"]))
397404

@@ -1935,6 +1942,7 @@ def _evaluate_instances(
19351942
metrics: Optional[list[types.MetricOrDict]] = None,
19361943
instance: Optional[types.EvaluationInstanceOrDict] = None,
19371944
metric_sources: Optional[list[types.MetricSourceOrDict]] = None,
1945+
allow_cross_region_model: Optional[bool] = None,
19381946
config: Optional[types.EvaluateInstancesConfigOrDict] = None,
19391947
) -> types.EvaluateInstancesResponse:
19401948
"""
@@ -1956,6 +1964,7 @@ def _evaluate_instances(
19561964
metrics=metrics,
19571965
instance=instance,
19581966
metric_sources=metric_sources,
1967+
allow_cross_region_model=allow_cross_region_model,
19591968
config=config,
19601969
)
19611970

@@ -3171,6 +3180,8 @@ def evaluate(
31713180
- evaluation_service_qps: The rate limit (queries per second) for
31723181
calls to the evaluation service. Defaults to 10. Increase this
31733182
value if your project has a higher EvaluateInstances API quota.
3183+
- allow_cross_region_model: Opt-in flag to authorize cross-region
3184+
routing for judge models.
31743185
**kwargs: Extra arguments to pass to evaluation, such as `agent_info`.
31753186
31763187
Returns:
@@ -3214,6 +3225,7 @@ def evaluate(
32143225
dest=config.dest,
32153226
location=location,
32163227
evaluation_service_qps=getattr(config, "evaluation_service_qps", None),
3228+
allow_cross_region_model=getattr(config, "allow_cross_region_model", None),
32173229
**kwargs,
32183230
)
32193231

@@ -4944,6 +4956,7 @@ async def _evaluate_instances(
49444956
metrics: Optional[list[types.MetricOrDict]] = None,
49454957
instance: Optional[types.EvaluationInstanceOrDict] = None,
49464958
metric_sources: Optional[list[types.MetricSourceOrDict]] = None,
4959+
allow_cross_region_model: Optional[bool] = None,
49474960
config: Optional[types.EvaluateInstancesConfigOrDict] = None,
49484961
) -> types.EvaluateInstancesResponse:
49494962
"""
@@ -4965,6 +4978,7 @@ async def _evaluate_instances(
49654978
metrics=metrics,
49664979
instance=instance,
49674980
metric_sources=metric_sources,
4981+
allow_cross_region_model=allow_cross_region_model,
49684982
config=config,
49694983
)
49704984

0 commit comments

Comments
 (0)