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
12 changes: 8 additions & 4 deletions agentplatform/_genai/_transformers.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,19 +72,23 @@ def t_metrics(
elif (
hasattr(metric, "remote_custom_function") and metric.remote_custom_function
):
metric_payload_item["custom_code_execution_spec"] = {
spec: dict[str, Any] = {
"evaluation_function": metric.remote_custom_function
}
if getattr(metric, "code_execution_region", None):
spec["code_execution_region"] = metric.code_execution_region
metric_payload_item["custom_code_execution_spec"] = spec
elif (
isinstance(metric, types.CodeExecutionMetric)
or (
isinstance(metric, types.Metric)
and isinstance(getattr(metric, "custom_function", None), str)
)
) and getattr(metric, "custom_function", None):
metric_payload_item["custom_code_execution_spec"] = {
"evaluation_function": metric.custom_function
}
spec = {"evaluation_function": metric.custom_function}
if getattr(metric, "code_execution_region", None):
spec["code_execution_region"] = metric.code_execution_region
metric_payload_item["custom_code_execution_spec"] = spec
# LLM-based metrics
elif hasattr(metric, "prompt_template") and metric.prompt_template:
llm_based_spec: dict[str, Any] = {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -43,40 +43,81 @@ def evaluate(instance):
],
)
def test_custom_code_execution(client, custom_metric):
"""Tests that custom code execution metric produces a correctly structured EvaluationResult."""
"""Tests that custom code execution metric produces a correctly structured EvaluationResult."""

prompts_df = pd.DataFrame(
prompts_df = pd.DataFrame(
{
"prompt": ["What is 2+2?", "What is 3+3?"],
"response": ["4", "5"],
"reference": ["4", "6"],
}
)

eval_dataset = types.EvaluationDataset(
eval_dataset = types.EvaluationDataset(
eval_dataset_df=prompts_df,
candidate_name="test_model",
)

evaluation_result = client.evals.evaluate(
evaluation_result = client.evals.evaluate(
dataset=eval_dataset,
metrics=[custom_metric],
)

assert isinstance(evaluation_result, types.EvaluationResult)

assert evaluation_result.summary_metrics is not None
assert evaluation_result.summary_metrics
for summary in evaluation_result.summary_metrics:
assert isinstance(summary, types.AggregatedMetricResult)
assert summary.metric_name == "my_custom_code_metric"

assert evaluation_result.eval_case_results is not None
assert evaluation_result.eval_case_results
for case_result in evaluation_result.eval_case_results:
assert isinstance(case_result, types.EvalCaseResult)
assert case_result.eval_case_index is not None
assert case_result.response_candidate_results is not None
assert isinstance(evaluation_result, types.EvaluationResult)

assert evaluation_result.summary_metrics is not None
assert evaluation_result.summary_metrics
for summary in evaluation_result.summary_metrics:
assert isinstance(summary, types.AggregatedMetricResult)
assert summary.metric_name == "my_custom_code_metric"

assert evaluation_result.eval_case_results is not None
assert evaluation_result.eval_case_results
for case_result in evaluation_result.eval_case_results:
assert isinstance(case_result, types.EvalCaseResult)
assert case_result.eval_case_index is not None
assert case_result.response_candidate_results is not None


def test_custom_code_execution_with_region(client):
"""Tests that code_execution_region is included in the custom code execution spec."""

prompts_df = pd.DataFrame({
"prompt": ["What is 2+2?", "What is 3+3?"],
"response": ["4", "5"],
"reference": ["4", "6"],
})

eval_dataset = types.EvaluationDataset(
eval_dataset_df=prompts_df,
candidate_name="test_model",
)

metric = types.Metric(
name="my_custom_code_metric",
remote_custom_function=CODE_SNIPPET,
code_execution_region="europe-west3",
)

evaluation_result = client.evals.evaluate(
dataset=eval_dataset,
metrics=[metric],
)

assert isinstance(evaluation_result, types.EvaluationResult)

assert evaluation_result.summary_metrics is not None
assert evaluation_result.summary_metrics
for summary in evaluation_result.summary_metrics:
assert isinstance(summary, types.AggregatedMetricResult)
assert summary.metric_name == "my_custom_code_metric"

assert evaluation_result.eval_case_results is not None
assert evaluation_result.eval_case_results
for case_result in evaluation_result.eval_case_results:
assert isinstance(case_result, types.EvalCaseResult)
assert case_result.eval_case_index is not None
assert case_result.response_candidate_results is not None


@pytest.mark.parametrize(
Expand Down
12 changes: 8 additions & 4 deletions vertexai/_genai/_transformers.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,19 +72,23 @@ def t_metrics(
elif (
hasattr(metric, "remote_custom_function") and metric.remote_custom_function
):
metric_payload_item["custom_code_execution_spec"] = {
spec: dict[str, Any] = {
"evaluation_function": metric.remote_custom_function
}
if getattr(metric, "code_execution_region", None):
spec["code_execution_region"] = metric.code_execution_region
metric_payload_item["custom_code_execution_spec"] = spec
elif (
isinstance(metric, types.CodeExecutionMetric)
or (
isinstance(metric, types.Metric)
and isinstance(getattr(metric, "custom_function", None), str)
)
) and getattr(metric, "custom_function", None):
metric_payload_item["custom_code_execution_spec"] = {
"evaluation_function": metric.custom_function
}
spec = {"evaluation_function": metric.custom_function}
if getattr(metric, "code_execution_region", None):
spec["code_execution_region"] = metric.code_execution_region
metric_payload_item["custom_code_execution_spec"] = spec
# LLM-based metrics
elif hasattr(metric, "prompt_template") and metric.prompt_template:
llm_based_spec: dict[str, Any] = {
Expand Down
14 changes: 14 additions & 0 deletions vertexai/_genai/evals.py
Original file line number Diff line number Diff line change
Expand Up @@ -206,6 +206,13 @@ def _CustomCodeExecutionSpec_from_vertex(
getv(from_object, ["evaluation_function"]),
)

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

return to_object


Expand All @@ -228,6 +235,13 @@ def _CustomCodeExecutionSpec_to_vertex(
getv(from_object, ["remote_custom_function"]),
)

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

return to_object


Expand Down
8 changes: 8 additions & 0 deletions vertexai/_genai/types/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -1816,6 +1816,10 @@ class Metric(_common.BaseModel):
default=None,
description="""The evaluation function for the custom code execution metric. This custom code is run remotely in the evaluation service.""",
)
code_execution_region: Optional[str] = Field(
default=None,
description="""Optional. The region to use for code execution. If set, the Code Execution Sandbox will be invoked in the specified region regardless of the request's originating region. Supported regions: us-central1, us-east1, us-east4, us-west1, us-west4, southamerica-east1, europe-west2, europe-west3, asia-east1, asia-south1, asia-southeast1. If unset, the request's originating region is used.""",
)
judge_model: Optional[str] = Field(
default=None, description="""The judge model for the metric."""
)
Expand Down Expand Up @@ -2164,6 +2168,10 @@ def evaluate(instance: dict[str, Any]) -> float:
Instance is the evaluation instance, any fields populated in the instance
are available to the function as instance[field_name].""",
)
code_execution_region: Optional[str] = Field(
default=None,
description="""Optional. The region to use for code execution. If set, the Code Execution Sandbox will be invoked in the specified region regardless of the request's originating region. Supported regions: us-central1, us-east1, us-east4, us-west1, us-west4, southamerica-east1, europe-west2, europe-west3, asia-east1, asia-south1, asia-southeast1. If unset, the request's originating region is used.""",
)


class CustomCodeExecutionSpecDict(TypedDict, total=False):
Expand Down
Loading