chore: import upstream snapshot with attribution
This commit is contained in:
@@ -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.
|
||||
@@ -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.
|
||||
@@ -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.
|
||||
|
||||

|
||||
|
||||
**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:
|
||||
|
||||

|
||||
|
||||
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 |
@@ -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()
|
||||
@@ -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": [
|
||||
""
|
||||
]
|
||||
},
|
||||
{
|
||||
"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",
|
||||
"\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
|
||||
}
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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"
|
||||
@@ -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.
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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('"', '\\"')
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user