chore: import upstream snapshot with attribution

This commit is contained in:
wehub-resource-sync
2026-07-13 13:30:30 +08:00
commit 914fea506e
2793 changed files with 802106 additions and 0 deletions
+29
View File
@@ -0,0 +1,29 @@
## Before you begin
### Sign our Contributor License Agreement
Contributions to this project must be accompanied by a
[Contributor License Agreement](https://cla.developers.google.com/about) (CLA).
You (or your employer) retain the copyright to your contribution; this simply
gives us permission to use and redistribute your contributions as part of the
project.
If you or your current employer have already signed the Google CLA (even if it
was for a different project), you probably don't need to do it again.
Visit <https://cla.developers.google.com/> to see your current agreements or to
sign a new one.
### Review our community guidelines
This project follows
[Google's Open Source Community Guidelines](https://opensource.google/conduct/).
## Contribution process
### Code reviews
All submissions, including submissions by project members, require review. We
use GitHub pull requests for this purpose. Consult
[GitHub Help](https://help.github.com/articles/about-pull-requests/) for more
information on using pull requests.
+202
View File
@@ -0,0 +1,202 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright [yyyy] [name of copyright owner]
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
+265
View File
@@ -0,0 +1,265 @@
# LLM EvalKit
## Summary
LLMEvalKit is a tool designed to help developers evaluate and improve the performance of Large Language Models (LLMs) on specific tasks. It provides a comprehensive workflow to create, test, and optimize prompts, manage datasets, and analyze evaluation results. With LLMEvalKit, developers can conduct both human and model-based evaluations, compare results, and use automated processes to refine prompts for better accuracy and relevance. This toolkit streamlines the iterative process of prompt engineering and evaluation, enabling developers to build more effective and reliable LLM-powered applications.
![Image](assets/image.gif)
**Authors: [Mike Santoro](https://github.com/Michael-Santoro), [Katherine Larson](https://github.com/kat-litinsky)**
## 🚀 Getting Started
There are two ways to work through a tutorial of this application one method is more stable one is less stable.
1. Scroll down to the Tutorial Section here.
2. Open this [notebook](https://github.com/GoogleCloudPlatform/generative-ai/blob/main/tools/llmevalkit/prompt-management-tutorial.ipynb) in colab running the application on a colab server.
## Overview
This tutorial provides a comprehensive guide to prompt engineering, covering the entire lifecycle from creation to evaluation and optimization. It's broken down into the following sections:
1. **Prompt Management:** This section focuses on the core tasks of creating, editing, and managing prompts. You can:
- **Create new prompts:** Define the prompt's name, text, the model it's designed for, and any system instructions.
- **Load and edit existing prompts:** Browse a library of saved prompts, load a specific version, and make modifications.
- **Test prompts:** Before saving, you can provide sample input and generate a response to see how the prompt performs.
- **Versioning:** Each time you save a change to a prompt, a new version is created, allowing you to track its evolution and compare different iterations.
2. **Dataset Creation:** A crucial part of prompt engineering is having good data to test and evaluate your prompts. This section allows you to:
- **Create new datasets:** A dataset is essentially a folder in Google Cloud Storage where you can group related files.
- **Upload data:** You can upload files in CSV, JSON, or JSONL format to your datasets. This data will be used for evaluating your prompts.
3. **Evaluation:** Once you have a prompt and a dataset, you need to see how well the prompt performs. The evaluation section helps you with this by:
- **Running evaluations:** You can select a prompt and a dataset and run an evaluation. This will generate responses from the model for each item in your dataset.
- **Human-in-the-loop rating:** For a more nuanced evaluation, you can manually review the model's responses and rate them.
- **Automated metrics:** The tutorial also supports automated evaluation metrics to get a quantitative measure of your prompt's performance.
4. **One-Click Refiner:** Instantly upgrade a draft prompt into a structured, production-ready instruction without managing any datasets. This is a quick way to apply prompt engineering best practices to your initial drafts.
5. **Performance Tuner:** Optimize your prompt's System Instructions using data-driven iteration to maximize metric performance. This tool uses hill-climbing algorithms to automatically refine prompts based on your evaluation metrics.
6. **Prompt Optimization:** This section helps you automatically improve your prompts using Agent Platform's prompt optimization capabilities. It provides a structured way to:
- **Configure and launch optimization jobs:** You can set up and run a job that will take your prompt and a dataset and try to find a better-performing version of the prompt.
7. **Prompt Optimization Results:** After an optimization job has run, this section allows you to:
- **View the results:** You can see the different prompt versions that the optimizer came up with and how they performed.
- **Compare versions:** The results are presented in a way that makes it easy to compare the different optimized prompts and choose the best one.
8. **Prompt Records:** This is a leaderboard that shows you the evaluation results of all your different prompt versions. It helps you to:
- **Track performance over time:** See how your prompts have improved with each new version.
- **Compare different prompts:** You can compare the performance of different prompts for the same task.
In summary, this tutorial provides a complete and integrated environment for all your prompt engineering needs, from initial creation to sophisticated optimization and evaluation.
## Tutorial: Step-by-Step
This section walks you through using the app.
### 0. Startup
First, clone the repository and set up the environment:
# Clone the repository
git clone https://github.com/GoogleCloudPlatform/generative-ai.git
# Navigate to the project directory
cd generative-ai/tools/llmevalkit
Next, `cp src/.env.example src/.env` open the file and set `BUCKET_NAME` and `PROJECT_ID`
# Authorize gcloud
`gcloud auth application-default login`
# Run the Streamlit application
`uv run streamlit run index.py`
### 1. Prompt Management
In the Prompt Name field enter:
```
math_prompt_test
```
In the Prompt Text field enter:
```
Problem: {{query}}
Image: {{image}} @@@image/jpeg
Answer: {{target}}
```
In the Model Name field enter:
```
gemini-2.5-flash
```
In the System Instructions field enter:
```
Solve the problem given the image.
```
Click `Save Prompt`
Copy this text for testing:
```
{"query": "Hint: Please answer the question and provide the correct option letter, e.g., A, B, C, D, at the end.\nQuestion: As shown in the figure, CD is the diameter of \u2299O, chord DE \u2225 OA, if the degree of \u2220D is 50.0, then the degree of \u2220C is ()", "Choices":"\n(A) 25\u00b0\n(B) 30\u00b0\n(C) 40\u00b0\n(D) 50\u00b0", "image": "gs://github-repo/prompts/prompt_optimizer/mathvista_dataset/images/643.jpg", "target": "25\u00b0"}
```
🖱️ Click `Generate`.
### 2. Dataset Creation
Download a copy of the dataset. Then upload this file in the application.
**Dataset Name:** `mathvista`
You can preview the dataset at the bottom of the page.
To download the dataset, run this command:
```bash
gsutil cp gs://github-repo/prompts/prompt_optimizer/mathvista_dataset/mathvista_input.jsonl .
```
### 3. Evaluation
*Note: Ensure you have completed Step 2 to create a dataset before proceeding here.*
We will now run an evaluation, prior to doing any tweaking to get a baseline.
- **Existing Dataset:** 'mathvista'
- **Dataset File:** 'mathvista_input.jsonl'
- **Number of Samples:** '100'
- **Ground Truth Column Name:** 'target'
- **Existing Prompt:** 'math_prompt_test'
- **Version:** '1'
**Note:** if the prompt is not in the list refresh the page.
Click Load Prompt, and Upload and Get Response... ⏰ Wait!!
Review the responses.
- **Model-Based:** 'question-answering-quality'
- **Model:** 'gemini-2.5-pro'
Launch the Eval... ⏰ Wait!!
View the Evaluation Results, and save to prompt records. This will save this initial version to the prompt records for the baseline.
### 4. One-Click Refiner
Use this for a quick, zero-data upgrade to your prompt.
- **Select Existing Prompt:** 'math_prompt_test'
- **Select Version:** '1'
🖱️ Click **Load Prompt**.
- **Target Model:** 'gemini-2.0-flash-001' (or your preferred model)
- **Tone:** 'Professional'
🖱️ Click **Auto-Suggest Directives**.
🖱️ Click **Optimize Now**.
Review the **Optimized Result** and the **Insights** (why it changed). If satisfied, click **Save as New Version**.
### (Optional) Run new Evaluation
Navigate back to **Evaluation** and run an evaluation similar to step 3, but load **Version 2** of the prompt.
### 5. Performance Tuner
Use this for data-driven optimization using a hill-climbing algorithm.
- **Select Dataset:** 'mathvista'
- **Select File:** 'mathvista_input.jsonl'
🖱️ Click **Load Dataset**.
- **Select Prompt:** 'math_prompt_test'
- **Select Version:** '1'
🖱️ Click **Load Prompt**.
- **Evaluation Metrics:** 'question_answering_correctness' (default)
- **Target Model:** 'gemini-2.5-flash'
🖱️ Click **Start Optimization Job**.
⏰ Wait!! This may take some time as it runs multiple iterations.
Once complete, click **Load Results** to see the **Score Jump** and the **Winner System Instruction**. You can test the winner on a blind case before clicking **Export Final Prompt** to save it as a new version.
### (Optional) Run new Evaluation
Navigate back to **Evaluation** and run an evaluation similar to step 3, but load **Version 2** of the prompt.
### 6. Prompt Optimization
🔧 Set-Up Prompt Optimization using Agent Platform's batch optimization service.
- **Target Model:** 'gemini-2.5-flash'
- **Existing Prompt:** 'math_prompt_test'
- **Version:** '1'
🖱️ Click Load Prompt.
- **Select Existing Dataset:** 'mathvista'
- **Select the File:** 'mathvista_input.jsonl'
🖱️ Click Load Dataset.
Preview the dataset.
🖱️ Click Start Optimization.
**Note:** If Interested in viewing the progress, Navigate to https://console.cloud.google.com/vertex-ai/training/custom-jobs
⏰ Wait!! This step will take about 20-min to run.
### 7. Prompt Optimization Results
View the Optimization Results.
The last run will be shown at the top of the screen. Pick this from the dropdown menu:
![image.png](assets/prompt_optimization_result.png)
Review the results and select the highest scoring version and copy the instruction.
### 8. Navigate Back to Prompt for New Version
Load your existing prompt from before.
📋 Paste your new instructions from the prompt optimizer, and save new version.
### 9. Run new Evaluation
Repeat step 3 with your new version.
### 10. View the Records
Navigate to the leaderboard and load the results.
## License
```
Copyright 2025 Google LLC
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
https://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language
```
Binary file not shown.

After

Width:  |  Height:  |  Size: 2.0 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.5 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 26 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 127 KiB

+54
View File
@@ -0,0 +1,54 @@
## Copyright 2025 Google LLC
## Licensed under the Apache License, Version 2.0 (the "License");
## you may not use this file except in compliance with the License.
## You may obtain a copy of the License at
## https://www.apache.org/licenses/LICENSE-2.0
## Unless required by applicable law or agreed to in writing, software
## distributed under the License is distributed on an "AS IS" BASIS,
## WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
## See the License for the specific language
"""The main landing page for the LLM EvalKit Streamlit application."""
import streamlit as st
from dotenv import load_dotenv
load_dotenv("src/.env")
def main() -> None:
"""Renders the main landing page of the application."""
st.set_page_config(
page_title="LLM EvalKit",
layout="wide",
initial_sidebar_state="expanded",
page_icon="assets/favicon.ico",
)
st.title("Welcome to the LLM EvalKit")
st.markdown(
"A suite of tools for managing, evaluating, and optimizing LLM prompts and datasets."
)
st.subheader("Getting Started")
st.markdown(
"""
This application helps you streamline your prompt engineering workflow.
Select a tool from the sidebar on the left to begin.
**Available Tools:**
* **Prompt Management:** Create, test, and manage your prompts.
* **Dataset Creation:** Create evaluation datasets from CSV files.
* **Simple Evaluation:** Run simple evaluations on your prompts.
* **Evaluation Human Judge:** Manually rate model responses for evaluation.
* **Prompt Optimization:** Optimize your prompts for better performance.
* **Prompt Optimization Results:** View the results of prompt optimization runs.
* **Prompt Records:** View and manage your prompt records.
"""
)
st.caption("LLM EvalKit | Home")
if __name__ == "__main__":
main()
@@ -0,0 +1,590 @@
# Copyright 2025 Google LLC
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
# https://www.apache.org/licenses/LICENSE-2.0
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language
"""Streamlit user interface for managing prompts in the LLM EvalKit.
This page provides a comprehensive interface for prompt engineering, allowing users
to create, load, edit, and test prompts that are stored and versioned in a
backend service (e.g., Google Cloud's Vertex AI Prompt Management).
The page is divided into two main sections:
1. **Create New Prompt**: A form to define a new prompt from scratch, including
its name, text, model, system instructions, and other metadata. Users can
test the prompt with sample input before saving it.
2. **Load & Edit Prompt**: A section to load existing prompts and their specific
versions. Users can modify the loaded prompt's details and save the changes
as a new version, facilitating iterative development and A/B testing.
Helper functions handle JSON parsing, data type conversions, and interactions
with the `gcp_prompt` object, which abstracts the backend communication.
"""
import json
import logging
from typing import Any
import streamlit as st
from dotenv import load_dotenv
from src.gcp_prompt import GcpPrompt as gcp_prompt
from vertexai.preview import prompts
# --- Initial Configuration ---
load_dotenv("src/.env")
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
# --- Constants ---
AVAILABLE_PROMPT_TASKS = [
"Classification",
"Summarization",
"Translation",
"Creative Writing",
"Q&A",
]
# --- Helper Functions ---
def _parse_json_input(json_string: str, field_name: str) -> dict[str, Any] | None:
"""Safely parses a JSON string from a text area.
Cleans the input string to handle common copy-paste errors and displays
an error in the Streamlit UI if parsing fails.
Args:
json_string: The raw string from a Streamlit text_area.
field_name: The user-facing name of the field for error messages.
Returns:
A dictionary if parsing is successful, otherwise None.
"""
if not json_string:
return None
try:
# Clean up common copy-paste issues like smart quotes and newlines
json_string_cleaned = (
json_string.replace("", "'")
.replace("\n", " ")
.replace("\t", " ")
.replace("\r", "")
)
return json.loads(json_string_cleaned)
except json.JSONDecodeError as e:
st.error(f"Invalid JSON format for {field_name}: {e}")
return None
def _apply_generation_config_typing(config: dict[str, Any]) -> dict[str, Any]:
"""Applies correct data types to generation config parameters.
Streamlit text inputs return strings, but the underlying API requires
specific types (e.g., float for temperature). This function converts
common configuration values to their expected types.
Args:
config: The generation configuration dictionary with string values.
Returns:
The configuration dictionary with values cast to the correct types.
"""
if "temperature" in config:
config["temperature"] = float(config["temperature"])
if "top_p" in config:
config["top_p"] = float(config["top_p"])
if "max_output_tokens" in config:
config["max_output_tokens"] = int(config["max_output_tokens"])
return config
# --- Handlers for "Create New Prompt" Tab ---
def _handle_save_new_prompt() -> None:
"""Validates inputs and saves a new prompt.
Retrieves all necessary data from the Streamlit session state for the
"Create New Prompt" tab, validates that required fields are filled,
constructs the prompt object, and calls the backend service to save it.
Displays success or error messages in the UI.
"""
required_fields = {
"new_prompt_name": "Prompt Name",
"new_prompt_data": "Prompt Text",
"new_model_name": "Model Name",
"new_system_instructions": "System Instructions",
}
for key, name in required_fields.items():
if not st.session_state.get(key):
st.warning(f"Please enter a value for {name}.")
return
prompt_obj = st.session_state.local_prompt
prompt_obj.prompt_to_run.prompt_name = st.session_state.new_prompt_name
prompt_obj.prompt_to_run.prompt_data = st.session_state.new_prompt_data
prompt_obj.prompt_to_run.model_name = st.session_state.new_model_name.strip()
prompt_obj.prompt_to_run.system_instruction = (
st.session_state.new_system_instructions
)
response_schema = _parse_json_input(
st.session_state.new_response_schema, "Response Schema"
)
generation_config = _parse_json_input(
st.session_state.new_generation_config, "Generation Config"
)
if generation_config:
generation_config = _apply_generation_config_typing(generation_config)
if response_schema:
generation_config["response_schema"] = response_schema
prompt_obj.prompt_meta["generation_config"] = generation_config
if response_schema:
prompt_obj.prompt_meta["response_schema"] = response_schema
prompt_obj.prompt_meta["meta_tags"] = st.session_state.new_meta_tags
try:
logger.info("Saving new prompt...")
prompt_meta_info = prompt_obj.save_prompt(check_existing=True)
logger.info("Prompt saved successfully: %s", prompt_meta_info)
st.success("Prompt saved successfully!")
except Exception as e:
logger.error("Failed to save prompt: %s", e, exc_info=True)
st.error(f"Failed to save prompt: {e}")
def _handle_generate_test_for_new() -> None:
"""Generates a test response for the new prompt form.
Takes the user-provided sample input and the current prompt configuration
from the "Create" tab, sends it to the model for a response, and displays
the output in the UI. This allows for quick testing before saving.
"""
user_input_str = st.session_state.new_sample_user_input
if not user_input_str:
st.warning("Please provide sample user input to generate a response.")
return
sample_user_input = _parse_json_input(user_input_str, "User Input")
if sample_user_input is None:
return
try:
prompt_obj = st.session_state.local_prompt
prompt_obj.prompt_to_run.prompt_data = st.session_state.new_prompt_data
prompt_obj.prompt_to_run.model_name = st.session_state.new_model_name.strip()
prompt_obj.prompt_to_run.system_instruction = (
st.session_state.new_system_instructions
)
prompt_obj.prompt_meta["sample_user_input"] = sample_user_input
with st.spinner("Generating response..."):
response = prompt_obj.generate_response(sample_user_input)
st.session_state.new_sample_output = response
st.success("Prompt response generated!")
except Exception as e:
logger.error("Error during test generation: %s", e, exc_info=True)
st.error(f"An error occurred during generation: {e}")
# --- Handlers for "Load & Edit Prompt" Tab ---
def _populate_ui_from_prompt() -> None:
"""Populates session state for UI widgets from the loaded prompt object.
After a prompt is loaded from the backend, this function takes the data
from the `gcp_prompt` object and sets the corresponding values in the
Streamlit session state. This updates the "Load & Edit" tab's input
widgets to display the loaded prompt's information.
"""
prompt_obj = st.session_state.local_prompt
st.session_state.edit_prompt_name = prompt_obj.prompt_to_run.prompt_name
st.session_state.edit_prompt_data = prompt_obj.prompt_to_run.prompt_data
st.session_state.edit_model_name = prompt_obj.prompt_to_run.model_name.split("/")[
-1
]
st.session_state.edit_system_instructions = (
prompt_obj.prompt_to_run.system_instruction
)
st.session_state.edit_response_schema = json.dumps(
prompt_obj.prompt_meta.get("response_schema", {}), indent=2
)
st.session_state.edit_generation_config = json.dumps(
prompt_obj.prompt_meta.get("generation_config", {}), indent=2
)
st.session_state.edit_meta_tags = prompt_obj.prompt_meta.get("meta_tags", [])
st.session_state.edit_sample_user_input = json.dumps(
prompt_obj.prompt_meta.get("sample_user_input", {}), indent=2
)
st.session_state.edit_sample_output = "" # Clear previous output
def _handle_load_prompt() -> None:
"""Loads the selected prompt and version and populates the UI.
Triggered by the 'Load Prompt' button. It retrieves the selected prompt
name and version from the UI, calls the backend to fetch the data,
and then uses `_populate_ui_from_prompt` to display it.
"""
if not st.session_state.get("selected_prompt") or not st.session_state.get(
"selected_version"
):
st.warning("Please select both a prompt and a version to load.")
return
prompt_name = st.session_state.selected_prompt
prompt_id = st.session_state.local_prompt.existing_prompts[prompt_name]
version_id = st.session_state.selected_version
try:
with st.spinner(f"Loading version '{version_id}' of prompt '{prompt_name}'..."):
st.session_state.local_prompt.load_prompt(
prompt_id, prompt_name, version_id
)
logger.info(
"Successfully loaded prompt '%s' version '%s'.", prompt_name, version_id
)
_populate_ui_from_prompt()
st.success(f"Loaded prompt '{prompt_name}' (Version: {version_id}).")
except Exception as e:
logger.error("Failed to load prompt: %s", e, exc_info=True)
st.error(f"Failed to load prompt: {e}")
def _handle_save_edited_prompt() -> None:
"""Validates inputs and saves the current prompt config as a new version.
Similar to saving a new prompt, but it takes the data from the "Edit" tab's
widgets. It saves the current configuration as a new version of the
already existing prompt.
"""
if not st.session_state.get("edit_prompt_name"):
st.warning("Cannot save. Please load a prompt first.")
return
required_fields = {
"edit_prompt_data": "Prompt Text",
"edit_model_name": "Model Name",
"edit_system_instructions": "System Instructions",
}
for key, name in required_fields.items():
if not st.session_state.get(key):
st.warning(f"Please ensure '{name}' is not empty.")
return
prompt_obj = st.session_state.local_prompt
prompt_obj.prompt_to_run.prompt_name = st.session_state.edit_prompt_name
prompt_obj.prompt_to_run.prompt_data = st.session_state.edit_prompt_data
prompt_obj.prompt_to_run.model_name = st.session_state.edit_model_name.strip()
prompt_obj.prompt_to_run.system_instruction = (
st.session_state.edit_system_instructions
)
response_schema = _parse_json_input(
st.session_state.edit_response_schema, "Response Schema"
)
generation_config = _parse_json_input(
st.session_state.edit_generation_config, "Generation Config"
)
if generation_config:
generation_config = _apply_generation_config_typing(generation_config)
if response_schema:
generation_config["response_schema"] = response_schema
prompt_obj.prompt_meta["generation_config"] = generation_config
if response_schema:
prompt_obj.prompt_meta["response_schema"] = response_schema
prompt_obj.prompt_meta["meta_tags"] = st.session_state.edit_meta_tags
try:
with st.spinner("Saving as new version..."):
prompt_meta_info = prompt_obj.save_prompt(check_existing=False)
logger.info("Prompt saved successfully: %s", prompt_meta_info)
st.success("Saved as a new version successfully!")
st.session_state.local_prompt.refresh_prompt_cache()
except Exception as e:
logger.error("Failed to save prompt: %s", e, exc_info=True)
st.error(f"Failed to save prompt: {e}")
def _handle_generate_test_for_edit() -> None:
"""Generates a test response for the edited prompt.
Allows users to test changes made in the "Edit" tab before saving them
as a new version. It uses the current values in the UI fields to generate
a response from the model.
"""
if not st.session_state.get("edit_prompt_name"):
st.warning("Please load a prompt before generating a response.")
return
user_input_str = st.session_state.get("edit_sample_user_input", "")
if not user_input_str:
st.warning("Please provide sample user input to generate a response.")
return
sample_user_input = _parse_json_input(user_input_str, "Sample User Input")
if sample_user_input is None:
return
try:
prompt_obj = st.session_state.local_prompt
prompt_obj.prompt_to_run.prompt_data = st.session_state.edit_prompt_data
prompt_obj.prompt_to_run.system_instruction = (
st.session_state.edit_system_instructions
)
prompt_obj.prompt_meta["sample_user_input"] = sample_user_input
with st.spinner("Generating response..."):
response = prompt_obj.generate_response(sample_user_input)
st.session_state.edit_sample_output = response
st.success("Prompt response generated!")
except Exception as e:
logger.error("Error during test generation: %s", e, exc_info=True)
st.error(f"An error occurred during generation: {e}")
# --- UI Rendering Functions ---
def render_create_tab() -> None:
"""Renders the UI components for the 'Create New Prompt' tab.
This function defines and lays out all the Streamlit widgets (text inputs,
buttons, etc.) for the prompt creation workflow.
"""
st.subheader("1. Define Prompt Details")
st.text_input(
"**Prompt Name**",
key="new_prompt_name",
placeholder="e.g., customer_sentiment_classifier_v1",
help="A unique name to identify your prompt.",
)
st.text_area(
"**Prompt Text**",
key="new_prompt_data",
height=150,
placeholder="e.g., Classify the sentiment of the following text: {customer_review}",
help="The core text of your prompt. Use curly braces `{}` for variables.",
)
st.text_input(
"**Model Name**",
key="new_model_name",
placeholder="gemini-2.5-pro-001",
help="The specific model version to use (e.g., gemini-2.5-pro).",
)
st.text_area(
"**System Instructions**",
key="new_system_instructions",
height=300,
placeholder="e.g., You are an expert in sentiment analysis...",
help="Optional instructions to guide the model's behavior.",
)
st.multiselect(
"**Prompt Task**",
options=AVAILABLE_PROMPT_TASKS,
key="new_meta_tags",
help="Select the most appropriate task type for this prompt.",
)
st.text_area(
"**Response Schema (JSON)**",
key="new_response_schema",
height=150,
placeholder='{\n "type": "object", ... \n}',
help="Define the desired JSON structure for the model's output.",
)
st.text_area(
"**Generation Config (JSON)**",
key="new_generation_config",
height=150,
placeholder='{\n "temperature": 0.2, ... \n}',
help="A dictionary of generation parameters.",
)
if st.button(
"Save Prompt", type="primary", use_container_width=True, key="save_new"
):
_handle_save_new_prompt()
st.divider()
st.subheader("2. Test Your Prompt")
st.markdown("You can test your prompt here before saving.")
st.text_area(
"**Sample User Input (JSON)**",
key="new_sample_user_input",
height=150,
placeholder='{\n "customer_review": "The product was amazing!"\n}',
help="A JSON object where keys match the variables in your prompt text.",
)
if st.button("Generate Test Response", use_container_width=True, key="test_new"):
_handle_generate_test_for_new()
st.text_area(
"**Test Output**",
key="new_sample_output",
height=150,
placeholder="The model's response will be displayed here.",
disabled=True,
)
def render_edit_tab() -> None:
"""Renders the UI components for the 'Load & Edit Prompt' tab.
This function defines and lays out all the Streamlit widgets for loading,
editing, and versioning existing prompts.
"""
st.subheader("1. Load Prompt")
if st.button("Refresh List"):
with st.spinner("Refreshing..."):
st.session_state.local_prompt.refresh_prompt_cache()
st.toast("Prompt list refreshed.")
col1, col2 = st.columns(2)
with col1:
selected_prompt_name = st.selectbox(
"Select Existing Prompt",
options=st.session_state.local_prompt.existing_prompts.keys(),
placeholder="Select Prompt...",
key="selected_prompt",
help="Choose the prompt you want to load.",
)
with col2:
versions = []
if selected_prompt_name:
try:
prompt_id = st.session_state.local_prompt.existing_prompts[
selected_prompt_name
]
versions = [v.version_id for v in prompts.list_versions(prompt_id)]
except Exception as e:
st.error(f"Could not fetch versions: {e}")
st.selectbox(
"Select Version",
options=versions,
placeholder="Select Version...",
key="selected_version",
help="Choose the specific version to load.",
)
st.button(
"Load Prompt",
on_click=_handle_load_prompt,
use_container_width=True,
type="primary",
)
st.divider()
st.subheader("2. Edit Prompt Details")
st.text_input("Prompt Name", key="edit_prompt_name", disabled=True)
st.text_area("Prompt Text", key="edit_prompt_data", height=150)
st.text_input("Model Name", key="edit_model_name")
st.text_area("System Instructions", key="edit_system_instructions", height=300)
st.multiselect("Prompt Task", options=AVAILABLE_PROMPT_TASKS, key="edit_meta_tags")
col_schema, col_config = st.columns(2)
with col_schema:
st.text_area("Response Schema (JSON)", key="edit_response_schema", height=200)
with col_config:
st.text_area(
"Generation Config (JSON)", key="edit_generation_config", height=200
)
if st.button(
"Save as New Version", type="primary", use_container_width=True, key="save_edit"
):
_handle_save_edited_prompt()
st.divider()
st.subheader("3. Test Your Prompt")
st.text_area("Sample User Input (JSON)", key="edit_sample_user_input", height=150)
if st.button("Generate Test Response", use_container_width=True, key="test_edit"):
_handle_generate_test_for_edit()
st.text_area(
"Test Output",
key="edit_sample_output",
height=150,
placeholder="The model's response will be displayed here.",
disabled=True,
)
# --- Main Application ---
def main() -> None:
"""Renders the main Prompt Management page.
Sets the page configuration, initializes the session state (including the
`gcp_prompt` object and UI field defaults), and renders the main title
and tabbed layout for creating and editing prompts.
"""
st.set_page_config(
layout="wide",
page_title="Prompt Management",
page_icon="assets/favicon.ico",
)
# Initialize session state object and UI fields
if "local_prompt" not in st.session_state:
st.session_state.local_prompt = gcp_prompt()
ui_fields = {
"new_prompt_name": "",
"new_prompt_data": "",
"new_model_name": "",
"new_system_instructions": "",
"new_response_schema": "",
"new_generation_config": "",
"new_meta_tags": [],
"new_sample_user_input": "",
"new_sample_output": "",
"edit_prompt_name": "",
"edit_prompt_data": "",
"edit_model_name": "",
"edit_system_instructions": "",
"edit_response_schema": "",
"edit_generation_config": "",
"edit_meta_tags": [],
"edit_sample_user_input": "",
"edit_sample_output": "",
}
for field, default_val in ui_fields.items():
if field not in st.session_state:
st.session_state[field] = default_val
st.title("Prompt Management")
st.markdown(
"Create new prompts or load, edit, and test existing ones from the Prompt Management service."
)
st.divider()
# Use st.radio to create stateful tabs that persist across reruns.
# This prevents the UI from resetting to the first tab on every interaction.
selected_tab = st.radio(
"Select Action",
["Create New Prompt", "Load & Edit Prompt"],
key="prompt_management_tab",
horizontal=True,
label_visibility="collapsed",
)
if selected_tab == "Create New Prompt":
render_create_tab()
elif selected_tab == "Load & Edit Prompt":
render_edit_tab()
st.caption("LLM EvalKit | Prompt Management")
if __name__ == "__main__":
main()
@@ -0,0 +1,250 @@
# Copyright 2025 Google LLC
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
# https://www.apache.org/licenses/LICENSE-2.0
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language
"""Streamlit page for creating and managing datasets in Google Cloud Storage."""
import logging
import os
import streamlit as st
from dotenv import load_dotenv
from google.cloud import storage
from streamlit.runtime.uploaded_file_manager import UploadedFile
# Load environment variables from .env file
load_dotenv("src/.env")
# Configure logging to the console
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
@st.cache_data(ttl=300)
def get_existing_datasets(
_storage_client: storage.Client, bucket_name: str
) -> list[str]:
"""Lists 'directories' in GCS under the 'datasets/' prefix.
These directories represent the existing datasets.
"""
if not bucket_name or not _storage_client:
return []
bucket = _storage_client.bucket(bucket_name)
prefix = "datasets/"
retrieved_prefixes = set()
try:
# Explicitly iterate through pages for robustness.
iterator = bucket.list_blobs(prefix=prefix, delimiter="/")
for page in iterator.pages:
retrieved_prefixes.update(page.prefixes)
# The retrieved prefixes are the "subdirectories".
# e.g., {'datasets/my_dataset_1/', 'datasets/my_dataset_2/'}
dir_names = []
for p in retrieved_prefixes:
# Extract 'my_dataset_1' from 'datasets/my_dataset_1/'
name = p[len(prefix) :].strip("/")
if name:
dir_names.append(name)
logger.info(f"Found datasets: {dir_names}")
return sorted(dir_names)
except Exception as e:
st.error(f"Error listing datasets from GCS: {e}")
logger.error(f"Error in get_existing_datasets: {e}", exc_info=True)
return []
def _handle_upload(
storage_client: storage.Client,
bucket_name: str,
dataset_name: str,
uploaded_file: UploadedFile,
) -> None:
"""Handles the logic of uploading a file to GCS."""
if not all([storage_client, bucket_name, dataset_name, uploaded_file]):
st.warning("Missing required information for upload.")
return
try:
file_name = uploaded_file.name
content_type = "text/plain" # Default
if file_name.endswith(".csv"):
content_type = "text/csv"
elif file_name.endswith(".json"):
content_type = "application/json"
elif file_name.endswith(".jsonl"):
content_type = "application/x-jsonlines"
blob_path = f"datasets/{dataset_name}/{uploaded_file.name}"
bucket = storage_client.bucket(bucket_name)
blob = bucket.blob(blob_path)
with st.spinner(f"Uploading '{uploaded_file.name}' to '{dataset_name}'..."):
blob.upload_from_string(uploaded_file.getvalue(), content_type=content_type)
st.success(
f"Successfully uploaded '{uploaded_file.name}' to dataset '{dataset_name}'!"
)
logger.info(f"Uploaded file to gs://{bucket_name}/{blob_path}")
# Clear the cache for get_existing_datasets to reflect the new dataset if created
get_existing_datasets.clear()
st.rerun()
except Exception as e:
st.error(f"Failed to upload file: {e}")
logger.error(f"Error during GCS upload: {e}", exc_info=True)
def _ensure_datasets_folder_exists(
storage_client: storage.Client, bucket_name: str
) -> None:
"""Ensures the 'datasets/' folder exists by creating a placeholder object if needed.
This helps it appear in the GCS UI even when empty.
"""
if not storage_client or not bucket_name:
return
try:
bucket = storage_client.bucket(bucket_name)
blob = bucket.blob("datasets/")
if not blob.exists():
blob.upload_from_string("", content_type="application/x-directory")
logger.info(
f"Created placeholder for 'datasets/' folder in bucket '{bucket_name}'."
)
except Exception as e:
# This is not a critical failure, so just log a warning.
logger.warning(f"Could not ensure 'datasets/' folder exists: {e}")
def main() -> None:
"""Renders the Dataset Creation page."""
st.set_page_config(
layout="wide", page_title="Dataset Management", page_icon="assets/favicon.ico"
)
# --- Initialize Session State & GCS Client ---
if "storage_client" not in st.session_state:
try:
st.session_state.storage_client = storage.Client()
except Exception as e:
st.error(f"Could not connect to Google Cloud Storage: {e}")
st.stop()
BUCKET_NAME = os.getenv("BUCKET")
if not BUCKET_NAME:
st.error("BUCKET environment variable is not set. Please configure it in .env.")
st.stop()
# Ensure the base 'datasets/' folder exists for UI consistency
_ensure_datasets_folder_exists(st.session_state.storage_client, BUCKET_NAME)
st.title("Dataset Management")
st.markdown(
"Create new datasets or upload files (CSV, JSON, or JSONL) to existing ones. "
"A 'Dataset' is a folder in your GCS bucket used to group related evaluation files."
)
st.divider()
# --- Section 1: Upload File ---
st.subheader("1. Upload a File")
existing_datasets = get_existing_datasets(
st.session_state.storage_client, BUCKET_NAME
)
# Let user choose whether to create a new dataset or add to an existing one
upload_mode = st.radio(
"Choose an action:",
("Create a new dataset", "Add to an existing dataset"),
key="upload_mode",
horizontal=True,
)
dataset_name = ""
if upload_mode == "Create a new dataset":
dataset_name = st.text_input(
"Enter a name for the new dataset:",
key="new_dataset_name",
help="Use a descriptive name, e.g., 'sentiment_analysis_v1'.",
)
else:
dataset_name = st.selectbox(
"Select an existing dataset:",
options=existing_datasets,
key="selected_dataset_for_upload",
help="Choose the dataset folder to upload your file into.",
index=None,
placeholder="Select a dataset...",
)
uploaded_file = st.file_uploader(
"Select a file to upload",
type=["csv", "json", "jsonl"],
key="file_uploader",
)
if st.button("Upload to Cloud Storage", type="primary", use_container_width=True):
if not dataset_name:
st.warning("Please provide or select a dataset name.")
elif not uploaded_file:
st.warning("Please select a file to upload.")
else:
_handle_upload(
st.session_state.storage_client,
BUCKET_NAME,
dataset_name,
uploaded_file,
)
st.divider()
# --- Section 2: View Existing Datasets ---
st.subheader("2. View Existing Datasets")
with st.expander("Browse datasets and their contents", expanded=True):
selected_dataset_to_view = st.selectbox(
"Select a dataset to view its contents:",
options=existing_datasets,
key="selected_dataset_for_view",
index=None,
placeholder="Select a dataset...",
)
if selected_dataset_to_view:
prefix = f"datasets/{selected_dataset_to_view}/"
blobs = st.session_state.storage_client.list_blobs(
BUCKET_NAME, prefix=prefix
)
filenames = [
os.path.basename(b.name)
for b in blobs
if b.name.endswith((".csv", ".json", ".jsonl"))
]
if filenames:
st.write(f"**Files in '{selected_dataset_to_view}':**")
st.text_area(
"Files",
value="\n".join(filenames),
height=150,
disabled=True,
label_visibility="collapsed",
)
else:
st.info(f"No files found in the '{selected_dataset_to_view}' dataset.")
st.caption("LLM EvalKit | Dataset Management")
if __name__ == "__main__":
main()
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,321 @@
# Copyright 2025 Google LLC
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
# https://www.apache.org/licenses/LICENSE-2.0
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language
"""Streamlit user interface for the One-Click Refiner.
This page provides an interface to instantly upgrade a draft prompt into a
structured, production-ready instruction without managing any datasets.
"""
import json
import logging
import streamlit as st
from dotenv import load_dotenv
from src.gcp_prompt import GcpPrompt as gcp_prompt
from vertexai.generative_models import GenerationConfig, GenerativeModel
from vertexai.preview import prompts
load_dotenv("src/.env")
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
# --- Prompt Templates ---
META_PROMPT_TEMPLATE = """You are an expert prompt engineer. Your goal is to improve the user's draft prompt and system instructions into highly structured, production-ready iterations.
Ensure you include and follow these directives:
{custom_directives}
Ensure tone relates to the optional requested Tone: {tone}.
CRITICAL REQUIREMENTS:
- You MUST preserve all variable placeholders exactly as they appear (e.g., `{{{{query}}}}`, `{{{{target}}}}`). Note: the draft prompt might use curly brackets like `{{variable}}`. Do NOT strip them.
- You MUST preserve any multimodal tags exactly as they appear (e.g., `@@@image/jpeg`). Do not alter or remove image attachments.
Draft System Instructions:
{draft_system_instructions}
Draft Prompt:
{draft_prompt}
You must respond in pure JSON format with exactly three keys:
1. "optimized_system_instruction": A single string containing the rewritten system instructions.
2. "optimized_prompt": A single string containing the fully rewritten structured prompt template.
3. "insights": A list of strings explaining exactly what you changed and why.
"""
SUGGEST_DIRECTIVES_PROMPT = """Analyze the following draft prompt and system instructions. Suggest 3-5 specific prompt engineering best practices that would improve it. Focus on structure, constraints, format, clarity, and safety.
Return ONLY a markdown list of suggestions suitable to be used as instructions for another LLM prompt engineer. Do not include introductory text.
Draft System Instructions:
{draft_system_instructions}
Draft Prompt:
{draft_prompt}
"""
def initialize_session_state() -> None:
"""Initializes needed session state variables."""
if "local_prompt" not in st.session_state:
st.session_state.local_prompt = gcp_prompt()
if "ocr_directives" not in st.session_state:
st.session_state.ocr_directives = "1. Add a clear Role definition.\n2. Add specific Context to constrain the generator.\n3. Clarify output format expectations."
if "opt_sys" not in st.session_state:
st.session_state.opt_sys = ""
if "opt_prompt" not in st.session_state:
st.session_state.opt_prompt = ""
if "ocr_insights" not in st.session_state:
st.session_state.ocr_insights = None
def _handle_load_prompt():
"""Loads the selected prompt and version into the gcp_prompt object."""
if not st.session_state.get("selected_prompt") or not st.session_state.get(
"selected_version"
):
st.warning("Please select both a prompt and a version to load.")
return
prompt_name = st.session_state.selected_prompt
prompt_id = st.session_state.local_prompt.existing_prompts[prompt_name]
version_id = st.session_state.selected_version
try:
with st.spinner(f"Loading version '{version_id}' of prompt '{prompt_name}'..."):
st.session_state.local_prompt.load_prompt(
prompt_id, prompt_name, version_id
)
st.success(f"Loaded prompt '{prompt_name}' (Version: {version_id}).")
# Clear previous optimizations
st.session_state.opt_sys = ""
st.session_state.opt_prompt = ""
st.session_state.ocr_insights = None
except Exception as e:
logger.error("Failed to load prompt: %s", e, exc_info=True)
st.error(f"Failed to load prompt: {e}")
def _handle_auto_suggest():
"""Calls Agent Platform to automatically suggest prompt engineering directives."""
sys_inst = st.session_state.local_prompt.prompt_to_run.system_instruction or "None"
prompt_data = st.session_state.local_prompt.prompt_to_run.prompt_data or "None"
model_name = st.session_state.get("ocr_target_model", "gemini-2.5-pro")
if not model_name:
model_name = "gemini-2.5-pro"
try:
model = GenerativeModel(model_name)
prompt_text = SUGGEST_DIRECTIVES_PROMPT.format(
draft_system_instructions=sys_inst, draft_prompt=prompt_data
)
with st.spinner("Analyzing prompt and generating suggestions..."):
response = model.generate_content(prompt_text)
st.session_state.ocr_directives = response.text
except Exception as e:
logger.error("Error auto-suggesting directives: %s", e, exc_info=True)
st.error(f"Failed to generate suggestions: {e}")
def _handle_optimize():
"""Optimizes the loaded prompt using the meta-prompt and custom directives."""
sys_inst = st.session_state.local_prompt.prompt_to_run.system_instruction or "None"
prompt_data = st.session_state.local_prompt.prompt_to_run.prompt_data or "None"
directives = st.session_state.get("ocr_directives", "")
tone = st.session_state.get("ocr_tone", "Professional")
model_name = st.session_state.get("ocr_target_model", "gemini-2.5-pro")
if not model_name:
model_name = "gemini-2.5-pro"
try:
model = GenerativeModel(model_name)
prompt_text = META_PROMPT_TEMPLATE.format(
custom_directives=directives,
tone=tone,
draft_system_instructions=sys_inst,
draft_prompt=prompt_data,
)
with st.spinner("Optimizing..."):
response = model.generate_content(
prompt_text,
generation_config=GenerationConfig(
temperature=0.4, response_mime_type="application/json"
),
)
# Parse response
try:
res_obj = json.loads(response.text)
st.session_state.opt_sys = res_obj.get(
"optimized_system_instruction", ""
)
st.session_state.opt_prompt = res_obj.get("optimized_prompt", "")
st.session_state.ocr_insights = res_obj.get("insights", [])
st.success("Optimization Complete!")
except json.JSONDecodeError as e:
st.error(f"Failed to parse optimization output as JSON: {e}")
logger.error("Raw response: %s", response.text)
except Exception as e:
logger.error("Error optimizing prompt: %s", e, exc_info=True)
st.error(f"Failed to optimize prompt: {e}")
def _handle_save_new_version():
"""Saves the optimized prompt to the backend registry as a new version."""
prompt_obj = st.session_state.local_prompt
if not prompt_obj.prompt_to_run.prompt_name:
st.warning("No prompt is currently loaded to save.")
return
prompt_obj.prompt_to_run.prompt_data = st.session_state.opt_prompt
prompt_obj.prompt_to_run.system_instruction = st.session_state.opt_sys
try:
with st.spinner("Saving as new version..."):
prompt_obj.save_prompt(check_existing=False)
st.success("Successfully saved new optimized version to registry!")
prompt_obj.refresh_prompt_cache()
except Exception as e:
logger.error("Failed to save new version: %s", e, exc_info=True)
st.error(f"Failed to save prompt: {e}")
def main():
"""Renders the One-Click Refiner page layout."""
st.set_page_config(
layout="wide", page_title="One-Click Refiner", page_icon="assets/favicon.ico"
)
initialize_session_state()
st.title("One-Click Refiner")
st.markdown(
"Instantly upgrade a draft prompt into a structured, production-ready instruction without managing any datasets."
)
st.divider()
# SECTION 1: Load Existing Prompt
st.subheader("1. Load Prompt")
if st.button("Refresh List"):
with st.spinner("Refreshing..."):
st.session_state.local_prompt.refresh_prompt_cache()
st.toast("Prompt list refreshed.")
col1, col2 = st.columns(2)
with col1:
selected_prompt_name = st.selectbox(
"Select Existing Prompt",
options=st.session_state.local_prompt.existing_prompts.keys(),
placeholder="Select Prompt...",
key="selected_prompt",
)
with col2:
versions = []
if selected_prompt_name:
try:
prompt_id = st.session_state.local_prompt.existing_prompts[
selected_prompt_name
]
versions = [v.version_id for v in prompts.list_versions(prompt_id)]
except Exception as e:
st.error(f"Could not fetch versions: {e}")
st.selectbox(
"Select Version",
options=versions,
placeholder="Select Version...",
key="selected_version",
)
st.button("Load Prompt", on_click=_handle_load_prompt, type="primary")
st.divider()
p_data = st.session_state.local_prompt.prompt_to_run.prompt_data
if p_data:
# SECTION 2: Configuration
st.subheader("2. Configuration")
c1, c2 = st.columns(2)
with c1:
current_model = st.session_state.local_prompt.prompt_to_run.model_name
if current_model and "/" in current_model:
current_model = current_model.split("/")[-1]
st.text_input(
"Target Model",
value=current_model if current_model else "gemini-2.0-flash-001",
key="ocr_target_model",
)
with c2:
st.selectbox(
"Tone",
options=[
"Professional",
"Creative",
"Concise",
"Assertive",
"Friendly",
"None",
],
key="ocr_tone",
)
st.markdown("**Optimization Directives**")
st.text_area(
"Modify the guidelines the optimizer should follow:",
key="ocr_directives",
height=120,
)
st.button("✨ Auto-Suggest Directives", on_click=_handle_auto_suggest)
st.button("🚀 Optimize Now", on_click=_handle_optimize, type="primary")
st.divider()
# SECTION 3: Review
st.subheader("3. Review")
rev_c1, rev_c2 = st.columns(2)
with rev_c1:
st.markdown("### Original Draft")
st.text_area(
"System Instructions",
value=st.session_state.local_prompt.prompt_to_run.system_instruction
or "",
disabled=True,
height=200,
key="org_sys",
)
st.text_area(
"Prompt Data",
value=p_data or "",
disabled=True,
height=200,
key="org_prompt",
)
with rev_c2:
st.markdown("### Optimized Result")
st.text_area("System Instructions", key="opt_sys", height=200)
st.text_area("Prompt Data", key="opt_prompt", height=200)
if st.session_state.ocr_insights:
with st.expander("💡 Why this changed (Insights)", expanded=True):
for insight in st.session_state.ocr_insights:
st.markdown(f"- {insight}")
st.divider()
st.subheader("4. Action")
st.button(
"Save as New Version", on_click=_handle_save_new_version, type="primary"
)
if __name__ == "__main__":
main()
@@ -0,0 +1,604 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Streamlit page for Performance Tuner (Prompt Optimization)."""
import json
import logging
import os
from argparse import Namespace
from datetime import datetime
import pandas as pd
import streamlit as st
from dotenv import load_dotenv
from etils import epath
from google.cloud import aiplatform, storage
from src import vapo_lib
from src.gcp_prompt import GcpPrompt as gcp_prompt
from vertexai.evaluation import MetricPromptTemplateExamples
from vertexai.preview import prompts
load_dotenv("src/.env")
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
TARGET_MODELS = [
"gemini-2.5-pro",
"gemini-2.5-flash",
"gemini-2.5-flash-lite",
"gemini-2.0-flash",
"gemini-2.0-flash-001",
"gemini-2.0-flash-lite",
"gemini-2.0-flash-lite-001",
]
def initialize_session_state() -> None:
if "op_id" not in st.session_state:
st.session_state.op_id = vapo_lib.get_id()
if "local_prompt" not in st.session_state:
st.session_state.local_prompt = gcp_prompt()
if "storage_client" not in st.session_state:
st.session_state["storage_client"] = storage.Client()
if "data_uris" not in st.session_state:
st.session_state["data_uris"] = refresh_bucket()
if "dataset" not in st.session_state:
st.session_state["dataset"] = None
if "cached_data_files" not in st.session_state:
st.session_state.cached_data_files = {}
if "last_selected_dataset_for_cache" not in st.session_state:
st.session_state.last_selected_dataset_for_cache = None
if "tuner_launched_job" not in st.session_state:
st.session_state.tuner_launched_job = None
if "tuner_run_uri" not in st.session_state:
st.session_state.tuner_run_uri = None
if "tuner_winning_template" not in st.session_state:
st.session_state.tuner_winning_template = None
def refresh_bucket() -> list[str]:
logger.info("Bucket: %s", os.getenv("BUCKET"))
bucket = st.session_state.storage_client.bucket(os.getenv("BUCKET"))
blobs = bucket.list_blobs()
data_uris = []
for i in blobs:
if i.name.split("/")[0] == "datasets" and (
i.name.endswith(".csv") or i.name.endswith(".jsonl")
):
data_uris.append(f"gs://{i.bucket.name}/{i.name}")
return data_uris
def get_optimization_args(
input_optimization_data_file_uri,
output_optimization_run_uri,
target_model,
selected_metrics,
target_qps=1.0,
optimizer_qps=1.0,
eval_qps=1.0,
data_limit=10,
):
response_schema_str = st.session_state.local_prompt.prompt_meta.get(
"response_schema", "{}"
)
try:
response_schema = (
json.loads(response_schema_str)
if isinstance(response_schema_str, str)
else response_schema_str
)
except json.JSONDecodeError:
response_schema = {}
response_mime_type = "application/json" if response_schema else "text/plain"
response_schema_arg = response_schema if response_schema else ""
has_multimodal = False
if (
st.session_state.dataset is not None
and "image" in st.session_state.dataset.columns
):
has_multimodal = True
metrics = (
selected_metrics if selected_metrics else ["question_answering_correctness"]
)
weights = [1.0 for _ in metrics]
return Namespace(
system_instruction=st.session_state.local_prompt.prompt_to_run.system_instruction,
prompt_template=(
f"{st.session_state.local_prompt.prompt_to_run.prompt_data}"
"\n\tAnswer: {target}"
),
target_model=target_model,
optimization_mode="instruction",
eval_metrics_types=metrics,
eval_metrics_weights=weights,
aggregation_type="weighted_sum",
input_data_path=input_optimization_data_file_uri,
output_path=f"gs://{output_optimization_run_uri}",
project=os.getenv("PROJECT_ID"),
num_steps=5,
num_demo_set_candidates=10,
demo_set_size=3,
target_model_location="us-central1",
source_model="",
source_model_location="",
target_model_qps=target_qps,
optimizer_model_qps=optimizer_qps,
eval_qps=eval_qps,
source_model_qps="",
response_mime_type=response_mime_type,
response_schema=response_schema_arg,
language="English",
placeholder_to_content=json.loads("{}"),
data_limit=data_limit,
translation_source_field_name="",
has_multimodal_inputs=has_multimodal,
)
def check_job_status(job_name: str, project_id: str, location: str) -> str:
client_options = {"api_endpoint": f"{location}-aiplatform.googleapis.com"}
client = aiplatform.gapic.JobServiceClient(client_options=client_options)
parent = f"projects/{project_id}/locations/{location}"
response = client.list_custom_jobs(parent=parent)
for job in response:
if job.display_name == job_name:
return job.state.name
return "NOT_FOUND"
def write_record(metrics, prompt_name, version, system_instruction):
bucket_name = os.getenv("BUCKET")
if not bucket_name:
return
record = {
"timestamp": datetime.now().isoformat(),
"prompt_name": prompt_name,
"prompt_version": version,
"system_instruction": system_instruction,
"scores": metrics,
}
blob_name = f"records/{prompt_name}_{version}_{datetime.now().strftime('%Y%m%d%H%M%S')}.json"
bucket = st.session_state.storage_client.bucket(bucket_name)
blob = bucket.blob(blob_name)
blob.upload_from_string(json.dumps([record], indent=2))
logger.info(f"Saved optimized record to gs://{bucket_name}/{blob_name}")
def main() -> None:
st.set_page_config(
layout="wide", page_title="Performance Tuner", page_icon="assets/favicon.ico"
)
initialize_session_state()
st.title("Performance Tuner")
st.markdown(
"Optimize your prompt's System Instructions using data-driven iteration to maximize metric performance."
)
# 1. & 2. Data Setup & Template Definition
st.header("1. Data & Prompt Setup")
col1, col2 = st.columns(2)
with col1:
st.subheader("Data Setup")
data_sets = list({i.split("/")[4] for i in st.session_state.data_uris})
st.selectbox("Select Dataset", options=[None, *data_sets], key="tuner_dataset")
if st.session_state.tuner_dataset:
if (
st.session_state.tuner_dataset
!= st.session_state.last_selected_dataset_for_cache
or st.session_state.tuner_dataset
not in st.session_state.cached_data_files
):
bucket = st.session_state.storage_client.bucket(os.getenv("BUCKET"))
prefix = f"datasets/{st.session_state.tuner_dataset}/"
blobs_iterator = bucket.list_blobs(prefix=prefix)
current_dataset_files = [
blob.name[len(prefix) :]
for blob in blobs_iterator
if (blob.name.endswith(".csv") or blob.name.endswith(".jsonl"))
and not blob.name.endswith("/")
]
st.session_state.cached_data_files[st.session_state.tuner_dataset] = (
sorted(set(current_dataset_files))
)
st.session_state.last_selected_dataset_for_cache = (
st.session_state.tuner_dataset
)
files = st.session_state.cached_data_files.get(
st.session_state.tuner_dataset, []
)
st.selectbox(
"Select File (.csv or .jsonl)", options=[None, *files], key="tuner_file"
)
if st.button("Load Dataset", key="tuner_load_data"):
if st.session_state.tuner_file:
gcs_uri = f"gs://{os.getenv('BUCKET')}/datasets/{st.session_state.tuner_dataset}/{st.session_state.tuner_file}"
if st.session_state.tuner_file.endswith(".jsonl"):
st.session_state.dataset = pd.read_json(gcs_uri, lines=True)
else:
st.session_state.dataset = pd.read_csv(gcs_uri)
st.success(f"Loaded {len(st.session_state.dataset)} rows.")
with col2:
st.subheader("Template Definition")
st.selectbox(
"Select Prompt",
options=st.session_state.local_prompt.existing_prompts.keys(),
placeholder="Select Prompt...",
key="tuner_prompt",
)
if st.session_state.tuner_prompt:
st.session_state.local_prompt.prompt_meta["name"] = (
st.session_state.tuner_prompt
)
versions = [
i.version_id
for i in prompts.list_versions(
st.session_state.local_prompt.existing_prompts[
st.session_state.tuner_prompt
]
)
]
st.selectbox(
"Select Version",
options=versions,
placeholder="Select Version...",
key="tuner_version",
)
if st.button("Load Prompt", key="tuner_load_prompt"):
if st.session_state.tuner_prompt and st.session_state.tuner_version:
st.session_state.local_prompt.load_prompt(
st.session_state.local_prompt.existing_prompts[
st.session_state.tuner_prompt
],
st.session_state.tuner_prompt,
st.session_state.tuner_version,
)
st.success("Prompt loaded successfully.")
if st.session_state.local_prompt.prompt_to_run.system_instruction:
with st.expander("View Loaded Prompt Details", expanded=False):
st.text_area(
"System Instruction",
st.session_state.local_prompt.prompt_to_run.system_instruction,
disabled=True,
height=100,
)
st.text_area(
"Prompt Template",
st.session_state.local_prompt.prompt_to_run.prompt_data,
disabled=True,
height=100,
)
st.divider()
# 3. Metric Selection
st.header("2. Metric Selection")
metric_names = MetricPromptTemplateExamples.list_example_metric_names()
computation_metrics = [
"bleu",
"rouge_1",
"rouge_2",
"rouge_l",
"rouge_l_sum",
"exact_match",
"question_answering_correctness",
]
all_metrics = list(set(metric_names + computation_metrics))
selected_metrics = st.multiselect(
"Select Evaluation Metrics for Optimization",
options=all_metrics,
default=["question_answering_correctness"],
key="tuner_metrics",
)
target_model = st.selectbox(
"Select Target Model", options=TARGET_MODELS, key="tuner_target_model"
)
with st.expander("Advanced Settings"):
st.session_state.tuner_target_qps = st.number_input(
"Target Model QPS",
min_value=0.1,
max_value=10.0,
value=1.0,
step=0.1,
key="tuner_target_qps_input",
)
st.session_state.tuner_optimizer_qps = st.number_input(
"Optimizer Model QPS",
min_value=0.1,
max_value=10.0,
value=1.0,
step=0.1,
key="tuner_optimizer_qps_input",
)
st.session_state.tuner_eval_qps = st.number_input(
"Evaluation QPS",
min_value=0.1,
max_value=10.0,
value=1.0,
step=0.1,
key="tuner_eval_qps_input",
)
st.session_state.tuner_data_limit = st.number_input(
"Data Limit (Sample Size)",
min_value=1,
max_value=1000,
value=10,
step=1,
key="tuner_data_limit_input",
)
st.divider()
# 4. Execution
st.header("3. Execution")
if st.button("Start Optimization Job", type="primary"):
if not st.session_state.dataset is not None:
st.error("Please load a dataset first.")
return
if not st.session_state.local_prompt.prompt_to_run.system_instruction:
st.error("Please load a prompt first.")
return
if not selected_metrics:
st.error("Please select at least one metric.")
return
with st.spinner("Initializing Job..."):
workspace_uri = (
epath.Path(os.getenv("BUCKET"))
/ "optimization"
/ st.session_state.op_id
)
input_data_uri = workspace_uri / "data"
workspace_uri.mkdir(parents=True, exist_ok=True)
input_data_uri.mkdir(parents=True, exist_ok=True)
output_optimization_data_uri = workspace_uri / "optimization_jobs"
job_name = f"{st.session_state.tuner_prompt}-{st.session_state.tuner_version}-{st.session_state.tuner_dataset}-{st.session_state.op_id}"
output_optimization_run_uri = str(output_optimization_data_uri / job_name)
input_optimization_data_file_uri = f"gs://{input_data_uri}/{job_name}.jsonl"
st.session_state.dataset.to_json(
str(input_optimization_data_file_uri), orient="records", lines=True
)
args = get_optimization_args(
input_optimization_data_file_uri,
output_optimization_run_uri,
target_model,
selected_metrics,
st.session_state.tuner_target_qps,
st.session_state.tuner_optimizer_qps,
st.session_state.tuner_eval_qps,
st.session_state.tuner_data_limit,
)
args_dict = vars(args)
config_file_uri = "gs://" + str(workspace_uri / "config" / "config.json")
with epath.Path(config_file_uri).open("w") as config_file:
json.dump(args_dict, config_file)
worker_pool_specs = [
{
"machine_spec": {"machine_type": "n1-standard-4"},
"replica_count": 1,
"container_spec": {
"image_uri": os.getenv("APD_CONTAINER_URI"),
"args": ["--config=" + config_file_uri],
},
}
]
custom_job = aiplatform.CustomJob(
display_name=job_name,
worker_pool_specs=worker_pool_specs,
staging_bucket=str(workspace_uri),
)
custom_job.run(service_account=os.getenv("APD_SERVICE_ACCOUNT"), sync=False)
st.session_state.tuner_launched_job = job_name
st.session_state.tuner_run_uri = f"gs://{os.getenv('BUCKET')}/optimization/{st.session_state.op_id}/optimization_jobs/{job_name}"
st.success(f"Started Optimization Job: {job_name}")
if st.session_state.tuner_launched_job:
st.info(f"Active Job Tracked: {st.session_state.tuner_launched_job}")
st.divider()
# 5. Results Report
st.header("4. Results Report")
if st.button("Load Results"):
if not st.session_state.tuner_launched_job:
st.warning("No optimization job has been launched in this session.")
else:
with st.spinner("Checking job status..."):
status = check_job_status(
st.session_state.tuner_launched_job,
os.getenv("PROJECT_ID"),
os.getenv("LOCATION"),
)
if status in ["JOB_STATE_PENDING", "JOB_STATE_RUNNING"]:
st.info(
f"Job is still {status.replace('JOB_STATE_', '')}. Please check back later. (Hill-climbing algorithms may take a while)."
)
elif status == "JOB_STATE_FAILED":
st.error(
"The optimization job failed. Check Agent Platform console logs."
)
elif status == "JOB_STATE_SUCCEEDED":
st.success("Job Complete! Processing results...")
try:
results_ui = vapo_lib.ResultsUI(st.session_state.tuner_run_uri)
if getattr(results_ui, "templates", None) and getattr(
results_ui, "eval_results", None
):
baseline = results_ui.templates[0]
winner = results_ui.templates[-1]
st.subheader("Score Jump")
mean_cols = [
c
for c in baseline.columns
if c.startswith("metrics.") and "/mean" in c
]
col_metrics = st.columns(min(len(mean_cols), 4) or 1)
final_scores = {}
for idx, m_col in enumerate(mean_cols):
b_val = (
float(baseline[m_col].iloc[0])
if m_col in baseline
else 0.0
)
w_val = (
float(winner[m_col].iloc[0])
if m_col in winner
else 0.0
)
diff = w_val - b_val
label = (
m_col.replace("metrics.", "")
.replace("/mean", "")
.title()
)
final_scores[label] = w_val
with col_metrics[idx % len(col_metrics)]:
st.metric(
label, f"{w_val:.3f}", delta=f"{diff:.3f}"
)
st.subheader("Winner System Instruction")
winning_system_text = (
winner["prompt"].iloc[0]
if "prompt" in winner
else "Unable to parse winning prompt."
)
# Usually the optimizer alters the instruction which is the "prompt" field in vapo_lib output.
st.text_area(
"Optimized Instruction", winning_system_text, height=150
)
st.session_state.tuner_winning_template = (
winning_system_text
)
st.session_state.tuner_final_scores = final_scores
else:
st.warning("Results found but could not be parsed.")
except Exception as e:
st.error(f"Error loading results: {e}")
else:
st.warning(f"Job state is currently: {status}")
st.divider()
# 6. Validation
st.header("5. Validation")
st.markdown("Test the best performing prompt on a blind test case.")
blind_test_json = st.text_area(
"Blind Test Case Input (JSON)",
placeholder='{"ticket_text": "I lost my password..."}',
height=100,
)
if st.button("Test Best Prompt"):
if not st.session_state.tuner_winning_template:
st.warning("Please load successful results first.")
elif not blind_test_json:
st.warning("Please provide a test case in JSON format.")
else:
try:
test_input = json.loads(blind_test_json)
# Setup temporary prompt object to run test
prompt_obj = st.session_state.local_prompt
prompt_obj.prompt_to_run.system_instruction = (
st.session_state.tuner_winning_template
)
# Keep original generation config/schema
with st.spinner("Generating Response..."):
res = prompt_obj.generate_response(test_input)
st.success("Evaluation complete.")
st.text_area("Validation Response", res, height=150, disabled=True)
except Exception as e:
st.error(f"Failed to generate test response: {e}")
st.divider()
# 7. Outcome
st.header("6. Outcome")
if st.button("Export Final Prompt & Save Report", type="primary"):
if not st.session_state.tuner_winning_template:
st.error("No winning template found. Please load results first.")
else:
try:
prompt_obj = st.session_state.local_prompt
prompt_obj.prompt_to_run.system_instruction = (
st.session_state.tuner_winning_template
)
# Save as new version
with st.spinner("Saving optimized prompt to registry..."):
prompt_obj.save_prompt(check_existing=False)
new_version = prompt_obj.prompt_to_run._version_id or "latest"
st.success(f"Successfully exported as version: {new_version}")
# Save evaluation records
if "tuner_final_scores" in st.session_state:
with st.spinner("Saving performance report..."):
write_record(
st.session_state.tuner_final_scores,
st.session_state.tuner_prompt,
new_version,
st.session_state.tuner_winning_template,
)
st.success("Performance report saved to GCS.")
except Exception as e:
st.error(f"Failed to save outcome: {e}")
if __name__ == "__main__":
main()
@@ -0,0 +1,456 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Streamlit page for running Vertex AI Prompt Optimization.
This script provides a user interface for:
- Loading existing prompts from Vertex AI Prompt Registry.
- Loading datasets from a Google Cloud Storage bucket.
- Generating baseline responses and evaluating them against a ground truth.
- Configuring and launching a Vertex AI CustomJob for prompt optimization.
- Displaying baseline evaluation results.
File Source:
https://github.com/GoogleCloudPlatform/generative-ai/blob/main/gemini/prompts/prompt_optimizer/vapo_lib.py
"""
import json
import logging
import os
from argparse import Namespace
import pandas as pd
import streamlit as st
from dotenv import load_dotenv
from etils import epath
from google.cloud import aiplatform, storage
from src import vapo_lib
from src.gcp_prompt import GcpPrompt as gcp_prompt
from vertexai.preview import prompts
load_dotenv("src/.env")
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
TARGET_MODELS = [
"gemini-2.5-pro",
"gemini-2.5-flash",
"gemini-2.0-flash-001",
"gemini-2.0-flash-lite-001",
"gemini-1.5-pro-002",
"gemini-1.5-flash-002",
]
def initialize_session_state() -> None:
"""Initializes the session state variables."""
if "op_id" not in st.session_state:
st.session_state.op_id = vapo_lib.get_id()
if "local_prompt" not in st.session_state:
st.session_state.local_prompt = gcp_prompt()
if "storage_client" not in st.session_state:
st.session_state["storage_client"] = storage.Client()
if "data_uris" not in st.session_state:
st.session_state["data_uris"] = refresh_bucket()
if "dataset" not in st.session_state:
st.session_state["dataset"] = None
if "cached_data_files" not in st.session_state:
st.session_state.cached_data_files = {}
if "last_selected_dataset_for_cache" not in st.session_state:
st.session_state.last_selected_dataset_for_cache = None
def refresh_bucket() -> list[str]:
"""Refreshes the list of available dataset URIs from the GCS bucket.
This function lists all blobs in the configured GCS bucket, filters for
CSV and JSONL files located within the 'datasets/' prefix, and constructs a list
of their full gs:// URI paths.
Returns:
A list of strings, where each string is a GCS URI to a dataset file.
"""
logger.info("Bucket: %s", os.getenv("BUCKET"))
bucket = st.session_state.storage_client.bucket(os.getenv("BUCKET"))
blobs = bucket.list_blobs()
data_uris = []
for i in blobs:
if i.name.split("/")[0] == "datasets" and (
i.name.endswith(".csv") or i.name.endswith(".jsonl")
):
data_uris.append(f"gs://{i.bucket.name}/{i.name}")
logger.info("Data URIs: %s", data_uris)
return data_uris
def prompt_selection() -> None:
"""Handles the prompt selection and loading."""
st.selectbox(
"Select Existing Prompt",
options=st.session_state.local_prompt.existing_prompts.keys(),
placeholder="Select Prompt...",
key="selected_prompt",
)
if st.session_state.selected_prompt:
logger.info("Prompt Meta: %s", st.session_state.local_prompt.prompt_meta)
st.session_state.local_prompt.prompt_meta["name"] = (
st.session_state.selected_prompt
)
versions = [
i.version_id
for i in prompts.list_versions(
st.session_state.local_prompt.existing_prompts[
st.session_state.selected_prompt
]
)
]
st.selectbox(
"Select Version",
options=versions,
placeholder="Select Version...",
key="selected_version",
)
st.button("Load Prompt", key="load_existing_prompt_button")
if st.session_state.load_existing_prompt_button:
logger.info(
"Selected Prompt ID: %s",
st.session_state.local_prompt.existing_prompts[
st.session_state.selected_prompt
],
)
logger.info("Version: %s", st.session_state.selected_version)
st.session_state.local_prompt.load_prompt(
st.session_state.local_prompt.existing_prompts[
st.session_state.selected_prompt
],
st.session_state.selected_prompt,
st.session_state.selected_version,
)
logger.info("Local Prompt Meta: %s", st.session_state.local_prompt.prompt_meta)
logger.info(
"Local Prompt Meta Dict Keys: %s",
st.session_state.local_prompt.prompt_meta.keys(),
)
st.session_state.prompt_name = (
st.session_state.local_prompt.prompt_to_run.prompt_name
)
st.session_state.prompt_data = (
st.session_state.local_prompt.prompt_to_run.prompt_data
)
st.session_state.model_name = (
st.session_state.local_prompt.prompt_to_run.model_name.split("/")[-1]
)
st.session_state.system_instructions = (
st.session_state.local_prompt.prompt_to_run.system_instruction
)
st.session_state.response_schema = json.dumps(
st.session_state.local_prompt.prompt_meta.get("response_schema", {})
)
st.session_state.generation_config = json.dumps(
st.session_state.local_prompt.prompt_meta.get("generation_config", {})
)
st.session_state.meta_tags = st.session_state.local_prompt.prompt_meta[
"meta_tags"
]
def dataset_selection() -> None:
"""Handles the dataset selection and loading."""
data_sets = list({i.split("/")[4] for i in st.session_state.data_uris})
logger.info("Data Sets: %s", data_sets)
st.selectbox(
"Select an Existing Dataset", options=[None, *data_sets], key="selected_dataset"
)
files_to_display_in_selectbox = []
if st.session_state.selected_dataset:
if (
st.session_state.selected_dataset
!= st.session_state.last_selected_dataset_for_cache
or st.session_state.selected_dataset
not in st.session_state.cached_data_files
):
logger.info(
"Cache miss or dataset changed for files. Fetching for: %s",
st.session_state.selected_dataset,
)
bucket = st.session_state.storage_client.bucket(os.getenv("BUCKET"))
prefix = f"datasets/{st.session_state.selected_dataset}/"
blobs_iterator = bucket.list_blobs(prefix=prefix)
current_dataset_files = []
for blob in blobs_iterator:
if (
blob.name.endswith(".csv") or blob.name.endswith(".jsonl")
) and not blob.name.endswith("/"):
filename = blob.name[len(prefix) :]
if filename:
current_dataset_files.append(filename)
st.session_state.cached_data_files[st.session_state.selected_dataset] = (
sorted(set(current_dataset_files))
)
st.session_state.last_selected_dataset_for_cache = (
st.session_state.selected_dataset
)
logger.info(
"Cached files for %s: %s",
st.session_state.selected_dataset,
st.session_state.cached_data_files[st.session_state.selected_dataset],
)
if "selected_file_from_dataset" in st.session_state:
st.session_state.selected_file_from_dataset = None
logger.info("Reset selected_file_from_dataset due to dataset change.")
files_to_display_in_selectbox = st.session_state.cached_data_files.get(
st.session_state.selected_dataset, []
)
st.selectbox(
"Select a file from this dataset:",
options=[None, *files_to_display_in_selectbox],
key="selected_file_from_dataset",
)
st.button("Load Dataset", key="load_existing_dataset_button")
if st.session_state.load_existing_dataset_button:
if not st.session_state.get("selected_dataset") or not st.session_state.get(
"selected_file_from_dataset"
):
st.warning("Please select a dataset and a file first.")
else:
gcs_uri = f"gs://{os.getenv('BUCKET')}/datasets/{st.session_state.selected_dataset}/{st.session_state.selected_file_from_dataset}"
logger.info("Loading file: %s", gcs_uri)
if st.session_state.selected_file_from_dataset.endswith(".jsonl"):
st.session_state.dataset = pd.read_json(gcs_uri, lines=True)
else:
st.session_state.dataset = pd.read_csv(gcs_uri)
if st.session_state.dataset is not None:
st.dataframe(st.session_state.dataset)
def get_optimization_args(
input_optimization_data_file_uri,
output_optimization_run_uri,
target_model,
target_qps=1.0,
optimizer_qps=1.0,
eval_qps=1.0,
data_limit=10,
):
"""Gets the arguments for the optimization job."""
response_schema_str = st.session_state.local_prompt.prompt_meta.get(
"response_schema", "{}"
)
try:
response_schema = (
json.loads(response_schema_str)
if isinstance(response_schema_str, str)
else response_schema_str
)
except json.JSONDecodeError:
response_schema = {}
if response_schema and response_schema != {}:
response_mime_type = "application/json"
response_schema_arg = response_schema
else:
response_mime_type = "text/plain"
response_schema_arg = ""
has_multimodal = False
if (
st.session_state.dataset is not None
and "image" in st.session_state.dataset.columns
):
has_multimodal = True
return Namespace(
system_instruction=st.session_state.local_prompt.prompt_to_run.system_instruction,
prompt_template=(
f"{st.session_state.local_prompt.prompt_to_run.prompt_data}"
"\n\tAnswer: {target}"
),
target_model=target_model,
optimization_mode="instruction",
eval_metrics_types=[
"question_answering_correctness",
],
eval_metrics_weights=[
1.0,
],
aggregation_type="weighted_sum",
input_data_path=input_optimization_data_file_uri,
output_path=f"gs://{output_optimization_run_uri}",
project=os.getenv("PROJECT_ID"),
num_steps=10,
num_demo_set_candidates=10,
demo_set_size=3,
target_model_location="us-central1",
source_model="",
source_model_location="",
target_model_qps=target_qps,
optimizer_model_qps=optimizer_qps,
eval_qps=eval_qps,
source_model_qps="",
response_mime_type=response_mime_type,
response_schema=response_schema_arg,
language="English",
placeholder_to_content=json.loads("{}"),
data_limit=data_limit,
translation_source_field_name="",
has_multimodal_inputs=has_multimodal,
)
def start_optimization() -> None:
"""Starts the optimization job."""
st.divider()
st.subheader("Run Optimization")
st.button("Start Optimization", key="start_optimization_button")
if st.session_state.start_optimization_button:
workspace_uri = (
epath.Path(os.getenv("BUCKET")) / "optimization" / st.session_state.op_id
)
logger.info("Workspace URI: %s", workspace_uri)
input_data_uri = epath.Path(workspace_uri) / "data"
logger.info("Input Data URI: %s", input_data_uri)
workspace_uri.mkdir(parents=True, exist_ok=True)
input_data_uri.mkdir(parents=True, exist_ok=True)
output_optimization_data_uri = epath.Path(workspace_uri) / "optimization_jobs"
logger.info("Output Data URI: %s", output_optimization_data_uri)
prompt_optimization_job = (
f"{st.session_state.selected_prompt}-"
f"{st.session_state.selected_version}-"
f"{st.session_state.selected_dataset}-"
f"{st.session_state.op_id}"
)
output_optimization_run_uri = str(
output_optimization_data_uri / prompt_optimization_job
)
input_optimization_data_file_uri = (
f"gs://{input_data_uri}/{prompt_optimization_job}.jsonl"
)
logger.info("Input Optimization Data URI: %s", input_optimization_data_file_uri)
if st.session_state.dataset is not None:
st.session_state.dataset.to_json(
str(input_optimization_data_file_uri), orient="records", lines=True
)
else:
st.error("Please load a dataset first.")
return
args = get_optimization_args(
input_optimization_data_file_uri,
output_optimization_run_uri,
st.session_state.target_model_optimization,
st.session_state.target_qps,
st.session_state.optimizer_qps,
st.session_state.eval_qps,
st.session_state.data_limit,
)
with st.expander("Prompt Optimization Config"):
st.json(vars(args))
args = vars(args)
config_file_uri = "gs://" + str(workspace_uri / "config" / "config.json")
with epath.Path(config_file_uri).open("w") as config_file:
json.dump(args, config_file)
config_file.close()
st.success(f"Successfully wrote config file to {config_file_uri}")
worker_pool_specs = [
{
"machine_spec": {
"machine_type": "n1-standard-4",
},
"replica_count": 1,
"container_spec": {
"image_uri": os.getenv("APD_CONTAINER_URI"),
"args": ["--config=" + config_file_uri],
},
}
]
custom_job = aiplatform.CustomJob(
display_name=prompt_optimization_job,
worker_pool_specs=worker_pool_specs,
staging_bucket=str(workspace_uri),
)
custom_job.run(service_account=os.getenv("APD_SERVICE_ACCOUNT"), sync=False)
st.success("Successfully Started Job!!")
def main() -> None:
"""Streamlit page for Prompt Optimization."""
st.set_page_config(
layout="wide", page_title="Prompt Optimization", page_icon="assets/favicon.ico"
)
initialize_session_state()
st.header("Prompt Optimization")
st.selectbox(
"Select Target Model for Optimization:",
options=TARGET_MODELS,
key="target_model_optimization",
)
with st.expander("Advanced Settings"):
st.session_state.target_qps = st.number_input(
"Target Model QPS", min_value=0.1, max_value=10.0, value=1.0, step=0.1
)
st.session_state.optimizer_qps = st.number_input(
"Optimizer Model QPS", min_value=0.1, max_value=10.0, value=1.0, step=0.1
)
st.session_state.eval_qps = st.number_input(
"Evaluation QPS", min_value=0.1, max_value=10.0, value=1.0, step=0.1
)
st.session_state.data_limit = st.number_input(
"Data Limit (Sample Size)", min_value=1, max_value=1000, value=10, step=1
)
prompt_selection()
dataset_selection()
start_optimization()
if __name__ == "__main__":
main()
@@ -0,0 +1,481 @@
## Copyright 2025 Google LLC
## Licensed under the Apache License, Version 2.0 (the "License");
## you may not use this file except in compliance with the License.
## You may obtain a copy of the License at
## https://www.apache.org/licenses/LICENSE-2.0
## Unless required by applicable law or agreed to in writing, software
## distributed under the License is distributed on an "AS IS" BASIS,
## WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
## See the License for the specific language
import json
import logging
import os
import pandas as pd
import streamlit as st
from dotenv import load_dotenv
from google.cloud import storage
from src import vapo_lib
# Load environment variables
load_dotenv("src/.env")
# Configure logging to the console
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
# --- Constants ---
BASE_OPTIMIZATION_PREFIX = "optimization/"
OPTIMIZATION_JOBS_SUBDIR = "optimization_jobs/"
from google.cloud import aiplatform
def list_custom_training_jobs(project_id: str, location: str):
"""Lists all custom training jobs and their statuses in a given project and location.
Args:
project_id: The Google Cloud project ID.
location: The region for the Agent Platform jobs, e.g., "us-central1".
Returns:
A list of dictionaries, where each dictionary contains details of a custom job.
"""
# Initialize the Agent Platform client
# The API endpoint is determined by the location
client_options = {"api_endpoint": f"{location}-aiplatform.googleapis.com"}
client = aiplatform.gapic.JobServiceClient(client_options=client_options)
# The parent resource path format
parent = f"projects/{project_id}/locations/{location}"
# Make the API request to list custom jobs
response = client.list_custom_jobs(parent=parent)
# Process the response and format the output
jobs_list = []
print(f"Fetching jobs from project '{project_id}' in '{location}'...")
for job in response:
job_info = {
"display_name": job.display_name,
"name": job.name,
"status": job.state.name, # .name gets the string representation of the enum
}
jobs_list.append(job_info)
print(f"Found {len(jobs_list)} jobs.")
return jobs_list
# --- Example Usage ---
if __name__ == "__main__":
# Replace with your project ID and desired location
PROJECT_ID = os.getenv("PROJECT_ID")
LOCATION = os.getenv("LOCATION")
# Ensure you have authenticated with Google Cloud CLI:
# gcloud auth application-default login
# And have the necessary permissions (e.g., "Agent Platform User" role)
try:
all_jobs = list_custom_training_jobs(project_id=PROJECT_ID, location=LOCATION)
# Print the results
if all_jobs:
print("\n--- Job Statuses ---")
for job in all_jobs:
print(f" - Name: {job['display_name']:<40} Status: {job['status']}")
print("--------------------\n")
else:
print("No custom jobs found.")
except Exception as e:
print(
"\nAn error occurred. Please ensure your project ID and location are correct,"
)
print(f"and that you have authenticated correctly. Error: {e}")
def safe_json_loads(s):
"""Safely loads a JSON string, returning the original value on failure."""
if not isinstance(s, str):
return s
try:
return json.loads(s)
except (json.JSONDecodeError, TypeError):
return s
@st.cache_data(ttl=300)
def list_gcs_directories(
bucket_name: str, prefix: str, _storage_client: storage.Client
) -> list[str]:
"""Lists 'directories' in GCS under a given prefix.
A 'directory' is inferred from the common prefixes of objects.
Caches the result for 5 minutes to improve performance.
"""
if not bucket_name:
st.warning("BUCKET environment variable is not set.")
return []
if not _storage_client:
st.warning("Storage client is not initialized.")
return []
bucket = _storage_client.bucket(bucket_name)
retrieved_prefixes = set()
try:
for page in bucket.list_blobs(prefix=prefix, delimiter="/").pages:
retrieved_prefixes.update(page.prefixes)
# The retrieved prefixes are the "subdirectories".
# e.g., for prefix 'optimization/', a retrieved prefix might be 'optimization/op_id/'.
# We want to extract just 'op_id'.
dir_names = []
for p in retrieved_prefixes:
name = p.replace(prefix, "").strip("/")
if name:
dir_names.append(name)
return sorted(set(dir_names))
except Exception as e:
st.error(
f"Error listing GCS directories under gs://{bucket_name}/{prefix}: {e}"
)
logger.error(
f"Error listing GCS directories under gs://{bucket_name}/{prefix}: {e}",
exc_info=True,
)
return []
def _display_interactive_results(results_ui: vapo_lib.ResultsUI) -> None:
"""Processes results from a VAPO run and displays them in an interactive
Streamlit UI with tabs for each prompt version.
"""
try:
if (
not hasattr(results_ui, "templates")
or not results_ui.templates
or not hasattr(results_ui, "eval_results")
):
logger.info(
"ResultsUI object does not have 'templates' or 'eval_results', or templates list is empty. Falling back."
)
st.info(
"No completed runs found yet in this directory. The evaluation might still be running or failed to produce results."
)
else:
processed_results_for_tabs = []
for i, template_summary_df in enumerate(results_ui.templates):
if (
not isinstance(template_summary_df, pd.DataFrame)
or template_summary_df.empty
):
logger.warning(
f"Template summary data at index {i} is not a non-empty DataFrame. Skipping."
)
continue
# Get the detailed results to perform the custom calculation
detailed_eval_df = pd.DataFrame()
if i < len(results_ui.eval_results) and isinstance(
results_ui.eval_results[i], pd.DataFrame
):
detailed_eval_df = results_ui.eval_results[i]
# Add a custom exact_match calculation. This is more robust than simple
# string comparison as it handles differences in JSON key order and whitespace.
if (
not detailed_eval_df.empty
and "ground_truth" in detailed_eval_df.columns
and "reference" in detailed_eval_df.columns
):
# Parse the JSON strings into Python objects before comparing.
parsed_ground_truths = detailed_eval_df["ground_truth"].apply(
safe_json_loads
)
parsed_references = detailed_eval_df["reference"].apply(
safe_json_loads
)
# Create a boolean series for the comparison
is_match = parsed_ground_truths.eq(parsed_references)
# Map boolean to 'yes'/'no' for display in the detailed table
detailed_eval_df["calculated_exact_match"] = is_match.map(
{True: "yes", False: "no"}
)
# Calculate the mean from the boolean series for the summary metric
new_exact_match_mean = is_match.mean()
template_summary_df["metrics.calculated_exact_match/mean"] = (
new_exact_match_mean
)
prompt_text = "Prompt text not found in template data."
if "prompt" in template_summary_df.columns:
prompt_text = template_summary_df["prompt"].iloc[0]
else:
logger.warning(
f"Column 'prompt' not found in template_summary_df at index {i}."
)
# Determine the primary score and build the tab name.
primary_score_label = "Score"
primary_score_value = "N/A"
if "metrics.calculated_exact_match/mean" in template_summary_df.columns:
primary_score_label = "Calculated Exact Match"
primary_score_value = template_summary_df[
"metrics.calculated_exact_match/mean"
].iloc[0]
else:
# Fallback to the first available metric
mean_metric_columns = [
col
for col in template_summary_df.columns
if col.startswith("metrics.") and "/mean" in col
]
if mean_metric_columns:
first_metric_col = mean_metric_columns[0]
primary_score_label = (
first_metric_col.replace("metrics.", "")
.replace("/mean", "")
.replace("_", " ")
.title()
)
primary_score_value = template_summary_df[
first_metric_col
].iloc[0]
# Build the tab name with all available metrics for a quick overview.
tab_name_metrics_parts = []
mean_metric_columns = [
col
for col in template_summary_df.columns
if col.startswith("metrics.") and "/mean" in col
]
for metric_col in mean_metric_columns:
metric_name_short = metric_col.replace("metrics.", "").replace(
"/mean", ""
)
metric_val = template_summary_df[metric_col].iloc[0]
if metric_name_short == "calculated_exact_match" and isinstance(
metric_val, float
):
tab_name_metrics_parts.append(
f"{metric_name_short}: {metric_val:.1%}"
)
else:
tab_name_metrics_parts.append(
f"{metric_name_short}: {metric_val:.3f}"
if isinstance(metric_val, float)
else f"{metric_name_short}: {metric_val}"
)
tab_name = f"Template {i}"
if tab_name_metrics_parts:
tab_name += f" ({', '.join(tab_name_metrics_parts)})"
current_summary_df_display = template_summary_df.copy()
if "prompt" in current_summary_df_display.columns:
current_summary_df_display = current_summary_df_display.drop(
columns=["prompt"]
)
processed_results_for_tabs.append(
{
"name": tab_name,
"template_text": prompt_text,
"primary_score_label": primary_score_label,
"primary_score_value": primary_score_value,
"summary_metrics_df": current_summary_df_display,
"detailed_eval_df": detailed_eval_df,
}
)
if (
processed_results_for_tabs
): # If we successfully processed data, show the new UI
st.write("### Interactive Prompt Versions")
tab_titles = [res["name"] for res in processed_results_for_tabs]
tabs = st.tabs(tab_titles)
for i, tab_content in enumerate(tabs):
with tab_content:
result_data = processed_results_for_tabs[i]
st.subheader("Prompt Template")
# Sanitize tab name for key
clean_key_name = "".join(
filter(str.isalnum, result_data["name"])
)
st.text_area(
"Template",
value=result_data["template_text"],
height=200,
disabled=True,
key=f"template_view_{clean_key_name}_{i}",
)
st.subheader("Primary Score")
score_val = result_data["primary_score_value"]
score_label = result_data["primary_score_label"]
if score_label == "Calculated Exact Match" and isinstance(
score_val, float
):
st.metric(label=score_label, value=f"{score_val:.2%}")
else:
st.metric(
label=score_label,
value=f"{score_val:.4f}"
if isinstance(score_val, float)
else str(score_val),
)
if not result_data["summary_metrics_df"].empty:
st.subheader("Summary Metrics (from templates.json)")
st.dataframe(result_data["summary_metrics_df"])
if not result_data["detailed_eval_df"].empty:
st.subheader(
"Detailed Evaluation Results (from eval_results.json)"
)
st.dataframe(result_data["detailed_eval_df"])
else:
st.caption(
"No detailed evaluation results available for this template."
)
else:
st.warning("No valid results could be processed for display.")
except Exception as e:
st.error(f"An error occurred while trying to display results: {e}")
logger.error(f"Error in results display section: {e}", exc_info=True)
st.markdown(
"For now, you can access the results directly at the GCS path shown above."
)
def main() -> None:
"""Renders the Streamlit page for viewing Prompt Optimization Results."""
st.set_page_config(
layout="wide",
page_title="Prompt Optimization Results",
page_icon="assets/favicon.ico",
)
st.header("Prompt Optimization Results Browser")
if "storage_client" not in st.session_state:
try:
st.session_state["storage_client"] = storage.Client()
logger.info("Storage client initialized.")
except Exception as e:
st.error(f"Failed to initialize Google Cloud Storage client: {e}")
logger.error(
f"Failed to initialize Google Cloud Storage client: {e}", exc_info=True
)
st.session_state["storage_client"] = None
return
bucket_name = os.getenv("BUCKET")
if not bucket_name:
st.error("BUCKET environment variable is not set. Please configure it in .env.")
return
# --- Step 1: Select Operation ID ---
op_ids = list_gcs_directories(
bucket_name, BASE_OPTIMIZATION_PREFIX, st.session_state.storage_client
)
if not op_ids:
st.info(
f"No optimization operation IDs found under gs://{bucket_name}/{BASE_OPTIMIZATION_PREFIX}"
)
return
if "op_id" in st.session_state and st.session_state.op_id:
st.caption(
f"Hint: The last optimization run you initiated had the ID: `{st.session_state.op_id}`."
)
selected_op_id = st.selectbox(
"Select an Operation ID:", options=[None, *op_ids], key="selected_op_id_results"
)
if not selected_op_id:
st.write("Please select an Operation ID to see its optimization job runs.")
return
st.divider()
# --- Step 2: Select Experiment Run ---
st.subheader(f"Optimization Job Runs for Operation ID: {selected_op_id}")
optimization_jobs_prefix = (
f"{BASE_OPTIMIZATION_PREFIX}{selected_op_id}/{OPTIMIZATION_JOBS_SUBDIR}"
)
experiment_runs = list_gcs_directories(
bucket_name, optimization_jobs_prefix, st.session_state.storage_client
)
if not experiment_runs:
st.info(
f"No completed optimization job runs found under gs://{bucket_name}/{optimization_jobs_prefix}"
)
return
selected_run = st.selectbox(
"Select an Optimization Job Run:",
options=[None, *experiment_runs],
key="selected_experiment_run",
)
if not selected_run:
st.write("Please select an optimization job run to view its results.")
return
st.divider()
# --- Step 3: Check Job Status and Display Results ---
st.subheader(f"Results for: {selected_run}")
project_id = os.getenv("PROJECT_ID")
location = os.getenv("LOCATION")
if not project_id or not location:
st.error("PROJECT_ID or REGION environment variables are not set.")
return
try:
jobs = list_custom_training_jobs(project_id=project_id, location=location)
job_status = "Not Found"
for job in jobs:
if job["display_name"] == selected_run:
job_status = job["status"]
break
st.info(f"Status for job '{selected_run}': **{job_status}**")
if job_status == "JOB_STATE_FAILED":
st.error(
"This optimization job has failed. Please check the logs in the Agent Platform console for more details."
)
return
if job_status not in ["JOB_STATE_SUCCEEDED", "JOB_STATE_CANCELLED"]:
st.warning(
f"Job is currently in status: {job_status}. Results may be incomplete."
)
except Exception as e:
st.error(f"Could not retrieve job status. Error: {e}")
logger.error(
f"Failed to retrieve job status for {selected_run}: {e}", exc_info=True
)
run_uri = f"gs://{bucket_name}/{optimization_jobs_prefix}{selected_run}"
st.info(f"Loading results from: {run_uri}")
results_ui = vapo_lib.ResultsUI(run_uri)
_display_interactive_results(results_ui)
if __name__ == "__main__":
main()
+129
View File
@@ -0,0 +1,129 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import json
import logging
import os
import pandas as pd
import streamlit as st
from dotenv import load_dotenv
from google.cloud import storage
load_dotenv("src/.env")
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
def load_records_from_gcs(bucket_name: str, prefix: str) -> pd.DataFrame:
"""Loads all JSON record files from a GCS prefix and returns a DataFrame."""
try:
storage_client = storage.Client()
bucket = storage_client.bucket(bucket_name)
blobs = bucket.list_blobs(prefix=prefix)
all_records = []
for blob in blobs:
if blob.name.endswith(".json"):
logger.info(f"Loading record from {blob.name}")
try:
record_data = json.loads(blob.download_as_string())
if isinstance(record_data, list):
all_records.extend(record_data)
else:
all_records.append(record_data)
except json.JSONDecodeError:
logger.warning(f"Could not decode JSON from {blob.name}")
except Exception as e:
logger.exception(f"Failed to process blob {blob.name}: {e}")
if not all_records:
st.warning(f"No JSON records found at gs://{bucket_name}/{prefix}")
return pd.DataFrame()
return pd.json_normalize(all_records)
except Exception as e:
st.error(f"Failed to load or parse records from GCS: {e}")
logger.error("Error loading records: %s", e, exc_info=True)
return pd.DataFrame()
def main() -> None:
"""Renders the Prompt Records Leaderboard page."""
st.set_page_config(
layout="wide",
page_title="Prompt Records Leaderboard",
page_icon="assets/favicon.ico",
)
st.header("Prompt Records Leaderboard")
st.markdown(
"This page allows you to view and compare the evaluation results of different prompt versions."
)
records_prefix = "records/"
if "leaderboard_df" not in st.session_state:
st.session_state.leaderboard_df = pd.DataFrame()
if st.button("Load/Refresh Leaderboard"):
with st.spinner("Loading records from GCS..."):
st.session_state.leaderboard_df = load_records_from_gcs(
os.getenv("BUCKET"), records_prefix
)
if not st.session_state.leaderboard_df.empty:
st.success("Leaderboard loaded successfully.")
else:
st.info("Leaderboard is empty or could not be loaded.")
if st.session_state.leaderboard_df.empty:
st.info("Click the button above to load the leaderboard data.")
return
st.divider()
prompt_names = st.session_state.leaderboard_df["prompt_name"].unique().tolist()
selected_prompt = st.selectbox(
"Select a Prompt to Compare Versions", options=[None, *prompt_names]
)
if selected_prompt:
st.subheader(f"Comparison for: {selected_prompt}")
prompt_df = st.session_state.leaderboard_df[
st.session_state.leaderboard_df["prompt_name"] == selected_prompt
].copy()
if prompt_df.empty:
st.info("No records found for the selected prompt.")
return
# Use prompt_df directly as it already contains flattened mean_scores columns
comparison_df = prompt_df.reset_index(drop=True)
# Clean up the view
display_columns = [
col
for col in comparison_df.columns
if col not in ["scores", "evaluation_data"]
and not (col.startswith("score.") and col[6:].isdigit())
]
st.dataframe(comparison_df[display_columns])
if __name__ == "__main__":
main()
@@ -0,0 +1,580 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "81450b47de75"
},
"outputs": [],
"source": [
"# Copyright 2025 Google LLC\n",
"#\n",
"# Licensed under the Apache License, Version 2.0 (the \"License\");\n",
"# you may not use this file except in compliance with the License.\n",
"# You may obtain a copy of the License at\n",
"#\n",
"# http://www.apache.org/licenses/LICENSE-2.0\n",
"#\n",
"# Unless required by applicable law or agreed to in writing, software\n",
"# distributed under the License is distributed on an \"AS IS\" BASIS,\n",
"# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n",
"# See the License for the specific language governing permissions and\n",
"# limitations under the License."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "a204d0ab284d"
},
"source": [
"# Tutorial for Running Prompt Management and Evaluation\n",
"\n",
"<table align=\"left\">\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://colab.research.google.com/github/GoogleCloudPlatform/generative-ai/blob/main/tools/llmevalkit/prompt-management-tutorial.ipynb\">\n",
" <img width=\"32px\" src=\"https://www.gstatic.com/pantheon/images/bigquery/welcome_page/colab-logo.svg\" alt=\"Google Colaboratory logo\"><br> Open in Colab\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fgenerative-ai%2Fmain%2Ftools%2Fllmevalkit%2Fprompt-management-tutorial.ipynb\">\n",
" <img width=\"32px\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" alt=\"Google Cloud Colab Enterprise logo\"><br> Open in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/workbench/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/generative-ai/main/tools/llmevalkit/prompt-management-tutorial.ipynb\">\n",
" <img src=\"https://www.gstatic.com/images/branding/gcpiconscolors/vertexai/v1/32px.svg\" alt=\"Vertex AI logo\"><br> Open in Vertex AI Workbench\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/generative-ai/blob/main/tools/llmevalkit/prompt-management-tutorial.ipynb\">\n",
" <img width=\"32px\" src=\"https://raw.githubusercontent.com/primer/octicons/refs/heads/main/icons/mark-github-24.svg\" alt=\"GitHub logo\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</table>\n",
"\n",
"<div style=\"clear: both;\"></div>\n",
"\n",
"<b>Share to:</b>\n",
"\n",
"<a href=\"https://www.linkedin.com/sharing/share-offsite/?url=https%3A//github.com/GoogleCloudPlatform/generative-ai/blob/main/tools/llmevalkit/prompt-management-tutorial.ipynb\" target=\"_blank\">\n",
" <img width=\"20px\" src=\"https://upload.wikimedia.org/wikipedia/commons/8/81/LinkedIn_icon.svg\" alt=\"LinkedIn logo\">\n",
"</a>\n",
"\n",
"<a href=\"https://bsky.app/intent/compose?text=https%3A//github.com/GoogleCloudPlatform/generative-ai/blob/main/tools/llmevalkit/prompt-management-tutorial.ipynb\" target=\"_blank\">\n",
" <img width=\"20px\" src=\"https://upload.wikimedia.org/wikipedia/commons/7/7a/Bluesky_Logo.svg\" alt=\"Bluesky logo\">\n",
"</a>\n",
"\n",
"<a href=\"https://twitter.com/intent/tweet?url=https%3A//github.com/GoogleCloudPlatform/generative-ai/blob/main/tools/llmevalkit/prompt-management-tutorial.ipynb\" target=\"_blank\">\n",
" <img width=\"20px\" src=\"https://upload.wikimedia.org/wikipedia/commons/5/5a/X_icon_2.svg\" alt=\"X logo\">\n",
"</a>\n",
"\n",
"<a href=\"https://reddit.com/submit?url=https%3A//github.com/GoogleCloudPlatform/generative-ai/blob/main/tools/llmevalkit/prompt-management-tutorial.ipynb\" target=\"_blank\">\n",
" <img width=\"20px\" src=\"https://redditinc.com/hubfs/Reddit%20Inc/Brand/Reddit_Logo.png\" alt=\"Reddit logo\">\n",
"</a>\n",
"\n",
"<a href=\"https://www.facebook.com/sharer/sharer.php?u=https%3A//github.com/GoogleCloudPlatform/generative-ai/blob/main/tools/llmevalkit/prompt-management-tutorial.ipynb\" target=\"_blank\">\n",
" <img width=\"20px\" src=\"https://upload.wikimedia.org/wikipedia/commons/5/51/Facebook_f_logo_%282019%29.svg\" alt=\"Facebook logo\">\n",
"</a>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "8aee03ecd776"
},
"source": [
"| Author(s) |\n",
"| --- |\n",
"| [Mike Santoro](https://github.com/Michael-Santoro) |"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "7cb0282fcd41"
},
"source": [
"## 1. Overview\n",
"\n",
"This tutorial provides a comprehensive guide to prompt engineering, covering the entire lifecycle from creation to evaluation and optimization. It's broken down into the following sections:\n",
"\n",
"1. **Prompt Management:** This section focuses on the core tasks of creating, editing, and managing prompts. You can: \n",
" - **Create new prompts:** Define the prompt's name, text, the model it's designed for, and any system instructions. \n",
" - **Load and edit existing prompts:** Browse a library of saved prompts, load a specific version, and make modifications.\n",
" - **Test prompts:** Before saving, you can provide sample input and generate a response to see how the prompt performs.\n",
" - **Versioning:** Each time you save a change to a prompt, a new version is created, allowing you to track its evolution and compare different iterations.\n",
"\n",
"2. **Dataset Creation:** A crucial part of prompt engineering is having good data to test and evaluate your prompts. This section allows you to:\n",
"\n",
" - **Create new datasets:** A dataset is essentially a folder in Google Cloud Storage where you can group related files.\n",
" - **Upload data:** You can upload files in CSV, JSON, or JSONL format to your datasets. This data will be used for evaluating your prompts.\n",
"\n",
"3. **Evaluation:** Once you have a prompt and a dataset, you need to see how well the prompt performs. The evaluation section helps you with this by:\n",
"\n",
" - **Running evaluations:** You can select a prompt and a dataset and run an evaluation. This will generate responses from the model for each item in your dataset.\n",
" - **Human-in-the-loop rating:** For a more nuanced evaluation, you can manually review the model's responses and rate them.\n",
" - **Automated metrics:** The tutorial also supports automated evaluation metrics to get a quantitative measure of your prompt's performance.\n",
"\n",
"4. **Prompt Optimization:** This section helps you automatically improve your prompts. It uses Vertex AI's prompt optimization capabilities to:\n",
"\n",
" - **Configure and launch optimization jobs:** You can set up and run a job that will take your prompt and a dataset and try to find a better-performing version of the prompt.\n",
"\n",
"5. **Prompt Optimization Results:** After an optimization job has run, this section allows you to:\n",
"\n",
" - **View the results:** You can see the different prompt versions that the optimizer came up with and how they performed.\n",
" - **Compare versions:** The results are presented in a way that makes it easy to compare the different optimized prompts and choose the best one.\n",
"\n",
"6. **Prompt Records:** This is a leaderboard that shows you the evaluation results of all your different prompt versions. It helps you to:\n",
"\n",
" - **Track performance over time:** See how your prompts have improved with each new version.\n",
" - **Compare different prompts:** You can compare the performance of different prompts for the same task.\n",
"\n",
"In summary, this tutorial provides a complete and integrated environment for all your prompt engineering needs, from initial creation to sophisticated optimization and evaluation.\n"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "bbd5f8c2144a"
},
"source": [
"## 2. Before you start"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "6b2e004206a7"
},
"source": [
"### Clone the GitHub Repo"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "bcafdf59c8b9"
},
"outputs": [],
"source": [
"! git clone https://github.com/GoogleCloudPlatform/generative-ai.git"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "36e497494e38"
},
"outputs": [],
"source": [
"! gcloud storage cp gs://github-repo/prompts/prompt_optimizer/mathvista_dataset/mathvista_input.jsonl mathvista_input.jsonl"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "7445a03507f7"
},
"source": [
"### Install Python Dependencies"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "44fee8ab2679"
},
"outputs": [],
"source": [
"% pip install -r requirements.txt"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "870814a62e87"
},
"source": [
"### Authenticate your notebook environment (Colab only)\n",
"\n",
"Authenticate your environment on Google Colab."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "9b6bfee6ba31"
},
"outputs": [],
"source": [
"import sys\n",
"\n",
"if \"google.colab\" in sys.modules:\n",
" from google.colab import auth\n",
"\n",
" auth.authenticate_user()"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "f5cee1b8b3ab"
},
"source": [
"### Alternative Authenticate"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "4632dbaa3b73"
},
"outputs": [],
"source": [
"# fmt: off\n",
"PROJECT_ID = \"[your-project-id]\" # @param {type: \"string\", placeholder: \"[your-project-id]\", isTemplate: true}\n",
"LOCATION = \"[your-project-region]\" # @param{type: \"string\", placeholder: \"[your-project-region]\", isTemplate: true}\n",
"# fmt: on\n",
"\n",
"! gcloud auth application-default login\n",
"! gcloud config set project {PROJECT_ID}"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "390a62d0e8de"
},
"source": [
"### Set Google Cloud project information\n",
"\n",
"**TO-DO: Check these APIs**\n",
"To get started using Vertex AI, you must have an existing Google Cloud project and [enable the following APIs](https://console.cloud.google.com/flows/enableapi?apiid=cloudresourcemanager.googleapis.com,aiplatform.googleapis.com,cloudfunctions.googleapis.com,run.googleapis.com).\n",
"\n",
"Learn more about [setting up a project and a development environment](https://cloud.google.com/vertex-ai/docs/start/cloud-environment)."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "ad1be5b6c02d"
},
"outputs": [],
"source": [
"! cp src/.env.example src/.env"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "94dd1d71983f"
},
"source": [
"### Copy sample.env and Modify\n",
"\n",
"- BUCKET_NAME - Pick an existing bucket or make a new one below\n",
"- PROJECT_ID\n",
"- SERVICE_ACCOUNT - Created Below\n",
"\n",
"\n",
"#### Create a New Bucket (Not Required if using existing)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "e27ab4bb8a87"
},
"outputs": [],
"source": [
"# fmt: off\n",
"BUCKET_NAME = \"[your-bucket-name]\" # @param {type: \"string\", placeholder: \"[your-bucket-name]\", isTemplate: true}\n",
"# fmt: on\n",
"\n",
"BUCKET_URI = f\"gs://{BUCKET_NAME}\"\n",
"\n",
"\n",
"! gcloud storage buckets create {BUCKET_URI} --location {LOCATION}"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "33cc66a9351d"
},
"source": [
"#### Create a Service Account"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "492567ee96db"
},
"outputs": [],
"source": [
"PROJECT_NUMBER = !gcloud projects describe {PROJECT_ID} --format=\"get(projectNumber)\"[0]\n",
"PROJECT_NUMBER = PROJECT_NUMBER[0]\n",
"SERVICE_ACCOUNT = f\"{PROJECT_NUMBER}-compute@developer.gserviceaccount.com\"\n",
"\n",
"for role in ['aiplatform.user', 'storage.objectAdmin']:\n",
"\n",
" ! gcloud projects add-iam-policy-binding {PROJECT_ID} \\\n",
" --member=serviceAccount:{SERVICE_ACCOUNT} \\\n",
" --role=roles/{role} --condition=None"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "4bc0ae8e0c5f"
},
"source": [
"## 3. Run the App"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "21fed89471c8"
},
"outputs": [],
"source": [
"! cd generative-ai/llmevalkit && streamlit run index.py & npx localtunnel --port 8501"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "90a521d73796"
},
"source": [
"Click the link and use just the external ip as the password.\n",
"\n",
"📝 **Note:** You can run `wget -q -O - https://loca.lt/mytunnelpassword` to get the external ip (i.e 35.194.128.20)\n",
"\n",
"📝 **Note:** If you are having issues displaying the app, clear your cache."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "1e31b6cc7645"
},
"source": [
"![image.png](assets/welcome_page.png)"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "9ca8d05072cd"
},
"source": [
"## 4. Work with the App"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "c7c11886dc17"
},
"source": [
"### 1. Prompt Management\n",
"\n",
"In the Prompt Name field enter:\n",
"\n",
"```\n",
"math_prompt_test\n",
"```\n",
"\n",
"In the Prompt Data field enter:\n",
"\n",
"```\n",
"Problem: {{query}}\n",
"Image: {{image}} @@@image/jpeg\n",
"Answer: {{target}}\n",
"```\n",
"\n",
"In the Model Name field enter:\n",
"```\n",
"gemini-2.0-flash-001\n",
"```\n",
"\n",
"In the System Instructions field enter:\n",
"```\n",
"Solve the problem given the image.\n",
"```\n",
"\n",
"Click `Save`\n",
"\n",
"Copy this text for testing:\n",
"\n",
"```\n",
"{\"query\": \"Hint: Please answer the question and provide the correct option letter, e.g., A, B, C, D, at the end.\\nQuestion: As shown in the figure, CD is the diameter of \\u2299O, chord DE \\u2225 OA, if the degree of \\u2220D is 50.0, then the degree of \\u2220C is ()\\nChoices:\\n(A) 25\\u00b0\\n(B) 30\\u00b0\\n(C) 40\\u00b0\\n(D) 50\\u00b0\", \"image\": \"gs://github-repo/prompts/prompt_optimizer/mathvista_dataset/images/643.jpg\", \"target\": \"25\\u00b0\"}\n",
"```\n",
"\n",
"🖱️ Click `Generate`.\n"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "306b19d374b6"
},
"source": [
"### 2. Dataset Creation\n",
"\n",
"Download a copy of the dataset. Then upload this file in the application.\n",
"\n",
"**Dataset Name:** `mathvista`\n",
"\n",
"You can preview the dataset at the bottom of the page."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "814fb0866bd5"
},
"outputs": [],
"source": [
"! gcloud storage cp gs://github-repo/prompts/prompt_optimizer/mathvista_dataset/mathvista_input.jsonl ."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "0b06fa624a7a"
},
"source": [
"### 3. Evaluation\n",
"\n",
"We will now run an evaluation, prior to doing any tweaking to get a baseline.\n",
"\n",
"- **Existing Dataset:** 'mathvista'\n",
"- **Dataset File:** 'mathvista_input.jsonl'\n",
"- **Number of Samples:** '100'\n",
"- **Ground Truth Column Name:** 'target'\n",
"- **Existing Prompt:** 'math_prompt_test'\n",
"- **Version:** '1'\n",
"\n",
"Click Load Prompt, and Upload and Get Response... ⏰ Wait!!\n",
"\n",
"Review the responses.\n",
"\n",
"- **Model-Based:** 'question-answering-quality'\n",
"\n",
"Launch the Eval... ⏰ Wait!!\n",
"\n",
"View the Evaluation Results, and save to prompt records. This will save this initial version to the prompt records for the baseline.\n"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "c76e0e5065fa"
},
"source": [
"### 4. Prompt Optimization\n",
"\n",
"🔧 Set-Up Prompt Optimization.\n",
"\n",
"- **Target Model:** 'gemini-2.0-flash-001'\n",
"- **Existing Prompt:** 'math_prompt_test'\n",
"- **Version:** '1'\n",
"\n",
"🖱️ Click Load Prompt.\n",
"\n",
"- **Select Existing Dataset:** 'mathvista'\n",
"- **Select the File:** 'mathvista_input.jsonl'\n",
"\n",
"🖱️ Click Load Dataset.\n",
"\n",
"Preview the dataset.\n",
"\n",
"🖱️ Click Start Optimization.\n",
"\n",
"**Note:** If Interested in viewing the progress, Navigate to https://console.cloud.google.com/vertex-ai/training/custom-jobs\n",
"\n",
"⏰ Wait!! This step will take about 20-min to run."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "9987b016703e"
},
"source": [
"### 5. Prompt Optimization Results\n",
"\n",
"View the Optimization Results.\n",
"\n",
"The last run will be shown at the top of the screen. Pick this from the dropdown menu: \n",
"\n",
"![image.png](assets/prompt_optimization_result.png)\n",
"\n",
"Review the results and select the highest scoring version and copy the instruction."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "f15bce27fc94"
},
"source": [
"### 6. Navigate Back to Prompt for New Version\n",
"\n",
"Load your existing prompt from before.\n",
"\n",
"📋 Paste your new instructions from the prompt optimizer, and save new version."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "68957d00c89e"
},
"source": [
"### 7. Run new Evaluation\n",
"\n",
"Repeat step 3 with your new version."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "c2e6c524d69a"
},
"source": [
"### 8. View the Records\n",
"\n",
"Navigate to the leaderboard and load the results."
]
}
],
"metadata": {
"colab": {
"name": "prompt-management-tutorial.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
+41
View File
@@ -0,0 +1,41 @@
[project]
name = "llmevalkit"
version = "0.1.0"
description = "LLM Evaluation Kit"
readme = "README.md"
requires-python = ">=3.10"
dependencies = [
"google-genai",
"vertexai",
"streamlit",
"pandas",
"google-cloud-storage",
"python-dotenv",
"pydantic",
"plotly",
"matplotlib",
"db-dtypes",
"nbformat>=4.2.0",
"streamlit-tags",
"google-cloud-pipeline-components",
"gcsfs",
"ipython",
"tqdm",
"tenacity",
"etils",
"importlib-resources",
"fsspec",
"ipywidgets",
"google-cloud-aiplatform",
"tensorflow",
"pylint",
"ipykernel"
]
[tool.uv]
index = [{ url = "https://pypi.org/simple", default = true }]
[dependency-groups]
dev = [
"pytest>=9.0.2",
]
+25
View File
@@ -0,0 +1,25 @@
ipykernel
google-cloud-storage
streamlit
pandas
plotly
matplotlib
db-dtypes
nbformat>=4.2.0
streamlit-tags
dotenv
google-cloud-pipeline-components
gcsfs
pipfile
ipython
asyncio
tqdm
tenacity
etils
importlib-resources
fsspec
ipywidgets
google-cloud-aiplatform
tensorflow
pylint
google-genai
+25
View File
@@ -0,0 +1,25 @@
# Project Configuration
BUCKET = ""
PROMPT_PREFIX = "prompts"
PROJECT_ID = ""
LOCATION = ""
START_PROMPT_PATH = "src/initial_local_prompt_empty.json"
## Constants for Eval
JITTER_AMOUNT = 0.1
PLOTLY_RENDERER = "colab"
DEFAULT_HUMAN_RATING_COL = "human_rating"
DEFAULT_SCORE_COL = "score"
EXPERIMENT_NAME = "eval-open-judge-test"
## Constants for Optimize
OPTIMIZATION_MODE = "instruction"
EVAL_METRIC = "exact_match"
NUM_INST_OPTIMIZATION_STEPS = 10
TARGET_MODEL_QPS = 3.0
EVAL_QPS = 3.0
RESPONSE_MIME_TYPE = "text/plain"
TARGET_LANGUAGE = "English"
SERVICE_ACCOUNT = ""
APD_CONTAINER_URI = "us-docker.pkg.dev/vertex-ai-restricted/builtin-algorithm/apd:preview_v1_0"
+13
View File
@@ -0,0 +1,13 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
+95
View File
@@ -0,0 +1,95 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""This module provides utility functions for managing and processing datasets
stored on Google Cloud Storage (GCS).
It includes functionalities to:
- List available datasets from a specified GCS bucket.
- Fetch and read CSV files from GCS into pandas DataFrames.
- Process raw data from CSVs to generate structured user prompts and expected
outcomes for model evaluation.
- A helper function to escape special characters in text to be used with the
Gemini API.
The primary purpose is to abstract the GCS interactions and data preprocessing
steps required for the prompt management and evaluation application.
"""
import io
import logging
import os
import pandas as pd
from dotenv import load_dotenv
from google.cloud import storage
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
load_dotenv("src/.env")
PREFIX = "datasets_meta"
def get_existing_datasets() -> list[str]:
"""Gets the list of existing dataset names from the GCS bucket.
Returns:
list[str]: A list of dataset names.
"""
storage_client = storage.Client()
blobs = storage_client.list_blobs(os.getenv("BUCKET_NAME"), prefix=PREFIX)
return [i.name.split("/")[-1] for i in blobs]
def escape_special_characters(text: str) -> str:
"""Escapes special characters for Gemini Flash API."""
if not isinstance(text, str):
return text
text = text.replace("\\", "\\")
text = text.replace("\n", "\n")
text = text.replace("\r", "\r")
text = text.replace("\t", "\t")
return text.replace('"', '"')
def process_csv_from_gcs(bucket_name: str, file_path: str) -> pd.DataFrame:
"""Reads a CSV file from GCS, processes each row to create user prompts
and expected results, and returns a pandas DataFrame.
Args:
bucket_name: Name of the GCS bucket.
file_path: Path to the CSV file within the bucket.
Returns:
pandas.DataFrame: DataFrame with original data and added user_prompt
and expected_result columns.
"""
try:
client = storage.Client()
bucket = client.get_bucket(bucket_name)
blob = bucket.blob(file_path)
content = blob.download_as_bytes()
return pd.read_csv(io.BytesIO(content))
except (OSError, ValueError) as e:
print(f"An error occurred: {e}")
return pd.DataFrame()
+401
View File
@@ -0,0 +1,401 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Provides tools for evaluating and visualizing model performance against human ratings.
This module contains functions to process evaluation data, which includes both
human-provided scores and model-generated scores. It prepares the data into a
pandas DataFrame, extracts key metrics like the confusion matrix, and generates
a series of Plotly visualizations to compare the two sets of scores.
The primary functions are:
- prepare_dataframe: Cleans and formats the raw evaluation data.
- extract_completeness_metrics: Pulls confusion matrix data from metrics logs.
- plot_distribution_comparison: Compares the distribution of human vs. model scores.
- plot_confusion_matrix: Creates a heatmap to show agreement and disagreement.
- plot_jitter_scatter: Visualizes the alignment of individual data points.
- run_visual_analysis: An orchestrator function that runs the full analysis and
returns a summary report along with the generated figures.
"""
import ast
import logging
import os
from typing import Any
import numpy as np
import pandas as pd
import plotly.graph_objects as go
from dotenv import load_dotenv
from plotly.subplots import make_subplots
load_dotenv("src/.env")
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
def format_user_content(user_content: str, tokenizer: Any, **kwargs: Any) -> str | None:
"""Applies tokenizer.apply_chat_template to user content string.
Assumes user_content is the text for the 'user' role.
"""
message = [
{"role": "user", "content": user_content},
]
kwargs.setdefault("tokenize", False)
kwargs.setdefault("add_generation_prompt", True)
try:
return tokenizer.apply_chat_template(message, **kwargs)
except Exception as e:
print(
f"Error applying chat template to content: '{user_content[:50]}...'. Error: {e}"
)
return None
def prepare_dataframe(
raw_data: Any,
human_col_name: str,
score_col_name: str,
) -> pd.DataFrame:
"""Prepares the DataFrame from raw data, renames columns, and converts types."""
try:
if isinstance(raw_data, pd.DataFrame):
df = raw_data.copy()
if len(df.columns) >= 2:
# If it's a DataFrame, check if desired columns exist, else use first two
if human_col_name not in df.columns or score_col_name not in df.columns:
original_cols = df.columns
df = df[
[original_cols[0], original_cols[1]]
].copy() # Use first two
df.columns = [human_col_name, score_col_name]
else:
# Ensure we only keep the needed columns if more exist
pass
if human_col_name not in df.columns or score_col_name not in df.columns:
original_cols = df.columns
df = df[[original_cols[0], original_cols[1]]].copy()
df.columns = [human_col_name, score_col_name]
else:
df = df[[human_col_name, score_col_name]].copy()
else:
print(
f"Warning: Input DataFrame has < 2 columns. Expected at least '{human_col_name}' and '{score_col_name}'."
)
return pd.DataFrame(columns=[human_col_name, score_col_name])
else:
df = pd.DataFrame(raw_data)
if df.empty:
print("Warning: Created empty DataFrame from raw_data.")
return pd.DataFrame(columns=[human_col_name, score_col_name])
if len(df.columns) >= 2:
df = df.iloc[:, :2]
df.columns = [human_col_name, score_col_name]
else:
print(
f"Warning: DataFrame from raw_data has < 2 columns. Cannot set '{human_col_name}' and '{score_col_name}'."
)
return pd.DataFrame(columns=[human_col_name, score_col_name])
df[human_col_name] = pd.to_numeric(df[human_col_name], errors="coerce")
df[score_col_name] = pd.to_numeric(df[score_col_name], errors="coerce")
df.dropna(subset=[human_col_name, score_col_name], inplace=True)
df[human_col_name] = df[human_col_name].astype(float)
df[score_col_name] = df[score_col_name].astype(float)
return df
except Exception as e:
print(f"Error preparing DataFrame: {e}")
return pd.DataFrame(columns=[human_col_name, score_col_name])
def extract_completeness_metrics(
metrics_data: list[dict[str, Any]] | None,
) -> tuple[list[list[Any]] | None, list[str] | None, list[float] | None]:
"""Extracts confusion matrix info from metrics data (expected at index 0)."""
if not metrics_data or not isinstance(metrics_data, list) or len(metrics_data) == 0:
print("Warning: Metrics data is empty or not a list.")
return None, None, None
try:
completeness_metrics = metrics_data[0]
if (
"confusion_matrix" not in completeness_metrics
or "confusion_matrix_labels" not in completeness_metrics
):
print(
"Warning: 'confusion_matrix' or 'confusion_matrix_labels' not in first metrics item."
)
return None, None, None
cm = completeness_metrics["confusion_matrix"]
cm_labels = completeness_metrics["confusion_matrix_labels"]
cm_labels_numeric = [float(cl) for cl in cm_labels]
return cm, cm_labels, cm_labels_numeric
except (KeyError, IndexError, ValueError, TypeError) as e:
print(f"Warning: Could not extract confusion matrix info from metrics: {e}")
return None, None, None
def plot_distribution_comparison(
df: pd.DataFrame,
human_col: str,
score_col: str,
) -> go.Figure:
"""Generates bar charts comparing distributions of human ratings and model scores."""
human_counts = df[human_col].value_counts().sort_index()
score_counts = df[score_col].value_counts().sort_index()
fig = make_subplots(
rows=1,
cols=2,
subplot_titles=("Human Rating Distribution", "Model Score Distribution"),
)
fig.add_trace(
go.Bar(
x=human_counts.index,
y=human_counts.values,
name="Human Rating",
marker_color="indianred",
),
row=1,
col=1,
)
fig.add_trace(
go.Bar(
x=score_counts.index,
y=score_counts.values,
name="Model Score",
marker_color="lightsalmon",
),
row=1,
col=2,
)
fig.update_layout(
title_text="Distribution of Human Ratings vs. Model Scores",
bargap=0.2,
xaxis1_title="Rating Value",
yaxis1_title="Count",
xaxis2_title="Score Value",
yaxis2_title="Count",
xaxis1_type="category",
xaxis2_type="category",
xaxis1={
"categoryorder": "array",
"categoryarray": sorted(human_counts.index.unique()),
},
xaxis2={
"categoryorder": "array",
"categoryarray": sorted(score_counts.index.unique()),
},
height=400,
)
return fig
def plot_confusion_matrix(
cm: list[list[Any]] | None, cm_labels: list[str] | None
) -> go.Figure | None:
"""Generates a heatmap for the confusion matrix."""
if cm is None or cm_labels is None:
print("Skipping confusion matrix plot: missing data.")
return None
fig = go.Figure(
data=go.Heatmap(
z=cm,
x=cm_labels,
y=cm_labels,
hoverongaps=False,
colorscale="Blues",
text=cm,
texttemplate="%{text}",
zmin=0,
)
)
fig.update_layout(
title="Confusion Matrix: Human Rating vs. Model Score (Completeness)",
xaxis_title="Predicted (Model Score)",
yaxis_title="True (Human Rating)",
yaxis={
"type": "category",
"categoryorder": "array",
"categoryarray": cm_labels,
},
xaxis={
"type": "category",
"categoryorder": "array",
"categoryarray": cm_labels,
},
height=600,
width=600,
)
return fig
def plot_jitter_scatter(
df: pd.DataFrame,
cm_labels: list[str] | None,
cm_labels_numeric: list[float] | None,
human_col: str,
score_col: str,
) -> go.Figure:
"""Generates a jitter scatter plot comparing individual scores and ratings."""
df_jitter = df[[human_col, score_col]].copy()
if cm_labels is None or cm_labels_numeric is None:
print(
"Using data range for jitter plot axes due to missing confusion matrix labels."
)
min_val: float = (
min(df_jitter[human_col].min(), df_jitter[score_col].min()) - 0.5
)
max_val: float = (
max(df_jitter[human_col].max(), df_jitter[score_col].max()) + 0.5
)
plot_range = [min_val, max_val]
tick_vals = sorted(
df_jitter[human_col].unique()
) # Use unique human ratings for ticks if available
tick_text = [str(int(v)) if v == int(v) else str(v) for v in tick_vals]
else:
plot_range = [min(cm_labels_numeric) - 0.5, max(cm_labels_numeric) + 0.5]
tick_vals = cm_labels_numeric
tick_text = cm_labels
df_jitter[f"{human_col}_jitter"] = df_jitter[human_col] + np.random.uniform(
-ast.literal_eval(os.getenv("JITTER_AMOUNT")),
ast.literal_eval(os.getenv("JITTER_AMOUNT")),
size=len(df_jitter),
)
df_jitter[f"{score_col}_jitter"] = df_jitter[score_col] + np.random.uniform(
-ast.literal_eval(os.getenv("JITTER_AMOUNT")),
ast.literal_eval(os.getenv("JITTER_AMOUNT")),
size=len(df_jitter),
)
fig = go.Figure()
fig.add_trace(
go.Scatter(
x=df_jitter[f"{score_col}_jitter"],
y=df_jitter[f"{human_col}_jitter"],
mode="markers",
marker={
"color": "rgba(0, 100, 200, 0.7)",
"size": 10,
"line": {"width": 1, "color": "DarkSlateGrey"},
},
text=[
f"HR: {hr:.1f}, Score: {s:.1f}"
for hr, s in zip(
df_jitter[human_col], df_jitter[score_col], strict=False
)
],
hoverinfo="text",
name="Ratings",
)
)
fig.add_trace(
go.Scatter(
x=plot_range,
y=plot_range,
mode="lines",
name="Ideal Alignment (Score = Human Rating)",
line={"color": "red", "dash": "dash"},
)
)
fig.update_layout(
title="Model Score vs. Human Rating (with Jitter)",
xaxis_title="Model Score (Jittered)",
yaxis_title="Human Rating (Jittered)",
xaxis={"range": plot_range, "tickvals": list(tick_vals), "ticktext": tick_text},
yaxis={"range": plot_range, "tickvals": list(tick_vals), "ticktext": tick_text},
width=600,
height=600,
showlegend=True,
hovermode="closest",
)
return fig
def run_visual_analysis(
df_data: Any,
metrics: list[dict[str, Any]] | None,
human_col_name: str,
score_col_name: str,
) -> tuple[str, go.Figure | None, go.Figure | None, go.Figure | None]:
"""Runs the visual analysis comparing model scores and human ratings.
Returns a markdown string and the three figure objects.
"""
result_str = ""
df = prepare_dataframe(
df_data, human_col_name=human_col_name, score_col_name=score_col_name
)
cm, cm_labels, cm_labels_numeric = extract_completeness_metrics(metrics)
if df.empty:
result_str = f"Could not create valid DataFrame with columns '{human_col_name}' and '{score_col_name}' from input data. Stopping analysis."
return result_str, None, None, None
if human_col_name not in df.columns or score_col_name not in df.columns:
result_str = f"Expected columns '{human_col_name}' and '{score_col_name}' not found in DataFrame. Stopping analysis."
return result_str, None, None, None
result_str += "# Visual Analysis: Model Score vs. Human Rating Alignment \n"
result_str += "## 1. Distributions Comparison\n"
fig_dist: go.Figure | None = plot_distribution_comparison(
df, human_col=human_col_name, score_col=score_col_name
)
if fig_dist:
result_str += "*Shows if the model score distribution mirrors human ratings.*\n"
else:
result_str += "*Could not generate distribution comparison plot.*\n"
result_str += "## 2. Confusion Matrix (Human vs. Model)\n"
fig_cm: go.Figure | None = plot_confusion_matrix(cm, cm_labels)
if fig_cm:
result_str += "*Visualizes agreement and disagreement between discrete human ratings and model scores. Ideal alignment is along the diagonal.*\n"
else:
result_str += "*Confusion Matrix could not be generated (check metrics data or ensure it's at index 0).**\n"
result_str += "## 3. Score vs. Rating Alignment (Jitter Plot)\n"
fig_scatter: go.Figure | None = plot_jitter_scatter(
df,
cm_labels,
cm_labels_numeric,
human_col=human_col_name,
score_col=score_col_name,
)
if fig_scatter:
result_str += "*Shows individual item alignment. Points close to the red dashed line indicate good agreement.*\n"
else:
result_str += "*Could not generate jitter scatter plot.*"
return result_str, fig_dist, fig_cm, fig_scatter
+280
View File
@@ -0,0 +1,280 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Manages the lifecycle of prompts using Google Cloud Platform services.
This module provides a `gcp_prompt` class that facilitates the creation,
storage, retrieval, and execution of prompts with Vertex AI and Cloud Storage.
It handles the interaction with the Vertex AI SDK to manage prompt versions
and uses Cloud Storage to persist prompt metadata.
Key functionalities include:
- Initializing a connection to Google Cloud Platform services (Vertex AI, Cloud Storage).
- Caching and refreshing a list of existing prompts.
- Saving new or updated prompts, including their metadata, to both the
Vertex AI prompt registry and a Cloud Storage bucket.
- Loading existing prompts from the registry and their associated metadata
from the bucket.
- Generating responses from a specified model using a loaded prompt and
a given set of variables.
- Helper functions for escaping special characters in prompt text.
The `model_response` Pydantic model defines the expected structure of the
response from the generative model.
"""
import json
import logging
import os
import warnings
from typing import Any
from dotenv import load_dotenv
from google import genai
from google.cloud import aiplatform, storage
from vertexai.generative_models import GenerationConfig
from vertexai.preview import prompts
from vertexai.preview.prompts import Prompt
# Suppress the UserWarning from vertexai.generative_models as we need to use it
# for Prompt Management until google-genai SDK supports it.
warnings.filterwarnings(
"ignore", category=UserWarning, module="vertexai.generative_models"
)
load_dotenv("src/.env")
# Configure logging to the console
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
aiplatform.init(
project=os.getenv("PROJECT_ID"),
location=os.getenv("LOCATION"),
staging_bucket=os.getenv("BUCKET"),
)
# --- Constants ---
PROMPT_PREFIX = os.getenv("PROMPT_PREFIX", "prompts_meta")
class GcpPrompt:
"""A wrapper class for the Vertex AI Prompt Management service."""
def __init__(self) -> None:
"""Initializes the GcpPrompt client."""
self.storage_client = storage.Client()
self.refresh_prompt_cache()
logger.info("Found %d existing prompts.", len(self.existing_prompts))
# Load a default/initial state for the prompt metadata.
self.prompt_meta: dict[str, Any] = {}
self.prompt_to_run: Prompt = Prompt()
self.refresh_bucket_cache()
def refresh_prompt_cache(self) -> None:
"""Refreshes the local cache of existing prompts from the service."""
self.existing_prompts = {p.display_name: p.prompt_id for p in prompts.list()}
def refresh_bucket_cache(self) -> None:
"""Refreshes the list of metadata files from GCS."""
blobs = self.storage_client.list_blobs(
os.getenv("BUCKET"), prefix=PROMPT_PREFIX
)
self.bucket_cache_index = [i.name for i in blobs]
logger.debug("Found %d metadata files in GCS.", len(self.bucket_cache_index))
def _get_metadata_blob_name(self) -> str:
"""Constructs the GCS blob name for the prompt metadata file."""
return (
f"{PROMPT_PREFIX}/{self.prompt_to_run.prompt_name}_"
f"{self.prompt_to_run._prompt_name}_{self.prompt_to_run._version_id}_"
f"{self.prompt_to_run._version_name}.json"
)
def save_prompt(self, check_existing: bool = False) -> str:
"""Saves the current prompt.
If check_existing is True, it ensures a prompt with the same name does not
already exist before creating version "1".
If check_existing is False, it creates a new version of an existing prompt.
Args:
check_existing: If True, raises an error if a prompt with the same
display name already exists.
Returns:
A string with the details of the saved prompt version.
"""
logger.info("Attempting to save prompt: %s", self.prompt_to_run.prompt_name)
if check_existing and self.prompt_to_run.prompt_name in self.existing_prompts:
raise ValueError(
f"Prompt with name '{self.prompt_to_run.prompt_name}' already exists. "
"To create a new version, load the prompt and save from the "
"'Existing Prompt' page."
)
if "generation_config" in self.prompt_meta and isinstance(
self.prompt_meta["generation_config"], dict
):
self.prompt_to_run.generation_config = GenerationConfig(
**self.prompt_meta["generation_config"]
)
elif not isinstance(self.prompt_to_run.generation_config, GenerationConfig):
logger.warning("No valid generation_config found. Using default.")
self.prompt_to_run.generation_config = GenerationConfig()
logger.debug("Prompt object being sent to SDK: %s", self.prompt_to_run)
self.prompt_to_run = prompts.create_version(prompt=self.prompt_to_run)
sdk_details = (
f"SDK Prompt Version Details received:\n"
f"- Resource Name: {self.prompt_to_run._prompt_name}\n"
f"- Version ID: {self.prompt_to_run._version_id}\n"
f"- Version Name: {self.prompt_to_run._version_name}\n"
f"- Display Name: {self.prompt_to_run.prompt_name}"
)
logger.info(sdk_details)
self.prompt_meta["name"] = self.prompt_to_run.prompt_name
self.write_to_bucket()
return sdk_details
def load_prompt(self, prompt_id: str, prompt_name: str, version_id: str) -> None:
"""Loads a specific version of a prompt and its associated metadata."""
self.prompt_to_run = prompts.get(prompt_id, version_id)
self.prompt_to_run.prompt_name = prompt_name
blob_name = self._get_metadata_blob_name()
none_blob_name = (
f"{PROMPT_PREFIX}/{None}_{None}_"
f"{self.prompt_to_run._version_id}_{self.prompt_to_run._version_name}.json"
)
self.refresh_bucket_cache()
bucket = self.storage_client.bucket(os.getenv("BUCKET"))
if blob_name in self.bucket_cache_index:
blob = bucket.blob(blob_name)
self.prompt_meta = json.loads(blob.download_as_string())
logger.info("Loaded metadata from %s", blob_name)
elif none_blob_name in self.bucket_cache_index:
logger.warning(
"Metadata not found at %s, using fallback %s", blob_name, none_blob_name
)
blob = bucket.blob(none_blob_name)
self.prompt_meta = json.loads(blob.download_as_string())
self.write_to_bucket()
else:
raise FileNotFoundError(
f"Could not find metadata file in GCS for prompt '{prompt_name}' "
f"version '{version_id}'. Looked for '{blob_name}' and '{none_blob_name}'."
)
def write_to_bucket(self) -> None:
"""Writes the current prompt_meta to a GCS blob."""
blob_name = self._get_metadata_blob_name()
bucket = self.storage_client.bucket(os.getenv("BUCKET"))
blob = bucket.blob(blob_name)
blob.upload_from_string(
json.dumps(self.prompt_meta, indent=2), content_type="application/json"
)
logger.info("Wrote metadata to gs://%s/%s", os.getenv("BUCKET"), blob_name)
def generate_response(self, variables: dict[str, Any]) -> str | None:
"""Generates a response from the currently loaded prompt and variables,
handling image uploads.
"""
if not self.prompt_to_run.prompt_data:
raise ValueError("Prompt data is not loaded. Cannot generate response.")
client = genai.Client(
vertexai=True,
project=os.getenv("PROJECT_ID"),
location=os.getenv("LOCATION"),
)
updated_contents: list[Any] = []
for key, value in variables.items():
logger.info("Iterating through variables")
if key == "image":
if "jpg" in value or "jpeg" in value:
mime_type = "image/jpeg"
elif "png" in value:
mime_type = "image/png"
else:
logger.warning(
f"Unsupported image type for {value}, attempting to process as is."
)
updated_contents.append(value)
continue
try:
# Create Image Part from URI
updated_contents.append(
genai.types.Part.from_uri(mime_type=mime_type, file_uri=value)
)
except Exception as e:
logger.error(
"Error creating Part from URI for variable '%s': %s",
key,
e,
exc_info=True,
)
updated_contents.append(value)
else:
updated_contents.append(value)
# Append the prompt text if it's not already included in variables
prompt_text = self.prompt_to_run.prompt_data
if not any(
isinstance(item, str) and item == prompt_text for item in updated_contents
):
updated_contents.append(prompt_text)
logger.info("Generating response with contents: %s", updated_contents)
model = self.prompt_to_run.model_name
response = client.models.generate_content(
model=model, contents=updated_contents
)
return response.text if response and response.text else None
def get_generation_config_dict(self):
"""Returns the generation configuration as a dictionary."""
return self.prompt_to_run.generation_config.to_dict()
def escape_special_characters(text: str) -> str:
"""Escapes special characters for embedding in a string.
Note: This function is also present in gcp_dataset.py and could be moved
to a shared utility module in the future.
"""
if not isinstance(text, str):
return text
text = text.replace("\\", "\\\\")
text = text.replace("\n", "\\n")
text = text.replace("\r", "\\r")
text = text.replace("\t", "\\t")
return text.replace('"', '\\"')
+27
View File
@@ -0,0 +1,27 @@
# Copyright 2024 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Utility functions and classes for the VAPO notebook.
This file imports all the functions and classes from the original vapo_lib.py file.
"""
import os
import sys
repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))
if repo_root not in sys.path:
sys.path.append(repo_root)
from gemini.prompts.prompt_optimizer.vapo_lib import * # noqa: F403
@@ -0,0 +1,30 @@
import os
import sys
from unittest.mock import patch
from streamlit.testing.v1 import AppTest
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
os.environ["PROJECT_ID"] = "test-project"
os.environ["LOCATION"] = "us-central1"
os.environ["BUCKET"] = "test-bucket"
@patch("vertexai.preview.prompts.list_versions")
@patch("src.gcp_prompt.GcpPrompt")
@patch("google.cloud.storage.Client")
@patch("vertexai.init")
def test_evaluation_load_prompt(mock_init, mock_storage, mock_gcp_prompt, mock_prompts):
# Setup mock behavior
mock_instance = mock_gcp_prompt.return_value
mock_instance.existing_prompts = {"test_prompt": "123"}
at = AppTest.from_file("pages/3_Evaluation.py")
at.run(timeout=30)
# Verify the warning is triggered
at.button(key="load_prompt_button").click().run(timeout=30)
warnings = [w.value for w in getattr(at, "warning", [])]
assert any("Please select a prompt before loading." in w for w in warnings)