mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat: Add a new option eval_storage_uri in adk web & adk eval to specify GCS bucket to store eval data
PiperOrigin-RevId: 774947795
This commit is contained in:
committed by
Copybara-Service
parent
120cbabeb2
commit
fa025d7559
@@ -31,12 +31,15 @@ import uvicorn
|
||||
from . import cli_create
|
||||
from . import cli_deploy
|
||||
from .. import version
|
||||
from ..evaluation.gcs_eval_set_results_manager import GcsEvalSetResultsManager
|
||||
from ..evaluation.gcs_eval_sets_manager import GcsEvalSetsManager
|
||||
from ..evaluation.local_eval_set_results_manager import LocalEvalSetResultsManager
|
||||
from ..sessions.in_memory_session_service import InMemorySessionService
|
||||
from .cli import run_cli
|
||||
from .cli_eval import MISSING_EVAL_DEPENDENCIES_MESSAGE
|
||||
from .fast_api import get_fast_api_app
|
||||
from .utils import envs
|
||||
from .utils import evals
|
||||
from .utils import logs
|
||||
|
||||
LOG_LEVELS = click.Choice(
|
||||
@@ -282,11 +285,21 @@ def cli_run(
|
||||
default=False,
|
||||
help="Optional. Whether to print detailed results on console or not.",
|
||||
)
|
||||
@click.option(
|
||||
"--eval_storage_uri",
|
||||
type=str,
|
||||
help=(
|
||||
"Optional. The evals storage URI to store agent evals,"
|
||||
" supported URIs: gs://<bucket name>."
|
||||
),
|
||||
default=None,
|
||||
)
|
||||
def cli_eval(
|
||||
agent_module_file_path: str,
|
||||
eval_set_file_path: tuple[str],
|
||||
eval_set_file_path: list[str],
|
||||
config_file_path: str,
|
||||
print_detailed_results: bool,
|
||||
eval_storage_uri: Optional[str] = None,
|
||||
):
|
||||
"""Evaluates an agent given the eval sets.
|
||||
|
||||
@@ -338,12 +351,33 @@ def cli_eval(
|
||||
root_agent = get_root_agent(agent_module_file_path)
|
||||
reset_func = try_get_reset_func(agent_module_file_path)
|
||||
|
||||
gcs_eval_sets_manager = None
|
||||
eval_set_results_manager = None
|
||||
if eval_storage_uri:
|
||||
gcs_eval_managers = evals.create_gcs_eval_managers_from_uri(
|
||||
eval_storage_uri
|
||||
)
|
||||
gcs_eval_sets_manager = gcs_eval_managers.eval_sets_manager
|
||||
eval_set_results_manager = gcs_eval_managers.eval_set_results_manager
|
||||
else:
|
||||
eval_set_results_manager = LocalEvalSetResultsManager(
|
||||
agents_dir=os.path.dirname(agent_module_file_path)
|
||||
)
|
||||
eval_set_file_path_to_evals = parse_and_get_evals_to_run(eval_set_file_path)
|
||||
eval_set_id_to_eval_cases = {}
|
||||
|
||||
# Read the eval_set files and get the cases.
|
||||
for eval_set_file_path, eval_case_ids in eval_set_file_path_to_evals.items():
|
||||
eval_set = load_eval_set_from_file(eval_set_file_path, eval_set_file_path)
|
||||
if gcs_eval_sets_manager:
|
||||
eval_set = gcs_eval_sets_manager._load_eval_set_from_blob(
|
||||
eval_set_file_path
|
||||
)
|
||||
if not eval_set:
|
||||
raise click.ClickException(
|
||||
f"Eval set {eval_set_file_path} not found in GCS."
|
||||
)
|
||||
else:
|
||||
eval_set = load_eval_set_from_file(eval_set_file_path, eval_set_file_path)
|
||||
eval_cases = eval_set.eval_cases
|
||||
|
||||
if eval_case_ids:
|
||||
@@ -378,16 +412,13 @@ def cli_eval(
|
||||
raise click.ClickException(MISSING_EVAL_DEPENDENCIES_MESSAGE)
|
||||
|
||||
# Write eval set results.
|
||||
local_eval_set_results_manager = LocalEvalSetResultsManager(
|
||||
agents_dir=os.path.dirname(agent_module_file_path)
|
||||
)
|
||||
eval_set_id_to_eval_results = collections.defaultdict(list)
|
||||
for eval_case_result in eval_results:
|
||||
eval_set_id = eval_case_result.eval_set_id
|
||||
eval_set_id_to_eval_results[eval_set_id].append(eval_case_result)
|
||||
|
||||
for eval_set_id, eval_case_results in eval_set_id_to_eval_results.items():
|
||||
local_eval_set_results_manager.save_eval_set_result(
|
||||
eval_set_results_manager.save_eval_set_result(
|
||||
app_name=os.path.basename(agent_module_file_path),
|
||||
eval_set_id=eval_set_id,
|
||||
eval_case_results=eval_case_results,
|
||||
@@ -444,6 +475,15 @@ def adk_services_options():
|
||||
),
|
||||
default=None,
|
||||
)
|
||||
@click.option(
|
||||
"--eval_storage_uri",
|
||||
type=str,
|
||||
help=(
|
||||
"Optional. The evals storage URI to store agent evals,"
|
||||
" supported URIs: gs://<bucket name>."
|
||||
),
|
||||
default=None,
|
||||
)
|
||||
@click.option(
|
||||
"--memory_service_uri",
|
||||
type=str,
|
||||
@@ -564,6 +604,7 @@ def fast_api_common_options():
|
||||
)
|
||||
def cli_web(
|
||||
agents_dir: str,
|
||||
eval_storage_uri: Optional[str] = None,
|
||||
log_level: str = "INFO",
|
||||
allow_origins: Optional[list[str]] = None,
|
||||
host: str = "127.0.0.1",
|
||||
@@ -616,6 +657,7 @@ def cli_web(
|
||||
session_service_uri=session_service_uri,
|
||||
artifact_service_uri=artifact_service_uri,
|
||||
memory_service_uri=memory_service_uri,
|
||||
eval_storage_uri=eval_storage_uri,
|
||||
allow_origins=allow_origins,
|
||||
web=True,
|
||||
trace_to_cloud=trace_to_cloud,
|
||||
@@ -654,6 +696,7 @@ def cli_web(
|
||||
)
|
||||
def cli_api_server(
|
||||
agents_dir: str,
|
||||
eval_storage_uri: Optional[str] = None,
|
||||
log_level: str = "INFO",
|
||||
allow_origins: Optional[list[str]] = None,
|
||||
host: str = "127.0.0.1",
|
||||
@@ -685,6 +728,7 @@ def cli_api_server(
|
||||
session_service_uri=session_service_uri,
|
||||
artifact_service_uri=artifact_service_uri,
|
||||
memory_service_uri=memory_service_uri,
|
||||
eval_storage_uri=eval_storage_uri,
|
||||
allow_origins=allow_origins,
|
||||
web=False,
|
||||
trace_to_cloud=trace_to_cloud,
|
||||
@@ -771,6 +815,15 @@ def cli_api_server(
|
||||
" version in the dev environment)"
|
||||
),
|
||||
)
|
||||
@click.option(
|
||||
"--eval_storage_uri",
|
||||
type=str,
|
||||
help=(
|
||||
"Optional. The evals storage URI to store agent evals,"
|
||||
" supported URIs: gs://<bucket name>."
|
||||
),
|
||||
default=None,
|
||||
)
|
||||
@adk_services_options()
|
||||
@deprecated_adk_services_options()
|
||||
@click.argument(
|
||||
@@ -797,6 +850,7 @@ def cli_deploy_cloud_run(
|
||||
session_service_uri: Optional[str] = None,
|
||||
artifact_service_uri: Optional[str] = None,
|
||||
memory_service_uri: Optional[str] = None,
|
||||
eval_storage_uri: Optional[str] = None,
|
||||
session_db_url: Optional[str] = None, # Deprecated
|
||||
artifact_storage_uri: Optional[str] = None, # Deprecated
|
||||
):
|
||||
|
||||
Reference in New Issue
Block a user