Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 21 additions & 16 deletions agentplatform/_genai/_evals_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -2886,6 +2886,21 @@ def _resolve_dataset_inputs(
return processed_eval_dataset, num_response_candidates


def _prebuilt_evaluation_run_metric(
resolved_metric: types.Metric,
) -> types.EvaluationRunMetric:
"""Builds the evaluation run metric for a resolved RubricMetric."""
if resolved_metric.name in _evals_constant.SUPPORTED_PREDEFINED_METRICS:
metric_config = t.t_metrics([resolved_metric])[0]
else:
metric_config = {
"predefined_metric_spec": {"metric_spec_name": resolved_metric.name}
}
return types.EvaluationRunMetric(
metric=resolved_metric.name, metric_config=metric_config
)


def _resolve_evaluation_run_metrics(
metrics: Union[list[types.EvaluationRunMetric], list[types.Metric]], api_client: Any
) -> list[types.EvaluationRunMetric]:
Expand All @@ -2903,14 +2918,7 @@ def _resolve_evaluation_run_metrics(
resolved_metric = metric_instance.resolve(api_client=api_client)
if resolved_metric.name:
resolved_metrics_list.append(
types.EvaluationRunMetric(
metric=resolved_metric.name,
metric_config=types.UnifiedMetric(
predefined_metric_spec=genai_types.PredefinedMetricSpec(
metric_spec_name=resolved_metric.name,
)
),
)
_prebuilt_evaluation_run_metric(resolved_metric)
)
except Exception as e:
logger.error(
Expand Down Expand Up @@ -2944,14 +2952,7 @@ def _resolve_evaluation_run_metrics(
)
if resolved_metric.name:
resolved_metrics_list.append(
types.EvaluationRunMetric(
metric=resolved_metric.name,
metric_config=types.UnifiedMetric(
predefined_metric_spec=genai_types.PredefinedMetricSpec(
metric_spec_name=resolved_metric.name,
)
),
)
_prebuilt_evaluation_run_metric(resolved_metric)
)
else:
raise TypeError(
Expand Down Expand Up @@ -3020,6 +3021,7 @@ def _execute_evaluation( # type: ignore[no-untyped-def]
dest: Optional[str] = None,
location: Optional[str] = None,
evaluation_service_qps: Optional[float] = None,
allow_cross_region_model: Optional[bool] = None,
**kwargs,
) -> types.EvaluationResult:
"""Evaluates a dataset using the provided metrics.
Expand All @@ -3035,6 +3037,8 @@ def _execute_evaluation( # type: ignore[no-untyped-def]
evaluation_service_qps: The rate limit (queries per second) for calls
to the evaluation service. Defaults to 10. Increase this value if
your project has a higher EvaluateInstances API quota.
allow_cross_region_model: Opt-in flag to authorize cross-region
routing for judge models.
**kwargs: Extra arguments to pass to evaluation, such as `agent_info`.

Returns:
Expand Down Expand Up @@ -3117,6 +3121,7 @@ def _execute_evaluation( # type: ignore[no-untyped-def]
evaluation_result = _evals_metric_handlers.compute_metrics_and_aggregate(
evaluation_run_config,
evaluation_service_qps=evaluation_service_qps,
allow_cross_region_model=allow_cross_region_model,
)
t2 = time.perf_counter()
logger.info("Evaluation took: %f seconds", t2 - t1)
Expand Down
26 changes: 21 additions & 5 deletions agentplatform/_genai/_evals_metric_handlers.py
Original file line number Diff line number Diff line change
Expand Up @@ -288,6 +288,7 @@ class MetricHandler(abc.ABC, Generic[T]):
def __init__(self, module: "evals.Evals", metric: T):
self.module = module
self.metric: T = metric
self.allow_cross_region_model: Optional[bool] = None

@property
@abc.abstractmethod
Expand Down Expand Up @@ -761,6 +762,7 @@ def get_metric_result(
lambda: self.module._evaluate_instances(
metrics=[self.metric],
instance=instance,
allow_cross_region_model=self.allow_cross_region_model,
),
self.metric_name,
)
Expand Down Expand Up @@ -981,15 +983,16 @@ def __init__(self, module: "evals.Evals", metric: types.Metric):
raise ValueError(
f"Metric '{self.metric.name}' is not a supported predefined metric."
)
if (
if self.metric.name.startswith("multi_turn") and (
self.metric.judge_model
or self.metric.judge_model_generation_config
or self.metric.judge_model_sampling_count
):
logger.warning(
"Autorater config settings (judge_model, "
"judge_model_generation_config, judge_model_sampling_count) "
"are ignored for predefined metric '%s'.",
"are ignored for multi-turn metric '%s'. Use "
"judge_model_step_configs to set its judges.",
self.metric.name,
)

Expand Down Expand Up @@ -1071,6 +1074,7 @@ def get_metric_result(
metrics=[self.metric],
instance=payload.get("instance"),
autorater_config=payload.get("autorater_config"),
allow_cross_region_model=self.allow_cross_region_model,
),
metric_name,
)
Expand Down Expand Up @@ -1318,6 +1322,7 @@ def get_metric_result(
metric_sources=[metric_source],
instance=payload.get("instance"),
autorater_config=payload.get("autorater_config"),
allow_cross_region_model=self.allow_cross_region_model,
),
metric_name,
)
Expand Down Expand Up @@ -1404,12 +1409,16 @@ def aggregate(


def get_handler_for_metric(
module: "evals.Evals", metric: types.Metric
module: "evals.Evals",
metric: types.Metric,
allow_cross_region_model: Optional[bool] = None,
) -> Union[MetricHandlerType, Any]:
"""Returns a metric handler for the given metric."""
for condition, handler_class in _METRIC_HANDLER_MAPPING:
if condition(metric): # type: ignore[no-untyped-call]
return handler_class(module=module, metric=metric)
handler = handler_class(module=module, metric=metric)
handler.allow_cross_region_model = allow_cross_region_model
return handler
raise ValueError(f"Unsupported metric: {metric.name}")


Expand Down Expand Up @@ -1548,6 +1557,7 @@ def _rate_limited_get_metric_result(
def compute_metrics_and_aggregate(
evaluation_run_config: EvaluationRunConfig,
evaluation_service_qps: Optional[float] = None,
allow_cross_region_model: Optional[bool] = None,
) -> types.EvaluationResult:
"""Computes metrics and aggregates them for a given evaluation run config.

Expand All @@ -1556,6 +1566,8 @@ def compute_metrics_and_aggregate(
evaluation_service_qps: Optional QPS limit for the evaluation service.
Defaults to _DEFAULT_EVAL_SERVICE_QPS (10). Users with higher
quotas can increase this value.
allow_cross_region_model: Opt-in flag to authorize cross-region
routing for judge models.
"""
metric_handlers = []
all_futures = []
Expand All @@ -1574,7 +1586,11 @@ def compute_metrics_and_aggregate(

for eval_metric in evaluation_run_config.metrics:
metric_handlers.append(
get_handler_for_metric(evaluation_run_config.evals_module, eval_metric)
get_handler_for_metric(
evaluation_run_config.evals_module,
eval_metric,
allow_cross_region_model=allow_cross_region_model,
)
)

eval_case_count = len(evaluation_run_config.dataset.eval_cases)
Expand Down
13 changes: 9 additions & 4 deletions agentplatform/_genai/_evals_metric_loaders.py
Original file line number Diff line number Diff line change
Expand Up @@ -215,8 +215,11 @@ def resolve(self, api_client: Any) -> "types.Metric":
if self._resolved_metric:
return self._resolved_metric

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

# Fallback to GCS loading for custom LLM-based Prebuilt Metrics
Expand All @@ -234,8 +238,9 @@ def resolve(self, api_client: Any) -> "types.Metric":
)
try:
gcs_metric = self._fetch_and_parse(api_client)
final_cache_key = f"{self.name}@{self.version}"
LazyLoadedPrebuiltMetric._cache[final_cache_key] = gcs_metric
if use_cache:
final_cache_key = f"{self.name}@{self.version}"
LazyLoadedPrebuiltMetric._cache[final_cache_key] = gcs_metric
self._resolved_metric = gcs_metric
return self._resolved_metric
except Exception as e:
Expand Down
6 changes: 5 additions & 1 deletion agentplatform/_genai/_transformers.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,10 +64,14 @@ def t_metrics(
elif (
metric_name and metric_name in _evals_constant.SUPPORTED_PREDEFINED_METRICS
):
metric_payload_item["predefined_metric_spec"] = {
predefined_spec: dict[str, Any] = {
"metric_spec_name": metric_name,
"metric_spec_parameters": metric.metric_spec_parameters,
}
step_configs = getv(metric, ["judge_model_step_configs"])
if step_configs:
predefined_spec["step_autorater_configs"] = step_configs
metric_payload_item["predefined_metric_spec"] = predefined_spec
# Custom Code Execution Metric
elif (
hasattr(metric, "remote_custom_function") and metric.remote_custom_function
Expand Down
14 changes: 14 additions & 0 deletions agentplatform/_genai/evals.py
Original file line number Diff line number Diff line change
Expand Up @@ -392,6 +392,13 @@ def _EvaluateInstancesRequestParameters_to_vertex(
],
)

if getv(from_object, ["allow_cross_region_model"]) is not None:
setv(
to_object,
["allowCrossRegionModel"],
getv(from_object, ["allow_cross_region_model"]),
)

if getv(from_object, ["config"]) is not None:
setv(to_object, ["config"], getv(from_object, ["config"]))

Expand Down Expand Up @@ -1935,6 +1942,7 @@ def _evaluate_instances(
metrics: Optional[list[types.MetricOrDict]] = None,
instance: Optional[types.EvaluationInstanceOrDict] = None,
metric_sources: Optional[list[types.MetricSourceOrDict]] = None,
allow_cross_region_model: Optional[bool] = None,
config: Optional[types.EvaluateInstancesConfigOrDict] = None,
) -> types.EvaluateInstancesResponse:
"""
Expand All @@ -1956,6 +1964,7 @@ def _evaluate_instances(
metrics=metrics,
instance=instance,
metric_sources=metric_sources,
allow_cross_region_model=allow_cross_region_model,
config=config,
)

Expand Down Expand Up @@ -3171,6 +3180,8 @@ def evaluate(
- evaluation_service_qps: The rate limit (queries per second) for
calls to the evaluation service. Defaults to 10. Increase this
value if your project has a higher EvaluateInstances API quota.
- allow_cross_region_model: Opt-in flag to authorize cross-region
routing for judge models.
**kwargs: Extra arguments to pass to evaluation, such as `agent_info`.

Returns:
Expand Down Expand Up @@ -3214,6 +3225,7 @@ def evaluate(
dest=config.dest,
location=location,
evaluation_service_qps=getattr(config, "evaluation_service_qps", None),
allow_cross_region_model=getattr(config, "allow_cross_region_model", None),
**kwargs,
)

Expand Down Expand Up @@ -4944,6 +4956,7 @@ async def _evaluate_instances(
metrics: Optional[list[types.MetricOrDict]] = None,
instance: Optional[types.EvaluationInstanceOrDict] = None,
metric_sources: Optional[list[types.MetricSourceOrDict]] = None,
allow_cross_region_model: Optional[bool] = None,
config: Optional[types.EvaluateInstancesConfigOrDict] = None,
) -> types.EvaluateInstancesResponse:
"""
Expand All @@ -4965,6 +4978,7 @@ async def _evaluate_instances(
metrics=metrics,
instance=instance,
metric_sources=metric_sources,
allow_cross_region_model=allow_cross_region_model,
config=config,
)

Expand Down
Loading
Loading