Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9cf8ce16fa | ||
|
|
4b983a2701 | ||
|
|
3b11c876bd | ||
|
|
cc0d791ef2 | ||
|
|
df83a345bb | ||
|
|
e6ded7beaa | ||
|
|
e64a4e89d5 | ||
|
|
7ac54985e4 | ||
|
|
6ce96a08d3 | ||
|
|
756711b3c9 | ||
|
|
a25d209139 | ||
|
|
59da536b9a | ||
|
|
1bc2839a2b | ||
|
|
0d5e268a1f | ||
|
|
1b019a76e4 | ||
|
|
87c1ed686a | ||
|
|
187fdc526c | ||
|
|
215c8eee3e | ||
|
|
f90cd0d6ed | ||
|
|
1985f06e99 | ||
|
|
ff428dc589 | ||
|
|
f6124370b0 | ||
|
|
b8822f5008 | ||
|
|
8976c57b9c | ||
|
|
3985da440e | ||
|
|
1c9092ced3 | ||
|
|
77b2af09ce | ||
|
|
c6d33c2a0d | ||
|
|
89772b320e | ||
|
|
a1f6d2c069 | ||
|
|
936a6adf77 | ||
|
|
37a85d53f4 | ||
|
|
0b0e362ab9 | ||
|
|
98103d462f | ||
|
|
9ea1cf3b86 | ||
|
|
8f3e6668e1 | ||
|
|
5d9853db5c | ||
|
|
003fb5121b | ||
|
|
a62695fb38 | ||
|
|
3c630fdbb8 | ||
|
|
6ca1d899d6 | ||
|
|
1894602fff | ||
|
|
31a52d6e92 | ||
|
|
0f9d9734c3 | ||
|
|
e85cf9a174 | ||
|
|
b4c0bbc1a0 | ||
|
|
24244351cd | ||
|
|
bf0e1300a9 | ||
|
|
cf048b6fe4 | ||
|
|
8c8820ecfa | ||
|
|
a1a52d8145 | ||
|
|
a06ce545e7 | ||
|
|
913780c4cb | ||
|
|
71be46e7d8 | ||
|
|
849e88a627 | ||
|
|
daf56bcd0b | ||
|
|
563f423b93 | ||
|
|
292e540e96 | ||
|
|
6c6a703c5a | ||
|
|
7ef83c6f73 | ||
|
|
aba6598109 | ||
|
|
88a6b8037e | ||
|
|
8845f7ab27 | ||
|
|
3b2e711a16 | ||
|
|
5c0629cdc7 | ||
|
|
e107d30807 | ||
|
|
b98ab36913 | ||
|
|
f1d90b5a71 | ||
|
|
f848db6132 | ||
|
|
cc9fffd945 | ||
|
|
7606a1de03 | ||
|
|
cca59aa753 | ||
|
|
a1907da27a | ||
|
|
e8cb7738d0 | ||
|
|
18e8d603de | ||
|
|
1bd5901fdb | ||
|
|
5b245024cd | ||
|
|
ba043c196c | ||
|
|
28ce8f6d7a | ||
|
|
3b5a8cad41 | ||
|
|
c6d7971bc9 | ||
|
|
dbe28965cb | ||
|
|
0d34d6bbea | ||
|
|
a1d898f35e | ||
|
|
772ee71bc3 | ||
|
|
425851cedc | ||
|
|
bcccbee164 | ||
|
|
5ae325528a | ||
|
|
86674effee | ||
|
|
8b4708c606 | ||
|
|
f3dd6cbca3 | ||
|
|
1f9e93993c | ||
|
|
062835174e | ||
|
|
cb4916f590 | ||
|
|
2933fe606b | ||
|
|
b468809df7 | ||
|
|
0417d8b9c4 | ||
|
|
7750e83fbb | ||
|
|
649800e646 | ||
|
|
42b35056fa | ||
|
|
24974eda95 | ||
|
|
ad41377783 | ||
|
|
36dea3ca01 | ||
|
|
b75b2ea4d7 | ||
|
|
28f7fc4445 | ||
|
|
bf2c1226fd | ||
|
|
2990c53292 | ||
|
|
531d9cfee0 | ||
|
|
ff18ec7af5 | ||
|
|
008eb409ef | ||
|
|
b2dba4b568 | ||
|
|
665547f790 | ||
|
|
5078c44eb8 | ||
|
|
da6e46531e | ||
|
|
0a4091a3b1 | ||
|
|
5afa83dd25 | ||
|
|
8cab85d6ad | ||
|
|
a7f3940635 | ||
|
|
5bb1a75a48 | ||
|
|
8fe4985aa8 | ||
|
|
996b6534d9 | ||
|
|
1061ae5348 | ||
|
|
4fac3a630f | ||
|
|
1f2db2903c | ||
|
|
912a52de70 | ||
|
|
5e29090e86 | ||
|
|
cd8fcd1839 | ||
|
|
e51075ec4b | ||
|
|
9709c0dddb | ||
|
|
a21ae41762 | ||
|
|
ff5ff7b609 | ||
|
|
23748f443e | ||
|
|
fc6b2167de | ||
|
|
e79a45358c | ||
|
|
20d19fb11c | ||
|
|
9c9f7a6e2a | ||
|
|
ede41c2115 | ||
|
|
ca53786c04 | ||
|
|
419f8310c9 | ||
|
|
996b690e03 | ||
|
|
5b9d04d63a | ||
|
|
9ed3c2f83d | ||
|
|
5fc93c8bdb | ||
|
|
3cf46226e9 | ||
|
|
439f6a0cae | ||
|
|
993898bb71 | ||
|
|
9e9e639375 | ||
|
|
633cf6a799 | ||
|
|
8b618bc455 | ||
|
|
6bbe3bcfe0 | ||
|
|
814827ac19 | ||
|
|
7a613785b9 | ||
|
|
100243e90a | ||
|
|
6424515b03 | ||
|
|
8d22b221b4 | ||
|
|
2a8ad7cdbb | ||
|
|
3239b301f2 | ||
|
|
acb10d14b8 | ||
|
|
778d145970 | ||
|
|
4d00356f4b | ||
|
|
0654305994 | ||
|
|
c53f392c5a | ||
|
|
f4b56e92ae | ||
|
|
45ec1cf18a | ||
|
|
f6c8bcf937 | ||
|
|
9e590d5a9f | ||
|
|
46e0ea4f1c | ||
|
|
300fce6b9f | ||
|
|
52e3066c38 | ||
|
|
27ebf52198 | ||
|
|
ca7d4e153e | ||
|
|
5b6c766629 | ||
|
|
5efa51206f | ||
|
|
e604a4d43e | ||
|
|
23af5373ec | ||
|
|
7577c0b1fc | ||
|
|
b648f9e73b | ||
|
|
966bbc49a7 | ||
|
|
f5d341ae45 | ||
|
|
70770a50c7 | ||
|
|
f181c39cbf | ||
|
|
a5637f87f2 | ||
|
|
0103299084 | ||
|
|
1867536d76 | ||
|
|
0edae683e7 | ||
|
|
8471b5cb6f | ||
|
|
87f540ac53 | ||
|
|
babeba9f02 | ||
|
|
85c649dd26 | ||
|
|
b7135ae1f0 | ||
|
|
2f5119a266 | ||
|
|
23e64ca76f | ||
|
|
0ba5a62cc9 | ||
|
|
0be2c6fd0c | ||
|
|
cef4928c49 | ||
|
|
9bb8107110 | ||
|
|
b075990d88 | ||
|
|
6a83c4c695 | ||
|
|
820c0f8db4 | ||
|
|
4ab197a4ba | ||
|
|
2a5877fbd1 | ||
|
|
090e1d9fee | ||
|
|
2e049d4830 | ||
|
|
e3320d2126 | ||
|
|
5fc0e03ca3 | ||
|
|
ff2a16237d | ||
|
|
d26f081642 | ||
|
|
8a0a39176c | ||
|
|
646532ea69 | ||
|
|
9e96a3da67 | ||
|
|
db34e1fbd5 | ||
|
|
82308acbac | ||
|
|
ee0ba75d1e | ||
|
|
d61aedc721 | ||
|
|
19f7f94af5 | ||
|
|
4e5ce9b226 | ||
|
|
99938244f4 | ||
|
|
0cc7be4a6a | ||
|
|
b6bde41850 | ||
|
|
52444a0933 | ||
|
|
bab9c398fd | ||
|
|
447affcc93 | ||
|
|
6132c37be0 | ||
|
|
8d7f59aeec | ||
|
|
bd327ad424 | ||
|
|
3b1fbdb382 | ||
|
|
954043a729 | ||
|
|
16ef9ee80e | ||
|
|
79301b4a4d | ||
|
|
0b38d02e6f | ||
|
|
c52ff25ba4 | ||
|
|
424400bace | ||
|
|
b1dfac2043 | ||
|
|
b81ffcddab | ||
|
|
21d8f144aa | ||
|
|
065a674305 | ||
|
|
571d498d08 | ||
|
|
e936882123 | ||
|
|
f754f99052 | ||
|
|
aa5523a5e9 | ||
|
|
81393ede1a | ||
|
|
5a1c0222da | ||
|
|
1f9326bd56 | ||
|
|
a94cae2e79 | ||
|
|
cb861713c8 | ||
|
|
3a55087789 | ||
|
|
d359b21f3e | ||
|
|
c7d4123b25 | ||
|
|
07a8bb2d0c | ||
|
|
f6b6f365b6 |
@@ -365,7 +365,7 @@ def process_and_execute_notebook(
|
||||
# Use gcloud to get tail
|
||||
try:
|
||||
result.error_message = subprocess.check_output(
|
||||
["gsutil", "cat", "-r", "-1000", log_file_uri], encoding="UTF-8"
|
||||
["gcloud", "storage", "cat", "--range", "-1000", log_file_uri], encoding="UTF-8"
|
||||
)
|
||||
except Exception as error:
|
||||
result.error_message = str(error)
|
||||
|
||||
@@ -56,8 +56,8 @@ def execute_notebook(
|
||||
print("\n=== DOWNLOAD EXECUTED NOTEBOOK ===\n")
|
||||
print(f"Please debug the executed notebook by downloading the executed notebook:")
|
||||
|
||||
print("Option 1. Using gsutil. Run the following command in your terminal.")
|
||||
print(f'\tgsutil cp "{output_file_or_uri}" .')
|
||||
print("Option 1. Using gcloud storage. Run the following command in your terminal.")
|
||||
print(f'\tgcloud storage cp "{output_file_or_uri}" .')
|
||||
|
||||
print("Option 2. Using this link.")
|
||||
print(f"\thttps://storage.googleapis.com/{output_file_or_uri[5:]}")
|
||||
|
||||
@@ -108,7 +108,7 @@ class VertexAIInstallProprocessor(Preprocessor):
|
||||
if "google-cloud-aiplatform" not in content:
|
||||
return content
|
||||
return (
|
||||
f"gsutil cp {self.vertex_ai_wheel} google-cloud-aiplatform.whl\n" +
|
||||
f"gcloud storage cp {self.vertex_ai_wheel} google-cloud-aiplatform.whl\n" +
|
||||
content.replace("google-cloud-aiplatform\n", "google-cloud-aiplatform.whl\n")
|
||||
.replace("google-cloud-aiplatform ", "google-cloud-aiplatform.whl ")
|
||||
)
|
||||
|
||||
@@ -15,7 +15,7 @@ def download_file(bucket_name: str, blob_name: str, destination_file: str) -> st
|
||||
remote_file_path = "".join(["gs://", "/".join([bucket_name, blob_name])])
|
||||
|
||||
subprocess.check_output(
|
||||
["gsutil", "cp", remote_file_path, destination_file], encoding="UTF-8"
|
||||
["gcloud", "storage", "cp", remote_file_path, destination_file], encoding="UTF-8"
|
||||
)
|
||||
|
||||
return destination_file
|
||||
@@ -27,7 +27,7 @@ def upload_file(
|
||||
) -> str:
|
||||
"""Copies a local file to a GCS path"""
|
||||
subprocess.check_output(
|
||||
["gsutil", "cp", local_file_path, remote_file_path], encoding="UTF-8"
|
||||
["gcloud", "storage", "cp", local_file_path, remote_file_path], encoding="UTF-8"
|
||||
)
|
||||
|
||||
return remote_file_path
|
||||
|
||||
@@ -7,11 +7,11 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version: '3.x'
|
||||
python-version: '3.12'
|
||||
- name: Fetch pull request branch
|
||||
uses: actions/checkout@v4
|
||||
uses: actions/checkout@v6
|
||||
with:
|
||||
fetch-depth: 0
|
||||
- name: Fetch base main branch
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
# 2. To lint specific notebooks:
|
||||
# docker run -v ${PWD}:/setup/app gcr.io/python-docs-samples-tests/notebook_linter:latest notebooks/1.ipynb notebooks/2.ipynb
|
||||
|
||||
FROM python:3.13
|
||||
FROM python:3.14
|
||||
|
||||
WORKDIR setup
|
||||
|
||||
|
||||
@@ -2,9 +2,9 @@ git+https://github.com/tensorflow/docs
|
||||
ipython
|
||||
jupyter
|
||||
nbconvert
|
||||
black==25.1.0
|
||||
pyupgrade==3.20.0
|
||||
isort==6.0.1
|
||||
black==26.5.1
|
||||
pyupgrade==3.21.2
|
||||
isort==8.0.1
|
||||
flake8==7.3.0
|
||||
nbqa==1.9.1
|
||||
|
||||
|
||||
@@ -1,12 +1,12 @@
|
||||
#  Google Cloud Vertex AI Samples
|
||||
|
||||
This repository contains notebooks, code samples, sample apps, and other resources that demonstrate how to use, develop and manage machine learning and generative AI workflows using Google Cloud Vertex AI.
|
||||
This repository contains notebooks, code samples, sample apps, skills, and other resources that demonstrate how to use, develop and manage machine learning and generative AI workflows using Google Cloud Vertex AI.
|
||||
|
||||
## Overview
|
||||
|
||||
[Vertex AI](https://cloud.google.com/vertex-ai) is a fully-managed, unified AI development platform for building and using generative AI. This repository is designed to help you get started with Vertex AI. Whether you're new to Vertex AI or an experienced ML practitioner, you'll find valuable resources here.
|
||||
|
||||
For more Vertex AI Generative AI notebook samples, please visit the Vertex AI [Generative AI](https://github.com/GoogleCloudPlatform/generative-ai) GitHub repository.
|
||||
⚠️ For more Vertex AI Generative AI notebook samples, please visit the Vertex AI [Generative AI](https://github.com/GoogleCloudPlatform/generative-ai) GitHub repository.
|
||||
|
||||
## Explore, learn and contribute
|
||||
|
||||
@@ -16,11 +16,11 @@ You can explore, learn, and contribute to this repository to unleash the full po
|
||||
|
||||
Explore this repository, follow the links in the header section of each of the notebooks to -
|
||||
|
||||
 Open and run the notebook in [Colab](https://colab.google/)\
|
||||
 Open and run the notebook in [Colab Enterprise](https://cloud.google.com/colab/docs/introduction)\
|
||||
 Open and run the notebook in [Vertex AI Workbench](https://cloud.google.com/vertex-ai/docs/workbench/introduction)\
|
||||
 View the notebook on Github
|
||||
|
||||
- Open and run the notebook in [Colab](https://colab.google/)
|
||||
- Open and run the notebook in [Colab Enterprise](https://cloud.google.com/colab/docs/introduction)
|
||||
- Open and run the notebook in [Vertex AI Workbench](https://cloud.google.com/vertex-ai/docs/workbench/introduction)
|
||||
- View the notebook on Github
|
||||
|
||||
### Contribute
|
||||
|
||||
See the [Contributing Guide](https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/master/CONTRIBUTING.md).
|
||||
@@ -35,7 +35,7 @@ To get started using Vertex AI, you must have a Google Cloud project.
|
||||
|
||||
## Repository structure
|
||||
|
||||
```bash
|
||||
```text
|
||||
├── notebooks
|
||||
│ ├── official - Notebooks demonstrating use of each Vertex AI service
|
||||
│ │ ├── automl
|
||||
@@ -45,7 +45,23 @@ To get started using Vertex AI, you must have a Google Cloud project.
|
||||
│ │ ├── model_garden
|
||||
│ │ ├── ...
|
||||
├── community-content - Sample code and tutorials contributed by the community
|
||||
|
||||
├── docs - Deep-dive documentation and advanced setup guides
|
||||
└── skills - Suite of AI Agent "Skills" for Vertex AI
|
||||
├── README.md # Developer guide for Vertex AI skills
|
||||
├── vertex-ai/ # Primary router for Vertex AI tasks
|
||||
│ └── SKILL.md # Entry point that routes across capabilities
|
||||
├── genai-sdk/ # Gemini API usage with Gen AI SDK
|
||||
│ └── SKILL.md # Guides for Python, JS/TS, Go, Java, C#
|
||||
├── vertex-deploy/ # Deploying models to Endpoints
|
||||
│ └── SKILL.md # Commands for open models & custom weights
|
||||
├── vertex-inference/ # Inferencing with GenAI models
|
||||
│ └── SKILL.md # Code samples for Gemini and OpenMaaS
|
||||
└── vertex-tuning/ # Secondary router for model fine-tuning
|
||||
├── SKILL.md # Router for tuning tasks
|
||||
├── gemini/ # Fine-tuning first-party Gemini models
|
||||
│ └── SKILL.md
|
||||
└── open-model/ # Fine-tuning third-party open models
|
||||
└── SKILL.md
|
||||
```
|
||||
## Examples
|
||||
|
||||
|
||||
@@ -148,7 +148,7 @@ implementation:
|
||||
|
||||
# Downloading the model archive from GCS
|
||||
# TODO: Fix gsutil bugs (requires project ID, has auth issues) and use gsutil instead.
|
||||
# gsutil cp "$model_archive_uri" "$model_archive_local_path"
|
||||
# gcloud storage cp "$model_archive_uri" "$model_archive_local_path"
|
||||
pip install google-cloud-storage
|
||||
python -c '
|
||||
import sys
|
||||
|
||||
@@ -24,12 +24,12 @@ implementation:
|
||||
|
||||
# Checking whether the URI points to a single blob, a directory or a URI pattern
|
||||
# URI points to a blob when that URI does not end with slash and listing that URI only yields the same URI
|
||||
if [[ "$uri" != */ ]] && (gsutil ls "$uri" | grep --fixed-strings --line-regexp "$uri"); then
|
||||
if [[ "$uri" != */ ]] && (gcloud storage ls "$uri" | grep --fixed-strings --line-regexp "$uri"); then
|
||||
mkdir -p "$(dirname "$output_path")"
|
||||
gsutil -m cp -r "$uri" "$output_path"
|
||||
gcloud storage cp --recursive "$uri" "$output_path"
|
||||
else
|
||||
mkdir -p "$output_path" # When source path is a directory, gsutil requires the destination to also be a directory
|
||||
gsutil -m rsync -r "$uri" "$output_path" # gsutil cp has different path handling than Linux cp. It always puts the source directory (name) inside the destination directory. gsutil rsync does not have that problem.
|
||||
gcloud storage rsync --recursive "$uri" "$output_path" # gsutil cp has different path handling than Linux cp. It always puts the source directory (name) inside the destination directory. gsutil rsync does not have that problem.
|
||||
fi
|
||||
- inputValue: GCS path
|
||||
- outputPath: Data
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
torch==2.8.0
|
||||
torch==2.13.0
|
||||
torchvision==0.9.1
|
||||
tensorboard==2.5.0
|
||||
@@ -1,3 +1,3 @@
|
||||
torch==2.7.0
|
||||
torch==2.13.0
|
||||
torchvision==0.9.1
|
||||
tensorboard==2.5.0
|
||||
@@ -110,7 +110,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil ls $gcs_output_uri_prefix"
|
||||
"! gcloud storage ls $gcs_output_uri_prefix"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -192,7 +192,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil cp -r $gcs_output_uri_prefix/model ./model_server/"
|
||||
"! gcloud storage cp --recursive $gcs_output_uri_prefix/model ./model_server/"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -556,7 +556,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil rm -rf $gcs_output_uri_prefix"
|
||||
"! gcloud storage rm --recursive --continue-on-error $gcs_output_uri_prefix"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
@@ -412,7 +412,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil ls $gcs_output_uri_prefix"
|
||||
"! gcloud storage ls $gcs_output_uri_prefix"
|
||||
]
|
||||
}
|
||||
],
|
||||
|
||||
@@ -77,4 +77,4 @@ echo "After the job is completed successfully, model files will be saved at $JOB
|
||||
|
||||
# # Verify the model was exported
|
||||
# echo "Verify the model was exported:"
|
||||
# gsutil ls ${JOB_DIR}/
|
||||
# gcloud storage ls ${JOB_DIR}/
|
||||
|
||||
@@ -34,4 +34,4 @@ RUN echo "service_envelope=json\n" "inference_address=http://0.0.0.0:${AIP_H
|
||||
USER model-server
|
||||
|
||||
# run Torchserve HTTP serve to respond to prediction requests
|
||||
CMD ["echo", "AIP_STORAGE_URI=${AIP_STORAGE_URI}", ";", "gsutil", "cp", "-r", "${AIP_STORAGE_URI}/${MODEL_NAME}.mar", "/home/model-server/model-store/", ";", "ls", "-ltr", "/home/model-server/model-store/", ";", "torchserve", "--start", "--ts-config=/home/model-server/config.properties", "--models", "${MODEL_NAME}=${MODEL_NAME}.mar", "--model-store", "/home/model-server/model-store"]
|
||||
CMD ["echo", "AIP_STORAGE_URI=${AIP_STORAGE_URI}", ";", "gcloud", "storage", "cp", "--recursive", "${AIP_STORAGE_URI}/${MODEL_NAME}.mar", "/home/model-server/model-store/", ";", "ls", "-ltr", "/home/model-server/model-store/", ";", "torchserve", "--start", "--ts-config=/home/model-server/config.properties", "--models", "${MODEL_NAME}=${MODEL_NAME}.mar", "--model-store", "/home/model-server/model-store"]
|
||||
|
||||
@@ -67,4 +67,4 @@ echo "After the job is completed successfully, model files will be saved at $JOB
|
||||
|
||||
# # Verify the model was exported
|
||||
# echo "Verify the model was exported:"
|
||||
# gsutil ls ${JOB_DIR}/
|
||||
# gcloud storage ls ${JOB_DIR}/
|
||||
|
||||
@@ -478,8 +478,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil mb -l $REGION $BUCKET_NAME"
|
||||
]
|
||||
"! gcloud storage buckets create --location $REGION $BUCKET_NAME" ]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
@@ -498,8 +497,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil ls -al $BUCKET_NAME"
|
||||
]
|
||||
"! gcloud storage ls --all-versions --long $BUCKET_NAME" ]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
@@ -582,8 +580,7 @@
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Download the sample data into your RAW_DATA_PATH\n",
|
||||
"! gsutil cp \"gs://cloud-samples-data/vertex-ai/community-content/tf_agents_bandits_movie_recommendation_with_kfp_and_vertex_sdk/u.data\" $RAW_DATA_PATH"
|
||||
]
|
||||
"! gcloud storage cp \"gs://cloud-samples-data/vertex-ai/community-content/tf_agents_bandits_movie_recommendation_with_kfp_and_vertex_sdk/u.data\" $RAW_DATA_PATH" ]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
@@ -1621,9 +1618,7 @@
|
||||
"! gcloud scheduler jobs delete $SIMULATOR_SCHEDULER_JOB --quiet\n",
|
||||
"\n",
|
||||
"# Delete Cloud Storage objects that were created.\n",
|
||||
"! gsutil -m rm -r $PIPELINE_ROOT\n",
|
||||
"! gsutil -m rm -r $TRAINING_ARTIFACTS_DIR"
|
||||
]
|
||||
"! gcloud storage rm --recursive $PIPELINE_ROOT\n", "! gcloud storage rm --recursive $TRAINING_ARTIFACTS_DIR" ]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
google-cloud-bigquery==2.20.0
|
||||
tensorflow==2.12.1
|
||||
pillow==10.3.0
|
||||
pillow==12.3.0
|
||||
tf-agents==0.8.0
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
google-cloud-pubsub==2.5.0
|
||||
pillow==10.3.0
|
||||
pillow==12.3.0
|
||||
tf-agents==0.8.0
|
||||
tensorflow==2.12.1
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
dataclasses==0.6
|
||||
google-cloud-aiplatform==1.8.1
|
||||
tensorflow==2.12.1
|
||||
pillow==10.3.0
|
||||
pillow==12.3.0
|
||||
tf-agents==0.8.0
|
||||
@@ -398,6 +398,7 @@
|
||||
"if not IS_GOOGLE_CLOUD_NOTEBOOK:\n",
|
||||
" if \"google.colab\" in sys.modules:\n",
|
||||
" from google.colab import auth as google_auth\n",
|
||||
"\n",
|
||||
" google_auth.authenticate_user()\n",
|
||||
"\n",
|
||||
" # If you are running this notebook locally, replace the string below with the\n",
|
||||
@@ -472,7 +473,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil mb -l $REGION $BUCKET_NAME"
|
||||
"! gcloud storage buckets create --location $REGION $BUCKET_NAME"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -492,7 +493,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil ls -al $BUCKET_NAME"
|
||||
"! gcloud storage ls --all-versions --long $BUCKET_NAME"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -565,7 +566,7 @@
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Copy the sample data into your DATA_PATH\n",
|
||||
"! gsutil cp \"gs://cloud-samples-data/vertex-ai/community-content/tf_agents_bandits_movie_recommendation_with_kfp_and_vertex_sdk/u.data\" $DATA_PATH"
|
||||
"! gcloud storage cp \"gs://cloud-samples-data/vertex-ai/community-content/tf_agents_bandits_movie_recommendation_with_kfp_and_vertex_sdk/u.data\" $DATA_PATH"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -579,11 +580,15 @@
|
||||
"# Set hyperparameters.\n",
|
||||
"BATCH_SIZE = 8 # @param {type:\"integer\"} Training and prediction batch size.\n",
|
||||
"TRAINING_LOOPS = 5 # @param {type:\"integer\"} Number of training iterations.\n",
|
||||
"STEPS_PER_LOOP = 2 # @param {type:\"integer\"} Number of driver steps per training iteration.\n",
|
||||
"STEPS_PER_LOOP = (\n",
|
||||
" 2 # @param {type:\"integer\"} Number of driver steps per training iteration.\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"# Set MovieLens simulation environment parameters.\n",
|
||||
"RANK_K = 20 # @param {type:\"integer\"} Rank for matrix factorization in the MovieLens environment; also the observation dimension.\n",
|
||||
"NUM_ACTIONS = 20 # @param {type:\"integer\"} Number of actions (movie items) to choose from.\n",
|
||||
"NUM_ACTIONS = (\n",
|
||||
" 20 # @param {type:\"integer\"} Number of actions (movie items) to choose from.\n",
|
||||
")\n",
|
||||
"PER_ARM = False # Use the non-per-arm version of the MovieLens environment.\n",
|
||||
"\n",
|
||||
"# Set agent parameters.\n",
|
||||
@@ -621,7 +626,8 @@
|
||||
"source": [
|
||||
"# Define RL environment.\n",
|
||||
"env = movielens_py_environment.MovieLensPyEnvironment(\n",
|
||||
" DATA_PATH, RANK_K, BATCH_SIZE, num_movies=NUM_ACTIONS, csv_delimiter=\"\\t\")\n",
|
||||
" DATA_PATH, RANK_K, BATCH_SIZE, num_movies=NUM_ACTIONS, csv_delimiter=\"\\t\"\n",
|
||||
")\n",
|
||||
"environment = tf_py_environment.TFPyEnvironment(env)\n",
|
||||
"\n",
|
||||
"# Define RL agent/algorithm.\n",
|
||||
@@ -631,7 +637,8 @@
|
||||
" tikhonov_weight=TIKHONOV_WEIGHT,\n",
|
||||
" alpha=AGENT_ALPHA,\n",
|
||||
" dtype=tf.float32,\n",
|
||||
" accepts_per_arm_features=PER_ARM)\n",
|
||||
" accepts_per_arm_features=PER_ARM,\n",
|
||||
")\n",
|
||||
"print(\"TimeStep Spec (for each batch):\\n\", agent.time_step_spec, \"\\n\")\n",
|
||||
"print(\"Action Spec (for each batch):\\n\", agent.action_spec, \"\\n\")\n",
|
||||
"print(\"Reward Spec (for each batch):\\n\", environment.reward_spec(), \"\\n\")\n",
|
||||
@@ -639,7 +646,8 @@
|
||||
"# Define RL metric.\n",
|
||||
"optimal_reward_fn = functools.partial(\n",
|
||||
" environment_utilities.compute_optimal_reward_with_movielens_environment,\n",
|
||||
" environment=environment)\n",
|
||||
" environment=environment,\n",
|
||||
")\n",
|
||||
"regret_metric = tf_bandit_metrics.RegretMetric(optimal_reward_fn)\n",
|
||||
"metrics = [regret_metric]"
|
||||
]
|
||||
@@ -704,35 +712,38 @@
|
||||
" if training_data_spec_transformation_fn is None:\n",
|
||||
" data_spec = agent.policy.trajectory_spec\n",
|
||||
" else:\n",
|
||||
" data_spec = training_data_spec_transformation_fn(\n",
|
||||
" agent.policy.trajectory_spec)\n",
|
||||
" replay_buffer = trainer.get_replay_buffer(data_spec, environment.batch_size,\n",
|
||||
" steps_per_loop)\n",
|
||||
" data_spec = training_data_spec_transformation_fn(agent.policy.trajectory_spec)\n",
|
||||
" replay_buffer = trainer.get_replay_buffer(\n",
|
||||
" data_spec, environment.batch_size, steps_per_loop\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" # `step_metric` records the number of individual rounds of bandit interaction;\n",
|
||||
" # that is, (number of trajectories) * batch_size.\n",
|
||||
" step_metric = tf_metrics.EnvironmentSteps()\n",
|
||||
" metrics = [\n",
|
||||
" tf_metrics.NumberOfEpisodes(),\n",
|
||||
" tf_metrics.AverageEpisodeLengthMetric(batch_size=environment.batch_size)\n",
|
||||
" tf_metrics.AverageEpisodeLengthMetric(batch_size=environment.batch_size),\n",
|
||||
" ]\n",
|
||||
" if additional_metrics:\n",
|
||||
" metrics += additional_metrics\n",
|
||||
"\n",
|
||||
" if isinstance(environment.reward_spec(), dict):\n",
|
||||
" metrics += [tf_metrics.AverageReturnMultiMetric(\n",
|
||||
" reward_spec=environment.reward_spec(),\n",
|
||||
" batch_size=environment.batch_size)]\n",
|
||||
" else:\n",
|
||||
" metrics += [\n",
|
||||
" tf_metrics.AverageReturnMetric(batch_size=environment.batch_size)]\n",
|
||||
" tf_metrics.AverageReturnMultiMetric(\n",
|
||||
" reward_spec=environment.reward_spec(), batch_size=environment.batch_size\n",
|
||||
" )\n",
|
||||
" ]\n",
|
||||
" else:\n",
|
||||
" metrics += [tf_metrics.AverageReturnMetric(batch_size=environment.batch_size)]\n",
|
||||
"\n",
|
||||
" # Store intermediate metric results, indexed by metric names.\n",
|
||||
" metric_results = defaultdict(list)\n",
|
||||
"\n",
|
||||
" if training_data_spec_transformation_fn is not None:\n",
|
||||
" def add_batch_fn(data): return replay_buffer.add_batch(training_data_spec_transformation_fn(data)) \n",
|
||||
" \n",
|
||||
"\n",
|
||||
" def add_batch_fn(data):\n",
|
||||
" return replay_buffer.add_batch(training_data_spec_transformation_fn(data))\n",
|
||||
"\n",
|
||||
" else:\n",
|
||||
" add_batch_fn = replay_buffer.add_batch\n",
|
||||
"\n",
|
||||
@@ -742,10 +753,12 @@
|
||||
" env=environment,\n",
|
||||
" policy=agent.collect_policy,\n",
|
||||
" num_steps=steps_per_loop * environment.batch_size,\n",
|
||||
" observers=observers)\n",
|
||||
" observers=observers,\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" training_loop = trainer.get_training_loop_fn(\n",
|
||||
" driver, replay_buffer, agent, steps_per_loop)\n",
|
||||
" driver, replay_buffer, agent, steps_per_loop\n",
|
||||
" )\n",
|
||||
" saver = policy_saver.PolicySaver(agent.policy)\n",
|
||||
"\n",
|
||||
" for _ in range(training_loops):\n",
|
||||
@@ -783,7 +796,8 @@
|
||||
" environment=environment,\n",
|
||||
" training_loops=TRAINING_LOOPS,\n",
|
||||
" steps_per_loop=STEPS_PER_LOOP,\n",
|
||||
" additional_metrics=metrics)\n",
|
||||
" additional_metrics=metrics,\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"tf.profiler.experimental.stop()"
|
||||
]
|
||||
@@ -1092,11 +1106,15 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"RUN_HYPERPARAMETER_TUNING = True # Execute hyperparameter tuning instead of regular training.\n",
|
||||
"RUN_HYPERPARAMETER_TUNING = (\n",
|
||||
" True # Execute hyperparameter tuning instead of regular training.\n",
|
||||
")\n",
|
||||
"TRAIN_WITH_BEST_HYPERPARAMETERS = False # Do not train.\n",
|
||||
"\n",
|
||||
"HPTUNING_RESULT_DIR = \"hptuning/\" # @param {type: \"string\"} Directory to store the best hyperparameter(s) in `BUCKET_NAME` and locally (temporarily).\n",
|
||||
"HPTUNING_RESULT_PATH = os.path.join(HPTUNING_RESULT_DIR, \"result.json\") # @param {type: \"string\"} Path to the file containing the best hyperparameter(s)."
|
||||
"HPTUNING_RESULT_PATH = os.path.join(\n",
|
||||
" HPTUNING_RESULT_DIR, \"result.json\"\n",
|
||||
") # @param {type: \"string\"} Path to the file containing the best hyperparameter(s)."
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -1124,7 +1142,7 @@
|
||||
" image_uri: str,\n",
|
||||
" args: List[str],\n",
|
||||
" location: str = \"us-central1\",\n",
|
||||
" api_endpoint: str = \"us-central1-aiplatform.googleapis.com\"\n",
|
||||
" api_endpoint: str = \"us-central1-aiplatform.googleapis.com\",\n",
|
||||
") -> None:\n",
|
||||
" \"\"\"Creates a hyperparameter tuning job using a custom container.\n",
|
||||
"\n",
|
||||
@@ -1197,8 +1215,8 @@
|
||||
"\n",
|
||||
" # Create job\n",
|
||||
" response = client.create_hyperparameter_tuning_job(\n",
|
||||
" parent=parent,\n",
|
||||
" hyperparameter_tuning_job=hyperparameter_tuning_job)\n",
|
||||
" parent=parent, hyperparameter_tuning_job=hyperparameter_tuning_job\n",
|
||||
" )\n",
|
||||
" job_id = response.name.split(\"/\")[-1]\n",
|
||||
" print(\"Job ID:\", job_id)\n",
|
||||
" print(\"Job config:\", response)\n",
|
||||
@@ -1242,7 +1260,8 @@
|
||||
" image_uri=f\"gcr.io/{PROJECT_ID}/{HPTUNING_TRAINING_CONTAINER}:latest\",\n",
|
||||
" args=args,\n",
|
||||
" location=REGION,\n",
|
||||
" api_endpoint=f\"{REGION}-aiplatform.googleapis.com\")"
|
||||
" api_endpoint=f\"{REGION}-aiplatform.googleapis.com\",\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -1292,7 +1311,8 @@
|
||||
" name = client.hyperparameter_tuning_job_path(\n",
|
||||
" project=project,\n",
|
||||
" location=location,\n",
|
||||
" hyperparameter_tuning_job=hyperparameter_tuning_job_id)\n",
|
||||
" hyperparameter_tuning_job=hyperparameter_tuning_job_id,\n",
|
||||
" )\n",
|
||||
" response = client.get_hyperparameter_tuning_job(name=name)\n",
|
||||
" return response"
|
||||
]
|
||||
@@ -1313,7 +1333,8 @@
|
||||
" location=REGION,\n",
|
||||
" api_endpoint=f\"{REGION}-aiplatform.googleapis.com\")\n",
|
||||
" if response.state.name == 'JOB_STATE_SUCCEEDED':\n",
|
||||
" print(\"Job succeeded.\\nJob Time:\", response.update_time - response.create_time)\n",
|
||||
" print(\"Job succeeded.\n",
|
||||
"Job Time:\", response.update_time - response.create_time)\n",
|
||||
" trials = response.trials\n",
|
||||
" print(\"Trials:\", trials)\n",
|
||||
" break\n",
|
||||
@@ -1348,8 +1369,8 @@
|
||||
"if trials:\n",
|
||||
" # Dict mapping from metric names to the best metric values seen so far\n",
|
||||
" best_objective_values = dict.fromkeys(\n",
|
||||
" [metric.metric_id for metric in trials[0].final_measurement.metrics],\n",
|
||||
" -np.inf)\n",
|
||||
" [metric.metric_id for metric in trials[0].final_measurement.metrics], -np.inf\n",
|
||||
" )\n",
|
||||
" # Dict mapping from metric names to a list of the best combination(s) of\n",
|
||||
" # hyperparameter(s). Each combination is a dict mapping from hyperparameter\n",
|
||||
" # names to their values.\n",
|
||||
@@ -1358,12 +1379,13 @@
|
||||
" # `final_measurement` and `parameters` are `RepeatedComposite` objects.\n",
|
||||
" # Reference the structure above to extract the value of your interest.\n",
|
||||
" for metric in trial.final_measurement.metrics:\n",
|
||||
" params = {\n",
|
||||
" param.parameter_id: param.value for param in trial.parameters}\n",
|
||||
" params = {param.parameter_id: param.value for param in trial.parameters}\n",
|
||||
" if metric.value > best_objective_values[metric.metric_id]:\n",
|
||||
" best_params[metric.metric_id] = [params]\n",
|
||||
" elif metric.value == best_objective_values[metric.metric_id]:\n",
|
||||
" best_params[param.parameter_id].append(params) # Handle cases where multiple hyperparameter values lead to the same performance.\n",
|
||||
" best_params[param.parameter_id].append(\n",
|
||||
" params\n",
|
||||
" ) # Handle cases where multiple hyperparameter values lead to the same performance.\n",
|
||||
" print(\"Best hyperparameter value(s):\")\n",
|
||||
" for metric, params in best_params.items():\n",
|
||||
" print(f\"Metric={metric}: {sorted(params)}\")\n",
|
||||
@@ -1443,7 +1465,9 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"PREDICTION_CONTAINER = \"prediction-custom-container\" # @param {type:\"string\"} Name of the container image."
|
||||
"PREDICTION_CONTAINER = (\n",
|
||||
" \"prediction-custom-container\" # @param {type:\"string\"} Name of the container image.\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -1475,7 +1499,7 @@
|
||||
" machineType: 'E2_HIGHCPU_8'\"\"\".format(\n",
|
||||
" PROJECT_ID=PROJECT_ID,\n",
|
||||
" PREDICTION_CONTAINER=PREDICTION_CONTAINER,\n",
|
||||
" ARTIFACTS_DIR=ARTIFACTS_DIR\n",
|
||||
" ARTIFACTS_DIR=ARTIFACTS_DIR,\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"with open(\"cloudbuild.yaml\", \"w\") as fp:\n",
|
||||
@@ -1592,8 +1616,12 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"RUN_HYPERPARAMETER_TUNING = False # Execute regular training instead of hyperparameter tuning.\n",
|
||||
"TRAIN_WITH_BEST_HYPERPARAMETERS = True # @param {type:\"bool\"} Whether to use learned hyperparameters in training."
|
||||
"RUN_HYPERPARAMETER_TUNING = (\n",
|
||||
" False # Execute regular training instead of hyperparameter tuning.\n",
|
||||
")\n",
|
||||
"TRAIN_WITH_BEST_HYPERPARAMETERS = (\n",
|
||||
" True # @param {type:\"bool\"} Whether to use learned hyperparameters in training.\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -1633,10 +1661,12 @@
|
||||
"job = aiplatform.CustomContainerTrainingJob(\n",
|
||||
" display_name=\"train-movielens\",\n",
|
||||
" container_uri=f\"gcr.io/{PROJECT_ID}/{HPTUNING_TRAINING_CONTAINER}:latest\",\n",
|
||||
" command=[\"python3\", \"-m\", \"src.training.task\"] + args, # Pass in training arguments, including hyperparameters.\n",
|
||||
" command=[\"python3\", \"-m\", \"src.training.task\"]\n",
|
||||
" + args, # Pass in training arguments, including hyperparameters.\n",
|
||||
" model_serving_container_image_uri=f\"gcr.io/{PROJECT_ID}/{PREDICTION_CONTAINER}:latest\",\n",
|
||||
" model_serving_container_predict_route=\"/predict\",\n",
|
||||
" model_serving_container_health_route=\"/health\")\n",
|
||||
" model_serving_container_health_route=\"/health\",\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"print(\"Training Spec:\", job._managed_model)\n",
|
||||
"\n",
|
||||
@@ -1645,7 +1675,8 @@
|
||||
" replica_count=1,\n",
|
||||
" machine_type=\"n1-standard-4\",\n",
|
||||
" accelerator_type=\"ACCELERATOR_TYPE_UNSPECIFIED\",\n",
|
||||
" accelerator_count=0)"
|
||||
" accelerator_count=0,\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -1784,7 +1815,7 @@
|
||||
"! gcloud ai models delete $model.name --quiet\n",
|
||||
"\n",
|
||||
"# Delete Cloud Storage objects that were created\n",
|
||||
"! gsutil -m rm -r $ARTIFACTS_DIR"
|
||||
"! gcloud storage rm --recursive $ARTIFACTS_DIR"
|
||||
]
|
||||
}
|
||||
],
|
||||
|
||||
@@ -324,7 +324,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil ls $gcs_output_uri_prefix"
|
||||
"! gcloud storage ls $gcs_output_uri_prefix"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -344,7 +344,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil rm -rf $gcs_output_uri_prefix"
|
||||
"! gcloud storage rm --recursive --continue-on-error $gcs_output_uri_prefix"
|
||||
]
|
||||
}
|
||||
],
|
||||
|
||||
@@ -328,7 +328,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil ls $gcs_output_uri_prefix"
|
||||
"! gcloud storage ls $gcs_output_uri_prefix"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -348,7 +348,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil rm -rf $gcs_output_uri_prefix"
|
||||
"! gcloud storage rm --recursive --continue-on-error $gcs_output_uri_prefix"
|
||||
]
|
||||
}
|
||||
],
|
||||
|
||||
@@ -341,7 +341,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil ls $gcs_output_uri_prefix"
|
||||
"! gcloud storage ls $gcs_output_uri_prefix"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -361,7 +361,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"! gsutil rm -rf $gcs_output_uri_prefix"
|
||||
"! gcloud storage rm --recursive --continue-on-error $gcs_output_uri_prefix"
|
||||
]
|
||||
}
|
||||
],
|
||||
|
||||
@@ -2,9 +2,9 @@ dllogger@git+https://github.com/NVIDIA/dllogger@v1.0.0
|
||||
|
||||
# Fixing these libraries versions to avoid conflicting or broken packages.
|
||||
immutabledict==4.2.1
|
||||
protobuf==4.25.8
|
||||
protobuf==5.29.6
|
||||
opencv-python-headless==4.11.0.86
|
||||
docutils==0.16
|
||||
urllib3==2.5.0
|
||||
urllib3==2.7.0
|
||||
google-cloud-storage==3.0.0
|
||||
retrying
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
absl-py==2.2.2
|
||||
annotated-types==0.7.0
|
||||
anyio==4.9.0
|
||||
black==25.1.0
|
||||
black==26.3.1
|
||||
cachetools==5.5.2
|
||||
certifi==2025.4.26
|
||||
charset-normalizer==3.4.2
|
||||
@@ -9,7 +9,7 @@ click==8.1.8
|
||||
docstring_parser==0.16
|
||||
google-api-core==2.24.2
|
||||
google-auth==2.40.1
|
||||
google-cloud-aiplatform==1.92.0
|
||||
google-cloud-aiplatform==1.133.0
|
||||
google-cloud-bigquery==3.31.0
|
||||
google-cloud-core==2.4.3
|
||||
google-cloud-resource-manager==1.14.2
|
||||
@@ -24,26 +24,26 @@ grpcio-status==1.71.0
|
||||
h11==0.16.0
|
||||
httpcore==1.0.9
|
||||
httpx==0.28.1
|
||||
idna==3.10
|
||||
idna==3.15
|
||||
mypy_extensions==1.1.0
|
||||
numpy==2.2.5
|
||||
packaging==25.0
|
||||
pathspec==0.12.1
|
||||
platformdirs==4.3.8
|
||||
proto-plus==1.26.1
|
||||
protobuf==5.29.4
|
||||
pyasn1==0.6.1
|
||||
protobuf==5.29.6
|
||||
pyasn1==0.6.4
|
||||
pyasn1_modules==0.4.2
|
||||
pydantic==2.11.4
|
||||
pydantic_core==2.33.2
|
||||
python-dateutil==2.9.0.post0
|
||||
pytz==2025.2
|
||||
requests==2.32.4
|
||||
requests==2.33.0
|
||||
rsa==4.9.1
|
||||
shapely==2.1.0
|
||||
six==1.17.0
|
||||
sniffio==1.3.1
|
||||
typing-inspection==0.4.0
|
||||
typing_extensions==4.13.2
|
||||
urllib3==2.4.0
|
||||
urllib3==2.7.0
|
||||
websockets==15.0.1
|
||||
@@ -66,7 +66,7 @@ mkdir -p "$local_folder"
|
||||
mkdir -p "$output_folder"
|
||||
|
||||
# Download the content from the GCS URI
|
||||
gsutil -m cp -r "$gcs_dataset_path"/* "$local_folder/"
|
||||
gcloud storage cp --recursive "$gcs_dataset_path"/* "$local_folder/"
|
||||
|
||||
# Process files in the local folder
|
||||
for file in "$local_folder"/*; do
|
||||
@@ -122,23 +122,23 @@ cp -r "$output_folder" "$images_folder"/images_2
|
||||
pushd "$images_folder"/images_2
|
||||
ls | xargs -P 8 -I {} mogrify -resize 50% {}
|
||||
popd
|
||||
gsutil -m cp -r "$images_folder"/images_2/* "$gcs_experiment_path"/data/images_2
|
||||
gcloud storage cp --recursive "$images_folder"/images_2/* "$gcs_experiment_path"/data/images_2
|
||||
|
||||
cp -r "$output_folder" "$images_folder"/images_4
|
||||
pushd "$images_folder"/images_4
|
||||
ls | xargs -P 8 -I {} mogrify -resize 25% {}
|
||||
popd
|
||||
gsutil -m cp -r "$images_folder"/images_4/* "$gcs_experiment_path"/data/images_4
|
||||
gcloud storage cp --recursive "$images_folder"/images_4/* "$gcs_experiment_path"/data/images_4
|
||||
|
||||
cp -r "$output_folder" "$images_folder"/images_8
|
||||
pushd "$images_folder"/images_8
|
||||
ls | xargs -P 8 -I {} mogrify -resize 12.5% {}
|
||||
popd
|
||||
gsutil -m cp "$images_folder"/images_8/* "$gcs_experiment_path"/data/images_8
|
||||
gcloud storage cp "$images_folder"/images_8/* "$gcs_experiment_path"/data/images_8
|
||||
|
||||
# Copy images and sparse reconstruction files to gcs experiment folder.
|
||||
gsutil -m cp "$images_folder"/images/* "$gcs_experiment_path"/data/images
|
||||
gsutil -m cp -r "$local_folder"/sparse "$gcs_experiment_path"/data
|
||||
gsutil -m cp "$local_folder"/database.db "$gcs_experiment_path"/data
|
||||
gcloud storage cp "$images_folder"/images/* "$gcs_experiment_path"/data/images
|
||||
gcloud storage cp --recursive "$local_folder"/sparse "$gcs_experiment_path"/data
|
||||
gcloud storage cp "$local_folder"/database.db "$gcs_experiment_path"/data
|
||||
|
||||
echo "Processing complete."
|
||||
@@ -99,14 +99,14 @@ create_dir_if_not_exists "$CHECKPOINTS_PATH"
|
||||
touch "$local_experiment_path/$exp_folder_name/log_render.txt"
|
||||
|
||||
# Copy experiment from GCS bucket to local
|
||||
gsutil -m cp -r "${args[-gcs_experiment_path]}/data" "$local_experiment_path/$exp_folder_name" || exit 1
|
||||
gsutil -m cp -r "${args[-gcs_experiment_path]}/checkpoints/${training_job_name}/*" "$CHECKPOINTS_PATH" || exit 1
|
||||
gcloud storage cp --recursive "${args[-gcs_experiment_path]}/data" "$local_experiment_path/$exp_folder_name" || exit 1
|
||||
gcloud storage cp --recursive "${args[-gcs_experiment_path]}/checkpoints/${training_job_name}/*" "$CHECKPOINTS_PATH" || exit 1
|
||||
|
||||
# Check and copy keyframes file.
|
||||
if [[ -n ${args[-gcs_keyframes_file]} ]]; then
|
||||
keyframes_file_basename=$(basename "${args[-gcs_keyframes_file]}")
|
||||
local_keyframes_file="$local_dataset_path/$keyframes_file_basename"
|
||||
gsutil cp "${args[-gcs_keyframes_file]}" "$local_keyframes_file" || exit 1
|
||||
gcloud storage cp "${args[-gcs_keyframes_file]}" "$local_keyframes_file" || exit 1
|
||||
echo "Local keyframe file: $local_keyframes_file"
|
||||
launch_rendering "$local_keyframes_file"
|
||||
else
|
||||
@@ -114,4 +114,4 @@ else
|
||||
fi
|
||||
|
||||
# Copy rendered data back to GCS.
|
||||
gsutil -m cp -r "$OUTPUT_RENDER_PATH" "${args[-gcs_experiment_path]}/render/${rendering_job_name}"
|
||||
gcloud storage cp --recursive "$OUTPUT_RENDER_PATH" "${args[-gcs_experiment_path]}/render/${rendering_job_name}"
|
||||
@@ -74,7 +74,7 @@ create_dir_if_not_exists "$local_experiment_path"
|
||||
create_dir_if_not_exists "$local_experiment_path/$scene_folder_name"
|
||||
|
||||
# Copy experiment from GCS bucket to local.
|
||||
gsutil -m cp -r "${gcs_experiment_path}/data" "$local_experiment_path/$scene_folder_name" || exit 1
|
||||
gcloud storage cp --recursive "${gcs_experiment_path}/data" "$local_experiment_path/$scene_folder_name" || exit 1
|
||||
|
||||
echo "GCS Experiment: $gcs_experiment_path"
|
||||
echo "Gin Config File: $gin_config_file"
|
||||
@@ -89,6 +89,6 @@ accelerate launch train.py --gin_configs="$gin_config_file" \
|
||||
--gin_bindings="Config.factor = ${factor}" \
|
||||
--gin_bindings="Config.max_steps = ${max_training_steps}"
|
||||
|
||||
gsutil -m rm -r "${gcs_experiment_path}/checkpoints/${training_job_name}"
|
||||
gsutil -m cp -r "$local_experiment_path/$scene_folder_name/config.gin" "${gcs_experiment_path}/${training_job_name}_config.gin"
|
||||
gsutil -m cp -r "$local_experiment_path/$scene_folder_name/checkpoints/*/*" "${gcs_experiment_path}/checkpoints/${training_job_name}"
|
||||
gcloud storage rm --recursive "${gcs_experiment_path}/checkpoints/${training_job_name}"
|
||||
gcloud storage cp --recursive "$local_experiment_path/$scene_folder_name/config.gin" "${gcs_experiment_path}/${training_job_name}_config.gin"
|
||||
gcloud storage cp --recursive "$local_experiment_path/$scene_folder_name/checkpoints/*/*" "${gcs_experiment_path}/checkpoints/${training_job_name}"
|
||||
@@ -102,10 +102,10 @@ def download_gcs_uri_to_local(
|
||||
if not os.path.exists(destination_dir):
|
||||
os.mkdir(destination_dir)
|
||||
subprocess.check_output([
|
||||
"gsutil",
|
||||
"-m",
|
||||
"gcloud",
|
||||
"storage",
|
||||
"cp",
|
||||
"-r",
|
||||
"--recursive",
|
||||
gcs_uri,
|
||||
destination_dir,
|
||||
])
|
||||
|
||||
@@ -12,7 +12,7 @@ bitsandbytes==0.43.2
|
||||
cloudml-hypertune==0.1.0.dev6
|
||||
datasets==2.20.0
|
||||
deepspeed==0.15.2
|
||||
diffusers==0.25.1
|
||||
diffusers==0.38.0
|
||||
evaluate==0.4.3
|
||||
fsspec==2024.3.1
|
||||
gcsfs==2024.3.1
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
# Agent Platform Training Clusters Blog Series
|
||||
|
||||
This directory contains deep-dive documentation, extended guides, and architectural references for Google Cloud Agent Platform Training Clusters.
|
||||
|
||||
## Contents
|
||||
- **`vertex-training-cluster/`**: Documentation and setup guides for configuring and managing Agent Platform Training Clusters.
|
||||
|
||||
## Blog Posts
|
||||
- [Model Distillation Best Practices](https://googlecloudplatform.github.io/vertex-ai-samples/vertex-training-cluster/model_distillation_best_practices): Explores off-policy model distillation, dataset curation, and hyperparameter scaling laws for training student models on Vertex AI.
|
||||
- [Forgetting Mitigation via Data Mixing](https://googlecloudplatform.github.io/vertex-ai-samples/vertex-training-cluster/forgetting_mitigation_data_mixing): Discusses catastrophic forgetting in model fine-tuning and how to mitigate it using multi-domain data mixing on Vertex AI.
|
||||
- [Multi-Turn Reinforcement Learning for τ²-bench](https://googlecloudplatform.github.io/vertex-ai-samples/vertex-training-cluster/multi_turn_reinforcement_learning_for_tau2_bench): Explores multi-turn RL training for tool-calling agents using GRPO on the τ²-bench customer service benchmark with NeMo RL.
|
||||
@@ -0,0 +1,448 @@
|
||||
<script type="text/javascript" async
|
||||
src="https://cdn.mathjax.org/mathjax/latest/MathJax.js?config=TeX-MML-AM_CHTML">
|
||||
</script><br><br>
|
||||
|
||||
# VTC Multi-Domain Dataset: Mitigating Catastrophic Forgetting with Data Mixing
|
||||
|
||||
**Author:** [Mayank Sharan](mailto:mayanksharan@google.com)
|
||||
|
||||
## Table of Contents
|
||||
|
||||
* [Intro](#intro)
|
||||
* [Background](#background)
|
||||
* [Dataset Curation](#dataset-selection)
|
||||
* [Forgetting Mitigation Best Practices](#forgetting-mitigation-best-practices)
|
||||
* [Experimental Setup](#experimental-setup)
|
||||
* [Mitigating Forgetting](#mitigating-forgetting)
|
||||
* [Mixing Ratios](#mixing-ratios)
|
||||
* [Different Starting Models](#different-starting-models)
|
||||
* [Acknowledgements](#acknowledgements)
|
||||
* [References](#references)
|
||||
|
||||
## Intro
|
||||
|
||||
In this entry of our blog series on model training best practices for Vertex AI Training Cluster (VTC) customers, we talk about catastrophic forgetting and how to mitigate it. We focus on tuning public models using supervised fine tuning (SFT) with a specialized domain dataset. With both open and closed source models performing well on general tasks the primary goal of training one's own models is to improve the performance on specialized tasks. This typically comes at the cost of the model forgetting general capabilities which can severely limit the utility of the trained model.
|
||||
|
||||
There are many possible interventions to limit forgetting, the most effective is mixing the target dataset with the actual dataset used in the model’s training. Since this is not available even for the most open source models, we have curated a multi-domain dataset that delivers the same benefits. This allows Vertex AI Training Cluster (VTC) customers to maintain and surpass frontier level model capabilities while training to further performance on specialized tasks.
|
||||
|
||||
<figure align="center" id="fig-teaser">
|
||||
<table align="center" width="80%">
|
||||
<tr>
|
||||
<td align="center" width="100%">
|
||||
<img src="images_data_mixing/teaser_forgetting.png" width="100%"><br>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
<figcaption align="left">
|
||||
<sub><b>Figure 1: Impact of mixing VTC Post Training dataset on Forgetting (8B model). </b> <i>Comparing SFT runs using only a specialized target dataset (MedMCQA) vs a mix of the target dataset and the VTC Post training dataset. Forgetting across all non-target domains is significantly mitigated with no performance loss on the target metric. Qwen3 Public here is the instruction tuned public Qwen3 8B model and the other two models are trained starting from the base Qwen3 8B model using only the target dataset and a mix of target dataset with the VTC dataset.</i></sub>
|
||||
</figcaption>
|
||||
</figure>
|
||||
|
||||
We provide a thorough set of experiments to serve as a guide for reducing forgetting while post training the Qwen3 open-weight thinking model family, beginning from their base pre-trained checkpoints. Furthermore, we demonstrate the value our datasets provide across model sizes often surpassing the performance of the official Qwen3 models while preserving performance on the specialized task (See [Figure 1](#fig-teaser)). The Qwen3 family was specifically chosen for this study because its diverse range of parameter counts and the availability of both pre-trained and post-trained checkpoints provide an ideal environment for high-fidelity scaling analysis.
|
||||
|
||||
To ensure our findings can be applied to a broad set of applications we validate our findings across five model sizes: 0.6B, 1.7B, 4B, 8B and 14B parameters. To support our VTC community in accelerating their own development, all code, datasets, and experiment configurations used in this blog are being made available for use in your training workloads.
|
||||
|
||||
## Background
|
||||
|
||||
Loss landscapes for neural networks have always been a complex multidimensional manifold rather than the simple convex ones that gradient descent is built for. Forgetting is a well known phenomenon in model customization, the first academically recorded instance being (McCloskey and Cohen, 1989) [<a href="#ref1">1</a>]. These manifolds have become even more complex with the introduction of Large Language Models where the number of parameters being optimized are typically in the billions. This makes it hard to mathematically grasp issues like forgetting. [Figure 2](#fig-loss-landscape) demonstrates a geometric understanding of why forgetting happens and how data mixing can mitigate it.
|
||||
|
||||
<figure align="center" id="fig-loss-landscape">
|
||||
<table align="center" width="80%">
|
||||
<tr>
|
||||
<td align="center" width="100%">
|
||||
<img src="images_data_mixing/background_loss_landscape.png" width="100%"><br>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
<figcaption align="left">
|
||||
<sub><b>Figure 2: Geometric Interpretation of Data Mixing to Mitigate Forgetting. </b> <i>Fine tuning objectives being meaningfully out of distribution from the pre-trained model often drives forgetting. Mixing in a dataset similar to the model distribution adjusts the objective enough to learn the new task without as much forgetting.</i></sub>
|
||||
</figcaption>
|
||||
</figure>
|
||||
|
||||
## Dataset Selection
|
||||
|
||||
Our primary requirements for a target dataset to run experiments to validate this were:
|
||||
|
||||
1. It should be out of distribution to cause forgetting
|
||||
2. It should have an evaluation metric that it directly improves
|
||||
3. It should be able to train the model to perform better than the counterpart generalist model
|
||||
|
||||
A good heuristic to determine where the data lies with respect to the model distribution is by calculating perplexity on samples from the dataset. Assuming
|
||||
- <span>$$X={x_1, x_2, \dots, x_N}$$</span> is a dataset sample represented as sequence of tokens
|
||||
- <span>$$P(x_i \mid x_{<i})$$</span> is the model likelihood of the i-th token given the sample till that token
|
||||
|
||||
Then the perplexity for this sample can be calculated as follows:
|
||||
|
||||
$$\begin{align*}
|
||||
& ppl(X) = \exp \left( -\frac{1}{N} \sum_{i=1}^{N} \log P(x_i \mid x_{<i}) \right) \\
|
||||
& = \exp \left( -\frac{1}{N} \log (\prod_{i=1}^{N} P(x_i \mid x_{<i})) \right)
|
||||
\end{align*} $$
|
||||
|
||||
The product form of the equation shows that this is a direct measure of the joint probability of this sequence of tokens according to the model. Since this computation has a balancing negative sign to account for the negative log value a lower joint probability results in a higher perplexity value and vice versa. We evaluated the following datasets as out-of-distribution candidates:
|
||||
|
||||
- [MedMCQA](https://huggingface.co/datasets/syz-ml2025/medmcqa) : Multiple Choice Questions (MCQ) dataset focusing on the medical domain
|
||||
- [BirdSQL](https://huggingface.co/datasets/birdsql/bird23-train-filtered) : Text to SQL generation dataset
|
||||
- [HardGen](https://huggingface.co/datasets/Bingguang/HardGen) : Function calling dataset
|
||||
|
||||
We also calculate perplexity on [OpenR1-Math-220k](https://huggingface.co/datasets/open-r1/OpenR1-Math-220k) to provide a reference as we expect this to be in distribution for the model given the Qwen3 models are particularly strong in the math domain.
|
||||
|
||||
<table id="tab-perplexity" style="margin-left:auto; margin-right:auto;">
|
||||
<thead>
|
||||
<tr>
|
||||
<th>Dataset \ Model</th>
|
||||
<th>Qwen3-0.6B</th>
|
||||
<th>Qwen3-8B</th>
|
||||
<th>Ours-0.6B</th>
|
||||
<th>Ours-8B</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr>
|
||||
<td>MedMCQA</td>
|
||||
<td>63.00</td>
|
||||
<td>66.00</td>
|
||||
<td>42.00</td>
|
||||
<td>20.75</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>BirdSQL</td>
|
||||
<td>38.50</td>
|
||||
<td>55.50</td>
|
||||
<td>45.50</td>
|
||||
<td>17.00</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>HardGen</td>
|
||||
<td>2.23</td>
|
||||
<td>2.03</td>
|
||||
<td>2.28</td>
|
||||
<td>1.79</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>OpenR1-Math</td>
|
||||
<td>8.63</td>
|
||||
<td>9.75</td>
|
||||
<td>6.44</td>
|
||||
<td>5.34</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
<caption style="text-align: left;"><b>Table 1:</b> Perplexity score analysis with the public instruction tuned Qwen3 models and Qwen3 base models trained using the VTC dataset (Ours) to identify a suitable target dataset.</caption>
|
||||
</table>
|
||||
|
||||
We see from [Table 1](#tab-perplexity) that OpenR1-Math-220k as we expected has low perplexity scores and HardGen shows an even lower perplexity score eliminating it from consideration. MedMCQA samples have high perplexity scores across all considered models. This dataset also has the advantage of a straightforward evaluation metric as we can use the validation split in the form of an MCQ verified evaluation.
|
||||
|
||||
Based on this analysis we choose MedMCQA as our target dataset for these experiments. Additionally, since we are training a thinking model and the dataset does not have thinking traces we use the Qwen3-235B model to inject thinking traces into the training samples.
|
||||
|
||||
## Forgetting Mitigation Best Practices
|
||||
|
||||
### Experimental Setup
|
||||
|
||||
#### Dataset Mixing
|
||||
|
||||
We tested the impact of how forgetting responds to mixing the base dataset in different ratios with the target dataset. The base dataset here refers to the multi domain SFT dataset we have developed (see our [distillation blog post](https://googlecloudplatform.github.io/vertex-ai-samples/vertex-training-cluster/model_distillation_best_practices) [<a href="#ref2">2</a>] for details of the generation process) that can replicate and on certain metrics beat the public Qwen3 models. The target dataset here refers to the MedMCQA dataset. It is important to understand that in all mixing scenarios where the target dataset is present we will use the complete target dataset as that is the reasonable course of action we expect any customer to take. This leads to the total number of samples used in training varying based on the mixing ratio.
|
||||
|
||||
We run 2 baseline experiments for each model size: using only the base dataset and only the target dataset. The mixing experiments are the base dataset being mixed in ratios of 0.9:0.1, 0.75:0.25 and 0.5:0.5. (0.9:0.1 means 90% of samples are from the base in-distribution dataset, while 10% are from the target out-of-distribution dataset.)
|
||||
|
||||
The base dataset is randomly subsampled for each of these experiments. For simpler reference and analysis let’s define a mixing ratio <span>$$0 \le \alpha < 1$$</span>, such that the final dataset mixture includes <span>$$ N'_{B} = \frac{\alpha}{1 - \alpha} N_T$$</span> samples from the base dataset where <span>$$N_T$$</span> is the number of samples in the target dataset. In each of these mixtures the complete target dataset is used, contributing <span>$$N_T$$</span> samples for a total training dataset size of <span>$$\frac{N_T}{1 - \alpha}$$</span>.
|
||||
|
||||
Since, our target dataset has 182,712 samples, this means that:
|
||||
|
||||
- 0.9:0.1 ratio (<span>$$\alpha = 0.9$$</span>) : Uses a total of 1,827,120 training samples
|
||||
- 0.75:0.25 ratio (<span>$$\alpha = 0.75$$</span>) : Uses a total of 730,849 training samples
|
||||
- 0.5:0.5 ratio (<span>$$\alpha = 0.5$$</span>) : Uses a total of 365,425 training samples
|
||||
|
||||
#### Evaluation
|
||||
|
||||
<table id="tab-eval-setup" style="margin-left:auto; margin-right:auto;">
|
||||
<thead>
|
||||
<tr>
|
||||
<th>Capabilities</th>
|
||||
<th>Benchmarks</th>
|
||||
<th># Test Samples</th>
|
||||
<th>Eval Metrics</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr>
|
||||
<td rowspan="7">Math</td>
|
||||
<td>AIME 24</td>
|
||||
<td>30</td>
|
||||
<td>pass@1 (average of 10)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>AIME 25</td>
|
||||
<td>30</td>
|
||||
<td>pass@1 (average of 10)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>BeyondAIME</td>
|
||||
<td>100</td>
|
||||
<td>pass@1 (average of 5)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>Math 500</td>
|
||||
<td>500</td>
|
||||
<td>pass@1</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>HMMT 25</td>
|
||||
<td>30</td>
|
||||
<td>pass@1 (average of 10)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>BRUMO 25</td>
|
||||
<td>30</td>
|
||||
<td>pass@1 (average of 10)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>CMIMC 25</td>
|
||||
<td>40</td>
|
||||
<td>pass@1 (average of 10)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td rowspan="3">Science</td>
|
||||
<td>GPQA</td>
|
||||
<td>448</td>
|
||||
<td>pass@1 (average of 5)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>MMLU</td>
|
||||
<td>14042</td>
|
||||
<td>pass@1</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>MMLU Pro</td>
|
||||
<td>12032</td>
|
||||
<td>pass@1</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td rowspan="2">Coding</td>
|
||||
<td>HumanEval</td>
|
||||
<td>164</td>
|
||||
<td>pass@1 (average of 5)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>LiveCodeBench v6</td>
|
||||
<td>175</td>
|
||||
<td>pass@1 (average of 5)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>Instruction Following</td>
|
||||
<td>IFEval</td>
|
||||
<td>541</td>
|
||||
<td>pass@1 (Strict Accuracy)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>Reasoning</td>
|
||||
<td>ARC-AGI 1</td>
|
||||
<td>400</td>
|
||||
<td>pass@1 (average of 5)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>Medical (Target domain)</td>
|
||||
<td>MedMCQA</td>
|
||||
<td>4183</td>
|
||||
<td>pass@1</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
<caption style="text-align: left;"><b>Table 2:</b> Comprehensive overview of task domains, evaluation benchmarks, and associated performance metrics.</caption>
|
||||
</table>
|
||||
|
||||
Our evaluation benchmarks and metrics are detailed in [Table 2](#tab-eval-setup). To ensure statistical reliability on smaller datasets, we report metrics averaged over multiple independent runs to mitigate variance. For each domain with multiple evaluations, we utilize the average score across the core benchmarks as our primary performance indicator. To maintain a consistent comparison, both our trained models and the official Qwen3 thinking models were evaluated using standardized sampling parameters — `Temperature=0.6`, `Top-P=0.95`, `Top-K=20` and `Max-tokens=32768` — aligning with the recommended [best practices](https://huggingface.co/Qwen/Qwen3-14B#best-practices) from the official Qwen3 model card.
|
||||
|
||||
Note that we have separated MedMCQA as a target metric instead of including it in the Science domain. This is to ensure clear outcomes from our experiments and to demonstrate impacts on model performance without any interference.
|
||||
|
||||
#### Training
|
||||
|
||||
##### Vertex AI Training Cluster
|
||||
|
||||
All experiments and results presented were orchestrated using the [Vertex AI Training Cluster (VTC)](https://docs.cloud.google.com/vertex-ai/docs/training/training-clusters/overview). VTC is a managed Google Cloud service designed to simplify and accelerate large-scale AI workloads. It provides a simple managed user experience that enables optimized GPU scheduling, automated fault tolerance, high hardware resiliency, quick start recipes and science tooling which drastically reduces the time from cluster setup to production training and speeds up experimentation.
|
||||
|
||||
##### Training Framework and Hyperparameters
|
||||
|
||||
We utilize NVIDIA [NeMo RL](https://github.com/NVIDIA-NeMo/RL), an open library from the [NVIDIA NeMo framework](https://github.com/NVIDIA-NeMo/) as the primary training library, leveraging the Megatron backend for distributed scaling. Models are initialized from a Qwen3 Base checkpoint and fine-tuned with a 32,768 context window on curated datasets. Optimization is handled via AdamW (<span>$$\beta_1=0.9$$</span>, <span>$$\beta_2=0.95$$</span>, weight decay=0.1) using a linear warmup and cosine decay schedule. All training is conducted using BF16 mixed precision. There are many model sizes and dataset mixes used in the experimentation so the maximum learning rate is guided by learning rate scaling laws (see [distillation blog post](https://googlecloudplatform.github.io/vertex-ai-samples/vertex-training-cluster/model_distillation_best_practices#hyperparameter-scaling) [<a href="#ref2">2</a>] for more) available as a part of VTC. The value is validated by testing slight adjustments from the recommended value for each dataset mixture.
|
||||
|
||||
### Mitigating Forgetting
|
||||
|
||||
All models in this experiment are trained starting from the Qwen3 base checkpoint. We explore the impact of dataset mixing by comparing the public Qwen3 instruction tuned model performance with our two baselines — model trained with only the target dataset and model trained only with the base dataset — and with a model trained using a 0.9 ratio mix.
|
||||
|
||||
<figure align="center" id="fig3_data_mixing">
|
||||
|
||||
<table align="center" width="100%">
|
||||
<tr>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig3_math.png" width="100%"><br>
|
||||
<sub><b>(a)</b> Math</sub>
|
||||
</td>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig3_science.png" width="100%"><br>
|
||||
<sub><b>(b)</b> Science</sub>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig3_coding.png" width="100%"><br>
|
||||
<sub><b>(c)</b> Coding</sub>
|
||||
</td>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig3_ifeval.png" width="100%"><br>
|
||||
<sub><b>(d)</b> IFEval</sub>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig3_arc_agi.png" width="100%"><br>
|
||||
<sub><b>(e)</b> ARC-AGI</sub>
|
||||
</td>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig3_medmcqa.png" width="100%"><br>
|
||||
<sub><b>(f)</b> MedMCQA</sub>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
<figcaption align="left">
|
||||
<sub><b>Figure 3: Performance with and without Data Mixing.</b> <i>A comparison across (a) Math, (b) Science, (c) Coding, (d) IFEval, (e) ARC-AGI and (f) MedMCQA benchmarks showing how data mixing impacts forgetting and performance on the target metric.</i></sub>
|
||||
</figcaption>
|
||||
|
||||
</figure>
|
||||
|
||||
[Figure 3](#fig3_data_mixing) shows that for all non-target metrics other than Science using just the target dataset shows significant forgetting. Math and ARC-AGI are almost completely forgotten for all model sizes up to 8B parameters. The mixed dataset recovers the performance to similar levels as the base dataset. The base dataset delivers performance comparable to the public model in all domains and significantly better on ARC-AGI.
|
||||
|
||||
The Science domain evaluations do not suffer severe forgetting likely because MedMCQA is very close to this domain. In fact, for the 8B and 14B sizes due to these transfer learning dynamics the <span>$$\alpha = 0.9$$</span> model outperforms both the public instruction-tuned and the base dataset (<span>$$\alpha = 1$$</span>) models.
|
||||
|
||||
Performance on the target metric of MedMCQA follows expected behavior with best results achieved by the model when trained only with the target dataset. It is important to note that the <span>$$\alpha = 0.9$$</span> model for all sizes is still significantly better than the public instruction-tuned and base dataset (<span>$$\alpha = 1$$</span>) model and for all sizes other than the 0.6B mostly maintains the performance gains of the target dataset (<span>$$\alpha = 0$$</span>) model.
|
||||
|
||||
#### Key Observations
|
||||
|
||||
Combining these conclusions we can see that mixing with our base dataset:
|
||||
|
||||
- Matches and outperforms the public instruction tuned model on general tasks.
|
||||
- Preserves the gains beyond the public model on target tasks.
|
||||
- Provides additional gains on tasks from a similar domain.
|
||||
|
||||
### Mixing Ratios
|
||||
|
||||
Now that we know that mixing the base dataset almost eliminates forgetting it is important to understand how performance changes for different mixing configurations. This is also important to examine as it determines training length and hence the cost. We will compare models trained only with the target dataset to models trained using dataset mixes with <span>$$\alpha = 0.5, 0.75, 0.9$$</span>. The ratio mentioned here refers to the proportion of the dataset from the base dataset.
|
||||
|
||||
<figure align="center" id="fig4_mixing_ratios">
|
||||
|
||||
<table align="center" width="100%">
|
||||
<tr>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig4_math.png" width="100%"><br>
|
||||
<sub><b>(a)</b> Math</sub>
|
||||
</td>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig4_science.png" width="100%"><br>
|
||||
<sub><b>(b)</b> Science</sub>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig4_coding.png" width="100%"><br>
|
||||
<sub><b>(c)</b> Coding</sub>
|
||||
</td>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig4_ifeval.png" width="100%"><br>
|
||||
<sub><b>(d)</b> IFEval</sub>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig4_arc_agi.png" width="100%"><br>
|
||||
<sub><b>(e)</b> ARC-AGI</sub>
|
||||
</td>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig4_medmcqa.png" width="100%"><br>
|
||||
<sub><b>(f)</b> MedMCQA</sub>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
<figcaption align="left">
|
||||
<sub><b>Figure 4: Performance across Mixing Ratios.</b> <i>A comparison across (a) Math, (b) Science, (c) Coding, (d) IFEval, (e) ARC-AGI and (f) MedMCQA benchmarks showing how dataset mixing ratios impact forgetting and performance on the target metric.</i></sub>
|
||||
</figcaption>
|
||||
|
||||
</figure>
|
||||
|
||||
[Figure 4](#fig4_mixing_ratios) shows that for all non-target metrics mixing helps achieve better performance than just using the target dataset even with a <span>$$\alpha = 0.5$$</span> mix. As expected the performance on non target metrics worsens as we lower the ratio of the base dataset. This effect is more pronounced in the smaller size models and for datasets like ARC-AGI where the mixed training provides a lot more gain. These patterns confirm that the gains on non target metrics are directly correlated to the base dataset.
|
||||
|
||||
The effect while present for Science domain metrics is much less pronounced due to the cross domain characteristics. Even with lower ratios the performance for models 4B and larger holds, confirming that our target dataset of MedMCQA here contributes to limiting forgetting for this domain.
|
||||
|
||||
The performance on the target metric, MedMCQA, stays mostly consistent with dips mostly when going from <span>$$\alpha = 0.75$$</span> mix to <span>$$\alpha = 0.5$$</span> mix. This aligns well as in all cases we are doing a complete epoch on the target dataset. The performance mostly holding at mixing ratios indicates that the tradeoff on the target metrics is relatively low even at an aggressive mixing ratio like 0.5.
|
||||
|
||||
#### Key Observations
|
||||
|
||||
The mixing ratio comparison shows us that:
|
||||
|
||||
- A mixing ratio of 0.9 is the best for achieving gains on target tasks and limiting forgetting.
|
||||
- A mixing ratio of even 0.5 limits forgetting well while only doubling the token budget compared to training without any mixing.
|
||||
|
||||
### Different Starting Models
|
||||
|
||||
We have trained all our models starting from Qwen3 base checkpoints. A natural question here might be: What happens if we train starting from the instruction tuned public Qwen3 checkpoints for our target task? In this section we examine this question and compare the instruction-tuned model tuned with the target dataset and an <span>$$\alpha = 0.9$$</span> mix to the instruction-tuned model itself and the base model tuned with an <span>$$\alpha = 0.9$$</span> mix.
|
||||
|
||||
<figure align="center" id="fig5_starting_models">
|
||||
|
||||
<table align="center" width="100%">
|
||||
<tr>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig5_math.png" width="100%"><br>
|
||||
<sub><b>(a)</b> Math</sub>
|
||||
</td>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig5_science.png" width="100%"><br>
|
||||
<sub><b>(b)</b> Science</sub>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig5_coding.png" width="100%"><br>
|
||||
<sub><b>(c)</b> Coding</sub>
|
||||
</td>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig5_ifeval.png" width="100%"><br>
|
||||
<sub><b>(d)</b> IFEval</sub>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig5_arc_agi.png" width="100%"><br>
|
||||
<sub><b>(e)</b> ARC-AGI</sub>
|
||||
</td>
|
||||
<td align="center" width="50%">
|
||||
<img src="images_data_mixing/fig5_medmcqa.png" width="100%"><br>
|
||||
<sub><b>(f)</b> MedMCQA</sub>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
<figcaption align="left">
|
||||
<sub><b>Figure 5: Performance across Starting Models.</b> <i>A comparison across (a) Math, (b) Science, (c) Coding, (d) IFEval, (e) ARC-AGI and (f) MedMCQA benchmarks showing how different starting models impact forgetting and performance on the target metric. Qwen 3 Public is the public instruction tuned Qwen3 model, α=0 (IT) and α=0.9 (IT) are the public instruction-tuned Qwen3 model trained only with the target dataset and the α=0.9 mixed dataset. α=0.9 (Base) is the base Qwen3 model trained on a 90% VTC dataset and 10% target dataset mix.</i></sub>
|
||||
</figcaption>
|
||||
|
||||
</figure>
|
||||
|
||||
In [Figure 5](#fig5_starting_models), among the non-target metrics other than science we see a common trend that starting with the IT model and using only the target dataset (<span>$$\alpha = 0$$</span>) shows severe forgetting. The base model and the instruction-tuned model trained using the <span>$$\alpha = 0.9$$</span> mix match or surpass the performance of the public model. This shows that starting with an instruction-tuned model while better than starting with the base model is still not a solution to forgetting. This also shows the high quality of our dataset that it can provide further gains on the public instruction-tuned model.
|
||||
|
||||
Science domain metrics show different trends based on the model size. The advantage of data mixing is much more apparent in 0.6B and 1.7B models. Overall though there are no disadvantages to mixing across all model sizes. The IT model demonstrating significant forgetting is a clear indication that cross domain characteristics of our target dataset are not enough to mitigate forgetting on its own.
|
||||
|
||||
The performance of the target metric, MedMCQA, shows no additional gain when we train using only the target dataset except for the 0.6B model, whether the starting model is a base model or the IT model. For all model sizes other than the 0.6B model we also see that the <span>$$\alpha = 0.9$$</span> mix trained model does not lose any meaningful performance compared to the target dataset only trained models. All the models trained using the target dataset clearly improve on the public model.
|
||||
|
||||
#### Key Observations
|
||||
|
||||
The comparison of different starting models shows us:
|
||||
|
||||
- Using the instruction-tuned model as the starting model is better than the Base model.
|
||||
- The IT model also shows catastrophic forgetting and loses performance on non target metrics.
|
||||
- The <span>$$\alpha = 0.9$$</span> mix avoids forgetting even with the instruction-tuned starting model showing its robustness.
|
||||
|
||||
## Acknowledgements
|
||||
|
||||
We would like to express our sincere gratitude to the NVIDIA NeMo RL team–specifically Terry Kong– for their invaluable support throughout this project.
|
||||
|
||||
We would also like to express our gratitude to our VTC teammates: Mohammadreza Mohseni, Weiran Zhao, Fei Xia, Youbao Tang, Xuehan Xiong, Joseph Pagadora, Jiuqiang Tang, Bo Wu, Lav Rai, and Minwoo Park for developing the underlying datasets, providing infrastructure support, feedback, and insightful discussions throughout the project. We also thank Ting Yu, Shengyang Dai, Peng Xu, and Saurabh Tiwary for their leadership and support.
|
||||
|
||||
## References
|
||||
|
||||
<a id="ref1"></a>[1] McCloskey, Michael, and Neal J. Cohen. "Catastrophic interference in connectionist networks: The sequential learning problem." Psychology of learning and motivation. Vol. 24. Academic Press, 1989. 109-165.
|
||||
|
||||
<a id="ref2"></a>[2] Google Cloud. "Model Distillation Best Practices." Vertex AI Training Cluster Samples. Google, 2026. https://googlecloudplatform.github.io/vertex-ai-samples/vertex-training-cluster/model_distillation_best_practices.
|
||||
|
After Width: | Height: | Size: 9.0 KiB |
|
After Width: | Height: | Size: 5.6 KiB |
|
After Width: | Height: | Size: 48 KiB |
|
After Width: | Height: | Size: 259 KiB |
|
After Width: | Height: | Size: 295 KiB |
|
After Width: | Height: | Size: 40 KiB |
|
After Width: | Height: | Size: 8.1 KiB |
|
After Width: | Height: | Size: 190 KiB |
|
After Width: | Height: | Size: 189 KiB |
|
After Width: | Height: | Size: 154 KiB |
|
After Width: | Height: | Size: 195 KiB |
|
After Width: | Height: | Size: 195 KiB |
|
After Width: | Height: | Size: 191 KiB |
|
After Width: | Height: | Size: 184 KiB |
|
After Width: | Height: | Size: 152 KiB |
|
After Width: | Height: | Size: 167 KiB |
|
After Width: | Height: | Size: 158 KiB |
|
After Width: | Height: | Size: 177 KiB |
|
After Width: | Height: | Size: 180 KiB |
|
After Width: | Height: | Size: 193 KiB |
|
After Width: | Height: | Size: 229 KiB |
|
After Width: | Height: | Size: 230 KiB |
|
After Width: | Height: | Size: 232 KiB |
|
After Width: | Height: | Size: 232 KiB |
|
After Width: | Height: | Size: 234 KiB |
|
After Width: | Height: | Size: 241 KiB |
|
After Width: | Height: | Size: 231 KiB |
|
After Width: | Height: | Size: 219 KiB |
|
After Width: | Height: | Size: 202 KiB |
|
After Width: | Height: | Size: 234 KiB |
|
After Width: | Height: | Size: 270 KiB |
|
After Width: | Height: | Size: 230 KiB |
|
After Width: | Height: | Size: 262 KiB |
|
After Width: | Height: | Size: 211 KiB |
|
After Width: | Height: | Size: 209 KiB |
|
After Width: | Height: | Size: 189 KiB |
|
After Width: | Height: | Size: 215 KiB |
|
After Width: | Height: | Size: 210 KiB |
|
After Width: | Height: | Size: 212 KiB |
|
After Width: | Height: | Size: 206 KiB |
|
After Width: | Height: | Size: 190 KiB |
|
After Width: | Height: | Size: 184 KiB |
|
After Width: | Height: | Size: 194 KiB |
|
After Width: | Height: | Size: 213 KiB |
|
After Width: | Height: | Size: 214 KiB |
|
After Width: | Height: | Size: 219 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 112 KiB |
|
After Width: | Height: | Size: 184 KiB |
|
After Width: | Height: | Size: 108 KiB |
|
After Width: | Height: | Size: 108 KiB |
|
After Width: | Height: | Size: 115 KiB |
|
After Width: | Height: | Size: 111 KiB |
|
After Width: | Height: | Size: 89 KiB |
|
After Width: | Height: | Size: 92 KiB |
|
After Width: | Height: | Size: 86 KiB |
|
After Width: | Height: | Size: 93 KiB |
|
After Width: | Height: | Size: 93 KiB |
|
After Width: | Height: | Size: 94 KiB |
|
After Width: | Height: | Size: 75 KiB |
|
After Width: | Height: | Size: 76 KiB |
|
After Width: | Height: | Size: 87 KiB |
|
After Width: | Height: | Size: 91 KiB |
|
After Width: | Height: | Size: 100 KiB |