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
645 changes: 645 additions & 0 deletions packages/gen/docs/gen_ai_hub/examples/prompt-optimization.ipynb

Large diffs are not rendered by default.

4 changes: 4 additions & 0 deletions packages/gen/gen_ai_hub/evaluations/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -181,6 +181,10 @@ def __init__(
self._validate_provider_params()
logger.info("Initialization of the client completed!")

@property
def _gen_ai_hub_proxy_client(self):
return self.__gen_ai_hub_proxy_client

def __repr__(self):
attrs = ", ".join(f"{k}={v!r}" for k, v in vars(self).items())
return f"{self.__class__.__name__}({attrs})"
Expand Down
4 changes: 4 additions & 0 deletions packages/gen/gen_ai_hub/optimizations/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
from .client import OptimizationClient
from .models import OptimizationRun, PromptOptimizationConfig, OptimizationResults

__all__ = ["OptimizationClient", "OptimizationRun", "PromptOptimizationConfig", "OptimizationResults"]
59 changes: 59 additions & 0 deletions packages/gen/gen_ai_hub/optimizations/client.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
"""Client for submitting and managing prompt optimization jobs on generative AI Hub."""
from gen_ai_hub.evaluations._internal._models import _AWSObjectStoreData
from gen_ai_hub.evaluations.client import EvaluationClient
from gen_ai_hub.evaluations.constants import DEFAULT_KEY
from gen_ai_hub.evaluations.exceptions.error_codes import ErrorCode
from gen_ai_hub.evaluations.helpers.collector import ValidationCollector
from gen_ai_hub.evaluations.utils.oss_secret_utils import fetch_object_store_secret_by_name
from gen_ai_hub.optimizations.models.optimization_config import PromptOptimizationConfig
from gen_ai_hub.optimizations.models.optimization_run import OptimizationRun
from gen_ai_hub.optimizations.optimization_flow import optimization_job_flow
from gen_ai_hub.optimizations.utils import validate_optimization_config


class OptimizationClient(EvaluationClient):
"""Client for running prompt optimization jobs against a target metric on generative AI Hub."""

def optimize(self, optimization_config: PromptOptimizationConfig) -> OptimizationRun:
"""Submit a prompt optimization job and return an OptimizationRun to track its progress."""
error_collector = ValidationCollector()
try:
if self.default_object_store_secret_name is None:
response = fetch_object_store_secret_by_name(
self.ai_core_client,
DEFAULT_KEY,
self.resource_group,
error_collector,
)
if response is None:
error_collector.add_error(
ErrorCode.MISSING_DEFAULT_OBJECT_STORE_SECRET_ERROR.value,
"Default Object Store secret is required to run optimize function. "
"Please use setup() function to create one!",
)

error_collector.raise_if_errors()

validate_optimization_config(
optimization_config,
self.ai_core_client,
self.resource_group,
error_collector,
)
error_collector.raise_if_errors()

object_store_credentials = _AWSObjectStoreData(
aws_access_key_id=self.aws_access_key_id,
aws_secret_access_key=self.aws_secret_access_key,
)
return optimization_job_flow(
optimization_config,
object_store_credentials,
self.ai_core_client,
self.resource_group,
error_collector,
proxy_client=self._gen_ai_hub_proxy_client,
)
except Exception as exc:
error_collector.raise_if_errors()
raise RuntimeError("Optimize function failed!") from exc
3 changes: 3 additions & 0 deletions packages/gen/gen_ai_hub/optimizations/constants.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
OPTIMIZATIONS_SCENARIO_ID = "genai-optimizations"
OPTIMIZATIONS_CONFIG_PREFIX_KEY = "optimization-config-"
OPTIMIZATIONS_ARTIFACT_KEY = "prompt-data"
5 changes: 5 additions & 0 deletions packages/gen/gen_ai_hub/optimizations/models/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
from .optimization_config import PromptOptimizationConfig
from .optimization_run import OptimizationRun
from .optimization_results import OptimizationResults

__all__ = ["PromptOptimizationConfig", "OptimizationRun", "OptimizationResults"]
102 changes: 102 additions & 0 deletions packages/gen/gen_ai_hub/optimizations/models/optimization_config.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,102 @@
"""PromptOptimizationConfig: configuration model for prompt optimization jobs."""
from typing import Dict, List, Optional


class PromptOptimizationConfig:
"""Configuration for a prompt optimization job.

:param target_prompt_mapping: Maps each target model (``name:version``) to the output
prompt name in the Prompt Registry.
:param target_models: List of models to optimize for, e.g. ``["gpt-4o:2024-11-20"]``.
:param base_prompt: Starting prompt template referenced as ``<scenario>/<name>:<version>``.
:param optimization_metric: System-defined metric to optimize for (e.g. ``"JSON_Match"``).
Mutually exclusive with ``custom_metric_id``.
:param artifact_id: AI Core artifact ID of a pre-uploaded dataset. Mutually exclusive with ``dataset_path``.
:param dataset: Relative file key within the artifact (required when using ``artifact_id``).
:param dataset_path: Local path to a labelled JSON dataset. The SDK uploads it automatically.
Mutually exclusive with ``artifact_id``.
:param base_model: Reference model used internally during optimization.
:param include_few_shot_examples: Whether to include few-shot examples in the optimized prompt
(default: ``False``).
:param custom_metric_id: ID of a custom metric to optimize for. Mutually exclusive with
``optimization_metric``.
:param maximize: Whether higher metric scores are better (default: ``True``).
:param correctness_cutoff: Score threshold for correctness classification.
:param prompt_template_scope: Scope for the output prompt template — ``"tenant"`` or
``"resourcegroup"`` (default: ``"tenant"``).
:param prototype_mode: Use as few as 3 samples for quick prototyping (default: ``False``).
:param train_dataset_config: Optional training dataset config. Requires ``test_dataset_config``.
:param test_dataset_config: Optional test dataset config. Required when ``train_dataset_config`` is provided.
:param model_params: JSON string mapping model IDs to parameter dicts (e.g. ``temperature``, ``max_tokens``).
:param variable_mapping: JSON string mapping prompt template variable names to dataset field names.
:param field_evaluation_metrics: JSON string mapping response format field names to their evaluation metrics
(e.g. ``'{"urgency": "ExactMatch", "sentiment": "LLMaaJ:Sem_Sim_1"}'``). Requires a ``response_format``
to be defined in the base prompt template. Mutually exclusive with ``optimization_metric`` and ``custom_metric_id``.
"""

def __init__(
self,
target_prompt_mapping: Dict[str, str],
target_models: List[str],
base_prompt: str,
optimization_metric: Optional[str] = None,
artifact_id: Optional[str] = None,
dataset: Optional[str] = None,
dataset_path: Optional[str] = None,
base_model: Optional[str] = "none",
include_few_shot_examples: Optional[bool] = False,
custom_metric_id: Optional[str] = None,
maximize: Optional[bool] = True,
correctness_cutoff: Optional[float] = None,
prompt_template_scope: Optional[str] = "tenant",
prototype_mode: Optional[bool] = False,
train_dataset_config=None,
test_dataset_config=None,
model_params: Optional[str] = None,
variable_mapping: Optional[str] = None,
field_evaluation_metrics: Optional[str] = None,
):
self.artifact_id = artifact_id
self.dataset = dataset
self.dataset_path = dataset_path
self.target_prompt_mapping = target_prompt_mapping
self.target_models = target_models
self.base_prompt = base_prompt
self.optimization_metric = optimization_metric
self.base_model = base_model
self.include_few_shot_examples = include_few_shot_examples
self.custom_metric_id = custom_metric_id
self.maximize = maximize
self.correctness_cutoff = correctness_cutoff
self.prompt_template_scope = prompt_template_scope
self.prototype_mode = prototype_mode
self.train_dataset_config = train_dataset_config
self.test_dataset_config = test_dataset_config
self.model_params = model_params
self.variable_mapping = variable_mapping
self.field_evaluation_metrics = field_evaluation_metrics
self._validate(
artifact_id, dataset, dataset_path,
optimization_metric, custom_metric_id,
train_dataset_config, test_dataset_config,
field_evaluation_metrics,
)

def _validate(self, artifact_id, dataset, dataset_path,
optimization_metric, custom_metric_id,
train_dataset_config, test_dataset_config,
field_evaluation_metrics=None):
if artifact_id is None and dataset_path is None:
raise ValueError("Either artifact_id or dataset_path must be provided.")
if artifact_id is not None and dataset_path is not None:
raise ValueError("Only one of artifact_id or dataset_path must be provided, not both.")
if artifact_id is not None and dataset is None:
raise ValueError("dataset (filename) must be provided when using artifact_id.")
if train_dataset_config is not None and test_dataset_config is None:
raise ValueError("test_dataset_config must be provided when train_dataset_config is provided.")
if optimization_metric is None and custom_metric_id is None and field_evaluation_metrics is None:
raise ValueError(
"At least one of optimization_metric, custom_metric_id, or field_evaluation_metrics must be provided."
)
if optimization_metric is not None and custom_metric_id is not None:
raise ValueError("Only one of optimization_metric or custom_metric_id must be provided, not both.")
Loading