Skip to content
Merged
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
221 changes: 221 additions & 0 deletions agentplatform/_genai/evals.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,29 @@
logger = logging.getLogger("agentplatform_genai.evals")


def _CreateEvaluationExperimentParameters_to_vertex(
from_object: Union[dict[str, Any], object],
parent_object: Optional[dict[str, Any]] = None,
) -> dict[str, Any]:
to_object: dict[str, Any] = {}
if getv(from_object, ["display_name"]) is not None:
setv(to_object, ["displayName"], getv(from_object, ["display_name"]))

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

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

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

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

return to_object


def _CreateEvaluationItemParameters_to_vertex(
from_object: Union[dict[str, Any], object],
parent_object: Optional[dict[str, Any]] = None,
Expand Down Expand Up @@ -1070,6 +1093,104 @@ def _UnifiedMetric_to_vertex(

class Evals(_api_module.BaseModule):

def create_evaluation_experiment(
self,
*,
display_name: Optional[str] = None,
labels: Optional[dict[str, str]] = None,
merge_strategy: Optional[types.EvaluationExperimentMergeStrategy] = None,
metadata: Optional[dict[str, Any]] = None,
config: Optional[types.CreateEvaluationExperimentConfigOrDict] = None,
) -> types.EvaluationExperiment:
"""
Creates an EvaluationExperiment.

Args:
display_name: The display name of the evaluation experiment.
labels: Labels for the evaluation experiment.
merge_strategy: Merge strategy for the evaluation experiment.
metadata: Metadata about the evaluation experiment, can be used by the
caller to store additional tracking information about the experiment.
config: Optional configuration for the create operation.

Returns:
The created evaluation experiment.

.. code-block:: python

eval_experiment = client.evals.create_evaluation_experiment(
display_name="my-experiment"
)

"""

parameter_model = types._CreateEvaluationExperimentParameters(
display_name=display_name,
labels=labels,
merge_strategy=merge_strategy,
metadata=metadata,
config=config,
)

request_url_dict: Optional[dict[str, str]]
if not self._api_client.vertexai:
raise ValueError(
"This method is only supported in Gemini Enterprise Agent Platform mode, not in Gemini Developer API mode."
)
else:
request_dict = _CreateEvaluationExperimentParameters_to_vertex(
parameter_model
)
request_url_dict = request_dict.get("_url")
if request_url_dict:
path = "evaluationExperiments".format_map(request_url_dict)
else:
path = "evaluationExperiments"

query_params = request_dict.get("_query")
if query_params:
path = f"{path}?{urlencode(query_params)}"
# TODO: remove the hack that pops config.
request_dict.pop("config", None)

http_options: Optional[types.HttpOptions] = None
if (
parameter_model.config is not None
and parameter_model.config.http_options is not None
):
http_options = parameter_model.config.http_options

request_dict = _common.convert_to_dict(request_dict)
request_dict = _common.encode_unserializable_types(request_dict)

response = self._api_client.request("post", path, request_dict, http_options)

response_dict = {} if not response.body else json.loads(response.body)

return_value = types.EvaluationExperiment._from_response(
response=response_dict,
kwargs=(
{
"config": {
"response_schema": getattr(
parameter_model.config, "response_schema", None
),
"response_json_schema": getattr(
parameter_model.config, "response_json_schema", None
),
"include_all_fields": getattr(
parameter_model.config, "include_all_fields", None
),
}
}
if getattr(parameter_model, "config", None)
else {}
),
)

self._api_client._verify_response(return_value)
return return_value

def _create_evaluation_item(
self,
*,
Expand Down Expand Up @@ -3403,6 +3524,106 @@ def delete_evaluation_metric(

class AsyncEvals(_api_module.BaseModule):

async def create_evaluation_experiment(
self,
*,
display_name: Optional[str] = None,
labels: Optional[dict[str, str]] = None,
merge_strategy: Optional[types.EvaluationExperimentMergeStrategy] = None,
metadata: Optional[dict[str, Any]] = None,
config: Optional[types.CreateEvaluationExperimentConfigOrDict] = None,
) -> types.EvaluationExperiment:
"""
Creates an EvaluationExperiment.

Args:
display_name: The display name of the evaluation experiment.
labels: Labels for the evaluation experiment.
merge_strategy: Merge strategy for the evaluation experiment.
metadata: Metadata about the evaluation experiment, can be used by the
caller to store additional tracking information about the experiment.
config: Optional configuration for the create operation.

Returns:
The created evaluation experiment.

.. code-block:: python

eval_experiment = client.evals.create_evaluation_experiment(
display_name="my-experiment"
)

"""

parameter_model = types._CreateEvaluationExperimentParameters(
display_name=display_name,
labels=labels,
merge_strategy=merge_strategy,
metadata=metadata,
config=config,
)

request_url_dict: Optional[dict[str, str]]
if not self._api_client.vertexai:
raise ValueError(
"This method is only supported in Gemini Enterprise Agent Platform mode, not in Gemini Developer API mode."
)
else:
request_dict = _CreateEvaluationExperimentParameters_to_vertex(
parameter_model
)
request_url_dict = request_dict.get("_url")
if request_url_dict:
path = "evaluationExperiments".format_map(request_url_dict)
else:
path = "evaluationExperiments"

query_params = request_dict.get("_query")
if query_params:
path = f"{path}?{urlencode(query_params)}"
# TODO: remove the hack that pops config.
request_dict.pop("config", None)

http_options: Optional[types.HttpOptions] = None
if (
parameter_model.config is not None
and parameter_model.config.http_options is not None
):
http_options = parameter_model.config.http_options

request_dict = _common.convert_to_dict(request_dict)
request_dict = _common.encode_unserializable_types(request_dict)

response = await self._api_client.async_request(
"post", path, request_dict, http_options
)

response_dict = {} if not response.body else json.loads(response.body)

return_value = types.EvaluationExperiment._from_response(
response=response_dict,
kwargs=(
{
"config": {
"response_schema": getattr(
parameter_model.config, "response_schema", None
),
"response_json_schema": getattr(
parameter_model.config, "response_json_schema", None
),
"include_all_fields": getattr(
parameter_model.config, "include_all_fields", None
),
}
}
if getattr(parameter_model, "config", None)
else {}
),
)

self._api_client._verify_response(return_value)
return return_value

async def _create_evaluation_item(
self,
*,
Expand Down
16 changes: 12 additions & 4 deletions agentplatform/_genai/types/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@
from .common import _CreateAgentEngineTaskRequestParameters
from .common import _CreateDatasetParameters
from .common import _CreateDatasetVersionParameters
from .common import _CreateEvaluationExperimentParameters
from .common import _CreateEvaluationItemParameters
from .common import _CreateEvaluationMetricParameters
from .common import _CreateEvaluationRunParameters
Expand Down Expand Up @@ -342,6 +343,9 @@
from .common import CreateDatasetVersionConfig
from .common import CreateDatasetVersionConfigDict
from .common import CreateDatasetVersionConfigOrDict
from .common import CreateEvaluationExperimentConfig
from .common import CreateEvaluationExperimentConfigDict
from .common import CreateEvaluationExperimentConfigOrDict
from .common import CreateEvaluationItemConfig
from .common import CreateEvaluationItemConfigDict
from .common import CreateEvaluationItemConfigOrDict
Expand Down Expand Up @@ -2013,6 +2017,12 @@
"ListAgentEngineTaskEventsResponse",
"ListAgentEngineTaskEventsResponseDict",
"ListAgentEngineTaskEventsResponseOrDict",
"CreateEvaluationExperimentConfig",
"CreateEvaluationExperimentConfigDict",
"CreateEvaluationExperimentConfigOrDict",
"EvaluationExperiment",
"EvaluationExperimentDict",
"EvaluationExperimentOrDict",
"CreateEvaluationItemConfig",
"CreateEvaluationItemConfigDict",
"CreateEvaluationItemConfigOrDict",
Expand Down Expand Up @@ -2337,9 +2347,6 @@
"GetEvaluationExperimentConfig",
"GetEvaluationExperimentConfigDict",
"GetEvaluationExperimentConfigOrDict",
"EvaluationExperiment",
"EvaluationExperimentDict",
"EvaluationExperimentOrDict",
"GetEvaluationMetricConfig",
"GetEvaluationMetricConfigDict",
"GetEvaluationMetricConfigOrDict",
Expand Down Expand Up @@ -3669,10 +3676,10 @@
"VersionState",
"QuotaState",
"FeedbackType",
"EvaluationExperimentMergeStrategy",
"EvaluationItemType",
"SamplingMethod",
"EvaluationRunState",
"EvaluationExperimentMergeStrategy",
"OptimizeTarget",
"MemoryMetadataMergeStrategy",
"GenerateMemoriesResponseGeneratedMemoryAction",
Expand Down Expand Up @@ -3708,6 +3715,7 @@
"_CreateAgentEngineTaskRequestParameters",
"_AppendAgentEngineTaskEventRequestParameters",
"_ListAgentEngineTaskEventsRequestParameters",
"_CreateEvaluationExperimentParameters",
"_CreateEvaluationItemParameters",
"_CreateEvaluationMetricParameters",
"_CreateEvaluationRunParameters",
Expand Down
Loading
Loading