Compare commits

...
Author SHA1 Message Date
Rayan DasoriyaandCopybara-Service 9cf8ce16fa Delete deprecated LoRA fine-tuning notebooks and related tutorials.
PiperOrigin-RevId: 976392412
2026-09-04 10:44:57 -07:00
Chun-Hsiang WangandGitHub 4b983a2701 feat: Claude Fable 5.1 Launch (#4581)
* feat: Claude Fable 5.1 Launch

* refactor: replace model/region if-elif chains with a dict lookup

Addresses review feedback on both Select Claude model cells. The mapping is
unchanged for all 20 models; only the lookup mechanism differs.

* chore: apply nbfmt

Runs the repo's own tensorflow-docs nbfmt over the notebook so the
'notebook format and lint' check passes.
2026-09-01 20:45:17 -04:00
Eric DongandGitHub 3b11c876bd fix: correct a typo in error message (#4577) 2026-08-25 17:03:21 -04:00
Mend RenovateandGitHub cc0d791ef2 Update dependency black to v26.5.1 (#4517) 2026-08-19 21:44:22 +00:00
Mend RenovateandGitHub df83a345bb Update dependency isort to v8 (#4444) 2026-08-19 20:52:37 +00:00
Mend RenovateandGitHub e6ded7beaa Update dependency pandas to v3.0.5 (#4491) 2026-08-19 20:51:20 +00:00
Mend RenovateandGitHub e64a4e89d5 chore(deps): update dependency google-cloud-aiplatform to v1.165.0 (#4457) 2026-08-19 20:50:48 +00:00
Mend RenovateandGitHub 7ac54985e4 chore(deps): update dependency smart_open to v8 (#4534) 2026-08-18 22:49:25 +00:00
Mend RenovateandGitHub 6ce96a08d3 Update dependency smart_open to v7.7.1 (#4494) 2026-08-18 21:16:42 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
756711b3c9 chore(deps): bump idna (#4518)
Bumps [idna](https://github.com/kjd/idna) from 3.10 to 3.15.
- [Release notes](https://github.com/kjd/idna/releases)
- [Changelog](https://github.com/kjd/idna/blob/master/HISTORY.md)
- [Commits](https://github.com/kjd/idna/compare/v3.10...v3.15)

---
updated-dependencies:
- dependency-name: idna
  dependency-version: '3.15'
  dependency-type: direct:production
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-18 21:15:21 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
a25d209139 chore(deps): bump torch (#4545)
Bumps [torch](https://github.com/pytorch/pytorch) from 2.8.0 to 2.13.0.
- [Release notes](https://github.com/pytorch/pytorch/releases)
- [Changelog](https://github.com/pytorch/pytorch/blob/main/RELEASE.md)
- [Commits](https://github.com/pytorch/pytorch/compare/v2.8.0...v2.13.0)

---
updated-dependencies:
- dependency-name: torch
  dependency-version: 2.13.0
  dependency-type: direct:production
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-18 21:14:39 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
59da536b9a chore(deps): bump pillow (#4548)
Bumps [pillow](https://github.com/python-pillow/Pillow) from 12.2.0 to 12.3.0.
- [Release notes](https://github.com/python-pillow/Pillow/releases)
- [Changelog](https://github.com/python-pillow/Pillow/blob/main/CHANGES.rst)
- [Commits](https://github.com/python-pillow/Pillow/compare/12.2.0...12.3.0)

---
updated-dependencies:
- dependency-name: pillow
  dependency-version: 12.3.0
  dependency-type: direct:production
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-18 21:13:58 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
1bc2839a2b chore(deps): bump pillow (#4568)
Bumps [pillow](https://github.com/python-pillow/Pillow) from 12.2.0 to 12.3.0.
- [Release notes](https://github.com/python-pillow/Pillow/releases)
- [Changelog](https://github.com/python-pillow/Pillow/blob/main/CHANGES.rst)
- [Commits](https://github.com/python-pillow/Pillow/compare/12.2.0...12.3.0)

---
updated-dependencies:
- dependency-name: pillow
  dependency-version: 12.3.0
  dependency-type: direct:production
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-18 21:13:27 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
0d5e268a1f chore(deps): bump urllib3 (#4512)
Bumps [urllib3](https://github.com/urllib3/urllib3) from 2.6.3 to 2.7.0.
- [Release notes](https://github.com/urllib3/urllib3/releases)
- [Changelog](https://github.com/urllib3/urllib3/blob/main/CHANGES.rst)
- [Commits](https://github.com/urllib3/urllib3/compare/2.6.3...2.7.0)

---
updated-dependencies:
- dependency-name: urllib3
  dependency-version: 2.7.0
  dependency-type: direct:production
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-18 21:12:21 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
1b019a76e4 Bump google-cloud-aiplatform (#4446)
Bumps [google-cloud-aiplatform](https://github.com/googleapis/python-aiplatform) from 1.92.0 to 1.133.0.
- [Release notes](https://github.com/googleapis/python-aiplatform/releases)
- [Changelog](https://github.com/googleapis/python-aiplatform/blob/main/CHANGELOG.md)
- [Commits](https://github.com/googleapis/python-aiplatform/compare/v1.92.0...v1.133.0)

---
updated-dependencies:
- dependency-name: google-cloud-aiplatform
  dependency-version: 1.133.0
  dependency-type: direct:production
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-18 21:11:56 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
87c1ed686a chore(deps): bump diffusers (#4510)
Bumps [diffusers](https://github.com/huggingface/diffusers) from 0.25.1 to 0.38.0.
- [Release notes](https://github.com/huggingface/diffusers/releases)
- [Commits](https://github.com/huggingface/diffusers/compare/v0.25.1...v0.38.0)

---
updated-dependencies:
- dependency-name: diffusers
  dependency-version: 0.38.0
  dependency-type: direct:production
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-18 21:11:22 +00:00
Mend RenovateandGitHub 187fdc526c Update dependency datasets to v5 (#4521) 2026-08-18 21:10:37 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
215c8eee3e chore(deps): bump urllib3 (#4513)
Bumps [urllib3](https://github.com/urllib3/urllib3) from 2.6.3 to 2.7.0.
- [Release notes](https://github.com/urllib3/urllib3/releases)
- [Changelog](https://github.com/urllib3/urllib3/blob/main/CHANGES.rst)
- [Commits](https://github.com/urllib3/urllib3/compare/2.6.3...2.7.0)

---
updated-dependencies:
- dependency-name: urllib3
  dependency-version: 2.7.0
  dependency-type: direct:production
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-18 21:10:01 +00:00
f90cd0d6ed Add AlphaFold 3 quickstart notebook (#4572)
* Add AlphaFold 3 quickstart notebook

* Update CODEOWNERS

---------

Co-authored-by: Amit Rai <raiamit@google.com>
2026-08-17 13:54:13 -07:00
Mend RenovateandGitHub 1985f06e99 Update dependency numpy to v2.5.2 (#4516) 2026-08-14 18:43:39 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
ff428dc589 chore(deps): bump torch (#4544)
Bumps [torch](https://github.com/pytorch/pytorch) from 2.7.0 to 2.13.0.
- [Release notes](https://github.com/pytorch/pytorch/releases)
- [Changelog](https://github.com/pytorch/pytorch/blob/main/RELEASE.md)
- [Commits](https://github.com/pytorch/pytorch/compare/v2.7.0...v2.13.0)

---
updated-dependencies:
- dependency-name: torch
  dependency-version: 2.13.0
  dependency-type: direct:production
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-14 18:41:44 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
f6124370b0 chore(deps): bump pillow (#4547)
Bumps [pillow](https://github.com/python-pillow/Pillow) from 12.1.1 to 12.3.0.
- [Release notes](https://github.com/python-pillow/Pillow/releases)
- [Changelog](https://github.com/python-pillow/Pillow/blob/main/CHANGES.rst)
- [Commits](https://github.com/python-pillow/Pillow/compare/12.1.1...12.3.0)

---
updated-dependencies:
- dependency-name: pillow
  dependency-version: 12.3.0
  dependency-type: direct:production
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-14 18:40:42 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
b8822f5008 chore(deps): bump pyasn1 (#4549)
Bumps [pyasn1](https://github.com/pyasn1/pyasn1) from 0.6.3 to 0.6.4.
- [Release notes](https://github.com/pyasn1/pyasn1/releases)
- [Changelog](https://github.com/pyasn1/pyasn1/blob/main/CHANGES.rst)
- [Commits](https://github.com/pyasn1/pyasn1/compare/v0.6.3...v0.6.4)

---
updated-dependencies:
- dependency-name: pyasn1
  dependency-version: 0.6.4
  dependency-type: direct:production
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-08-14 18:39:56 +00:00
Dustin LuongandCopybara-Service 8976c57b9c Update the Kimi-K3 deployment notebook image URI.
PiperOrigin-RevId: 964692807
2026-08-14 07:39:13 -07:00
gmaninatarajanandGitHub 3985da440e fix: Updated new whl file with SDK update to add interval_variants parameter to score_ism_variants() (#4565)
* fix: Updated new whl file with SDK update to add interval_variants parameter to score_ism_variants()

* fix: updating the whl file download cell
2026-08-11 19:56:05 -04:00
Damodar PanigrahiGitHubgemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
1c9092ced3 refactor - restructure the notebook (#4564)
* refactor - restructure the notebook

* Update notebooks/community/weathernext/CUSTOM_INPUTS_GUIDE.md

Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>

* Update notebooks/community/weathernext/weathernext_2_ic_pc.ipynb

Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>

* Update notebooks/community/weathernext/weathernext_2_dws.ipynb

Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>

---------

Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-08-07 20:53:54 +00:00
Damodar PanigrahiandGitHub 77b2af09ce feat: WN2 with GPU GA (#4563) 2026-08-07 17:27:35 +00:00
genquan9andGitHub c6d33c2a0d Add tau2-bench RL blog post to docs README (#4561) 2026-08-05 23:18:11 +00:00
genquan9andGitHub 89772b320e Fix inline math rendering: use span+1798467 for GitHub Pages MathJax (#4560) 2026-08-05 17:47:23 +00:00
genquan9andGitHub a1f6d2c069 Fix LaTeX rendering for Pass Rate formula (#4559)
Replace underscores in \text{num\_pass} with spaces to avoid
LaTeX math mode errors on GitHub rendering.
2026-08-05 17:30:29 +00:00
genquan9andGitHub 936a6adf77 Add multi-turn RL for tau2-bench technical report (#4558)
* Add multi-turn RL for tau2-bench technical report

Add technical report documenting multi-turn reinforcement learning
training pipeline for tau2-bench customer service benchmark, including
GRPO training, data synthesis pipeline, and evaluation results.

* Fix deprecated MathJax CDN and broken anchor link

- Remove deprecated cdn.mathjax.org script tag (GitHub renders LaTeX natively)
- Fix broken ToC anchor from #2-bench to #tau2-bench
2026-08-05 16:25:05 +00:00
Dustin LuongandCopybara-Service 37a85d53f4 No public description
MG_DOCKER_CODES_PIPER_ORIGIN_REV_ID: 958290655
2026-08-03 04:07:45 -07:00
Dustin LuongandCopybara-Service 0b0e362ab9 Add Kimi K3 Model Garden deployment notebook
PiperOrigin-RevId: 958290655
2026-08-03 04:06:44 -07:00
Damodar PanigrahiandGitHub 98103d462f test (#4554) 2026-07-30 23:37:05 +00:00
Oleh PrypinandCopybara-Service 9ea1cf3b86 No public description
MG_DOCKER_CODES_PIPER_ORIGIN_REV_ID: 955252640
2026-07-28 07:46:15 -07:00
Sam-DecigaandGitHub 8f3e6668e1 feat: Claude Opus 5 Launch (#4552) 2026-07-26 10:04:43 -04:00
Tianzi CaiandGitHub 5d9853db5c Fix formatting in Anthropic Claude intro notebook 2026-07-22 20:59:25 -07:00
Tianzi CaiandGitHub 003fb5121b Remove unused httpx imports and related comments 2026-07-22 20:56:48 -07:00
Tianzi CaiandGitHub a62695fb38 Update image URL and request handling in notebook (#4551)
* Update image URL and request handling in notebook

* Remove Colab link markdown cell

Removed markdown cell with Colab link from the notebook.

* Remove unused import
2026-07-22 23:52:09 +00:00
Sam-DecigaandGitHub 3c630fdbb8 feat: Claude-Sonnet5-Launch (#4537) 2026-06-30 16:31:27 -04:00
Damodar PanigrahiandGitHub 6ca1d899d6 feat: wn2 doc polished (#4531) 2026-06-23 23:22:47 +00:00
Damodar PanigrahiandGitHub 1894602fff feat: wn2 notebook (#4530) 2026-06-23 21:59:17 +00:00
Damodar PanigrahiandGitHub 31a52d6e92 feat: WeatherNext IC (#4523)
* feat: WeatherNext IC

* Fix: Replace weathernext_2_ic_early_access_program.ipynb symlink with actual notebook file

* fix: Replace Vertex Jobs with Gemini Enterprise Agent Platform Jobs in WeatherNext notebook

* fix: Correct typos, broken links, and apply linter formatting
2026-06-11 19:49:13 +00:00
0f9d9734c3 feat: Claude Fable 5 Launch (#4522)
Co-authored-by: Holt Skinner <13262395+holtskinner@users.noreply.github.com>
2026-06-09 14:47:46 -04:00
Vertex MG TeamandCopybara-Service e85cf9a174 Update link to Cloud Quotas page to correct location
PiperOrigin-RevId: 926490177
2026-06-03 23:21:57 -07:00
Sam-DecigaandGitHub b4c0bbc1a0 feat: Ant-Opus4.8 Launch (#4520) 2026-05-28 14:43:47 -04:00
54 changed files with 5499 additions and 19160 deletions
+2 -2
View File
@@ -2,9 +2,9 @@ git+https://github.com/tensorflow/docs
ipython
jupyter
nbconvert
black==26.3.1
black==26.5.1
pyupgrade==3.21.2
isort==7.0.0
isort==8.0.1
flake8==7.3.0
nbqa==1.9.1
@@ -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
@@ -1,4 +1,4 @@
google-cloud-bigquery==2.20.0
tensorflow==2.12.1
pillow==12.2.0
pillow==12.3.0
tf-agents==0.8.0
@@ -1,4 +1,4 @@
google-cloud-pubsub==2.5.0
pillow==12.2.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==12.1.1
pillow==12.3.0
tf-agents==0.8.0
@@ -5,6 +5,6 @@ immutabledict==4.2.1
protobuf==5.29.6
opencv-python-headless==4.11.0.86
docutils==0.16
urllib3==2.6.3
urllib3==2.7.0
google-cloud-storage==3.0.0
retrying
@@ -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,7 +24,7 @@ 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
@@ -32,7 +32,7 @@ pathspec==0.12.1
platformdirs==4.3.8
proto-plus==1.26.1
protobuf==5.29.6
pyasn1==0.6.3
pyasn1==0.6.4
pyasn1_modules==0.4.2
pydantic==2.11.4
pydantic_core==2.33.2
@@ -45,5 +45,5 @@ six==1.17.0
sniffio==1.3.1
typing-inspection==0.4.0
typing_extensions==4.13.2
urllib3==2.6.3
urllib3==2.7.0
websockets==15.0.1
@@ -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
+1
View File
@@ -8,3 +8,4 @@ This directory contains deep-dive documentation, extended guides, and architectu
## 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.
Binary file not shown.

After

Width:  |  Height:  |  Size: 300 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 258 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 174 KiB

@@ -0,0 +1,724 @@
<script type="text/javascript" async
src="https://cdn.mathjax.org/mathjax/latest/MathJax.js?config=TeX-MML-AM_CHTML">
</script><br><br>
# Multi-Turn Reinforcement Learning for &tau;<sup>2</sup>-bench
**Authors:** [Fei Xia](mailto:feixia@google.com), [Genquan Duan](mailto:genquan@google.com), [Youbao Tang](mailto:tangyoubao@google.com), [Jingya Liu](mailto:leyajiu@google.com), [Jiuqiang Tang](mailto:jqtang@google.com), [Xuehan Xiong](mailto:xxman@google.com)
## Table of Contents
* [Intro](#intro)
* [Background](#background)
* [Multi-Turn Tool-Calling Agents](#multi-turn-tool-calling-agents)
* [GRPO](#grpo)
* [&tau;<sup>2</sup>-bench](#tau2-bench)
* [Training Pipeline](#training-pipeline)
* [Training Framework](#training-framework)
* [User Simulator](#user-simulator)
* [Training Data Synthesis](#training-data-synthesis)
* [Experiments](#experiments)
* [Setup](#setup)
* [Main Results](#main-results)
* [Training Curves](#training-curves)
* [Ablation Studies](#ablation-studies)
* [More Analysis](#more-analysis)
* [Key Takeaways](#key-takeaways)
* [Acknowledgements](#acknowledgements)
## Intro
This blog is the third installment of our blog series dedicated to model training best practices for Managed Training Cluster (MTC) customers. Building on the [off-policy distillation methodology](./model_distillation_best_practices.md) covered in the first installment, this article explores how **reinforcement learning (RL)** can further improve tool-calling agent capabilities through direct environment interaction and reward optimization.
Training tool-calling agents with RL on multi-turn tasks is heavily constrained by sparse outcome rewards and complex credit assignment across extended dialogues. In this blog, we leverage [&tau;<sup>2</sup>-bench](https://github.com/sierra-research/tau2-bench) to evaluate agent capabilities across realistic retail, airline, and telecom customer service domains. Our training architecture employs the [NeMo RL](https://github.com/NVIDIA-NeMo/RL) framework paired with the Group Relative Policy Optimization (GRPO) algorithm. In this setup, the policy model (agent) learns optimal dialogue and tool-utilization strategies by interacting with a dedicated user simulator model powered by separate LLM endpoints, while an automated verifier evaluates final task completion. To establish a strong baseline, we synthesized data using open-source models ([GLM-4.7](https://huggingface.co/zai-org/GLM-4.7-FP8)) to boost our Supervised Fine-Tuning (SFT) checkpoints from 65.5% to 70.2% on the &tau;<sup>2</sup>-bench evaluation dataset.
To support our MTC community in accelerating their own development, we release our complete synthetic datasets, codebase, and training recipes to enable reproducible RL pipelines.
## Background
### Multi-Turn Tool-Calling Agents
Multi-Turn Tool-Calling Agents are autonomous architectures that interact with external functions or APIs over extended, iterative dialogues to solve complex, multi-step tasks. Instead of generating a final answer in a single pass, these agents alternate between reasoning, executing a tool, processing the tool's output, and planning their next move over several sequential rounds. At each turn <span>$$t$$</span>, the agent maintains an internal state consisting of the initial user query <span>$$q$$</span>, the hidden text history <span>$$h_t$$</span>, and a list of all prior tool executions and results <span>$$z_0, \dots, z_{t-1}$$</span>:
$$s_t = (q, h_t, z_0, \dots, z_{t-1})$$
Using this state, the agent's policy executes a classic Observation &rarr; Planning &rarr; Action loop:
* **Planning:** The agent decides whether it has enough information to answer the user or if it needs to invoke an external tool.
* **Action (Tool Invocation):** It generates a structured API call (e.g., JSON parameters) targeting a specific tool.
* **Observation (Execution):** The environment runs the API, captures the output, and appends it back into the agent's context window as a new message turn.
* **Iterate or Terminate:** The loop repeats until the agent determines it has solved the problem and yields a final answer.
### GRPO
Group Relative Policy Optimization (GRPO) normalizes rewards within groups of <span>$$G$$</span> rollouts per prompt. The Group Relative Advantage is calculated as <span>$$A_i = \frac{R_i - \bar{R}}{\sigma_R}$$</span>, where:
* <span>$$A_i$$</span>: The relative advantage of the <span>$$i$$</span>-th output in the group.
* <span>$$R_i$$</span>: The absolute reward score given to the <span>$$i$$</span>-th output.
* <span>$$\bar{R}$$</span>: The mean reward across all outputs in the sampled group (<span>$$G$$</span>): <span>$$\bar{R} = \frac{1}{G} \sum_{j=1}^G R_j$$</span>
* <span>$$\sigma_R$$</span>: The standard deviation of the rewards within the group: <span>$$\sigma_R = \sqrt{\frac{1}{G} \sum_{j=1}^G (R_j - \bar{R})^2}$$</span>
We apply the [decoupled clipped objective](https://arxiv.org/pdf/2110.00641):
$$L^{\text{CLIP}}_{\text{decoupled}}(\theta) := \hat{\mathbb{E}}_t \left[ \frac{\pi_{\theta_{\text{prox}}}(a_t \mid s_t)}{\pi_{\theta_{\text{behav}}}(a_t \mid s_t)} \min \left( r_t(\theta)\hat{A}_t, \text{clip}\left(r_t(\theta), 1-\epsilon, 1+\epsilon\right)\hat{A}_t \right) \right]$$
where <span>$$\hat{A}_t$$</span> is an estimator of the advantage at timestep <span>$$t$$</span>, <span>$$\hat{\mathbb{E}}_t[\dots]$$</span> indicates the empirical average over a finite batch of timesteps <span>$$t$$</span>, and the probability ratio <span>$$r_t(\theta)$$</span> is defined as <span>$$r_t(\theta) := \frac{\pi_\theta(a_t \mid s_t)}{\pi_{\theta_{\text{prox}}}(a_t \mid s_t)}$$</span>.
### &tau;<sup>2</sup>-bench
[&tau;<sup>2</sup>-bench](https://github.com/sierra-research/tau2-bench), developed by Sierra Research, is an open-source evaluation framework designed to test LLM-based autonomous agents in realistic customer service environments. While the original benchmark focused on agents working entirely on their own, &tau;<sup>2</sup>-bench introduces a shared action space where the AI agent and a simulated user must collaborate to solve problems. It tests agents across complex, multi-step tasks in industries like retail, airlines, telecom, and banking knowledge.
#### Reward
For any given task scenario, the overall reward for a completed interaction sequence is binary, <span>$$R_{\text{episode}} \in \{0, 1\}$$</span>. To achieve a perfect reward of 1, the agent must simultaneously clear two distinct evaluation layers: State-Based Verification and Action-Based Verification:
$$R_{\text{episode}}=\mathbf{1}(\text{State Verified}) \times \mathbf{1}(\text{Actions Verified})$$
**State-Based Verification:** The state of the environment is represented as a database state, <span>$$S_{\text{db}}$$</span>. At the beginning of a task, the database is initialized to a specific state, <span>$$S_{\text{db}}^{\text{init}}$$</span>. The user simulator interacts with the agent to achieve an underlying goal state. At the end of the conversation, the evaluation engine extracts the final database state, <span>$$S_{\text{db}}^{\text{final}}$$</span>, and compares it against the pre-annotated ground-truth expected state, <span>$$S_{\text{db}}^{\text{target}}$$</span>.
$$\mathbf{1}(\text{State Verified}) = \begin{cases} 1 & \text{if } S_{\text{db}}^{\text{final}} = S_{\text{db}}^{\text{target}} \\ 0 & \text{otherwise} \end{cases}$$
This ensures that regardless of the exact phrasing or natural language drift during the conversation, the structural side-effects of the agent's tool executions match the exact user intent.
**Action-Based Verification:** Even if the final database matches the target state, the agent must not violate organizational logic or safety guidelines along the way. The evaluation engine validates the trajectory's sequence of actions against a set of constraints:
* **Policy Adherence:** The agent must respect conditional boundaries (e.g., checking user ID before pulling records or refusing to apply a discount if the user is ineligible).
* **Structural Correctness:** The agent cannot execute invalid combinations of tools, such as firing multiple database mutations in parallel when the system guidelines demand single, sequential turn boundaries.
$$\mathbf{1}(\text{Actions Verified}) = \begin{cases} 1 & \text{if } \forall a_t \in \tau, \mathcal{C}_{\text{policy}}(a_t) = \text{True} \\ 0 & \text{otherwise} \end{cases}$$
Where <span>$$\tau$$</span> is the trajectory history and <span>$$\mathcal{C}_{\text{policy}}$$</span> maps an action to its validity given the policy document.
#### Metric
Because LLM-based agents are inherently stochastic, evaluating a task a single time can lead to misleading variance in performance numbers. The fundamental metric reported on the benchmark leaderboards is Pass<sup>1</sup>. It represents the expected success rate across the evaluation dataset when running exactly one trial per task scenario. Given a dataset of <span>$$N$$</span> unique task descriptions, Pass<sup>1</sup> is computed as:
$$\text{Pass}^1 = \frac{1}{N} \sum_{i=1}^{N} R_{\text{episode}}^{(i)}$$
We report Pass<sup>1</sup> with 4 trials in the evaluation below.
## Training Pipeline
### Training Framework
**Training Framework and System Architecture**
We utilize NVIDIA [NeMo RL](https://github.com/NVIDIA-NeMo/RL) as the primary training framework. We implement the &tau;<sup>2</sup>-bench sandbox environment inside NVIDIA [NeMo Gym](https://github.com/NVIDIA-NeMo/Gym), which provides a unified interface for building and scaling reinforcement learning environments and is seamlessly integrated with the NeMo RL library for RL training runs.
<figure align="center" id="fig-architecture">
<table align="center" width="90%">
<tr>
<td align="center" width="100%">
<img src="images_tau2/rl_tau2_architecture.png" width="100%"><br>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 1: RL Training System Architecture.</b> <i>The system partitions workloads across three execution domains&mdash;a CPU VM, a CPU cluster for environment execution, and a GPU cluster for training/sampling&mdash;so that each scales independently and GPUs stay saturated on training and generation.</i></sub>
</figcaption>
</figure>
We train on &tau;<sup>2</sup>-bench, a customer-service simulation benchmark spanning the airline, retail, and telecom domains. Each task instantiates a tool-augmented dialogue between a policy agent (the model under training) and an LLM-driven user simulator, grounded in a domain policy document and a per-domain tool/API suite. An episode is a multi-turn loop; at each turn the agent either replies to the user in natural language or issues a tool call against the domain backend, and the environment advances the user-simulator state, returns tool results and the user's next message. Rewards are produced by &tau;<sup>2</sup>'s built-in verifier against each task's expected outcome, yielding the per-episode scalar that drives GRPO.
The system architecture deliberately partitions the workload across three execution domains&mdash;a CPU VM, a CPU cluster for environment execution, and a GPU cluster for the trainer/sampler&mdash;so that each scales independently and the GPUs stay saturated on the only work that needs them: training and generation. As shown in [Figure 1](#fig-architecture), a single Driver Program on the CPU VM owns the training loop and hosts two cooperating components.
The first is the **Training Service Client**, which talks to the MTC Training Service on the GPU cluster and provisions two modules&mdash;a policy Trainer and a rollout Sampler&mdash;colocated to share GPUs or disaggregated for async workload. The client issues train / compute_logprobs calls to the Trainer and pulls generations from the Sampler, and after each update synchronizes policy weights Trainer&rarr;Sampler over a dedicated weights group so the next round of rollouts is on-policy.
The second component is the **Rollout Proxy & Trajectory Manager**. Rather than letting environment code call the Sampler directly, all generation is funneled through an OpenAI-compatible `/chat/completions` proxy that fronts the Sampler endpoint. This buys three things at once: (i) environment code stays a stock LLM client&mdash;the Episode Worker on the CPU cluster runs an unmodified &tau;<sup>2</sup> AgentGymEnv and reaches the model through a standard LiteLLM/OpenAI client pointed at the proxy URL; and (ii) because every agent turn transits the proxy, the Trajectory Manager records token-faithful prompt/completion segments and logprobs as they are generated, so trajectories are reconstructed exactly for the GRPO update instead of being re-tokenized after the fact.
This separation is what lets the environment tier scale horizontally and independently of the GPUs. Environment execution runs as a fleet of Ray actors on the CPU cluster, fanned out by the EnvRolloutDispatcher across two pools&mdash;a train pool and an eval pool&mdash;pinned to their respective Ray workergroups with the &tau;<sup>2</sup> data corpus baked into the worker image. Each step dispatches `num_prompts × repeat_n` episodes onto the train pool, all of them generating concurrently against the shared Sampler through the rollout proxy; the driver then filters failed and length-truncated trajectories, computes leave-one-out GRPO advantages within each prompt group, applies a clipped policy-gradient update on the Trainer, and syncs weights back to the Sampler before the next step. Evaluation runs periodically on the eval pool, and best-N checkpoint retention is keyed on the eval reward. The net effect is that slow, CPU-bound, highly parallel environment simulation is kept off the GPU critical path, while the GPU cluster does nothing but generate and train.
### User Simulator
Unlike passive benchmarks where the user is merely a text prompt, &tau;<sup>2</sup>-bench introduces a dual-control architecture. The User Simulator functions as an active environment entity. To eliminate the chaotic hallucinations common in pure LLM simulations, &tau;<sup>2</sup>-bench tightly couples the user's behavior to the actual underlying state machine. The user cannot magically fix a setting or misrepresent device states; they must be accurately guided by the RL agent's communication policy, making coordination and explicit user-modeling a strict requirement for policy success. The user simulator endpoints use vLLM or SGLang with OpenAI-compatible formats.
### Training Data Synthesis
To train our RL agent within &tau;<sup>2</sup>-bench's dual-control environment, we developed an efficient data synthesis pipeline to produce high-quality training data for three customer-service domains: Telecom, Retail, and Airline. The pipeline uses an LLM to generate tasks, then iteratively refines and verifies them through multiple stages to ensure solvability and correctness, and finally converts the verified rollout results into training data.
<figure align="center" id="fig-pipeline">
<table align="center" width="90%">
<tr>
<td align="center" width="100%">
<img src="images_tau2/rl_tau2_data_pipeline.png" width="100%"><br>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 2: Training Data Synthesis Pipeline.</b> <i>The pipeline generates task bundles, refines them through crash-fixing and solvability checks, verifies across multiple rollouts, and exports categorized training data.</i></sub>
</figcaption>
</figure>
The pipeline ([Figure 2](#fig-pipeline)) comprises the following stages:
* **Task Generation:** The process begins by prompting a large language model to generate a self-contained "Task Bundle". Each bundle contains a simulated database state, a concrete user scenario, and a list of machine-verifiable evaluation criteria. To prevent the LLM from generating repetitive tasks, a unique diversity seed is constructed for each call by randomly sampling:
* *User Profiles:* Names, addresses, and contact info.
* *Difficulty Levels:* Controlling the expected length and complexity (Easy, Medium, Hard).
* *Scenarios:* Specific problems mapped from domain pools (e.g., billing disputes, cancellations, or connectivity issues).
* **Task Refinement:**
* *Rollout Refinement (Crash Fixing):* Every task runs once in a live simulator. Tasks that crash are captured, and their stack tracebacks are sent back to the LLM for automated repair up to 3 rounds.
* *Ground-Truth (GT) Refinement (Solvability):* A specialized "Golden Agent" with perfect knowledge of the correct resolution path attempts each task. If this expert agent cannot achieve a perfect reward (reward=1.0), the task's database state or evaluation criteria are fundamentally misaligned and are sent back to the LLM to be repaired. If the expert fails to solve the task after 2 rounds, then the task is marked as failed to check ground truth.
* **Task Verification:** The pipeline verifies each task across 16 independent, stochastic rollouts with standard agents. This stage calculates a statistical Pass Rate for each task to evaluate solvability: <span>$$\text{Pass Rate} = \frac{\text{num pass}}{\text{num trials}}$$</span>. If a task is unsolvable by standard agents and has a 0% pass rate, then the task is marked as failed to check ground truth.
* **Failure Refinement and Re-verify:** Rather than discarding failed tasks entirely, the pipeline takes a "fix the test, not the code" approach. The LLM reviews the best recorded trajectory and only modifies evaluation criteria to make them solvable but still meaningful. Refined tasks are verified again and merged with previously verified results.
* **Task Export:** Generated tasks are categorized into difficulty buckets based on their statistical pass rates: easy (9&ndash;12 correct rollouts), medium (5&ndash;8 correct rollouts), and hard (1&ndash;4 correct rollouts). Tasks with 13&ndash;16 correct rollouts are excluded because they are already well-solved and provide limited training signal.
We used [GLM-4.7-FP8](https://huggingface.co/zai-org/GLM-4.7-FP8) and achieved the following synthesized data distribution:
<table id="tab-synth-data" style="margin-left:auto; margin-right:auto;">
<thead>
<tr>
<th>Difficulty</th>
<th>Airline</th>
<th>Retail</th>
<th>Telecom</th>
<th>Total Tasks</th>
</tr>
</thead>
<tbody>
<tr>
<td>Easy</td>
<td>170 (36.9%)</td>
<td>255 (55.3%)</td>
<td>36 (7.8%)</td>
<td>461</td>
</tr>
<tr>
<td>Medium</td>
<td>190 (54.6%)</td>
<td>143 (41.1%)</td>
<td>15 (4.3%)</td>
<td>348</td>
</tr>
<tr>
<td>Hard</td>
<td>346 (46.1%)</td>
<td>388 (51.7%)</td>
<td>16 (2.1%)</td>
<td>750</td>
</tr>
</tbody>
<caption style="text-align: left;"><b>Table 1:</b> Distribution of synthesized training data across domains and difficulty levels.</caption>
</table>
## Experiments
### Setup
#### User Simulator
The selected user simulator model for training and evaluation is [GLM-5-FP8](https://console.cloud.google.com/vertex-ai/publishers/zai-org/model-garden/glm-5). The user simulator endpoints can be deployed locally or in Vertex AI Model Garden. For easy reproduction, we provide sample scripts to deploy GLM-5-FP8 locally in clusters as well.
While our offline task generation pipeline utilized GLM-4.7 to efficiently scale the synthesis and verification of thousands of scenarios, utilizing a more powerful model as the live user simulator is essential to mitigate negative impacts on RL training stability. Specifically, GLM-5 outperforms GLM-4.7 in this role, providing a more robust and strictly compliant simulation environment. Furthermore, this decoupling mitigates self-reinforcing biases by ensuring the policy agent does not merely overfit to the linguistic quirks of the model used to generate its training data.
#### Training Configuration
* **Checkpoint:** Our SFT checkpoints were fine-tuned from Qwen3-8B, as described in the [Model Distillation Best Practices](./model_distillation_best_practices.md) blog.
* **Training Data:** Our synthesized data described [above](#training-data-synthesis).
* **Hyperparameters:**
<table id="tab-hyperparams" style="margin-left:auto; margin-right:auto;">
<thead>
<tr>
<th>Parameter</th>
<th>Value</th>
</tr>
</thead>
<tbody>
<tr>
<td>Prompts per step</td>
<td>64</td>
</tr>
<tr>
<td>Generations per prompt</td>
<td>16</td>
</tr>
<tr>
<td>Global batch size</td>
<td>1024</td>
</tr>
<tr>
<td>Max turns</td>
<td>40</td>
</tr>
<tr>
<td>Optimizer</td>
<td>Adam</td>
</tr>
<tr>
<td>Max num steps</td>
<td>150</td>
</tr>
<tr>
<td>Temperature</td>
<td>1.0</td>
</tr>
</tbody>
<caption style="text-align: left;"><b>Table 2:</b> Training hyperparameters for RL experiments.</caption>
</table>
#### Evaluation
We use &tau;<sup>2</sup>-bench (v2) as our evaluation dataset. The &tau;<sup>2</sup>-bench community mainly reports Pass<sup>1</sup> with 4 trials and averages across three different domains. The same models may produce different results across runs&mdash;this variance is by design in &tau;<sup>2</sup>-bench. Due to limited resources, we report the mean and standard deviation for the main results from 5 runs, and only report results from one run in ablation studies. Please refer to the [Background](#metric) section for a description of the evaluation metrics, and to [the original paper](https://arxiv.org/abs/2506.07982) for more details.
### Main Results
We compare our SFT and RL models against state-of-the-art models:
<table id="tab-main-results" style="margin-left:auto; margin-right:auto;">
<thead>
<tr>
<th>Model</th>
<th>Setup</th>
<th>Stage</th>
<th>Retail</th>
<th>Airline</th>
<th>Telecom</th>
<th>Avg</th>
</tr>
</thead>
<tbody>
<tr>
<td>Qwen3-8B-Base</td>
<td>Qwen3 official pre-trained checkpoint</td>
<td>Pre-trained</td>
<td>6.1</td>
<td>39.0</td>
<td>15.4</td>
<td>20.2</td>
</tr>
<tr>
<td>Qwen3-8B</td>
<td>Qwen3 official post-trained checkpoint</td>
<td>Post-trained</td>
<td>50.7</td>
<td>30.0</td>
<td>45.8</td>
<td>42.2</td>
</tr>
<tr>
<td>Qwen3-235B-A22B-Thinking-2507</td>
<td>Qwen3 official flagship post-trained model</td>
<td>Post-trained</td>
<td>72.1</td>
<td>56.5</td>
<td>73.2</td>
<td>67.3</td>
</tr>
<tr>
<td><b>Cirrus-Agent-SFT 8B [Ours]</b></td>
<td>Cirrus-0.5 8B, SFT with tool use data and rejection sampling</td>
<td>SFT</td>
<td>67.4 &plusmn; 3.0</td>
<td>55.5 &plusmn; 3.3</td>
<td>73.5 &plusmn; 1.3</td>
<td>65.5 &plusmn; 1.5</td>
</tr>
<tr>
<td><b>Cirrus-Agent-RL 8B [Ours]</b></td>
<td>RL based on Cirrus-Agent-SFT 8B</td>
<td>RL</td>
<td>68.1 &plusmn; 0.8</td>
<td>56.8 &plusmn; 3.2</td>
<td>85.9 &plusmn; 2.2</td>
<td>70.2 &plusmn; 1.4</td>
</tr>
</tbody>
<caption style="text-align: left;"><b>Table 3:</b> Comparison of SFT and RL models against state-of-the-art models on &tau;<sup>2</sup>-bench (Pass<sup>1</sup> with 4 trials, averaged over 5 runs for our models).</caption>
</table>
Key observations from our main results:
* **RL models improve Pass<sup>1</sup> from 65.5 to 70.2 (+4.7) overall** and achieve a massive improvement on telecom tasks from 73.5 to 85.9 (+12.4), confidently demonstrating that RL helps improve model performance.
* **On retail tasks**, RL models improve Pass<sup>1</sup> slightly (+0.7, within noise), but variance collapses from &plusmn;3.0 to &plusmn;0.8 (~73% reduction). This dramatic variance reduction means that while RL did not make the model more accurate on average, it made it far more consistent and predictable.
* **On airline tasks**, both the variances of SFT and RL models are large (~3) and the improvements of RL models are minor (+1.3, within noise).
#### Evaluation Details for SFT and RL Models
For better reproduction and understanding of evaluation results, here are detailed per-run results and a suggested interpretation guide. The evaluated RL model was trained with all synthetic data.
<table id="tab-eval-details" style="margin-left:auto; margin-right:auto;">
<thead>
<tr>
<th>Model</th>
<th>#Run</th>
<th>Retail</th>
<th>Airline</th>
<th>Telecom</th>
<th>Avg</th>
</tr>
</thead>
<tbody>
<tr>
<td rowspan="7"><b>Cirrus-Agent-SFT 8B</b></td>
<td>Run 1</td>
<td>71.3</td>
<td>51.0</td>
<td>73.0</td>
<td>65.1</td>
</tr>
<tr>
<td>Run 2</td>
<td>69.7</td>
<td>59.0</td>
<td>75.0</td>
<td>67.9</td>
</tr>
<tr>
<td>Run 3</td>
<td>66.2</td>
<td>58.5</td>
<td>72.1</td>
<td>65.6</td>
</tr>
<tr>
<td>Run 4</td>
<td>63.8</td>
<td>55.0</td>
<td>72.5</td>
<td>63.8</td>
</tr>
<tr>
<td>Run 5</td>
<td>66.2</td>
<td>54.0</td>
<td>74.8</td>
<td>65.0</td>
</tr>
<tr>
<td><i>x&#772;</i></td>
<td><i>67.4</i></td>
<td><i>55.5</i></td>
<td><i>73.5</i></td>
<td><i>65.5</i></td>
</tr>
<tr>
<td><i>&sigma;<sub>SFT</sub></i></td>
<td><i>3.0</i></td>
<td><i>3.3</i></td>
<td><i>1.3</i></td>
<td><i>1.5</i></td>
</tr>
<tr>
<td rowspan="7"><b>Cirrus-Agent-RL 8B</b></td>
<td>Run 1</td>
<td>67.5</td>
<td>62.5</td>
<td>87.1</td>
<td>72.4</td>
</tr>
<tr>
<td>Run 2</td>
<td>68.6</td>
<td>55.0</td>
<td>82.7</td>
<td>68.8</td>
</tr>
<tr>
<td>Run 3</td>
<td>69.1</td>
<td>55.0</td>
<td>84.6</td>
<td>69.6</td>
</tr>
<tr>
<td>Run 4</td>
<td>67.8</td>
<td>56.0</td>
<td>88.2</td>
<td>70.7</td>
</tr>
<tr>
<td>Run 5</td>
<td>67.3</td>
<td>55.5</td>
<td>86.8</td>
<td>69.9</td>
</tr>
<tr>
<td><i>x&#772;</i></td>
<td><i>68.1</i></td>
<td><i>56.8</i></td>
<td><i>85.9</i></td>
<td><i>70.2</i></td>
</tr>
<tr>
<td><i>&sigma;<sub>RL</sub></i></td>
<td><i>0.8</i></td>
<td><i>3.2</i></td>
<td><i>2.2</i></td>
<td><i>1.4</i></td>
</tr>
<tr style="border-top: 2px solid;">
<td colspan="2"><b>&Delta;x&#772;</b></td>
<td>0.7</td>
<td>1.3</td>
<td>12.4</td>
<td>4.7</td>
</tr>
<tr>
<td colspan="2"><b>&sigma;<sub>combined</sub></b></td>
<td>3.1</td>
<td>4.6</td>
<td>2.6</td>
<td>2.0</td>
</tr>
<tr>
<td colspan="2"><b>Significance</b></td>
<td>0.2</td>
<td>0.3</td>
<td>4.8</td>
<td>2.3</td>
</tr>
</tbody>
<caption style="text-align: left;"><b>Table 4:</b> Detailed per-run evaluation results for SFT and RL models. &sigma;<sub>combined</sub> is defined as &radic;(&sigma;<sub>SFT</sub>&sup2; + &sigma;<sub>RL</sub>&sup2;). Significance is &Delta;x&#772; / &sigma;<sub>combined</sub>.</caption>
</table>
**Suggested Interpretation Guide:**
The significance of overall (2.3&times;) and telecom (4.8&times;) results confidently demonstrates that RL improves performance.
<table id="tab-significance" style="margin-left:auto; margin-right:auto;">
<thead>
<tr>
<th>Significance Level</th>
<th>Sigma</th>
<th>Interpretation</th>
</tr>
</thead>
<tbody>
<tr>
<td>Very High</td>
<td>&gt;3&sigma;</td>
<td>Definitive effect</td>
</tr>
<tr>
<td>High</td>
<td>&gt;2&sigma;</td>
<td>Statistically significant</td>
</tr>
<tr>
<td>Moderate</td>
<td>1&sigma;&ndash;2&sigma;</td>
<td>Suggestive but inconclusive</td>
</tr>
<tr>
<td>Low</td>
<td>&lt;1&sigma;</td>
<td>Within random variation</td>
</tr>
</tbody>
<caption style="text-align: left;"><b>Table 5:</b> Significance level interpretation guide.</caption>
</table>
### Training Curves
<figure align="center" id="fig-training-curve">
<table align="center" width="80%">
<tr>
<td align="center" width="100%">
<img src="images_tau2/rl_tau2_training_curve.png" width="100%"><br>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 3: RL Training Reward Curve.</b> <i>Example training reward curve showing the progression of the GRPO optimization over training steps.</i></sub>
</figcaption>
</figure>
### Ablation Studies
We performed ablation studies on different learning rates, KL penalties, and data combinations. Due to limited resources, we only report Pass<sup>1</sup> with 4 trials from a single run.
<table id="tab-ablation" style="margin-left:auto; margin-right:auto;">
<thead>
<tr>
<th>Data</th>
<th>Step</th>
<th>LR</th>
<th>KL</th>
<th>Retail</th>
<th>Airline</th>
<th>Telecom</th>
<th>Avg</th>
</tr>
</thead>
<tbody>
<tr>
<td>Easy</td>
<td>70</td>
<td>1.0E-6</td>
<td>n/a</td>
<td>70.6</td>
<td>58.5</td>
<td>84.9</td>
<td>71.3</td>
</tr>
<tr>
<td>Easy</td>
<td>70</td>
<td>5.0E-7</td>
<td>n/a</td>
<td>68.2</td>
<td>57.5</td>
<td>79.2</td>
<td>68.3</td>
</tr>
<tr>
<td>Easy</td>
<td>70</td>
<td>1.5E-6</td>
<td>n/a</td>
<td>67.3</td>
<td>58.0</td>
<td>88.2</td>
<td>71.1</td>
</tr>
<tr>
<td>Easy</td>
<td>75</td>
<td>2.0E-6</td>
<td>n/a</td>
<td>72.4</td>
<td>56.0</td>
<td>84.9</td>
<td>71.1</td>
</tr>
<tr>
<td>Easy</td>
<td>135</td>
<td>1.0E-6</td>
<td>0.01</td>
<td>69.1</td>
<td>58.0</td>
<td>85.0</td>
<td>70.7</td>
</tr>
<tr>
<td>Easy</td>
<td>140</td>
<td>1.0E-6</td>
<td>0.02</td>
<td>67.3</td>
<td>58.0</td>
<td>81.4</td>
<td>68.9</td>
</tr>
<tr>
<td>Easy</td>
<td>115</td>
<td>1.0E-6</td>
<td>0.05</td>
<td>71.7</td>
<td>58.0</td>
<td>83.3</td>
<td>71.0</td>
</tr>
<tr>
<td>Easy</td>
<td>105</td>
<td>1.0E-6</td>
<td>0.1</td>
<td>69.5</td>
<td>58.5</td>
<td>82.9</td>
<td>70.3</td>
</tr>
<tr>
<td>Easy+Medium</td>
<td>45</td>
<td>2.0E-6</td>
<td>n/a</td>
<td>69.3</td>
<td>59.0</td>
<td>84.2</td>
<td>70.9</td>
</tr>
<tr>
<td>Easy+Medium+Hard</td>
<td>50</td>
<td>2.0E-6</td>
<td>n/a</td>
<td>67.5</td>
<td>62.5</td>
<td>87.1</td>
<td>72.4</td>
</tr>
</tbody>
<caption style="text-align: left;"><b>Table 6:</b> Ablation study results across learning rates, KL penalties, and data combinations (Pass<sup>1</sup> with 4 trials, single run).</caption>
</table>
Key observations from the ablation studies:
* Using the easy data, models trained with learning rates 1.0E-6, 1.5E-6, and 2.0E-6, or KL penalty 0.05, achieved similar results and outperformed other configurations.
* Mixing easy and medium data produced similar results to using only easy data.
* **Mixing easy, medium, and hard data yielded the best results**, outperforming both easy-only and easy+medium configurations.
* The best results occurred after training 45&ndash;75 steps (approximately 2&ndash;10 epochs) for training without KL. Training may overfit to the training data when running for additional steps.
## More Analysis
**Failure Patterns.** In the evaluation dataset, there are tasks with simple tool-call sequences&mdash;simple state toggles and straightforward procedures&mdash;such as all telecom tasks and partial airline/retail tasks. Other tasks require correct multi-step tool-call chains with multi-entity reasoning and constraints, such as the majority of airline/retail tasks. SFT models generally understand what to do and maintain strong user communication, but sometimes struggle to execute the correct tool-call sequences. RL models directly optimize tool-calling behavior through reward signals, improving performance overall, but exhibit some common failure patterns:
* **Skipped tool calls:** The model converses correctly but omits necessary actions (e.g., `modify_pending_order_items`, `get_reservation_details`), resulting in the database not being updated correctly.
* **Incorrect tool parameters:** The model calls the correct tools but with wrong arguments (e.g., wrong item IDs, order IDs), leaving the database in the wrong state.
* **Over-action:** Instead of refusing disallowed operations or escalating to a human agent (`transfer_to_human_agents`), the model proceeds with actions that should be declined, becoming more "action-biased."
**Data Paradox.** Telecom has 10&times; less training data than airline and retail, but achieves significantly better performance:
<table id="tab-data-paradox" style="margin-left:auto; margin-right:auto;">
<thead>
<tr>
<th>Domain</th>
<th>% of Training Data</th>
<th>Pass<sup>1</sup></th>
</tr>
</thead>
<tbody>
<tr>
<td>Retail</td>
<td>50.4%</td>
<td>68.1</td>
</tr>
<tr>
<td>Airline</td>
<td>45.3%</td>
<td>56.8</td>
</tr>
<tr>
<td>Telecom</td>
<td>4.3%</td>
<td>85.9</td>
</tr>
</tbody>
<caption style="text-align: left;"><b>Table 7:</b> The data paradox&mdash;telecom achieves the highest performance despite having the least training data.</caption>
</table>
This telecom performance advantage is likely driven by a more deterministic tool graph, structured slot-filling parameters, and lower linguistic variance from the simulator compared to the other more open-ended domains. We analyze airline and retail failures further:
* There are many airline failures for complex tasks (3+ actions), indicating that trained models should improve their ability to chain multi-step workflows.
* The retail failures are more long-tail in nature&mdash;various small failures where trained models make occasional mistakes on many different actions.
**Known Issues for Airline and Retail Evaluations.** The community has been invaluable in identifying issues&mdash;from annotation errors to underspecified tasks&mdash;in the original airline and retail domains. 50+ tasks were fixed in [&tau;<sup>3</sup>-bench releases](https://taubench.com/blog/tau3-task-fixes.html).
**Top Directions for Addressing Remaining Error Patterns:**
1. Add action-sequence SFT pre-training before RL to learn tool-calling patterns, which may accelerate RL convergence.
2. Enable light reward shaping (e.g., 0.15 format weight) to provide learning signal on total failures instead of pure 0 reward.
3. Use &tau;<sup>3</sup>-bench as evaluations.
## Key Takeaways
Thanks for reading. We hope this RL training framework and these insights help you build better tool-calling agents on Managed Training Clusters.
* **Performance Gains from RL:** RL training increases the overall Pass<sup>1</sup> success rate from 65.5 to 70.2 (+4.7), highlighted by a massive +12.4 performance boost on telecom tasks.
* **Variance Reduction in Retail:** While average performance gains on retail tasks are minor, RL reduces variance by roughly 73% (from &plusmn;3.0 to &plusmn;0.8), ensuring much more consistent and predictable agent behavior.
* **The Data Paradox:** Despite having 10&times; less training data than other domains, telecom achieves the highest performance (85.9 Pass<sup>1</sup>), demonstrating that domain clarity and data quality are far more critical than raw quantity.
* **Actionable Future Directions:** To address complex workflow failures and long-tail action errors, future iterations should incorporate action-sequence SFT pre-training to accelerate RL convergence and implement light reward shaping to provide a stronger learning signal.
## Acknowledgements
We would like to express our sincere gratitude to the NVIDIA NeMo RL team for their invaluable support throughout this project.
We would also like to express our gratitude to our MTC teammates: Mohammadreza Mohseni, Weiran Zhao, and Bo Wu for their infrastructure support, feedback, and insightful discussions throughout the project. We also thank Ting Yu, Shengyang Dai, Peng Xu, and Aparna Ramani for their leadership and support.
+4
View File
@@ -28,7 +28,11 @@
/vertex_endpoints/optimized_tensorflow_runtime @vlasenkoalexey
/notebooks/community/alphagenome/cloudai_alphagenome_vai_quickstart.ipynb @dpanigra
/notebooks/community/alphagenome/cloudai_alphagenome_finetune.ipynb @dpanigra
/notebooks/community/alphafold3/cloudai_alphafold3_vai_quickstart.ipynb @raiamitgit
/notebooks/community/weathernext/weathernext_2_early_access_program.ipynb @dpanigra
/notebooks/community/weathernext/weathernext_2_ic_early_access_program.ipynb @dpanigra
/notebooks/community/weathernext/weathernext_2_dws.ipynb @dpanigra
/notebooks/community/weathernext/weathernext_2_ic_pc.ipynb @dpanigra
/notebooks/community/ml_ops/stage2/get_started_with_visionapi_and_automl.ipynb @mansari
/notebooks/community/neo4j/graph_paysim.ipynb @benofben @laeg
/notebooks/community/ml_ops/stage1/get_started_with_visionapi_and_vertex_datasets.ipynb @mansari
+35
View File
@@ -0,0 +1,35 @@
# AlphaFold 3
[**Overview**](#overview) | [**Use cases**](#use-cases) | [**Documentation**](#documentation) | [**Prerequisites**](#prerequisites) | [**Quick start**](#quick-start)
## Overview
AlphaFold 3 is a revolutionary model developed by Google DeepMind and Isomorphic Labs that predicts the 3D structures and interactions of proteins, DNA, RNA, ligands, and chemical modifications.
By modeling these molecules and their interactions together in a unified diffusion-based architecture, AlphaFold 3 provides a comprehensive view of cellular machinery, enabling researchers to understand biological processes at atomic resolution.
AlphaFold 3 is available for commercial use on [Gemini Enterprise Agent Platform](https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/open-models/alphafold-3).
## Use cases
* **Protein-Ligand Interaction Prediction**: Model the binding of small molecule ligands to proteins, enabling drug discovery and development.
* **Nucleic Acid Interaction Prediction**: Predict the complex structures of proteins interacting with DNA and RNA sequences.
* **Chemical Modifications**: Predict structures containing modified residues, ions, and covalent linkages.
* **Antibody-Antigen Modeling**: Map the 3D structures of antibody-antigen complexes to support therapeutic antibody design.
## Documentation
The examples provided here demonstrate how to deploy and use AlphaFold 3 on Gemini Enterprise Agent Platform.
### Links
* Read the [Nature journal paper](https://doi.org/10.1038/s41586-024-07487-w)
* Read the [Google DeepMind blog post](https://blog.google/technology/ai/google-deepmind-isomorphic-alphafold-3-ai-model/)
* Explore the [AlphaFold Server](https://alphafoldserver.com/welcome)
* View the open-source code and non-commercial weights on [GitHub](https://github.com/google-deepmind/alphafold3)
## Prerequisites
To deploy and use AlphaFold 3 on Vertex AI:
1. **Request Access**: Submit the [AlphaFold 3 Request Form](https://console.cloud.google.com/vertex-ai/publishers/google/model-garden/alphafold3-request) and work with your Google Cloud account team for commercial subscription allowlisting.
2. **Hardware Quota**: Deployments require an `a3-highgpu-1g` machine type (1x NVIDIA H100 80GB GPU) with 750 GB Local SSD provisioned for database caching.
3. **Endpoint Configuration**: Deploy the model to a Dedicated Endpoint and configure the inference timeout to 3,600 seconds.
## Quick start
| Notebook | Description | Links |
| :--- | :--- | :--- |
| [AlphaFold 3 Quickstart](cloudai_alphafold3_vai_quickstart.ipynb) | End-to-end protein-ligand docking prediction (KRAS G12C covalent complex with Sotorasib), output handling, and 3D visualization. | <a href="https://colab.research.google.com/github/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/alphafold3/cloudai_alphafold3_vai_quickstart.ipynb"><img src="https://www.gstatic.com/pantheon/images/bigquery/welcome_page/colab-logo.svg" alt="Open in Colab" height="20"></a> |
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
@@ -233,7 +233,7 @@ def download_image(url: str) -> str:
base64 encoded image.
"""
response = requests.get(url)
return Image.open(io.BytesIO(response.content)) # pytype: disable=bad-return-type # pillow-102-upgrade
return Image.open(io.BytesIO(response.content))
def resize_image(image: Any, new_width: int = 1000) -> Any:
@@ -543,7 +543,8 @@ def get_quota_id(
"NVIDIA_H100_80GB": "H100GPUs",
"NVIDIA_H100_MEGA_80GB": "H100MEGAGPUs",
"NVIDIA_H200_141GB": "H200GPUs",
"NVIDIA_GB200": "B200GPUs",
"NVIDIA_GB200": "GB200GPUs",
"NVIDIA_B200": "B200GPUs",
"NVIDIA_TESLA_T4": "T4GPUs",
"NVIDIA_RTX_PRO_6000": "RTXPRO6000GPUs",
"TPU_7x": "7XTPU",
@@ -233,7 +233,7 @@ def download_image(url: str) -> str:
base64 encoded image.
"""
response = requests.get(url)
return Image.open(io.BytesIO(response.content)) # pytype: disable=bad-return-type # pillow-102-upgrade
return Image.open(io.BytesIO(response.content))
def resize_image(image: Any, new_width: int = 1000) -> Any:
@@ -543,7 +543,8 @@ def get_quota_id(
"NVIDIA_H100_80GB": "H100GPUs",
"NVIDIA_H100_MEGA_80GB": "H100MEGAGPUs",
"NVIDIA_H200_141GB": "H200GPUs",
"NVIDIA_GB200": "B200GPUs",
"NVIDIA_GB200": "GB200GPUs",
"NVIDIA_B200": "B200GPUs",
"NVIDIA_TESLA_T4": "T4GPUs",
"NVIDIA_RTX_PRO_6000": "RTXPRO6000GPUs",
"TPU_7x": "7XTPU",
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
@@ -1,406 +1,406 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "B8S-yo8qTIcO"
},
"outputs": [],
"source": [
"# Copyright 2026 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",
"# https://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": "MTRywGxLTZfU"
},
"source": [
"# Vertex AI Model Garden - Gemma Evaluation\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_gemma_evaluation.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gemma_evaluation.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "2CXS0vZfT8_7"
},
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates evaluating pre-trained and instruction-tuned Gemma models in Vertex AI.\n",
"\n",
"### Objective\n",
"\n",
"- Evaluate pre-trained and instruction-tuned Gemma model on any of the benchmark datasets\n",
"- Clean up the resources\n",
"\n",
"| Models |\n",
"| :- |\n",
"| [google/gemma-2b](https://huggingface.co/google/gemma-2b)\n",
"| [google/gemma-2b-it](https://huggingface.co/google/gemma-2b-it)\n",
"| [google/gemma-7b](https://huggingface.co/google/gemma-7b)\n",
"| [google/gemma-7b-it](https://huggingface.co/google/gemma-7b-it)\n",
"| [google/gemma-1.1-2b-it](https://huggingface.co/google/gemma-1.1-2b-it)\n",
"| [google/gemma-1.1-7b-it](https://huggingface.co/google/gemma-1.1-7b-it)\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "HCY8PGrFUbT1"
},
"source": [
"## Run the notebook"
]
},
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "81CC3tL1T_TL"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. **[Optional]** [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 3. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# Import the necessary packages\n",
"\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"import datetime\n",
"import importlib\n",
"import os\n",
"import uuid\n",
"\n",
"from google.cloud import aiplatform\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.notebooks.community.model_garden.docker_source_codes.notebook_util.common_util\"\n",
")\n",
"\n",
"models, endpoints = {}, {}\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"if not REGION:\n",
" if not os.environ.get(\"GOOGLE_CLOUD_REGION\"):\n",
" raise ValueError(\n",
" \"REGION must be set. See\"\n",
" \" https://cloud.google.com/vertex-ai/docs/general/locations for\"\n",
" \" available cloud locations.\"\n",
" )\n",
" REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, change the value yourself below.\n",
"now = datetime.datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}\"\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"MODEL_BUCKET = os.path.join(BUCKET_URI, \"gemma\")\n",
"\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default SERVICE_ACCOUNT.\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/storage.admin\"\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/aiplatform.user\""
]
},
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "pNHMbjr0UjrK"
},
"outputs": [],
"source": [
"# @title Evaluate Gemma models\n",
"\n",
"# @markdown This section demonstrates how to evaluate the Gemma models with and without finetuned LoRA adapters using EleutherAI's [Language Model Evaluation Harness (lm-evaluation-harness)](https://github.com/EleutherAI/lm-evaluation-harness) with Vertex CustomJob. Refer the peak GPU memory usage for serving and adjust the machine type, accelerator type and accelerator count accordingly.\n",
"\n",
"# @markdown You must provide a Hugging Face User Access Token (with read access) to access the Gemma models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"HF_TOKEN = \"\" # @param {type:\"string\", isTemplate:true}\n",
"\n",
"# @markdown This example uses the dataset [HellaSwag](https://arxiv.org/abs/1905.07830). All supported tasks are listed in [this task table](https://github.com/EleutherAI/lm-evaluation-harness/blob/master/docs/task_table.md).\n",
"# @markdown Set evaluation dataset.\n",
"eval_dataset = \"hellaswag\" # @param {type:\"string\"}\n",
"\n",
"# Worker pool spec.\n",
"# Find Vertex AI supported accelerators and regions in:\n",
"# https://cloud.google.com/vertex-ai/docs/training/configure-compute\n",
"\n",
"\n",
"# Setup evaluation job.\n",
"# @markdown Set the base model id.\n",
"base_model_id = \"google/gemma-1.1-2b-it\" # @param[\"google/gemma-2b\", \"google/gemma-2b-it\", \"google/gemma-7b\", \"google/gemma-7b-it\", \"google/gemma-1.1-2b-it\", \"google/gemma-1.1-7b-it\"] {isTemplate:true}\n",
"job_name = common_util.get_job_name_with_datetime(prefix=\"gemma-eval\")\n",
"eval_output_dir = os.path.join(MODEL_BUCKET, job_name)\n",
"eval_output_dir_gcsfuse = eval_output_dir.replace(\"gs://\", \"/gcs/\")\n",
"\n",
"# @markdown Set the accelerator type.\n",
"accelerator_type = \"NVIDIA_L4\" # @param[\"NVIDIA_TESLA_V100\", \"NVIDIA_L4\", \"NVIDIA_TESLA_A100\"]\n",
"\n",
"# @markdown To evaluate a PEFT-finetuned model, enter the PEFT output directory to the LoRA adapter below.\n",
"# @markdown Otherwise, leave it empty.\n",
"# @markdown See the [finetuning notebook](https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gemma_finetuning_on_vertex.ipynb) for more details.\n",
"# @markdown Set the PEFT output directory.\n",
"peft_output_dir = \"\" # @param {type:\"string\"}\n",
"peft_output_dir_gcsfuse = peft_output_dir.replace(\"gs://\", \"/gcs/\")\n",
"\n",
"if accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 1\n",
"elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-standard-8\"\n",
" accelerator_count = 2\n",
"elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-8\"\n",
" accelerator_count = 1\n",
"else:\n",
" print(f\"Unsupported accelerator type: {accelerator_type}\")\n",
"\n",
"replica_count = 1\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=True,\n",
")\n",
"\n",
"# Prepare evaluation command that runs the evaluation harness.\n",
"# Set `trust_remote_code = True` because evaluating the model requires\n",
"# executing code from the model repository.\n",
"# Set `use_accelerate = True` to enable evaluation across multiple GPUs.\n",
"eval_command = [\n",
" \"lm_eval\",\n",
" \"--model\",\n",
" \"hf\",\n",
" \"--tasks\",\n",
" f\"{eval_dataset}\",\n",
" \"--output_path\",\n",
" f\"{eval_output_dir_gcsfuse}\",\n",
"]\n",
"\n",
"if peft_output_dir_gcsfuse:\n",
" eval_command += [\n",
" \"--model_args\",\n",
" f\"pretrained={base_model_id},peft={peft_output_dir_gcsfuse},trust_remote_code=True,parallelize=True\",\n",
" ]\n",
"else:\n",
" eval_command += [\n",
" \"--model_args\",\n",
" f\"pretrained={base_model_id},trust_remote_code=True,parallelize=True\",\n",
" ]\n",
"\n",
"\n",
"# The evaluation docker image.\n",
"EVAL_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-lm-evaluation-harness:20241016_0934_RC00\"\n",
"\n",
"# Pass evaluation arguments and launch job.\n",
"worker_pool_specs = [\n",
" {\n",
" \"machine_spec\": {\n",
" \"machine_type\": machine_type,\n",
" \"accelerator_type\": accelerator_type,\n",
" \"accelerator_count\": accelerator_count,\n",
" },\n",
" \"replica_count\": replica_count,\n",
" \"disk_spec\": {\n",
" \"boot_disk_size_gb\": 500,\n",
" },\n",
" \"container_spec\": {\n",
" \"image_uri\": EVAL_DOCKER_URI,\n",
" \"env\": [\n",
" {\n",
" \"name\": \"HF_TOKEN\",\n",
" \"value\": HF_TOKEN,\n",
" }\n",
" ],\n",
" \"command\": eval_command,\n",
" \"args\": [],\n",
" },\n",
" }\n",
"]\n",
"\n",
"eval_job = aiplatform.CustomJob(\n",
" display_name=job_name,\n",
" worker_pool_specs=worker_pool_specs,\n",
" base_output_dir=eval_output_dir,\n",
")\n",
"\n",
"eval_job.run()\n",
"\n",
"print(\"Evaluation results were saved in:\", eval_output_dir)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "CVBxGpwWU3kY"
},
"outputs": [],
"source": [
"# @title Fetch and print evaluation results\n",
"import json\n",
"import re\n",
"\n",
"from google.cloud import storage\n",
"\n",
"# Fetch evaluation results.\n",
"storage_client = storage.Client()\n",
"BUCKET_NAME = BUCKET_URI.split(\"gs://\")[1]\n",
"bucket = storage_client.get_bucket(BUCKET_NAME)\n",
"\n",
"blobs = [b.name for b in bucket.list_blobs()]\n",
"\n",
"result_file_path = None\n",
"for file_path in filter(re.compile(\".*/*.json\").match, blobs):\n",
" result_file_path = file_path\n",
" print(f\"Found result file: {file_path}\")\n",
"\n",
"if result_file_path is None:\n",
" raise ValueError(\"No result file found.\")\n",
"\n",
"blob = bucket.blob(result_file_path)\n",
"raw_result = blob.download_as_string()\n",
"\n",
"# Print evaluation results.\n",
"result = json.loads(raw_result)\n",
"result_formatted = json.dumps(result, indent=2)\n",
"print(f\"Evaluation result:\\n{result_formatted}\")"
]
},
{
"cell_type": "markdown",
"execution_count": null,
"metadata": {
"id": "unjukbcjEBOd"
},
"outputs": [],
"source": [
"## Clean up resources"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "qWN3cl_VU7pa"
},
"outputs": [],
"source": [
"# Delete evaluation job.\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_URI\n",
" # Uncomment below to delete all artifacts\n",
" # !gsutil -m rm -r $STAGING_BUCKET $MODEL_BUCKET $EXPERIMENT_BUCKET\n",
"\n",
"eval_job.delete()"
]
}
],
"metadata": {
"colab": {
"name": "model_garden_gemma_evaluation.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
"cells": [
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "B8S-yo8qTIcO"
},
"outputs": [],
"source": [
"# Copyright 2026 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",
"# https://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."
]
},
"nbformat": 4,
"nbformat_minor": 0
{
"cell_type": "markdown",
"metadata": {
"id": "MTRywGxLTZfU"
},
"source": [
"# Vertex AI Model Garden - Gemma Evaluation\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_gemma_evaluation.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gemma_evaluation.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "2CXS0vZfT8_7"
},
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates evaluating pre-trained and instruction-tuned Gemma models in Vertex AI.\n",
"\n",
"### Objective\n",
"\n",
"- Evaluate pre-trained and instruction-tuned Gemma model on any of the benchmark datasets\n",
"- Clean up the resources\n",
"\n",
"| Models |\n",
"| :- |\n",
"| [google/gemma-2b](https://huggingface.co/google/gemma-2b)\n",
"| [google/gemma-2b-it](https://huggingface.co/google/gemma-2b-it)\n",
"| [google/gemma-7b](https://huggingface.co/google/gemma-7b)\n",
"| [google/gemma-7b-it](https://huggingface.co/google/gemma-7b-it)\n",
"| [google/gemma-1.1-2b-it](https://huggingface.co/google/gemma-1.1-2b-it)\n",
"| [google/gemma-1.1-7b-it](https://huggingface.co/google/gemma-1.1-7b-it)\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "HCY8PGrFUbT1"
},
"source": [
"## Run the notebook"
]
},
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "81CC3tL1T_TL"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. **[Optional]** [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 3. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# Import the necessary packages\n",
"\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"import datetime\n",
"import importlib\n",
"import os\n",
"import uuid\n",
"\n",
"from google.cloud import aiplatform\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.notebooks.community.model_garden.docker_source_codes.notebook_util.common_util\"\n",
")\n",
"\n",
"models, endpoints = {}, {}\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"if not REGION:\n",
" if not os.environ.get(\"GOOGLE_CLOUD_REGION\"):\n",
" raise ValueError(\n",
" \"REGION must be set. See\"\n",
" \" https://cloud.google.com/vertex-ai/docs/general/locations for\"\n",
" \" available cloud locations.\"\n",
" )\n",
" REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, change the value yourself below.\n",
"now = datetime.datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}\"\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"MODEL_BUCKET = os.path.join(BUCKET_URI, \"gemma\")\n",
"\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default SERVICE_ACCOUNT.\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/storage.admin\"\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/aiplatform.user\""
]
},
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "pNHMbjr0UjrK"
},
"outputs": [],
"source": [
"# @title Evaluate Gemma models\n",
"\n",
"# @markdown This section demonstrates how to evaluate the Gemma models with and without finetuned LoRA adapters using EleutherAI's [Language Model Evaluation Harness (lm-evaluation-harness)](https://github.com/EleutherAI/lm-evaluation-harness) with Vertex CustomJob. Refer the peak GPU memory usage for serving and adjust the machine type, accelerator type and accelerator count accordingly.\n",
"\n",
"# @markdown You must provide a Hugging Face User Access Token (with read access) to access the Gemma models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"HF_TOKEN = \"\" # @param {type:\"string\", isTemplate:true}\n",
"\n",
"# @markdown This example uses the dataset [HellaSwag](https://arxiv.org/abs/1905.07830). All supported tasks are listed in [this task table](https://github.com/EleutherAI/lm-evaluation-harness/blob/master/docs/task_table.md).\n",
"# @markdown Set evaluation dataset.\n",
"eval_dataset = \"hellaswag\" # @param {type:\"string\"}\n",
"\n",
"# Worker pool spec.\n",
"# Find Vertex AI supported accelerators and regions in:\n",
"# https://cloud.google.com/vertex-ai/docs/training/configure-compute\n",
"\n",
"\n",
"# Setup evaluation job.\n",
"# @markdown Set the base model id.\n",
"base_model_id = \"google/gemma-1.1-2b-it\" # @param[\"google/gemma-2b\", \"google/gemma-2b-it\", \"google/gemma-7b\", \"google/gemma-7b-it\", \"google/gemma-1.1-2b-it\", \"google/gemma-1.1-7b-it\"] {isTemplate:true}\n",
"job_name = common_util.get_job_name_with_datetime(prefix=\"gemma-eval\")\n",
"eval_output_dir = os.path.join(MODEL_BUCKET, job_name)\n",
"eval_output_dir_gcsfuse = eval_output_dir.replace(\"gs://\", \"/gcs/\")\n",
"\n",
"# @markdown Set the accelerator type.\n",
"accelerator_type = \"NVIDIA_L4\" # @param[\"NVIDIA_TESLA_V100\", \"NVIDIA_L4\", \"NVIDIA_TESLA_A100\"]\n",
"\n",
"# @markdown To evaluate a PEFT-finetuned model, enter the PEFT output directory to the LoRA adapter below.\n",
"# @markdown Otherwise, leave it empty.\n",
"# @markdown See the [finetuning notebook](https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gemma_finetuning_on_vertex.ipynb) for more details.\n",
"# @markdown Set the PEFT output directory.\n",
"peft_output_dir = \"\" # @param {type:\"string\"}\n",
"peft_output_dir_gcsfuse = peft_output_dir.replace(\"gs://\", \"/gcs/\")\n",
"\n",
"if accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 1\n",
"elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-standard-8\"\n",
" accelerator_count = 2\n",
"elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-8\"\n",
" accelerator_count = 1\n",
"else:\n",
" print(f\"Unsupported accelerator type: {accelerator_type}\")\n",
"\n",
"replica_count = 1\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=True,\n",
")\n",
"\n",
"# Prepare evaluation command that runs the evaluation harness.\n",
"# Set `trust_remote_code = True` because evaluating the model requires\n",
"# executing code from the model repository.\n",
"# Set `use_accelerate = True` to enable evaluation across multiple GPUs.\n",
"eval_command = [\n",
" \"lm_eval\",\n",
" \"--model\",\n",
" \"hf\",\n",
" \"--tasks\",\n",
" f\"{eval_dataset}\",\n",
" \"--output_path\",\n",
" f\"{eval_output_dir_gcsfuse}\",\n",
"]\n",
"\n",
"if peft_output_dir_gcsfuse:\n",
" eval_command += [\n",
" \"--model_args\",\n",
" f\"pretrained={base_model_id},peft={peft_output_dir_gcsfuse},trust_remote_code=True,parallelize=True\",\n",
" ]\n",
"else:\n",
" eval_command += [\n",
" \"--model_args\",\n",
" f\"pretrained={base_model_id},trust_remote_code=True,parallelize=True\",\n",
" ]\n",
"\n",
"\n",
"# The evaluation docker image.\n",
"EVAL_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-lm-evaluation-harness:20241016_0934_RC00\"\n",
"\n",
"# Pass evaluation arguments and launch job.\n",
"worker_pool_specs = [\n",
" {\n",
" \"machine_spec\": {\n",
" \"machine_type\": machine_type,\n",
" \"accelerator_type\": accelerator_type,\n",
" \"accelerator_count\": accelerator_count,\n",
" },\n",
" \"replica_count\": replica_count,\n",
" \"disk_spec\": {\n",
" \"boot_disk_size_gb\": 500,\n",
" },\n",
" \"container_spec\": {\n",
" \"image_uri\": EVAL_DOCKER_URI,\n",
" \"env\": [\n",
" {\n",
" \"name\": \"HF_TOKEN\",\n",
" \"value\": HF_TOKEN,\n",
" }\n",
" ],\n",
" \"command\": eval_command,\n",
" \"args\": [],\n",
" },\n",
" }\n",
"]\n",
"\n",
"eval_job = aiplatform.CustomJob(\n",
" display_name=job_name,\n",
" worker_pool_specs=worker_pool_specs,\n",
" base_output_dir=eval_output_dir,\n",
")\n",
"\n",
"eval_job.run()\n",
"\n",
"print(\"Evaluation results were saved in:\", eval_output_dir)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "CVBxGpwWU3kY"
},
"outputs": [],
"source": [
"# @title Fetch and print evaluation results\n",
"import json\n",
"import re\n",
"\n",
"from google.cloud import storage\n",
"\n",
"# Fetch evaluation results.\n",
"storage_client = storage.Client()\n",
"BUCKET_NAME = BUCKET_URI.split(\"gs://\")[1]\n",
"bucket = storage_client.get_bucket(BUCKET_NAME)\n",
"\n",
"blobs = [b.name for b in bucket.list_blobs()]\n",
"\n",
"result_file_path = None\n",
"for file_path in filter(re.compile(\".*/*.json\").match, blobs):\n",
" result_file_path = file_path\n",
" print(f\"Found result file: {file_path}\")\n",
"\n",
"if result_file_path is None:\n",
" raise ValueError(\"No result file found.\")\n",
"\n",
"blob = bucket.blob(result_file_path)\n",
"raw_result = blob.download_as_string()\n",
"\n",
"# Print evaluation results.\n",
"result = json.loads(raw_result)\n",
"result_formatted = json.dumps(result, indent=2)\n",
"print(f\"Evaluation result:\\n{result_formatted}\")"
]
},
{
"cell_type": "markdown",
"execution_count": null,
"metadata": {
"id": "unjukbcjEBOd"
},
"outputs": [],
"source": [
"## Clean up resources"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "qWN3cl_VU7pa"
},
"outputs": [],
"source": [
"# Delete evaluation job.\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_URI\n",
" # Uncomment below to delete all artifacts\n",
" # !gsutil -m rm -r $STAGING_BUCKET $MODEL_BUCKET $EXPERIMENT_BUCKET\n",
"\n",
"eval_job.delete()"
]
}
],
"metadata": {
"colab": {
"name": "model_garden_gemma_evaluation.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
@@ -1,964 +0,0 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "7d9bbf86da5e"
},
"outputs": [],
"source": [
"# Copyright 2026 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",
"# https://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": "99c1c3fc2ca5"
},
"source": [
"# Vertex AI Model Garden - Gemma Finetuning\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_gemma_finetuning_on_vertex.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gemma_finetuning_on_vertex.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "3de7470326a2"
},
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates finetuning and deploying Gemma models with [Vertex AI Custom Training Job](https://cloud.google.com/vertex-ai/docs/training/create-custom-job). All of the examples in this notebook use parameter efficient finetuning methods [PEFT (LoRA)](https://github.com/huggingface/peft) to reduce training and storage costs. LoRA (Low-Rank Adaptation) is one approach of Parameter Efficient FineTuning (PEFT), where pretrained model weights are frozen and rank decomposition matrices representing the change in model weights are trained during finetuning. Read more about LoRA in the following publication: [Hu, E.J., Shen, Y., Wallis, P., Allen-Zhu, Z., Li, Y., Wang, S., Wang, L. and Chen, W., 2021. Lora: Low-rank adaptation of large language models. *arXiv preprint arXiv:2106.09685*](https://arxiv.org/abs/2106.09685).\n",
"\n",
"\n",
"After tuning, we can deploy models on Vertex.\n",
"\n",
"\n",
"### Objective\n",
"\n",
"- Finetune and deploy Gemma models with Vertex AI Custom Training Jobs.\n",
"- Send prediction requests to your finetuned Gemma model.\n",
"\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "264c07757582"
},
"source": [
"## Before you begin"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "BNAlJh_pGxbL"
},
"outputs": [],
"source": [
"# @title Install Python Packages for Finetuning\n",
"\n",
"# @markdown 1. Install google-cloud-aiplatform package and restart the session if instructed.\n",
"! pip install --upgrade --quiet google-cloud-aiplatform==1.130.0\n",
"\n",
"# @markdown 2. Install packages to validate dataset with template.\n",
"! pip install --upgrade --quiet accelerate==0.31.0\n",
"! pip install --upgrade --quiet transformers==4.43.1\n",
"! pip install --upgrade --quiet datasets==2.19.2\n",
"\n",
"# Load local tensorboard.\n",
"%load_ext tensorboard"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "8CQcnBfWvc-f"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. For finetuning, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Frestricted_image_training_nvidia_a100_80gb_gpus)** to check if your project already has the required 8 Nvidia A100 80 GB GPUs in the us-central1 region. If yes, then run this notebook in the us-central1 region. If you do not have 8 Nvidia A100 80 GPUs or have more GPU requirements than this, then schedule your job with Nvidia H100 GPUs via Dynamic Workload Scheduler using [these instructions](https://cloud.google.com/vertex-ai/docs/training/schedule-jobs-dws). For Dynamic Workload Scheduler, check the [us-central1](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) or [europe-west4](https://console.cloud.google.com/iam-admin/quotas?location=europe-west4&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) quota for Nvidia H100 GPUs. If you do not have enough GPUs, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request quota.\n",
"\n",
"# @markdown 3. For serving, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_l4_gpus)** to check if your project already has the required 1 L4 GPU in the us-central1 region. If yes, then run this notebook in the us-central1 region. If you need more L4 GPUs for your project, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request more. Alternatively, if you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus).\n",
"\n",
"# @markdown > | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
"# @markdown | a2-ultragpu-1g | 1 NVIDIA_A100_80GB | us-central1, us-east4, europe-west4, asia-southeast1, us-east4 |\n",
"# @markdown | a3-highgpu-2g | 2 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-4g | 4 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"\n",
"# @markdown 4. **[Optional]** [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 5. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# Import the necessary packages\n",
"! rm -rf vertex-ai-samples && git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"! cd vertex-ai-samples && git reset --hard 7ae13b346a72ee2a2dc8152dd40c6ddd72d6c810\n",
"\n",
"import datetime\n",
"import importlib\n",
"import os\n",
"import uuid\n",
"from typing import Tuple\n",
"\n",
"from google.cloud import aiplatform\n",
"from google.cloud.aiplatform.compat.types import \\\n",
" custom_job as gca_custom_job_compat\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.notebooks.community.model_garden.docker_source_codes.notebook_util.common_util\"\n",
")\n",
"\n",
"models, endpoints = {}, {}\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"if not REGION:\n",
" if not os.environ.get(\"GOOGLE_CLOUD_REGION\"):\n",
" raise ValueError(\n",
" \"REGION must be set. See\"\n",
" \" https://cloud.google.com/vertex-ai/docs/general/locations for\"\n",
" \" available cloud locations.\"\n",
" )\n",
" REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, change the value yourself below.\n",
"now = datetime.datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}\"\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"MODEL_BUCKET = os.path.join(BUCKET_URI, \"gemma\")\n",
"\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default SERVICE_ACCOUNT.\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/storage.admin\"\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/aiplatform.user\"\n",
"\n",
"# @markdown ## Access Gemma Models\n",
"# @markdown For GPU based finetuning and serving, choose between accessing Gemma models on [Hugging Face](https://huggingface.co/)\n",
"# @markdown or Vertex AI as described below.\n",
"\n",
"# @markdown If you already obtained access to Gemma models on [Hugging Face](https://huggingface.co/), you can load models from there.\n",
"# @markdown Alternatively, you can also load the original Gemma models for finetuning and serving from Vertex AI after accepting the agreement.\n",
"\n",
"# @markdown **Select and fill one of the three following sections.**\n",
"LOAD_MODEL_FROM = \"Hugging Face\" # @param [\"Hugging Face\", \"Google Cloud\"] {isTemplate:true}\n",
"\n",
"# @markdown ---\n",
"\n",
"# @markdown ### Access Gemma models on Hugging Face for GPU based finetuning and serving\n",
"# @markdown You must provide a Hugging Face User Access Token (with read access) to access the Gemma models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"\n",
"HF_TOKEN = \"\" # @param {type:\"string\", isTemplate:true}\n",
"if LOAD_MODEL_FROM == \"Hugging Face\":\n",
" assert (\n",
" HF_TOKEN\n",
" ), \"Provide a read HF_TOKEN to load models from Hugging Face, or select a different model source.\"\n",
"\n",
"# @markdown *--- Or ---*\n",
"# @markdown ### Access Gemma models on Vertex AI for GPU based finetuning and serving\n",
"# @markdown Accept the model agreement to access the models:\n",
"# @markdown 1. Open the [Gemma model card](https://console.cloud.google.com/vertex-ai/publishers/google/model-garden/335) from [Vertex AI Model Garden](https://cloud.google.com/model-garden).\n",
"# @markdown 2. Review the agreement on the model card page.\n",
"# @markdown 3. After accepting the agreement of Gemma, a `https://` link containing Gemma pretrained and finetuned models will be shared.\n",
"# @markdown 4. Paste the link in the `VERTEX_MODEL_GARDEN_GEMMA` field below.\n",
"# @markdown **Note:** This will unzip and copy the Gemma model artifacts to your Cloud Storage bucket, which will take around 1 hour.\n",
"\n",
"VERTEX_AI_MODEL_GARDEN_GEMMA = \"\" # @param {type:\"string\", isTemplate:true}\n",
"\n",
"if LOAD_MODEL_FROM == \"Google Cloud\":\n",
" assert (\n",
" VERTEX_AI_MODEL_GARDEN_GEMMA\n",
" ), \"Accept the agreement of Gemma in Vertex AI Model Garden and get the URL to Gemma model artifacts, or select a different model source.\"\n",
"\n",
" # Only use the last part in case a full command is pasted.\n",
" signed_url = VERTEX_AI_MODEL_GARDEN_GEMMA.split(\" \")[-1].strip('\"')\n",
"\n",
" ! mkdir -p ./gemma\n",
" ! curl -X GET \"{signed_url}\" | tar -xzvf - -C ./gemma/\n",
" ! gsutil -m cp -R ./gemma/* {MODEL_BUCKET}\n",
"\n",
" model_path_prefix = MODEL_BUCKET\n",
" HF_TOKEN = \"\"\n",
"else:\n",
" model_path_prefix = \"google/\"\n",
"\n",
"conversion_job = None"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "cb56d402e84a"
},
"source": [
"## Finetune with HuggingFace PEFT and Deploy with vLLM on GPUs"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "KwAW99YZHTdy"
},
"outputs": [],
"source": [
"# @title Set dataset\n",
"\n",
"# @markdown Use the Vertex AI SDK to create and run the custom training jobs.\n",
"\n",
"# @markdown This notebook uses [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) dataset as an example.\n",
"# @markdown You can set `dataset_name` to any existing [Hugging Face dataset](https://huggingface.co/datasets) name, and set `instruct_column_in_dataset` to the name of the dataset column containing training data. The [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) has only one column `text`, and therefore we set `instruct_column_in_dataset` to `text` in this notebook.\n",
"\n",
"# @markdown ### (Optional) Prepare a custom JSONL dataset for finetuning\n",
"\n",
"# @markdown You can prepare a JSONL file where each line is a valid JSON string as your custom training dataset. For example, here is one line from the [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) dataset:\n",
"# @markdown ```\n",
"# @markdown {\"text\": \"### Human: Hola### Assistant: \\u00a1Hola! \\u00bfEn qu\\u00e9 puedo ayudarte hoy?\"}\n",
"# @markdown ```\n",
"\n",
"# @markdown The JSON object has a key `text`, which should match `instruct_column_in_dataset`; The value should be one training data point, i.e. a string. After you prepared your JSONL file, you can either upload it to [Hugging Face datasets](https://huggingface.co/datasets) or [Google Cloud Storage](https://cloud.google.com/storage).\n",
"\n",
"# @markdown - To upload a JSONL dataset to [Hugging Face datasets](https://huggingface.co/datasets), follow the instructions on [Uploading Datasets](https://huggingface.co/docs/hub/en/datasets-adding). Then, set `dataset_name` to the name of your newly created dataset on Hugging Face.\n",
"\n",
"# @markdown - To upload a JSONL dataset to [Google Cloud Storage](https://cloud.google.com/storage), follow the instructions on [Upload objects from a filesystem](https://cloud.google.com/storage/docs/uploading-objects). Then, set `dataset_name` to the `gs://` URI to your JSONL file. For example: `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`.\n",
"\n",
"# @markdown Optionally update the `instruct_column_in_dataset` field below if your JSON objects use a key other than the default `text`.\n",
"\n",
"# @markdown ### (Optional) Format your data with custom JSON template\n",
"\n",
"# @markdown Sometimes, your dataset might have multiple text columns and you want to construct the training data with a template. You can prepare a JSON template in the following format:\n",
"\n",
"# @markdown ```\n",
"# @markdown {\n",
"# @markdown \"description\": \"Template that accepts text-bison format.\",\n",
"# @markdown \"source\": \"https://cloud.google.com/vertex-ai/generative-ai/docs/models/tune-text-models-supervised#dataset-format\",\n",
"# @markdown \"prompt_input\": \"\\n\\n<|start_header_id|>user<|end_header_id|>\\n\\n{input_text}<|eot_id|>\\n\\n<|start_header_id|>assistant<|end_header_id|>\\n\\n{output_text}<|eot_id|>\",\n",
"# @markdown \"instruction_separator\": \"<|start_header_id|>user<|end_header_id|>\\n\\n\",\n",
"# @markdown \"response_separator\": \"<|start_header_id|>assistant<|end_header_id|>\\n\\n\"\n",
"# @markdown }\n",
"# @markdown ```\n",
"\n",
"# @markdown As an example, the template above can be used to format the following training data (this line comes from `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`):\n",
"\n",
"# @markdown ```\n",
"# @markdown {\"input_text\":\"TRANSCRIPT: \\nREASON FOR EVALUATION:,\\n\\n LABEL:\",\"output_text\":\"Chiropractic\"}\n",
"# @markdown ```\n",
"\n",
"# @markdown This example template simply concatenates `input_text` with `output_text` with some special tokens in between.\n",
"# @markdown\n",
"# @markdown To try such custom dataset, you can make the following changes:\n",
"# @markdown 1. Set `template` to `llama3-text-bison`\n",
"# @markdown 1. Set `train_dataset_name` to `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`\n",
"# @markdown 1. Set `train_split_name` to `train`\n",
"# @markdown 1. Set `eval_dataset_name` to `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_eval_sample.jsonl`\n",
"# @markdown 1. Set `eval_split_name` to `train` (**NOT** `test`)\n",
"# @markdown 1. Set `instruct_column_in_dataset` as `input_text`.\n",
"\n",
"# Template name or gs:// URI to a custom template.\n",
"template = \"openassistant-guanaco\" # @param {type:\"string\"}\n",
"\n",
"# Hugging Face dataset name or gs:// URI to a custom JSONL dataset.\n",
"train_dataset_name = \"timdettmers/openassistant-guanaco\" # @param {type:\"string\"}\n",
"train_split_name = \"train\" # @param {type:\"string\"}\n",
"eval_dataset_name = \"timdettmers/openassistant-guanaco\" # @param {type:\"string\"}\n",
"eval_split_name = \"test\" # @param {type:\"string\"}\n",
"\n",
"# Name of the dataset column containing training text input.\n",
"instruct_column_in_dataset = \"text\" # @param {type:\"string\"}"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "SdiyOeyFGxbM"
},
"outputs": [],
"source": [
"# @title Set model\n",
"\n",
"# @markdown Select a model variant of Gemma 2.\n",
"base_model_id = \"gemma-2b\" # @param[\"gemma-2b\", \"gemma-2b-it\", \"gemma-7b\", \"gemma-7b-it\", \"gemma-1.1-2b-it\", \"gemma-1.1-7b-it\"] {isTemplate:true}\n",
"pretrained_model_id = os.path.join(model_path_prefix, base_model_id)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "R5PcRc0MGxbM"
},
"outputs": [],
"source": [
"# @title Validate Dataset with Template\n",
"\n",
"# @markdown This section validates the train and eval datasets with the template before starting the fine tuning process.\n",
"\n",
"import transformers\n",
"\n",
"dataset_validation_util = importlib.import_module(\n",
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.dataset_validation_util\"\n",
")\n",
"\n",
"if dataset_validation_util.is_gcs_path(pretrained_model_id):\n",
" # Download tokenizer.\n",
" ! mkdir tokenizer\n",
" ! gsutil cp {pretrained_model_id}/tokenizer.json ./tokenizer\n",
" ! gsutil cp {pretrained_model_id}/config.json ./tokenizer\n",
" tokenizer_path = \"./tokenizer\"\n",
" access_token = \"\"\n",
"else:\n",
" tokenizer_path = pretrained_model_id\n",
" access_token = HF_TOKEN\n",
"\n",
"tokenizer = transformers.AutoTokenizer.from_pretrained(\n",
" tokenizer_path,\n",
" trust_remote_code=False,\n",
" use_fast=True,\n",
" token=access_token,\n",
")\n",
"\n",
"# Validate the train dataset.\n",
"dataset_validation_util.validate_dataset_with_template(\n",
" dataset_name=train_dataset_name,\n",
" split=train_split_name,\n",
" input_column=instruct_column_in_dataset,\n",
" template=template,\n",
" use_multiprocessing=False,\n",
" tokenizer=tokenizer,\n",
")\n",
"\n",
"# Validate the eval dataset.\n",
"dataset_validation_util.validate_dataset_with_template(\n",
" dataset_name=eval_dataset_name,\n",
" split=eval_split_name,\n",
" input_column=instruct_column_in_dataset,\n",
" template=template,\n",
" use_multiprocessing=False,\n",
" tokenizer=tokenizer,\n",
")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "ivVGS9dHXPOz"
},
"outputs": [],
"source": [
"# @title Finetune\n",
"# @markdown This section demonstrates how to finetune the Gemma model and merge the finetuned LoRA adapter with the base model on Vertex AI. It uses the Vertex AI SDK to create and run the custom training jobs.\n",
"\n",
"# @markdown The training job takes approximately between 10 to 20 mins to set-up. Once done, the training job is expected to take around 20 mins with the default configuration. To find the training time, throughput, and memory usage of your training job, you can go to the training logs and check the log line of the last training epoch.\n",
"\n",
"# @markdown **Note**:\n",
"# @markdown 1. We recommend setting `finetuning_precision_mode` to `4bit` because it enables using fewer hardware resources for finetuning.\n",
"# @markdown 1. If `max_steps > 0`, it takes precedence over `epochs`. One can set a small `max_steps` value to quickly check the pipeline.\n",
"\n",
"# Add labels for the finetuning job.\n",
"labels = {\n",
" \"mg-source\": \"notebook\",\n",
" \"mg-notebook-name\": \"model_garden_gemma_finetuning_on_vertex.ipynb\".split(\".\")[0],\n",
"}\n",
"\n",
"labels[\"mg-tune\"] = \"publishers-google-models-gemma\"\n",
"versioned_model_id = base_model_id.lower().replace(\".\", \"-\")\n",
"labels[\"versioned-mg-tune\"] = f\"{labels['mg-tune']}-{versioned_model_id}\"\n",
"\n",
"# @markdown Accelerator type to use for training.\n",
"training_accelerator_type = \"NVIDIA_A100_80GB\" # @param [\"NVIDIA_A100_80GB\", \"NVIDIA_H100_80GB\"]\n",
"\n",
"# The pre-built training docker image.\n",
"if training_accelerator_type == \"NVIDIA_A100_80GB\":\n",
" repo = \"us-docker.pkg.dev/vertex-ai-restricted\"\n",
" is_restricted_image = True\n",
" is_dynamic_workload_scheduler = False\n",
" dws_kwargs = {}\n",
"else:\n",
" repo = \"us-docker.pkg.dev/vertex-ai\"\n",
" is_restricted_image = False\n",
" is_dynamic_workload_scheduler = True\n",
" dws_kwargs = {\n",
" \"max_wait_duration\": 1800, # 30 minutes\n",
" \"scheduling_strategy\": gca_custom_job_compat.Scheduling.Strategy.FLEX_START,\n",
" }\n",
"\n",
"TRAIN_DOCKER_URI = (\n",
" f\"{repo}/vertex-vision-model-garden-dockers/pytorch-peft-train:stable_20240909\"\n",
")\n",
"\n",
"# Worker pool spec.\n",
"if training_accelerator_type == \"NVIDIA_A100_80GB\":\n",
" per_node_accelerator_count = 8\n",
" training_machine_type = \"a2-ultragpu-8g\"\n",
"elif training_accelerator_type == \"NVIDIA_H100_80GB\":\n",
" per_node_accelerator_count = 8\n",
" training_machine_type = \"a3-highgpu-8g\"\n",
"else:\n",
" raise ValueError(\n",
" f\"Recommended machine settings not found for: {training_accelerator_type}. To use another accelerator type, edit this code block to pass in an appropriate `training_machine_type`, `training_accelerator_type`, and `per_node_accelerator_count` to the deploy_model_vllm function by clicking `Show Code` and then modifying the code.\"\n",
" )\n",
"\n",
"# @markdown Batch size for finetuning.\n",
"per_device_train_batch_size = 1 # @param{type:\"integer\"}\n",
"# @markdown Number of updates steps to accumulate the gradients for, before performing a backward/update pass.\n",
"gradient_accumulation_steps = 4 # @param{type:\"integer\"}\n",
"# @markdown Maximum sequence length.\n",
"max_seq_length = 4096 # @param{type:\"integer\"}\n",
"# @markdown Setting a positive `max_steps` here will override `num_epochs`.\n",
"max_steps = -1 # @param{type:\"integer\"}\n",
"num_epochs = 1.0 # @param{type:\"number\"}\n",
"# @markdown Precision mode for finetuning.\n",
"finetuning_precision_mode = \"4bit\" # @param [\"4bit\", \"8bit\", \"float16\"]\n",
"# @markdown Learning rate.\n",
"learning_rate = 5e-5 # @param{type:\"number\"}\n",
"# @markdown The scheduler type to use.\n",
"lr_scheduler_type = \"cosine\" # @param{type:\"string\"}\n",
"# @markdown LoRA parameters.\n",
"lora_rank = 16 # @param{type:\"integer\"}\n",
"lora_alpha = 32 # @param{type:\"integer\"}\n",
"lora_dropout = 0.05 # @param{type:\"number\"}\n",
"# Activates gradient checkpointing for the current model (may be referred to as activation checkpointing or checkpoint activations in other frameworks).\n",
"enable_gradient_checkpointing = True\n",
"# Attention implementation to use in the model.\n",
"attn_implementation = \"eager\"\n",
"# The optimizer for which to schedule the learning rate.\n",
"optimizer = \"paged_adamw_32bit\"\n",
"# Define the proportion of training to be dedicated to a linear warmup where learning rate gradually increases.\n",
"warmup_ratio = \"0.01\"\n",
"# The list or string of integrations to report the results and logs to.\n",
"report_to = \"tensorboard\"\n",
"# Number of updates steps before two checkpoint saves.\n",
"save_steps = 10\n",
"# Number of update steps between two logs.\n",
"logging_steps = save_steps\n",
"# Train precision of the model.\n",
"train_precision = \"bfloat16\"\n",
"\n",
"replica_count = 1\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=training_accelerator_type,\n",
" accelerator_count=per_node_accelerator_count * replica_count,\n",
" is_for_training=True,\n",
" is_restricted_image=is_restricted_image,\n",
" is_dynamic_workload_scheduler=is_dynamic_workload_scheduler,\n",
")\n",
"\n",
"job_name = common_util.get_job_name_with_datetime(\"gemma-lora-train\")\n",
"\n",
"base_output_dir = os.path.join(STAGING_BUCKET, job_name)\n",
"# Create a GCS folder to store the LORA adapter.\n",
"lora_output_dir = os.path.join(base_output_dir, \"adapter\")\n",
"# Create a GCS folder to store the merged model with the base model and the\n",
"# finetuned LORA adapter.\n",
"merged_model_output_dir = os.path.join(base_output_dir, \"merged-model\")\n",
"\n",
"eval_args = [\n",
" f\"--eval_dataset_path={eval_dataset_name}\",\n",
" f\"--eval_column={instruct_column_in_dataset}\",\n",
" f\"--eval_template={template}\",\n",
" f\"--eval_split={eval_split_name}\",\n",
" f\"--eval_steps={save_steps}\",\n",
" \"--eval_tasks=builtin_eval\",\n",
" \"--eval_metric_name=loss\",\n",
"]\n",
"\n",
"train_job_args = [\n",
" \"--config_file=vertex_vision_model_garden_peft/deepspeed_zero2_8gpu.yaml\",\n",
" \"--task=instruct-lora\",\n",
" \"--completion_only=True\",\n",
" f\"--pretrained_model_id={pretrained_model_id}\",\n",
" f\"--dataset_name={train_dataset_name}\",\n",
" f\"--train_split_name={train_split_name}\",\n",
" f\"--instruct_column_in_dataset={instruct_column_in_dataset}\",\n",
" f\"--output_dir={lora_output_dir}\",\n",
" f\"--merge_base_and_lora_output_dir={merged_model_output_dir}\",\n",
" f\"--per_device_train_batch_size={per_device_train_batch_size}\",\n",
" f\"--gradient_accumulation_steps={gradient_accumulation_steps}\",\n",
" f\"--lora_rank={lora_rank}\",\n",
" f\"--lora_alpha={lora_alpha}\",\n",
" f\"--lora_dropout={lora_dropout}\",\n",
" f\"--max_steps={max_steps}\",\n",
" f\"--max_seq_length={max_seq_length}\",\n",
" f\"--learning_rate={learning_rate}\",\n",
" f\"--lr_scheduler_type={lr_scheduler_type}\",\n",
" f\"--precision_mode={finetuning_precision_mode}\",\n",
" f\"--train_precision={train_precision}\",\n",
" f\"--enable_gradient_checkpointing={enable_gradient_checkpointing}\",\n",
" f\"--num_epochs={num_epochs}\",\n",
" f\"--attn_implementation={attn_implementation}\",\n",
" f\"--optimizer={optimizer}\",\n",
" f\"--warmup_ratio={warmup_ratio}\",\n",
" f\"--report_to={report_to}\",\n",
" f\"--logging_output_dir={base_output_dir}\",\n",
" f\"--save_steps={save_steps}\",\n",
" f\"--logging_steps={logging_steps}\",\n",
" f\"--template={template}\",\n",
" f\"--huggingface_access_token={HF_TOKEN}\",\n",
"] + eval_args\n",
"\n",
"# Pass training arguments and launch job.\n",
"train_job = aiplatform.CustomContainerTrainingJob(\n",
" display_name=job_name,\n",
" container_uri=TRAIN_DOCKER_URI,\n",
" labels=labels,\n",
")\n",
"\n",
"print(\"Running training job with args:\")\n",
"print(\" \\\\\\n\".join(train_job_args))\n",
"train_job.run(\n",
" args=train_job_args,\n",
" replica_count=replica_count,\n",
" machine_type=training_machine_type,\n",
" accelerator_type=training_accelerator_type,\n",
" accelerator_count=per_node_accelerator_count,\n",
" boot_disk_size_gb=500,\n",
" service_account=SERVICE_ACCOUNT,\n",
" base_output_dir=base_output_dir,\n",
" sync=False, # Non-blocking call to run.\n",
" **dws_kwargs,\n",
")\n",
"\n",
"# Wait until resource has been created.\n",
"train_job.wait_for_resource_creation()\n",
"\n",
"print(\"LoRA adapter will be saved in:\", lora_output_dir)\n",
"print(\"Trained and merged models will be saved in:\", merged_model_output_dir)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "FvJV3FwFGxbM"
},
"outputs": [],
"source": [
"# @title Run TensorBoard\n",
"# @markdown This section shows how to launch TensorBoard in a [Cloud Shell](https://cloud.google.com/shell/docs).\n",
"# @markdown 1. Click the Cloud Shell icon(![terminal](https://github.com/google/material-design-icons/blob/master/png/action/terminal/materialicons/24dp/1x/baseline_terminal_black_24dp.png?raw=true)) on the top right to open the Cloud Shell.\n",
"# @markdown 2. Copy the `tensorboard` command shown below by running this cell.\n",
"# @markdown 3. Paste and run the command in the Cloud Shell to launch TensorBoard.\n",
"# @markdown 4. Once the command runs (You may have to click `Authorize` if prompted), click the link starting with `http://localhost`.\n",
"\n",
"# @markdown Note: You may need to wait around 10 minutes after the job starts in order for the TensorBoard logs to be written to the GCS bucket.\n",
"print(f\"Command to copy: tensorboard --logdir {base_output_dir}/logs\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "qmHW6m8xG_4U"
},
"outputs": [],
"source": [
"# @title Deploy\n",
"\n",
"# @markdown This section uploads the model to Model Registry and deploys it on the Endpoint. It takes 15 minutes to 1 hour to finish.\n",
"\n",
"# The pre-built serving docker image for vLLM.\n",
"VLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-vllm-serve:20240815_1634_RC00\"\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"\n",
"def get_deploy_source() -> str:\n",
" \"\"\"Gets deploy_source string based on running environment.\"\"\"\n",
" vertex_product = os.environ.get(\"VERTEX_PRODUCT\", \"\")\n",
" if vertex_product == \"COLAB_ENTERPRISE\":\n",
" return \"notebook_colab_enterprise\"\n",
" elif vertex_product == \"WORKBENCH_INSTANCE\":\n",
" return \"notebook_workbench\"\n",
" else:\n",
" # Legacy workbench, legacy colab, or other custom environments.\n",
" return \"notebook_environment_unspecified\"\n",
"\n",
"\n",
"def deploy_model_vllm(\n",
" model_name: str,\n",
" model_id: str,\n",
" publisher: str,\n",
" publisher_model_id: str,\n",
" service_account: str,\n",
" base_model_id: str = None,\n",
" machine_type: str = \"g2-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
" gpu_memory_utilization: float = 0.9,\n",
" max_model_len: int = 4096,\n",
" dtype: str = \"auto\",\n",
" enable_trust_remote_code: bool = False,\n",
" enforce_eager: bool = False,\n",
" enable_lora: bool = False,\n",
" enable_chunked_prefill: bool = False,\n",
" enable_prefix_cache: bool = False,\n",
" host_prefix_kv_cache_utilization_target: float = 0.0,\n",
" max_loras: int = 1,\n",
" max_cpu_loras: int = 8,\n",
" use_dedicated_endpoint: bool = False,\n",
" max_num_seqs: int = 256,\n",
" model_type: str = None,\n",
" enable_llama_tool_parser: bool = False,\n",
" is_spot: bool = False,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys trained models with vLLM into Vertex AI.\"\"\"\n",
" endpoint = aiplatform.Endpoint.create(\n",
" display_name=f\"{model_name}-endpoint\",\n",
" dedicated_endpoint_enabled=use_dedicated_endpoint,\n",
" )\n",
"\n",
" if not base_model_id:\n",
" base_model_id = model_id\n",
"\n",
" # See https://docs.vllm.ai/en/latest/models/engine_args.html for a list of possible arguments with descriptions.\n",
" vllm_args = [\n",
" \"python\",\n",
" \"-m\",\n",
" \"vllm.entrypoints.api_server\",\n",
" \"--host=0.0.0.0\",\n",
" \"--port=8080\",\n",
" f\"--model={model_id}\",\n",
" f\"--tensor-parallel-size={accelerator_count}\",\n",
" \"--swap-space=16\",\n",
" f\"--max-model-len={max_model_len}\",\n",
" f\"--dtype={dtype}\",\n",
" f\"--max-loras={max_loras}\",\n",
" f\"--max-cpu-loras={max_cpu_loras}\",\n",
" f\"--max-num-seqs={max_num_seqs}\",\n",
" \"--disable-log-stats\",\n",
" ]\n",
"\n",
" if gpu_memory_utilization:\n",
" vllm_args.append(f\"--gpu-memory-utilization={gpu_memory_utilization}\")\n",
"\n",
" if enable_trust_remote_code:\n",
" vllm_args.append(\"--trust-remote-code\")\n",
"\n",
" if enforce_eager:\n",
" vllm_args.append(\"--enforce-eager\")\n",
"\n",
" if enable_lora:\n",
" vllm_args.append(\"--enable-lora\")\n",
"\n",
" if enable_chunked_prefill:\n",
" vllm_args.append(\"--enable-chunked-prefill\")\n",
"\n",
" if enable_prefix_cache:\n",
" vllm_args.append(\"--enable-prefix-caching\")\n",
"\n",
" if 0 < host_prefix_kv_cache_utilization_target < 1:\n",
" vllm_args.append(\n",
" f\"--host-prefix-kv-cache-utilization-target={host_prefix_kv_cache_utilization_target}\"\n",
" )\n",
"\n",
" if model_type:\n",
" vllm_args.append(f\"--model-type={model_type}\")\n",
"\n",
" if enable_llama_tool_parser:\n",
" vllm_args.append(\"--enable-auto-tool-choice\")\n",
" vllm_args.append(\"--tool-call-parser=vertex-llama-3\")\n",
"\n",
" env_vars = {\n",
" \"MODEL_ID\": base_model_id,\n",
" \"DEPLOY_SOURCE\": \"notebook\",\n",
" }\n",
"\n",
" # HF_TOKEN is not a compulsory field and may not be defined.\n",
" try:\n",
" if HF_TOKEN:\n",
" env_vars[\"HF_TOKEN\"] = HF_TOKEN\n",
" except NameError:\n",
" pass\n",
"\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=VLLM_DOCKER_URI,\n",
" serving_container_args=vllm_args,\n",
" serving_container_ports=[8080],\n",
" serving_container_predict_route=\"/generate\",\n",
" serving_container_health_route=\"/ping\",\n",
" serving_container_environment_variables=env_vars,\n",
" serving_container_shared_memory_size_mb=(16 * 1024), # 16 GB\n",
" serving_container_deployment_timeout=7200,\n",
" model_garden_source_model_name=(\n",
" f\"publishers/{publisher}/models/{publisher_model_id}\"\n",
" ),\n",
" )\n",
" print(\n",
" f\"Deploying {model_name} on {machine_type} with {accelerator_count} {accelerator_type} GPU(s).\"\n",
" )\n",
" model.deploy(\n",
" endpoint=endpoint,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" deploy_request_timeout=1800,\n",
" service_account=service_account,\n",
" spot=is_spot,\n",
" system_labels={\n",
" \"NOTEBOOK_NAME\": \"model_garden_gemma_finetuning_on_vertex.ipynb\",\n",
" \"NOTEBOOK_ENVIRONMENT\": get_deploy_source(),\n",
" },\n",
" )\n",
" print(\"endpoint_name:\", endpoint.name)\n",
"\n",
" return model, endpoint\n",
"\n",
"\n",
"if train_job.end_time is None:\n",
" print(\"Waiting for the training job to finish...\")\n",
" train_job.wait()\n",
" print(\"The training job has finished.\")\n",
"\n",
"print(\"Deploying models in: \", merged_model_output_dir)\n",
"\n",
"# Find Vertex AI prediction supported accelerators and regions in [here](https://cloud.google.com/vertex-ai/docs/predictions/configure-compute).\n",
"# Sets 1 L4 (24G) to deploy Gemma models.\n",
"serve_machine_type = \"g2-standard-12\"\n",
"serve_accelerator_type = \"NVIDIA_L4\"\n",
"serve_accelerator_count = 1\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=serve_accelerator_type,\n",
" accelerator_count=serve_accelerator_count,\n",
" is_for_training=False,\n",
")\n",
"\n",
"# Note that a larger max_model_len will require more GPU memory.\n",
"max_model_len = 2048\n",
"\n",
"models[\"vllm_gpu\"], endpoints[\"vllm_gpu\"] = deploy_model_vllm(\n",
" model_name=common_util.get_job_name_with_datetime(prefix=\"gemma-vllm-serve\"),\n",
" base_model_id=f\"google/{base_model_id}\",\n",
" publisher=\"google\",\n",
" publisher_model_id=\"gemma\",\n",
" model_id=merged_model_output_dir,\n",
" service_account=SERVICE_ACCOUNT,\n",
" machine_type=serve_machine_type,\n",
" accelerator_type=serve_accelerator_type,\n",
" accelerator_count=serve_accelerator_count,\n",
" max_model_len=max_model_len,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"print(\"endpoint_name:\", endpoints[\"vllm_gpu\"].name)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "2UYUNn60G_4U"
},
"outputs": [],
"source": [
"# @title Predict\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint with text prompts.\n",
"\n",
"# @markdown Here we use an example from the [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) to show the finetuning outcome:\n",
"\n",
"# @markdown ```\n",
"# @markdown ### Human: How would the Future of AI in 10 Years look?### Assistant: Predicting the future is always a challenging task, but here are some possible ways that AI could evolve over the next 10 years: Continued advancements in deep learning: Deep learning has been one of the main drivers of recent AI breakthroughs, and we can expect continued advancements in this area. This may include improvements to existing algorithms, as well as the development of new architectures that are better suited to specific types of data and tasks. Increased use of AI in healthcare: AI has the potential to revolutionize healthcare, by improving the accuracy of diagnoses, developing new treatments, and personalizing patient care. We can expect to see continued investment in this area, with more healthcare providers and researchers using AI to improve patient outcomes. Greater automation in the workplace: Automation is already transforming many industries, and AI is likely to play an increasingly important role in this process. We can expect to see more jobs being automated, as well as the development of new types of jobs that require a combination of human and machine skills. More natural and intuitive interactions with technology: As AI becomes more advanced, we can expect to see more natural and intuitive ways of interacting with technology. This may include voice and gesture recognition, as well as more sophisticated chatbots and virtual assistants. Increased focus on ethical considerations: As AI becomes more powerful, there will be a growing need to consider its ethical implications. This may include issues such as bias in AI algorithms, the impact of automation on employment, and the use of AI in surveillance and policing. Overall, the future of AI in 10 years is likely to be shaped by a combination of technological advancements, societal changes, and ethical considerations. While there are many exciting possibilities for AI in the future, it will be important to carefully consider its potential impact on society and to work towards ensuring that its benefits are shared fairly and equitably.\n",
"# @markdown ```\n",
"\n",
"# @markdown Additionally, you can moderate the generated text with Vertex AI. See [Moderate text documentation](https://cloud.google.com/natural-language/docs/moderating-text) for more details.\n",
"\n",
"# Loads an existing endpoint instance using the endpoint name:\n",
"# - Using `endpoint_name = endpoint.name` allows us to get the\n",
"# endpoint name of the endpoint `endpoint` created in the cell\n",
"# above.\n",
"# - Alternatively, you can set `endpoint_name = \"1234567890123456789\"` to load\n",
"# an existing endpoint with the ID 1234567890123456789.\n",
"# You may uncomment the code below to load an existing endpoint.\n",
"\n",
"# endpoint_name = \"\" # @param {type:\"string\"}\n",
"# aip_endpoint_name = (\n",
"# f\"projects/{PROJECT_ID}/locations/{REGION}/endpoints/{endpoint_name}\"\n",
"# )\n",
"# endpoint = aiplatform.Endpoint(aip_endpoint_name)\n",
"\n",
"prompt = \"How would the Future of AI in 10 Years look?\" # @param {type: \"string\"}\n",
"max_tokens = 128 # @param {type:\"integer\"}\n",
"temperature = 1.0 # @param {type:\"number\"}\n",
"top_p = 0.9 # @param {type:\"number\"}\n",
"top_k = 1 # @param {type:\"integer\"}\n",
"\n",
"# Overrides max_tokens and top_k parameters during inferences.\n",
"# If you encounter the issue like `ServiceUnavailable: 503 Took too long to respond when processing`,\n",
"# you can reduce the max length, such as set max_tokens as 20.\n",
"instances = [\n",
" {\n",
" \"prompt\": f\"### Human: {prompt}### Assistant: \",\n",
" \"max_tokens\": max_tokens,\n",
" \"temperature\": temperature,\n",
" \"top_p\": top_p,\n",
" \"top_k\": top_k,\n",
" },\n",
"]\n",
"response = endpoints[\"vllm_gpu\"].predict(\n",
" instances=instances, use_dedicated_endpoint=use_dedicated_endpoint\n",
")\n",
"\n",
"for prediction in response.predictions:\n",
" print(prediction)"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "af21a3cff1e0"
},
"source": [
"## Clean up resources"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "911406c1561e"
},
"outputs": [],
"source": [
"# Delete the train job.\n",
"train_job.delete()\n",
"\n",
"# Delete the conversion job.\n",
"if conversion_job:\n",
" conversion_job.delete()\n",
"\n",
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continuous charges that may incur.\n",
"\n",
"# Undeploy model and delete endpoint.\n",
"for endpoint in endpoints.values():\n",
" endpoint.delete(force=True)\n",
"\n",
"# Delete models.\n",
"for model in models.values():\n",
" model.delete()\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_NAME"
]
}
],
"metadata": {
"colab": {
"name": "model_garden_gemma_finetuning_on_vertex.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
@@ -106,7 +106,7 @@
"# @markdown\n",
"# @markdown By default, the quota for TPU deployment `Custom model serving TPU v5e cores per region` is 4. This will be lifted to 16 in the future.\n",
"# @markdown Verify that you have the appropriate TPU quota for your chosen configuration (e.g., 1, 4, 8, or 16 cores) in the selected region.\n",
"# @markdown You can request for higher TPU quota following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota)."
"# @markdown You can request for higher TPU quota following the instructions at [\"Request a quota adjustment\"](https://cloud.google.com/docs/quotas/view-manage#requesting_higher_quota)."
]
},
{
@@ -104,7 +104,7 @@
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 4. If you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus). You can request for quota following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota).\n",
"# @markdown 4. If you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus). You can request for quota following the instructions at [\"Request a quota adjustment\"](https://cloud.google.com/docs/quotas/view-manage#requesting_higher_quota).\n",
"\n",
"# @markdown > | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
@@ -1,374 +1,373 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "7d9bbf86da5e"
},
"outputs": [],
"source": [
"# Copyright 2023 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",
"# https://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": "99c1c3fc2ca5"
},
"source": [
"# Vertex AI Model Garden - Falcon Evaluation\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_pytorch_falcon_evaluation.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_pytorch_falcon_evaluation.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "3de7470326a2"
},
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates evaluating a pre-trained or a PEFT-finetuned Falcon Instruct models in Vertex AI.\n",
"\n",
"### Objective\n",
"\n",
"- Evaluate a pre-trained or a PEFT-finetuned Falcon model on any of the benchmark datasets\n",
"- Clean up the resources\n",
"\n",
"| Models |\n",
"| :- |\n",
"| [tiiuae/falcon-7b-instruct](https://huggingface.co/tiiuae/falcon-7b-instruct)\n",
"| [tiiuae/falcon-40b-instruct](https://huggingface.co/tiiuae/falcon-40b-instruct)\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing) and [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "HsAZ1ozfRQt7"
},
"source": [
"## Run the notebook"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "855d6b96f291"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. [Optional] [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"# Import the necessary packages\n",
"import os\n",
"from datetime import datetime\n",
"\n",
"from google.cloud import aiplatform\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, please change the value yourself below.\n",
"now = datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" # Create a unique GCS bucket for this notebook, if not specified by the user\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}\"\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"EXPERIMENT_BUCKET = os.path.join(BUCKET_URI, \"peft\")\n",
"MODEL_BUCKET = os.path.join(EXPERIMENT_BUCKET, \"model\")\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default BUCKET_URI and SERVICE_ACCOUNT if they were not specified by the user.\n",
"SERVICE_ACCOUNT = None\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"\n",
"\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# The evaluation docker image.\n",
"EVAL_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-lm-evaluation-harness:20231011_0934_RC00\"\n",
"\n",
"# Define common functions\n",
"\n",
"\n",
"def get_job_name_with_datetime(prefix: str) -> str:\n",
" \"\"\"Gets the job name with date time when triggering training or deployment\n",
" jobs in Vertex AI.\n",
" \"\"\"\n",
" return prefix + datetime.now().strftime(\"_%Y%m%d_%H%M%S\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "g0t0RBixIw0P"
},
"outputs": [],
"source": [
"# @title Evaluate PEFT-finetuned Falcon Instruct models\n",
"\n",
"# @markdown This section demonstrates how to evaluate the Falcon Instruct models fintuned with PEFT LoRA using EleutherAI's [Language Model Evaluation Harness (lm-evaluation-harness)](https://github.com/EleutherAI/lm-evaluation-harness) with Vertex CustomJob. Please reference the peak GPU memory usage for serving and adjust the machine type, accelerator type and accelerator count accordingly.\n",
"\n",
"# @markdown This example uses the dataset [TruthfulQA](https://arxiv.org/abs/2109.07958). All supported tasks are listed in [this task table](https://github.com/EleutherAI/lm-evaluation-harness/blob/master/docs/task_table.md).\n",
"# @markdown Set evaluation dataset.\n",
"eval_dataset = \"truthfulqa_mc\" # @param {type:\"string\"}\n",
"\n",
"# Worker pool spec.\n",
"# Find Vertex AI supported accelerators and regions in:\n",
"# https://cloud.google.com/vertex-ai/docs/training/configure-compute\n",
"\n",
"\n",
"# Setup evaluation job.\n",
"# @markdown Set the base model id.\n",
"base_model_id = \"tiiuae/falcon-7b-instruct\" # @param [\"tiiuae/falcon-7b-instruct\", \"tiiuae/falcon-40b-instruct\"]\n",
"job_name = get_job_name_with_datetime(prefix=\"falcon-instruct-peft-eval\")\n",
"eval_output_dir = os.path.join(MODEL_BUCKET, job_name)\n",
"eval_output_dir_gcsfuse = eval_output_dir.replace(\"gs://\", \"/gcs/\")\n",
"\n",
"# @markdown Sets V100 (16G) to evaluate `tiiuae/falcon-7b-instruct` or `tiiuae/falcon-40b-instruct`.\n",
"# @markdown If A100 is not available, you may evaluate tiiuae/falcon-40b-instruct with\n",
"# @markdown multiple V100s. Please keep in mind that the efficiency of evaluating with\n",
"# @markdown multiple V100s is inferior to that of evaluating with A100s.\n",
"\n",
"# @markdown Set the accelerator type.\n",
"accelerator_type = \"NVIDIA_TESLA_V100\" # @param[\"NVIDIA_TESLA_V100\", \"NVIDIA_L4\", \"NVIDIA_TESLA_A100\", \"NVIDIA_TESLA_A100_80G\"]\n",
"\n",
"\n",
"# @markdown To evaluate a PEFT-finetuned model, enter the PEFT output directory below.\n",
"# @markdown Otherwise, leave it empty.\n",
"# @markdown See the finetuning notebook for more details: https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_pytorch_llama2_peft_finetuning.ipynb\n",
"peft_output_dir = \"\" # @param {type:\"string\"}\n",
"peft_output_dir_gcsfuse = peft_output_dir.replace(\"gs://\", \"/gcs/\")\n",
"\n",
"if \"7b\" in base_model_id:\n",
" # For models containing '7b', set configurations based on the accelerator type provided.\n",
" if accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 1\n",
" elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-standard-8\"\n",
" accelerator_count = 1\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-8\"\n",
" accelerator_count = 1\n",
" else:\n",
" print(f\"Unsupported accelerator type: {accelerator_type}\")\n",
"elif \"40b\" in base_model_id:\n",
" # For models containing '40b', set configurations based on the accelerator type provided.\n",
" if accelerator_type == \"NVIDIA_TESLA_A100_80GB\":\n",
" machine_type = \"a2-ultragpu-1g\"\n",
" accelerator_count = 2\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-48\"\n",
" accelerator_count = 4\n",
" elif (\n",
" accelerator_type == \"NVIDIA_TESLA_V100\"\n",
" ): # Assuming V100 can be used as a fallback for 40b models\n",
" machine_type = \"n1-standard-8\"\n",
" accelerator_count = 8\n",
" else:\n",
" print(f\"Unsupported accelerator type: {accelerator_type}\")\n",
"else:\n",
" print(\"The base_model_id does not specify a recognized model version.\")\n",
"\n",
"replica_count = 1\n",
"\n",
"\n",
"# Prepare evaluation command that runs the evaluation harness.\n",
"# Set `trust_remote_code = True` because evaluating the model requires\n",
"# executing code from the model repository.\n",
"# Set `use_accelerate = True` to enable evaluation across multiple GPUs.\n",
"eval_command = [\n",
" \"python\",\n",
" \"main.py\",\n",
" \"--model\",\n",
" \"hf-causal-experimental\",\n",
" \"--tasks\",\n",
" f\"{eval_dataset}\",\n",
" \"--output_path\",\n",
" f\"{eval_output_dir_gcsfuse}\",\n",
"]\n",
"\n",
"if peft_output_dir_gcsfuse:\n",
" eval_command += [\n",
" \"--model_args\",\n",
" f\"pretrained={base_model_id},peft={peft_output_dir_gcsfuse},trust_remote_code=True,use_accelerate=True,device_map_option=auto\",\n",
" ]\n",
"else:\n",
" eval_command += [\n",
" \"--model_args\",\n",
" f\"pretrained={base_model_id},trust_remote_code=True,use_accelerate=True,device_map_option=auto\",\n",
" ]\n",
"\n",
"\n",
"# Pass evaluation arguments and launch job.\n",
"worker_pool_specs = [\n",
" {\n",
" \"machine_spec\": {\n",
" \"machine_type\": machine_type,\n",
" \"accelerator_type\": accelerator_type,\n",
" \"accelerator_count\": accelerator_count,\n",
" },\n",
" \"replica_count\": replica_count,\n",
" \"disk_spec\": {\n",
" \"boot_disk_size_gb\": 500,\n",
" },\n",
" \"container_spec\": {\n",
" \"image_uri\": EVAL_DOCKER_URI,\n",
" \"command\": eval_command,\n",
" \"args\": [],\n",
" },\n",
" }\n",
"]\n",
"\n",
"# Submit evaluation custom job.\n",
"eval_job = aiplatform.CustomJob(\n",
" display_name=job_name,\n",
" worker_pool_specs=worker_pool_specs,\n",
" base_output_dir=eval_output_dir,\n",
")\n",
"\n",
"eval_job.run()\n",
"\n",
"print(\"Evaluation results were saved in:\", eval_output_dir)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "1f15ed6d375a"
},
"outputs": [],
"source": [
"# @title Fetch and print evaluation results\n",
"import json\n",
"\n",
"from google.cloud import storage\n",
"\n",
"# Fetch evaluation results.\n",
"storage_client = storage.Client()\n",
"BUCKET_NAME = BUCKET_URI.split(\"gs://\")[1]\n",
"bucket = storage_client.get_bucket(BUCKET_NAME)\n",
"RESULT_FILE_PATH = eval_output_dir[len(BUCKET_URI) + 1 :]\n",
"blob = bucket.blob(RESULT_FILE_PATH)\n",
"raw_result = blob.download_as_string()\n",
"\n",
"# Print evaluation results.\n",
"result = json.loads(raw_result)\n",
"result_formatted = json.dumps(result, indent=2)\n",
"print(f\"Evaluation result:\\n{result_formatted}\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "911406c1561e"
},
"outputs": [],
"source": [
"# @title Clean up resources\n",
"# Delete evaluation job.\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_URI\n",
"\n",
"eval_job.delete()"
]
}
],
"metadata": {
"colab": {
"name": "model_garden_pytorch_falcon_evaluation.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
"cells": [
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "7d9bbf86da5e"
},
"outputs": [],
"source": [
"# Copyright 2023 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",
"# https://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."
]
},
"nbformat": 4,
"nbformat_minor": 0
{
"cell_type": "markdown",
"metadata": {
"id": "99c1c3fc2ca5"
},
"source": [
"# Vertex AI Model Garden - Falcon Evaluation\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_pytorch_falcon_evaluation.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_pytorch_falcon_evaluation.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "3de7470326a2"
},
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates evaluating a pre-trained or a PEFT-finetuned Falcon Instruct models in Vertex AI.\n",
"\n",
"### Objective\n",
"\n",
"- Evaluate a pre-trained or a PEFT-finetuned Falcon model on any of the benchmark datasets\n",
"- Clean up the resources\n",
"\n",
"| Models |\n",
"| :- |\n",
"| [tiiuae/falcon-7b-instruct](https://huggingface.co/tiiuae/falcon-7b-instruct)\n",
"| [tiiuae/falcon-40b-instruct](https://huggingface.co/tiiuae/falcon-40b-instruct)\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing) and [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "HsAZ1ozfRQt7"
},
"source": [
"## Run the notebook"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "855d6b96f291"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. [Optional] [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"# Import the necessary packages\n",
"import os\n",
"from datetime import datetime\n",
"\n",
"from google.cloud import aiplatform\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, please change the value yourself below.\n",
"now = datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" # Create a unique GCS bucket for this notebook, if not specified by the user\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}\"\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"EXPERIMENT_BUCKET = os.path.join(BUCKET_URI, \"peft\")\n",
"MODEL_BUCKET = os.path.join(EXPERIMENT_BUCKET, \"model\")\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default BUCKET_URI and SERVICE_ACCOUNT if they were not specified by the user.\n",
"SERVICE_ACCOUNT = None\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"\n",
"\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# The evaluation docker image.\n",
"EVAL_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-lm-evaluation-harness:20231011_0934_RC00\"\n",
"\n",
"# Define common functions\n",
"\n",
"\n",
"def get_job_name_with_datetime(prefix: str) -> str:\n",
" \"\"\"Gets the job name with date time when triggering training or deployment\n",
" jobs in Vertex AI.\n",
" \"\"\"\n",
" return prefix + datetime.now().strftime(\"_%Y%m%d_%H%M%S\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "g0t0RBixIw0P"
},
"outputs": [],
"source": [
"# @title Evaluate PEFT-finetuned Falcon Instruct models\n",
"\n",
"# @markdown This section demonstrates how to evaluate the Falcon Instruct models fintuned with PEFT LoRA using EleutherAI's [Language Model Evaluation Harness (lm-evaluation-harness)](https://github.com/EleutherAI/lm-evaluation-harness) with Vertex CustomJob. Please reference the peak GPU memory usage for serving and adjust the machine type, accelerator type and accelerator count accordingly.\n",
"\n",
"# @markdown This example uses the dataset [TruthfulQA](https://arxiv.org/abs/2109.07958). All supported tasks are listed in [this task table](https://github.com/EleutherAI/lm-evaluation-harness/blob/master/docs/task_table.md).\n",
"# @markdown Set evaluation dataset.\n",
"eval_dataset = \"truthfulqa_mc\" # @param {type:\"string\"}\n",
"\n",
"# Worker pool spec.\n",
"# Find Vertex AI supported accelerators and regions in:\n",
"# https://cloud.google.com/vertex-ai/docs/training/configure-compute\n",
"\n",
"\n",
"# Setup evaluation job.\n",
"# @markdown Set the base model id.\n",
"base_model_id = \"tiiuae/falcon-7b-instruct\" # @param [\"tiiuae/falcon-7b-instruct\", \"tiiuae/falcon-40b-instruct\"]\n",
"job_name = get_job_name_with_datetime(prefix=\"falcon-instruct-peft-eval\")\n",
"eval_output_dir = os.path.join(MODEL_BUCKET, job_name)\n",
"eval_output_dir_gcsfuse = eval_output_dir.replace(\"gs://\", \"/gcs/\")\n",
"\n",
"# @markdown Sets V100 (16G) to evaluate `tiiuae/falcon-7b-instruct` or `tiiuae/falcon-40b-instruct`.\n",
"# @markdown If A100 is not available, you may evaluate tiiuae/falcon-40b-instruct with\n",
"# @markdown multiple V100s. Please keep in mind that the efficiency of evaluating with\n",
"# @markdown multiple V100s is inferior to that of evaluating with A100s.\n",
"\n",
"# @markdown Set the accelerator type.\n",
"accelerator_type = \"NVIDIA_TESLA_V100\" # @param[\"NVIDIA_TESLA_V100\", \"NVIDIA_L4\", \"NVIDIA_TESLA_A100\", \"NVIDIA_TESLA_A100_80G\"]\n",
"\n",
"\n",
"# @markdown To evaluate a PEFT-finetuned model, enter the PEFT output directory below.\n",
"# @markdown Otherwise, leave it empty.\n",
"peft_output_dir = \"\" # @param {type:\"string\"}\n",
"peft_output_dir_gcsfuse = peft_output_dir.replace(\"gs://\", \"/gcs/\")\n",
"\n",
"if \"7b\" in base_model_id:\n",
" # For models containing '7b', set configurations based on the accelerator type provided.\n",
" if accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 1\n",
" elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-standard-8\"\n",
" accelerator_count = 1\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-8\"\n",
" accelerator_count = 1\n",
" else:\n",
" print(f\"Unsupported accelerator type: {accelerator_type}\")\n",
"elif \"40b\" in base_model_id:\n",
" # For models containing '40b', set configurations based on the accelerator type provided.\n",
" if accelerator_type == \"NVIDIA_TESLA_A100_80GB\":\n",
" machine_type = \"a2-ultragpu-1g\"\n",
" accelerator_count = 2\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-48\"\n",
" accelerator_count = 4\n",
" elif (\n",
" accelerator_type == \"NVIDIA_TESLA_V100\"\n",
" ): # Assuming V100 can be used as a fallback for 40b models\n",
" machine_type = \"n1-standard-8\"\n",
" accelerator_count = 8\n",
" else:\n",
" print(f\"Unsupported accelerator type: {accelerator_type}\")\n",
"else:\n",
" print(\"The base_model_id does not specify a recognized model version.\")\n",
"\n",
"replica_count = 1\n",
"\n",
"\n",
"# Prepare evaluation command that runs the evaluation harness.\n",
"# Set `trust_remote_code = True` because evaluating the model requires\n",
"# executing code from the model repository.\n",
"# Set `use_accelerate = True` to enable evaluation across multiple GPUs.\n",
"eval_command = [\n",
" \"python\",\n",
" \"main.py\",\n",
" \"--model\",\n",
" \"hf-causal-experimental\",\n",
" \"--tasks\",\n",
" f\"{eval_dataset}\",\n",
" \"--output_path\",\n",
" f\"{eval_output_dir_gcsfuse}\",\n",
"]\n",
"\n",
"if peft_output_dir_gcsfuse:\n",
" eval_command += [\n",
" \"--model_args\",\n",
" f\"pretrained={base_model_id},peft={peft_output_dir_gcsfuse},trust_remote_code=True,use_accelerate=True,device_map_option=auto\",\n",
" ]\n",
"else:\n",
" eval_command += [\n",
" \"--model_args\",\n",
" f\"pretrained={base_model_id},trust_remote_code=True,use_accelerate=True,device_map_option=auto\",\n",
" ]\n",
"\n",
"\n",
"# Pass evaluation arguments and launch job.\n",
"worker_pool_specs = [\n",
" {\n",
" \"machine_spec\": {\n",
" \"machine_type\": machine_type,\n",
" \"accelerator_type\": accelerator_type,\n",
" \"accelerator_count\": accelerator_count,\n",
" },\n",
" \"replica_count\": replica_count,\n",
" \"disk_spec\": {\n",
" \"boot_disk_size_gb\": 500,\n",
" },\n",
" \"container_spec\": {\n",
" \"image_uri\": EVAL_DOCKER_URI,\n",
" \"command\": eval_command,\n",
" \"args\": [],\n",
" },\n",
" }\n",
"]\n",
"\n",
"# Submit evaluation custom job.\n",
"eval_job = aiplatform.CustomJob(\n",
" display_name=job_name,\n",
" worker_pool_specs=worker_pool_specs,\n",
" base_output_dir=eval_output_dir,\n",
")\n",
"\n",
"eval_job.run()\n",
"\n",
"print(\"Evaluation results were saved in:\", eval_output_dir)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "1f15ed6d375a"
},
"outputs": [],
"source": [
"# @title Fetch and print evaluation results\n",
"import json\n",
"\n",
"from google.cloud import storage\n",
"\n",
"# Fetch evaluation results.\n",
"storage_client = storage.Client()\n",
"BUCKET_NAME = BUCKET_URI.split(\"gs://\")[1]\n",
"bucket = storage_client.get_bucket(BUCKET_NAME)\n",
"RESULT_FILE_PATH = eval_output_dir[len(BUCKET_URI) + 1 :]\n",
"blob = bucket.blob(RESULT_FILE_PATH)\n",
"raw_result = blob.download_as_string()\n",
"\n",
"# Print evaluation results.\n",
"result = json.loads(raw_result)\n",
"result_formatted = json.dumps(result, indent=2)\n",
"print(f\"Evaluation result:\\n{result_formatted}\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "911406c1561e"
},
"outputs": [],
"source": [
"# @title Clean up resources\n",
"# Delete evaluation job.\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_URI\n",
"\n",
"eval_job.delete()"
]
}
],
"metadata": {
"colab": {
"name": "model_garden_pytorch_falcon_evaluation.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
@@ -1,605 +0,0 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "7d9bbf86da5e"
},
"outputs": [],
"source": [
"# Copyright 2024 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",
"# https://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": "99c1c3fc2ca5"
},
"source": [
"# Vertex AI Model Garden - Falcon Instruct (PEFT Finetuning)\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_pytorch_falcon_instruct_finetuning.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_pytorch_falcon_instruct_finetuning.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "3de7470326a2"
},
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates finetuning and deploying Falcon Instruct models with performance efficient finetuning libraries ([PEFT](https://github.com/huggingface/peft)) Falcon Instruct models in Vertex AI.\n",
"\n",
"### Objective\n",
"\n",
"- Finetune and deploy Falcon Instruct models with PEFT\n",
"- Cleanup the resources used\n",
"\n",
"| Models | LoRA |\n",
"| :- | :- |\n",
"| [tiiuae/falcon-7b-instruct](https://huggingface.co/tiiuae/falcon-7b-instruct) | Y |\n",
"| [tiiuae/falcon-40b-instruct](https://huggingface.co/tiiuae/falcon-40b-instruct) | Y |\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing) and [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "264c07757582"
},
"source": [
"## Run the notebook"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "855d6b96f291"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. [Optional] [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"# Import the necessary packages\n",
"import os\n",
"from datetime import datetime\n",
"from typing import Tuple\n",
"\n",
"from google.cloud import aiplatform\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, please change the value yourself below.\n",
"now = datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" # Create a unique GCS bucket for this notebook if not specified\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}\"\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"EXPERIMENT_BUCKET = os.path.join(BUCKET_URI, \"peft\")\n",
"DATA_BUCKET = os.path.join(EXPERIMENT_BUCKET, \"data\")\n",
"MODEL_BUCKET = os.path.join(EXPERIMENT_BUCKET, \"model\")\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default BUCKET_URI and SERVICE_ACCOUNT if they were not specified by the user.\n",
"SERVICE_ACCOUNT = None\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"# Create a unique GCS bucket for this notebook, if not specified by the user.\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}\"\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" shell_output = ! gsutil ls -Lb {BUCKET_URI} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"\n",
"# The pre-built training and serving docker images.\n",
"TRAIN_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-peft-train:20231222_0936_RC00\"\n",
"PREDICTION_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-peft-serve:20231129_0948_RC00\"\n",
"VLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-vllm-serve:20240410_0916_RC00\"\n",
"\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"\n",
"def create_name_with_datetime(prefix: str) -> str:\n",
" \"\"\"Creates a name with date time when triggering training or deployment\n",
" jobs in Vertex AI.\n",
" \"\"\"\n",
" return prefix + datetime.now().strftime(\"_%Y%m%d_%H%M%S\")\n",
"\n",
"\n",
"def deploy_model(\n",
" model_name: str,\n",
" base_model_id: str,\n",
" finetuned_lora_model_path: str,\n",
" service_account: str,\n",
" task: str,\n",
" machine_type: str = \"n1-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_TESLA_V100\",\n",
" accelerator_count: int = 1,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys trained models into Vertex AI.\"\"\"\n",
" endpoint = aiplatform.Endpoint.create(display_name=f\"{model_name}-endpoint\")\n",
" serving_env = {\n",
" \"BASE_MODEL_ID\": base_model_id,\n",
" \"TASK\": task,\n",
" \"DEPLOY_SOURCE\": \"notebook\",\n",
" }\n",
" if finetuned_lora_model_path:\n",
" serving_env[\"FINETUNED_LORA_MODEL_PATH\"] = finetuned_lora_model_path\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=PREDICTION_DOCKER_URI,\n",
" serving_container_ports=[7080],\n",
" serving_container_predict_route=\"/predictions/peft_serving\",\n",
" serving_container_health_route=\"/ping\",\n",
" serving_container_environment_variables=serving_env,\n",
" model_garden_source_model_name=\"publishers/tiiuae/models/falcon-instruct-7b-peft\"\n",
" )\n",
" model.deploy(\n",
" endpoint=endpoint,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" deploy_request_timeout=1800,\n",
" service_account=service_account,\n",
" system_labels={\n",
" \"NOTEBOOK_NAME\": \"model_garden_pytorch_falcon_instruct_finetuning.ipynb\"\n",
" },\n",
" )\n",
" return model, endpoint\n",
"\n",
"\n",
"def deploy_model_vllm(\n",
" model_name: str,\n",
" model_id: str,\n",
" service_account: str,\n",
" machine_type: str = \"n1-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_TESLA_V100\",\n",
" accelerator_count: int = 1,\n",
" quantization_method: str = \"\",\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys trained models with vLLM into Vertex AI.\"\"\"\n",
" endpoint = aiplatform.Endpoint.create(display_name=f\"{model_name}-endpoint\")\n",
"\n",
" vllm_args = [\n",
" \"--host=0.0.0.0\",\n",
" \"--port=7080\",\n",
" f\"--model={model_id}\",\n",
" f\"--tensor-parallel-size={accelerator_count}\",\n",
" \"--swap-space=16\",\n",
" \"--gpu-memory-utilization=0.9\",\n",
" \"--disable-log-stats\",\n",
" \"--dtype=float16\",\n",
" \"--trust-remote-code\",\n",
" ]\n",
" if quantization_method:\n",
" vllm_args.append(f\"--quantization={quantization_method}\")\n",
"\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=VLLM_DOCKER_URI,\n",
" serving_container_command=[\"python\", \"-m\", \"vllm.entrypoints.api_server\"],\n",
" serving_container_args=vllm_args,\n",
" serving_container_ports=[7080],\n",
" serving_container_predict_route=\"/generate\",\n",
" serving_container_health_route=\"/ping\",\n",
" model_garden_source_model_name=\"publishers/tiiuae/models/falcon-instruct-7b-peft\"\n",
" )\n",
"\n",
" model.deploy(\n",
" endpoint=endpoint,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" deploy_request_timeout=1800,\n",
" service_account=service_account,\n",
" )\n",
" return model, endpoint"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "65467b361315"
},
"outputs": [],
"source": [
"# @title Finetune and deploy Falcon Instruct models with PEFT\n",
"\n",
"# @markdown This section demonstrates how to finetune and deploy Falcon Instruct models with PEFT LoRA.\n",
"\n",
"# @markdown The peak GPU memory usages are ~11G and ~34G for finetuning LoRA models for [tiiuae/falcon-7b-instruct](https://huggingface.co/tiiuae/falcon-7b-instruct), and [tiiuae/falcon-40b-instruct](https://huggingface.co/tiiuae/falcon-40b-instruct) separately with default training parameters and the example dataset. Falcon-7b-instruct can be finetuned on 1 P100/V100 and falcon-40b-instruct can be finetuned on 1 A100 (40G).\n",
"\n",
"# @markdown This example uses the dataset [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco). You can either use a [dataset from huggingface](https://huggingface.co/datasets) or a custom JSONL dataset in [Vertex text model dataset format](https://cloud.google.com/vertex-ai/docs/generative-ai/models/tune-text-models-supervised#dataset-format) stored in Cloud Storage. The `template` parameter is optional.\n",
"\n",
"# @markdown To use a custom dataset, you should supply a `gs://` URI to a JSONL file in [Vertex text model dataset format](https://cloud.google.com/vertex-ai/docs/generative-ai/models/tune-text-models-supervised#dataset-format) in the `dataset_name` below.\n",
"\n",
"# @markdown For example, here is one data point from the sample dataset `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`:\n",
"\n",
"# @markdown ```json\n",
"# @markdown {\"input_text\":\"TRANSCRIPT: \\nREASON FOR EVALUATION:,\\n\\n LABEL:\",\"output_text\":\"Chiropractic\"}\n",
"# @markdown ```\n",
"\n",
"# @markdown To use this sample dataset that contains `input_text` and `output_text` fields, set `dataset_name` to `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl` and `template` to `vertex_sample`. For advanced usage with custom datatset fields, see [the template example](https://github.com/tloen/alpaca-lora/blob/main/templates/alpaca.json) and supply your own JSON template as `gs://` URIs.\n",
"\n",
"# @markdown Set the base model id.\n",
"base_model_id = \"tiiuae/falcon-7b-instruct\" # @param [\"tiiuae/falcon-7b-instruct\", \"tiiuae/falcon-40b-instruct\"]\n",
"\n",
"# @markdown Set the accelerator type.\n",
"accelerator_type = \"NVIDIA_TESLA_V100\" # @param[\"NVIDIA_TESLA_V100\", \"NVIDIA_L4\", \"NVIDIA_TESLA_A100\", \"NVIDIA_TESLA_A100_80G\"]\n",
"\n",
"# Huggingface dataset name or gs:// URI to a custom JSONL dataset.\n",
"dataset_name = \"timdettmers/openassistant-guanaco\" # @param {type:\"string\"}\n",
"# Optional. Template name or gs:// URI to a custom template.\n",
"template = \"\" # @param {type:\"string\"}\n",
"\n",
"# @markdown Set the number of steps in the finetuning job.\n",
"max_steps = 10 # @param {type:\"integer\"}\n",
"\n",
"# Worker pool spec.\n",
"# Find Vertex AI supported accelerators and regions in:\n",
"# https://cloud.google.com/vertex-ai/docs/training/configure-compute\n",
"\n",
"if \"7b\" in base_model_id:\n",
" # Uses V100 (16G) to finetune falcon-7b-instruct.\n",
" if accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-standard-8\"\n",
" accelerator_count = 2\n",
" # Uses L4 (24G) to finetune falcon-7b-instruct.\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-24\"\n",
" accelerator_count = 2\n",
" else:\n",
" raise ValueError(\n",
" f\"Recommended GPU setting not found for: {accelerator_type} and {base_model_name}.\"\n",
" )\n",
"elif \"40b\" in base_model_id:\n",
" if accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 1\n",
" elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-standard-16\"\n",
" accelerator_count = 4\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-24\"\n",
" accelerator_count = 2\n",
" else:\n",
" raise ValueError(\n",
" f\"Recommended GPU setting not found for: {accelerator_type} and {base_model_name}.\"\n",
" )\n",
"\n",
"replica_count = 1\n",
"\n",
"# Add labels for the finetuning job.\n",
"labels = {\n",
" \"mg-source\": \"notebook\",\n",
" \"mg-notebook-name\": (\n",
" \"model_garden_pytorch_falcon_instruct_finetuning.ipynb\".split(\".\")[0]\n",
" ),\n",
"}\n",
"\n",
"labels[\"mg-tune\"] = \"publishers/tiiuae/models/falcon\"\n",
"versioned_model_id = base_model_id.split(\"/\")[1].replace(\"_\", \"-\")\n",
"labels[\"versioned-mg-tune\"] = f\"{labels['mg-tune']}-{versioned_model_id}\"\n",
"\n",
"# Setup training job.\n",
"job_name = create_name_with_datetime(\"falcon-finetune-train\")\n",
"train_job = aiplatform.CustomContainerTrainingJob(\n",
" display_name=job_name,\n",
" container_uri=TRAIN_DOCKER_URI,\n",
" labels=labels,\n",
")\n",
"\n",
"# Create a GCS folder to store the LORA adapter.\n",
"finetune_dir = create_name_with_datetime(\"falcon-finetune\")\n",
"finetune_output_dir = os.path.join(MODEL_BUCKET, finetune_dir)\n",
"finetune_output_dir_gcsfuse = finetune_output_dir.replace(\"gs://\", \"/gcs/\")\n",
"\n",
"# Create a GCS folder to store the merged model with the base model and the\n",
"# finetuned LORA adapter.\n",
"merged_model_dir = create_name_with_datetime(\"falcon-merged-model\")\n",
"merged_model_output_dir = os.path.join(MODEL_BUCKET, merged_model_dir)\n",
"merged_model_output_dir_gcsfuse = merged_model_output_dir.replace(\"gs://\", \"/gcs/\")\n",
"\n",
"# Pass training arguments and launch job.\n",
"train_job.run(\n",
" args=[\n",
" \"--task=instruct-lora\",\n",
" f\"--pretrained_model_id={base_model_id}\",\n",
" f\"--dataset_name={dataset_name}\",\n",
" f\"--output_dir={finetune_output_dir_gcsfuse}\",\n",
" f\"--merge_base_and_lora_output_dir={merged_model_output_dir_gcsfuse}\",\n",
" \"--lora_rank=16\",\n",
" \"--lora_alpha=32\",\n",
" \"--lora_dropout=0.05\",\n",
" \"--warmup_steps=10\",\n",
" f\"--max_steps={max_steps}\",\n",
" \"--learning_rate=2e-4\",\n",
" f\"--template={template}\",\n",
" \"--per_device_train_batch_size=1\",\n",
" ],\n",
" replica_count=replica_count,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" boot_disk_size_gb=500,\n",
")\n",
"\n",
"print(\"The finetuned model can be found at: \", finetune_output_dir)\n",
"print(\n",
" \"The finetuned model merged with the base model can be found at: \",\n",
" merged_model_output_dir,\n",
")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "bf55e38815dc"
},
"outputs": [],
"source": [
"# @title Deploy to endpoint\n",
"\n",
"# @markdown This section uploads the model to Model Registry and deploys it on the Endpoint.\n",
"\n",
"# @markdown The model deployment step will take 15 minutes to 40 minutes to complete.\n",
"\n",
"# @markdown The peak GPU memory usages for [tiiuae/falcon-7b-instruct](https://huggingface.co/tiiuae/falcon-7b-instruct), and [tiiuae/falcon-40b-instruct](https://huggingface.co/tiiuae/falcon-40b-instruct) with LoRA weights are ~15.5G and ~84G separately with the default settings. Please adjust the machine type, accelerator type and accelerator count accordingly. We use V100 in deployments as an example. Note that V100 serving generally offers better throughput and latency performance than L4 serving, while L4 serving is generally more cost efficient than V100 serving. The serving efficiency of V100 and L4 GPUs is inferior to that of A100 GPUs, but V100 and L4 GPUs are nevertheless good serving solutions if you do not have A100 quota.\n",
"\n",
"# Find Vertex AI supported accelerators and regions in:\n",
"# https://cloud.google.com/vertex-ai/docs/predictions/configure-compute\n",
"\n",
"\n",
"# @markdown Set the base model id.\n",
"base_model_id = \"tiiuae/falcon-7b-instruct\" # @param [\"tiiuae/falcon-7b-instruct\", \"tiiuae/falcon-40b-instruct\"]\n",
"\n",
"# @markdown Set the accelerator type.\n",
"accelerator_type = \"NVIDIA_TESLA_V100\" # @param[\"NVIDIA_TESLA_V100\", \"NVIDIA_L4\", \"NVIDIA_TESLA_A100\", \"NVIDIA_TESLA_A100_80G\"]\n",
"\n",
"\n",
"if \"7b\" in base_model_id:\n",
" if accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 1\n",
" if accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-standard-8\"\n",
" accelerator_count = 1\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-8\"\n",
" accelerator_count = 1\n",
" else:\n",
" raise ValueError(\n",
" f\"Recommended GPU setting not found for: {accelerator_type} and {base_model_id}.\"\n",
" )\n",
"elif \"40b\" in base_model_id:\n",
" if accelerator_type == \"NVIDIA_TESLA_A100_80GB\":\n",
" machine_type = \"a2-ultragpu-1g\"\n",
" accelerator_count = 2\n",
" elif accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 4\n",
" elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-standard-8\"\n",
" accelerator_count = 8\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-48\"\n",
" accelerator_count = 4\n",
" else:\n",
" raise ValueError(\n",
" f\"Recommended GPU setting not found for: {accelerator_type} and {base_model_id}.\"\n",
" )\n",
"\n",
"if base_model_id == \"tiiuae/falcon-7b-instruct\":\n",
" model, endpoint = deploy_model(\n",
" model_name=create_name_with_datetime(prefix=\"falcon-instruct-serve\"),\n",
" base_model_id=base_model_id,\n",
" finetuned_lora_model_path=os.path.join(\n",
" finetune_output_dir, f\"checkpoint-{max_steps}\"\n",
" ), # This will avoid override finetuning models.\n",
" service_account=SERVICE_ACCOUNT,\n",
" task=\"instruct-lora\",\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" )\n",
"else:\n",
" model, endpoint = deploy_model_vllm(\n",
" model_name=create_name_with_datetime(prefix=\"falcon-instruct-vllm\"),\n",
" model_id=merged_model_output_dir,\n",
" service_account=SERVICE_ACCOUNT,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" )\n",
"\n",
"print(\"endpoint_name:\", endpoint.name)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "4ab04da3ec9a"
},
"outputs": [],
"source": [
"# @markdown NOTE: After the deployment succeeds, the base model weights will be downloaded on the fly from the original location and LoRA model weights will be downloaded from the GCS bucket used in training above. Thus, an additional 10-30 minutes of waiting time is needed **after** the above model deployment step succeeds and before you can run the next step below. Otherwise you might see a `ServiceUnavailable: 503 502:Bad Gateway` error when you send requests to the endpoint.\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint with text prompts.\n",
"\n",
"# @markdown Example:\n",
"\n",
"# @markdown ```\n",
"# @markdown Human: What is a car?\n",
"# @markdown Assistant: A car, or a motor car, is a road-connected human-transportation system used to move people or goods from one place to another. The term also encompasses a wide range of vehicles, including motorboats, trains, and aircrafts. Cars typically have four wheels, a cabin for passengers, and an engine or motor. They have been around since the early 19th century and are now one of the most popular forms of transportation, used for daily commuting, shopping, and other purposes.\n",
"# @markdown ```\n",
"\n",
"# @markdown Additionally, you can moderate the generated text with Vertex AI. See [Moderate text documentation](https://cloud.google.com/natural-language/docs/moderating-text) for more details.\n",
"\n",
"# Loads an existing endpoint instance using the endpoint name:\n",
"# - Using `endpoint_name = endpoint.name` allows us to get the\n",
"# endpoint name of the endpoint `endpoint` created in the cell\n",
"# above.\n",
"# - Alternatively, you can set `endpoint_name = \"1234567890123456789\"` to load\n",
"# an existing endpoint with the ID 1234567890123456789.\n",
"# You may uncomment the code below to load an existing endpoint.\n",
"\n",
"# endpoint_name = endpoint.name\n",
"# # endpoint_name = \"\" # @param {type:\"string\"}\n",
"# aip_endpoint_name = (\n",
"# f\"projects/{PROJECT_ID}/locations/{REGION}/endpoints/{endpoint_name}\"\n",
"# )\n",
"# endpoint = aiplatform.Endpoint(aip_endpoint_name)\n",
"\n",
"prompt = \"What is a car?\" # @param {type:\"string\"}\n",
"max_tokens = 50 # @param {type:\"integer\"}\n",
"temperature = 1.0 # @param {type:\"number\"}\n",
"top_p = 1.0 # @param {type:\"number\"}\n",
"top_k = 10 # @param {type:\"number\"}\n",
"\n",
"\n",
"instances = [\n",
" {\n",
" \"prompt\": prompt,\n",
" \"max_tokens\": max_tokens,\n",
" \"temperature\": temperature,\n",
" \"top_p\": top_p,\n",
" \"top_k\": top_k,\n",
" },\n",
"]\n",
"response = endpoint.predict(instances=instances)\n",
"\n",
"for prediction in response.predictions:\n",
" print(prediction)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "911406c1561e"
},
"outputs": [],
"source": [
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continouous charges that may incur.\n",
"\n",
"# @title Clean up resources\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_URI\n",
"\n",
"# Delete custom train, quantization, and evaluation jobs.\n",
"train_job.delete()\n",
"\n",
"# Undeploy models and delete endpoints.\n",
"endpoint.delete(force=True)\n",
"\n",
"# Delete models.\n",
"model.delete()"
]
}
],
"metadata": {
"colab": {
"name": "model_garden_pytorch_falcon_instruct_finetuning.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
@@ -1,966 +0,0 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "7d9bbf86da5e"
},
"outputs": [],
"source": [
"# Copyright 2026 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",
"# https://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": "99c1c3fc2ca5"
},
"source": [
"# Vertex AI Model Garden - Gemma Finetuning (PEFT + vLLM)\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_pytorch_gemma_peft_finetuning_hf.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_pytorch_gemma_peft_finetuning_hf.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "3de7470326a2"
},
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates finetuning and deploying Gemma models with [Vertex AI Custom Training Job](https://cloud.google.com/vertex-ai/docs/training/create-custom-job). Using Vertex AI Pipelines is the quickest way to start finetuning Gemma models, while using a Vertex AI Custom Training Job allows for a higher level of customization and control over the finetuning job. All of the examples in this notebook use parameter efficient finetuning methods [PEFT](https://github.com/huggingface/peft) to reduce training and storage costs.\n",
"\n",
"This notebook deploys the model with the [vLLM](https://github.com/vllm-project/vllm) docker and uses [Text moderation APIs](https://cloud.google.com/natural-language/docs/moderating-text) to analyze predictions against a list of safety attributes.\n",
"\n",
"\n",
"### Objective\n",
"\n",
"- Finetune and deploy Gemma models with a Vertex AI Custom Training Job.\n",
"- Send prediction requests to your finetuned Gemma model.\n",
"\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"* Cloud NL APIs\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), [Cloud NL API pricing](https://cloud.google.com/natural-language/pricing) and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "264c07757582"
},
"source": [
"## Before you begin"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "wQf_xXTXzjaS"
},
"outputs": [],
"source": [
"# @title Install Python Packages for Finetuning\n",
"\n",
"# @markdown 1. Install google-cloud-aiplatform package and restart the session if instructed.\n",
"! pip install --upgrade --quiet google-cloud-aiplatform==1.130.0\n",
"\n",
"# @markdown 2. Install packages to validate dataset with template.\n",
"! pip install --upgrade --quiet accelerate==0.31.0\n",
"! pip install --upgrade --quiet transformers==4.43.1\n",
"! pip install --upgrade --quiet datasets==2.19.2\n",
"\n",
"# Load local tensorboard.\n",
"%load_ext tensorboard"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "B9p8QmmcD_OP"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. For finetuning, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Frestricted_image_training_nvidia_a100_80gb_gpus)** to check if your project already has the required 8 Nvidia A100 80 GB GPUs in the us-central1 region. If yes, then run this notebook in the us-central1 region. If you do not have 8 Nvidia A100 80 GPUs or have more GPU requirements than this, then schedule your job with Nvidia H100 GPUs via Dynamic Workload Scheduler using [these instructions](https://cloud.google.com/vertex-ai/docs/training/schedule-jobs-dws). For Dynamic Workload Scheduler, check the [us-central1](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) or [europe-west4](https://console.cloud.google.com/iam-admin/quotas?location=europe-west4&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) quota for Nvidia H100 GPUs. If you do not have enough GPUs, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request quota.\n",
"\n",
"# @markdown 3. For serving, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_l4_gpus)** to check if your project already has the required 1 L4 GPU in the us-central1 region. If yes, then run this notebook in the us-central1 region. If you need more L4 GPUs for your project, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request more. Alternatively, if you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus).\n",
"\n",
"# @markdown > | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
"# @markdown | a2-ultragpu-1g | 1 NVIDIA_A100_80GB | us-central1, us-east4, europe-west4, asia-southeast1, us-east4 |\n",
"# @markdown | a3-highgpu-2g | 2 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-4g | 4 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"\n",
"# @markdown 4. **[Optional]** [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 5. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# Import the necessary packages\n",
"! rm -rf vertex-ai-samples && git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"! cd vertex-ai-samples && git reset --hard 0727e19520cf7957bceb701c248221bd3dbe4f1f\n",
"\n",
"import datetime\n",
"import importlib\n",
"import os\n",
"import uuid\n",
"from typing import Tuple\n",
"\n",
"from google.cloud import aiplatform, language\n",
"from google.cloud.aiplatform.compat.types import \\\n",
" custom_job as gca_custom_job_compat\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.notebooks.community.model_garden.docker_source_codes.notebook_util.common_util\"\n",
")\n",
"\n",
"models, endpoints = {}, {}\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"if not REGION:\n",
" if not os.environ.get(\"GOOGLE_CLOUD_REGION\"):\n",
" raise ValueError(\n",
" \"REGION must be set. See\"\n",
" \" https://cloud.google.com/vertex-ai/docs/general/locations for\"\n",
" \" available cloud locations.\"\n",
" )\n",
" REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, change the value yourself below.\n",
"now = datetime.datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}\"\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"MODEL_BUCKET = os.path.join(BUCKET_URI, \"gemma\")\n",
"\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default SERVICE_ACCOUNT.\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/storage.admin\"\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/aiplatform.user\"\n",
"\n",
"# @markdown ## Access Gemma Models\n",
"\n",
"# @markdown Provide a Hugging Face User Access Token (read) to access the Gemma models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"HF_TOKEN = \"\" # @param {type:\"string\", isTemplate:true}\n",
"assert HF_TOKEN, \"Provide a read HF_TOKEN to load models from Hugging Face.\"\n",
"\n",
"\n",
"def moderate_text(text: str) -> language.ModerateTextResponse:\n",
" \"\"\"Calls Vertex AI APIs to analyze text moderations.\"\"\"\n",
" client = language.LanguageServiceClient()\n",
" document = language.Document(\n",
" content=text,\n",
" type_=language.Document.Type.PLAIN_TEXT,\n",
" )\n",
" return client.moderate_text(document=document)\n",
"\n",
"\n",
"def show_text_moderation(text: str, response: language.ModerateTextResponse) -> None:\n",
" \"\"\"Shows text moderation results.\"\"\"\n",
" import pandas as pd\n",
"\n",
" def confidence(category: language.ClassificationCategory) -> float:\n",
" return category.confidence\n",
"\n",
" columns = [\"category\", \"confidence\"]\n",
" categories = sorted(response.moderation_categories, key=confidence, reverse=True)\n",
" data = ((category.name, category.confidence) for category in categories)\n",
" df = pd.DataFrame(columns=columns, data=data)\n",
"\n",
" print(f\"Text analyzed:\\n{text}\")\n",
" print(df.to_markdown(index=False, tablefmt=\"presto\", floatfmt=\".0%\"))"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "Pq4iF00YG_4T"
},
"source": [
"## Finetune with Vertex AI Custom Training Jobs\n",
"\n",
"This section demonstrates how to finetune and deploy Gemma models with PEFT LoRA on Vertex AI Custom Training Jobs. LoRA (Low-Rank Adaptation) is one approach of PEFT (Parameter Efficient FineTuning), where pretrained model weights are frozen and rank decomposition matrices representing the change in model weights are trained during finetuning. Read more about LoRA in the following publication: [Hu, E.J., Shen, Y., Wallis, P., Allen-Zhu, Z., Li, Y., Wang, S., Wang, L. and Chen, W., 2021. Lora: Low-rank adaptation of large language models. *arXiv preprint arXiv:2106.09685*](https://arxiv.org/abs/2106.09685)."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "FD1TvpYZzjaS"
},
"outputs": [],
"source": [
"# @title Set dataset\n",
"\n",
"# @markdown Use the Vertex AI SDK to create and run the custom training jobs.\n",
"\n",
"# @markdown This notebook uses [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) dataset as an example.\n",
"# @markdown You can set `dataset_name` to any existing [Hugging Face dataset](https://huggingface.co/datasets) name, and set `instruct_column_in_dataset` to the name of the dataset column containing training data. The [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) has only one column `text`, and therefore we set `instruct_column_in_dataset` to `text` in this notebook.\n",
"\n",
"# @markdown ### (Optional) Prepare a custom JSONL dataset for finetuning\n",
"\n",
"# @markdown You can prepare a JSONL file where each line is a valid JSON string as your custom training dataset. For example, here is one line from the [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) dataset:\n",
"# @markdown ```\n",
"# @markdown {\"text\": \"### Human: Hola### Assistant: \\u00a1Hola! \\u00bfEn qu\\u00e9 puedo ayudarte hoy?\"}\n",
"# @markdown ```\n",
"\n",
"# @markdown The JSON object has a key `text`, which should match `instruct_column_in_dataset`; The value should be one training data point, i.e. a string. After you prepared your JSONL file, you can either upload it to [Hugging Face datasets](https://huggingface.co/datasets) or [Google Cloud Storage](https://cloud.google.com/storage).\n",
"\n",
"# @markdown - To upload a JSONL dataset to [Hugging Face datasets](https://huggingface.co/datasets), follow the instructions on [Uploading Datasets](https://huggingface.co/docs/hub/en/datasets-adding). Then, set `dataset_name` to the name of your newly created dataset on Hugging Face.\n",
"\n",
"# @markdown - To upload a JSONL dataset to [Google Cloud Storage](https://cloud.google.com/storage), follow the instructions on [Upload objects from a filesystem](https://cloud.google.com/storage/docs/uploading-objects). Then, set `dataset_name` to the `gs://` URI to your JSONL file. For example: `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`.\n",
"\n",
"# @markdown Optionally update the `instruct_column_in_dataset` field below if your JSON objects use a key other than the default `text`.\n",
"\n",
"# @markdown ### (Optional) Format your data with custom JSON template\n",
"\n",
"# @markdown Sometimes, your dataset might have multiple text columns and you want to construct the training data with a template. You can prepare a JSON template in the following format:\n",
"\n",
"# @markdown ```\n",
"# @markdown {\n",
"# @markdown \"description\": \"Template that accepts text-bison format.\",\n",
"# @markdown \"source\": \"https://cloud.google.com/vertex-ai/generative-ai/docs/models/tune-text-models-supervised#dataset-format\",\n",
"# @markdown \"prompt_input\": \"\\n\\n<|start_header_id|>user<|end_header_id|>\\n\\n{input_text}<|eot_id|>\\n\\n<|start_header_id|>assistant<|end_header_id|>\\n\\n{output_text}<|eot_id|>\",\n",
"# @markdown \"instruction_separator\": \"<|start_header_id|>user<|end_header_id|>\\n\\n\",\n",
"# @markdown \"response_separator\": \"<|start_header_id|>assistant<|end_header_id|>\\n\\n\"\n",
"# @markdown }\n",
"# @markdown ```\n",
"\n",
"# @markdown As an example, the template above can be used to format the following training data (this line comes from `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`):\n",
"\n",
"# @markdown ```\n",
"# @markdown {\"input_text\":\"TRANSCRIPT: \\nREASON FOR EVALUATION:,\\n\\n LABEL:\",\"output_text\":\"Chiropractic\"}\n",
"# @markdown ```\n",
"\n",
"# @markdown This example template simply concatenates `input_text` with `output_text` with some special tokens in between.\n",
"# @markdown\n",
"# @markdown To try such custom dataset, you can make the following changes:\n",
"# @markdown 1. Set `template` to `llama3-text-bison`\n",
"# @markdown 1. Set `train_dataset_name` to `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`\n",
"# @markdown 1. Set `train_split_name` to `train`\n",
"# @markdown 1. Set `eval_dataset_name` to `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_eval_sample.jsonl`\n",
"# @markdown 1. Set `eval_split_name` to `train` (**NOT** `test`)\n",
"# @markdown 1. Set `instruct_column_in_dataset` as `input_text`.\n",
"\n",
"# Template name or gs:// URI to a custom template.\n",
"template = \"openassistant-guanaco\" # @param {type:\"string\"}\n",
"\n",
"# Hugging Face dataset name or gs:// URI to a custom JSONL dataset.\n",
"train_dataset_name = \"timdettmers/openassistant-guanaco\" # @param {type:\"string\"}\n",
"train_split_name = \"train\" # @param {type:\"string\"}\n",
"eval_dataset_name = \"timdettmers/openassistant-guanaco\" # @param {type:\"string\"}\n",
"eval_split_name = \"test\" # @param {type:\"string\"}\n",
"\n",
"# Name of the dataset column containing training text input.\n",
"instruct_column_in_dataset = \"text\" # @param {type:\"string\"}"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "wBxu3rEyzjaT"
},
"outputs": [],
"source": [
"# @title Set model\n",
"\n",
"# @markdown Select a model variant of Gemma.\n",
"base_model_id = \"gemma-1.1-2b-it\" # @param[\"gemma-2b\", \"gemma-2b-it\", \"gemma-7b\", \"gemma-7b-it\", \"gemma-1.1-2b-it\", \"gemma-1.1-7b-it\"] {isTemplate:true}\n",
"pretrained_model_id = os.path.join(\"google/\", base_model_id)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "jIEDYAnPzjaT"
},
"outputs": [],
"source": [
"# @title Validate Dataset with Template\n",
"\n",
"# @markdown This section validates the train and eval datasets with the template before starting the fine tuning process.\n",
"\n",
"import transformers\n",
"\n",
"dataset_validation_util = importlib.import_module(\n",
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.dataset_validation_util\"\n",
")\n",
"\n",
"if dataset_validation_util.is_gcs_path(pretrained_model_id):\n",
" # Download tokenizer.\n",
" ! mkdir tokenizer\n",
" ! gsutil cp {pretrained_model_id}/tokenizer.json ./tokenizer\n",
" ! gsutil cp {pretrained_model_id}/config.json ./tokenizer\n",
" tokenizer_path = \"./tokenizer\"\n",
" access_token = \"\"\n",
"else:\n",
" tokenizer_path = pretrained_model_id\n",
" access_token = HF_TOKEN\n",
"\n",
"tokenizer = transformers.AutoTokenizer.from_pretrained(\n",
" tokenizer_path,\n",
" trust_remote_code=False,\n",
" use_fast=True,\n",
" token=access_token,\n",
")\n",
"\n",
"# Validate the train dataset.\n",
"dataset_validation_util.validate_dataset_with_template(\n",
" dataset_name=train_dataset_name,\n",
" split=train_split_name,\n",
" input_column=instruct_column_in_dataset,\n",
" template=template,\n",
" use_multiprocessing=False,\n",
" tokenizer=tokenizer,\n",
")\n",
"\n",
"# Validate the eval dataset.\n",
"dataset_validation_util.validate_dataset_with_template(\n",
" dataset_name=eval_dataset_name,\n",
" split=eval_split_name,\n",
" input_column=instruct_column_in_dataset,\n",
" template=template,\n",
" use_multiprocessing=False,\n",
" tokenizer=tokenizer,\n",
")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "rU7ekq-0zjaT"
},
"outputs": [],
"source": [
"# @title Finetune\n",
"# @markdown This section demonstrates how to finetune the Gemma model and merge the finetuned LoRA adapter with the base model on Vertex AI. It uses the Vertex AI SDK to create and run the custom training jobs.\n",
"\n",
"# @markdown The training job takes approximately between 10 to 20 mins to set-up. Once done, the training job is expected to take around 20 mins with the default configuration. To find the training time, throughput, and memory usage of your training job, you can go to the training logs and check the log line of the last training epoch.\n",
"\n",
"# @markdown **Note**:\n",
"# @markdown 1. We recommend setting `finetuning_precision_mode` to `4bit` because it enables using fewer hardware resources for finetuning.\n",
"# @markdown 1. If `max_steps > 0`, it takes precedence over `epochs`. One can set a small `max_steps` value to quickly check the pipeline.\n",
"\n",
"# @markdown Accelerator type to use for training.\n",
"training_accelerator_type = \"NVIDIA_A100_80GB\" # @param [\"NVIDIA_A100_80GB\", \"NVIDIA_H100_80GB\"]\n",
"\n",
"# The pre-built training docker image.\n",
"if training_accelerator_type == \"NVIDIA_A100_80GB\":\n",
" repo = \"us-docker.pkg.dev/vertex-ai-restricted\"\n",
" is_restricted_image = True\n",
" is_dynamic_workload_scheduler = False\n",
" dws_kwargs = {}\n",
"else:\n",
" repo = \"us-docker.pkg.dev/vertex-ai\"\n",
" is_restricted_image = False\n",
" is_dynamic_workload_scheduler = True\n",
" dws_kwargs = {\n",
" \"max_wait_duration\": 1800, # 30 minutes\n",
" \"scheduling_strategy\": gca_custom_job_compat.Scheduling.Strategy.FLEX_START,\n",
" }\n",
"\n",
"TRAIN_DOCKER_URI = (\n",
" f\"{repo}/vertex-vision-model-garden-dockers/pytorch-peft-train:stable_20240909\"\n",
")\n",
"\n",
"# Worker pool spec.\n",
"if training_accelerator_type == \"NVIDIA_A100_80GB\":\n",
" per_node_accelerator_count = 8\n",
" training_machine_type = \"a2-ultragpu-8g\"\n",
"elif training_accelerator_type == \"NVIDIA_H100_80GB\":\n",
" per_node_accelerator_count = 8\n",
" training_machine_type = \"a3-highgpu-8g\"\n",
"else:\n",
" raise ValueError(\n",
" f\"Recommended machine settings not found for: {training_accelerator_type}. To use another accelerator type, edit this code block to pass in an appropriate `training_machine_type`, `training_accelerator_type`, and `per_node_accelerator_count` to the deploy_model_vllm function by clicking `Show Code` and then modifying the code.\"\n",
" )\n",
"\n",
"# @markdown Batch size for finetuning.\n",
"per_device_train_batch_size = 1 # @param{type:\"integer\"}\n",
"# @markdown Number of updates steps to accumulate the gradients for, before performing a backward/update pass.\n",
"gradient_accumulation_steps = 4 # @param{type:\"integer\"}\n",
"# @markdown Maximum sequence length.\n",
"max_seq_length = 4096 # @param{type:\"integer\"}\n",
"# @markdown Setting a positive `max_steps` here will override `num_epochs`.\n",
"max_steps = -1 # @param{type:\"integer\"}\n",
"num_epochs = 1.0 # @param{type:\"number\"}\n",
"# @markdown Precision mode for finetuning.\n",
"finetuning_precision_mode = \"4bit\" # @param [\"4bit\", \"8bit\", \"float16\"]\n",
"# @markdown Learning rate.\n",
"learning_rate = 5e-5 # @param{type:\"number\"}\n",
"# @markdown The scheduler type to use.\n",
"lr_scheduler_type = \"cosine\" # @param{type:\"string\"}\n",
"# @markdown LoRA parameters.\n",
"lora_rank = 16 # @param{type:\"integer\"}\n",
"lora_alpha = 32 # @param{type:\"integer\"}\n",
"lora_dropout = 0.05 # @param{type:\"number\"}\n",
"# Activates gradient checkpointing for the current model (may be referred to as activation checkpointing or checkpoint activations in other frameworks).\n",
"enable_gradient_checkpointing = True\n",
"# Attention implementation to use in the model.\n",
"attn_implementation = \"eager\"\n",
"# The optimizer for which to schedule the learning rate.\n",
"optimizer = \"paged_adamw_32bit\"\n",
"# Define the proportion of training to be dedicated to a linear warmup where learning rate gradually increases.\n",
"warmup_ratio = \"0.01\"\n",
"# The list or string of integrations to report the results and logs to.\n",
"report_to = \"tensorboard\"\n",
"# Number of updates steps before two checkpoint saves.\n",
"save_steps = 10\n",
"# Number of update steps between two logs.\n",
"logging_steps = save_steps\n",
"# Train precision of the model.\n",
"train_precision = \"bfloat16\"\n",
"\n",
"replica_count = 1\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=training_accelerator_type,\n",
" accelerator_count=per_node_accelerator_count * replica_count,\n",
" is_for_training=True,\n",
" is_restricted_image=is_restricted_image,\n",
" is_dynamic_workload_scheduler=is_dynamic_workload_scheduler,\n",
")\n",
"\n",
"job_name = common_util.get_job_name_with_datetime(\"gemma-lora-train\")\n",
"\n",
"base_output_dir = os.path.join(STAGING_BUCKET, job_name)\n",
"# Create a GCS folder to store the LORA adapter.\n",
"lora_output_dir = os.path.join(base_output_dir, \"adapter\")\n",
"# Create a GCS folder to store the merged model with the base model and the\n",
"# finetuned LORA adapter.\n",
"merged_model_output_dir = os.path.join(base_output_dir, \"merged-model\")\n",
"\n",
"# Add labels for the finetuning job.\n",
"labels = {\n",
" \"mg-source\": \"notebook\",\n",
" \"mg-notebook-name\": \"model_garden_pytorch_gemma_peft_finetuning_hf.ipynb\".split(\n",
" \".\"\n",
" )[0],\n",
"}\n",
"\n",
"labels[\"mg-tune\"] = \"publishers-google-models-gemma\"\n",
"versioned_model_id = base_model_id.lower().replace(\".\", \"-\")\n",
"labels[\"versioned-mg-tune\"] = f\"{labels['mg-tune']}-{versioned_model_id}\"\n",
"\n",
"eval_args = [\n",
" f\"--eval_dataset_path={eval_dataset_name}\",\n",
" f\"--eval_column={instruct_column_in_dataset}\",\n",
" f\"--eval_template={template}\",\n",
" f\"--eval_split={eval_split_name}\",\n",
" f\"--eval_steps={save_steps}\",\n",
" \"--eval_tasks=builtin_eval\",\n",
" \"--eval_metric_name=loss\",\n",
"]\n",
"\n",
"train_job_args = [\n",
" \"--config_file=vertex_vision_model_garden_peft/deepspeed_zero2_8gpu.yaml\",\n",
" \"--task=instruct-lora\",\n",
" \"--completion_only=True\",\n",
" f\"--pretrained_model_id={pretrained_model_id}\",\n",
" f\"--dataset_name={train_dataset_name}\",\n",
" f\"--train_split_name={train_split_name}\",\n",
" f\"--instruct_column_in_dataset={instruct_column_in_dataset}\",\n",
" f\"--output_dir={lora_output_dir}\",\n",
" f\"--merge_base_and_lora_output_dir={merged_model_output_dir}\",\n",
" f\"--per_device_train_batch_size={per_device_train_batch_size}\",\n",
" f\"--gradient_accumulation_steps={gradient_accumulation_steps}\",\n",
" f\"--lora_rank={lora_rank}\",\n",
" f\"--lora_alpha={lora_alpha}\",\n",
" f\"--lora_dropout={lora_dropout}\",\n",
" f\"--max_steps={max_steps}\",\n",
" f\"--max_seq_length={max_seq_length}\",\n",
" f\"--learning_rate={learning_rate}\",\n",
" f\"--lr_scheduler_type={lr_scheduler_type}\",\n",
" f\"--precision_mode={finetuning_precision_mode}\",\n",
" f\"--train_precision={train_precision}\",\n",
" f\"--enable_gradient_checkpointing={enable_gradient_checkpointing}\",\n",
" f\"--num_epochs={num_epochs}\",\n",
" f\"--attn_implementation={attn_implementation}\",\n",
" f\"--optimizer={optimizer}\",\n",
" f\"--warmup_ratio={warmup_ratio}\",\n",
" f\"--report_to={report_to}\",\n",
" f\"--logging_output_dir={base_output_dir}\",\n",
" f\"--save_steps={save_steps}\",\n",
" f\"--logging_steps={logging_steps}\",\n",
" f\"--template={template}\",\n",
" f\"--huggingface_access_token={HF_TOKEN}\",\n",
"] + eval_args\n",
"\n",
"# Pass training arguments and launch job.\n",
"train_job = aiplatform.CustomContainerTrainingJob(\n",
" display_name=job_name,\n",
" container_uri=TRAIN_DOCKER_URI,\n",
" labels=labels,\n",
")\n",
"\n",
"print(\"Running training job with args:\")\n",
"print(\" \\\\\\n\".join(train_job_args))\n",
"train_job.run(\n",
" args=train_job_args,\n",
" replica_count=replica_count,\n",
" machine_type=training_machine_type,\n",
" accelerator_type=training_accelerator_type,\n",
" accelerator_count=per_node_accelerator_count,\n",
" boot_disk_size_gb=500,\n",
" service_account=SERVICE_ACCOUNT,\n",
" base_output_dir=base_output_dir,\n",
" sync=False, # Non-blocking call to run.\n",
" **dws_kwargs,\n",
")\n",
"\n",
"# Wait until resource has been created.\n",
"train_job.wait_for_resource_creation()\n",
"\n",
"print(\"LoRA adapter will be saved in:\", lora_output_dir)\n",
"print(\"Trained and merged models will be saved in:\", merged_model_output_dir)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "NdBE5YabzjaT"
},
"outputs": [],
"source": [
"# @title Run TensorBoard\n",
"# @markdown This section shows how to launch TensorBoard in a [Cloud Shell](https://cloud.google.com/shell/docs).\n",
"# @markdown 1. Click the Cloud Shell icon(![terminal](https://github.com/google/material-design-icons/blob/master/png/action/terminal/materialicons/24dp/1x/baseline_terminal_black_24dp.png?raw=true)) on the top right to open the Cloud Shell.\n",
"# @markdown 2. Copy the `tensorboard` command shown below by running this cell.\n",
"# @markdown 3. Paste and run the command in the Cloud Shell to launch TensorBoard.\n",
"# @markdown 4. Once the command runs (You may have to click `Authorize` if prompted), click the link starting with `http://localhost`.\n",
"\n",
"# @markdown Note: You may need to wait around 10 minutes after the job starts in order for the TensorBoard logs to be written to the GCS bucket.\n",
"print(f\"Command to copy: tensorboard --logdir {base_output_dir}/logs\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "qmHW6m8xG_4U"
},
"outputs": [],
"source": [
"# @title Deploy with vLLM\n",
"# @markdown This section uploads the model to Model Registry and deploys it on the Endpoint. It takes 15 minutes to 1 hour to finish.\n",
"\n",
"if train_job.end_time is None:\n",
" print(\"Waiting for the training job to finish...\")\n",
" train_job.wait()\n",
" print(\"The training job has finished.\")\n",
"\n",
"print(\"Deploying models in:\", merged_model_output_dir)\n",
"\n",
"# The pre-built serving docker image for vLLM.\n",
"VLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-vllm-serve:20241010_0916_RC00\"\n",
"\n",
"# Find Vertex AI prediction supported accelerators and regions in [here](https://cloud.google.com/vertex-ai/docs/predictions/configure-compute).\n",
"# Sets 1 L4 (24G) to deploy Gemma models.\n",
"serve_machine_type = \"g2-standard-12\"\n",
"serve_accelerator_type = \"NVIDIA_L4\"\n",
"serve_accelerator_count = 1\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=serve_accelerator_type,\n",
" accelerator_count=serve_accelerator_count,\n",
" is_for_training=False,\n",
")\n",
"\n",
"# Note that a larger max_model_len will require more GPU memory.\n",
"max_model_len = 2048\n",
"\n",
"\n",
"def get_deploy_source() -> str:\n",
" \"\"\"Gets deploy_source string based on running environment.\"\"\"\n",
" vertex_product = os.environ.get(\"VERTEX_PRODUCT\", \"\")\n",
" if vertex_product == \"COLAB_ENTERPRISE\":\n",
" return \"notebook_colab_enterprise\"\n",
" elif vertex_product == \"WORKBENCH_INSTANCE\":\n",
" return \"notebook_workbench\"\n",
" else:\n",
" # Legacy workbench, legacy colab, or other custom environments.\n",
" return \"notebook_environment_unspecified\"\n",
"\n",
"\n",
"def deploy_model_vllm(\n",
" model_name: str,\n",
" model_id: str,\n",
" publisher: str,\n",
" publisher_model_id: str,\n",
" service_account: str,\n",
" base_model_id: str = None,\n",
" machine_type: str = \"g2-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
" gpu_memory_utilization: float = 0.9,\n",
" max_model_len: int = 4096,\n",
" dtype: str = \"auto\",\n",
" enable_trust_remote_code: bool = False,\n",
" enforce_eager: bool = False,\n",
" enable_lora: bool = False,\n",
" enable_chunked_prefill: bool = False,\n",
" enable_prefix_cache: bool = False,\n",
" host_prefix_kv_cache_utilization_target: float = 0.0,\n",
" max_loras: int = 1,\n",
" max_cpu_loras: int = 8,\n",
" use_dedicated_endpoint: bool = False,\n",
" max_num_seqs: int = 256,\n",
" model_type: str = None,\n",
" enable_llama_tool_parser: bool = False,\n",
" is_spot: bool = False,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys trained models with vLLM into Vertex AI.\"\"\"\n",
" endpoint = aiplatform.Endpoint.create(\n",
" display_name=f\"{model_name}-endpoint\",\n",
" dedicated_endpoint_enabled=use_dedicated_endpoint,\n",
" )\n",
"\n",
" if not base_model_id:\n",
" base_model_id = model_id\n",
"\n",
" # See https://docs.vllm.ai/en/latest/models/engine_args.html for a list of possible arguments with descriptions.\n",
" vllm_args = [\n",
" \"python\",\n",
" \"-m\",\n",
" \"vllm.entrypoints.api_server\",\n",
" \"--host=0.0.0.0\",\n",
" \"--port=8080\",\n",
" f\"--model={model_id}\",\n",
" f\"--tensor-parallel-size={accelerator_count}\",\n",
" \"--swap-space=16\",\n",
" f\"--max-model-len={max_model_len}\",\n",
" f\"--dtype={dtype}\",\n",
" f\"--max-loras={max_loras}\",\n",
" f\"--max-cpu-loras={max_cpu_loras}\",\n",
" f\"--max-num-seqs={max_num_seqs}\",\n",
" \"--disable-log-stats\",\n",
" ]\n",
"\n",
" if gpu_memory_utilization:\n",
" vllm_args.append(f\"--gpu-memory-utilization={gpu_memory_utilization}\")\n",
"\n",
" if enable_trust_remote_code:\n",
" vllm_args.append(\"--trust-remote-code\")\n",
"\n",
" if enforce_eager:\n",
" vllm_args.append(\"--enforce-eager\")\n",
"\n",
" if enable_lora:\n",
" vllm_args.append(\"--enable-lora\")\n",
"\n",
" if enable_chunked_prefill:\n",
" vllm_args.append(\"--enable-chunked-prefill\")\n",
"\n",
" if enable_prefix_cache:\n",
" vllm_args.append(\"--enable-prefix-caching\")\n",
"\n",
" if 0 < host_prefix_kv_cache_utilization_target < 1:\n",
" vllm_args.append(\n",
" f\"--host-prefix-kv-cache-utilization-target={host_prefix_kv_cache_utilization_target}\"\n",
" )\n",
"\n",
" if model_type:\n",
" vllm_args.append(f\"--model-type={model_type}\")\n",
"\n",
" if enable_llama_tool_parser:\n",
" vllm_args.append(\"--enable-auto-tool-choice\")\n",
" vllm_args.append(\"--tool-call-parser=vertex-llama-3\")\n",
"\n",
" env_vars = {\n",
" \"MODEL_ID\": base_model_id,\n",
" \"DEPLOY_SOURCE\": \"notebook\",\n",
" }\n",
"\n",
" # HF_TOKEN is not a compulsory field and may not be defined.\n",
" try:\n",
" if HF_TOKEN:\n",
" env_vars[\"HF_TOKEN\"] = HF_TOKEN\n",
" except NameError:\n",
" pass\n",
"\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=VLLM_DOCKER_URI,\n",
" serving_container_args=vllm_args,\n",
" serving_container_ports=[8080],\n",
" serving_container_predict_route=\"/generate\",\n",
" serving_container_health_route=\"/ping\",\n",
" serving_container_environment_variables=env_vars,\n",
" serving_container_shared_memory_size_mb=(16 * 1024), # 16 GB\n",
" serving_container_deployment_timeout=7200,\n",
" model_garden_source_model_name=(\n",
" f\"publishers/{publisher}/models/{publisher_model_id}\"\n",
" ),\n",
" )\n",
" print(\n",
" f\"Deploying {model_name} on {machine_type} with {accelerator_count} {accelerator_type} GPU(s).\"\n",
" )\n",
" model.deploy(\n",
" endpoint=endpoint,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" deploy_request_timeout=1800,\n",
" service_account=service_account,\n",
" spot=is_spot,\n",
" system_labels={\n",
" \"NOTEBOOK_NAME\": \"model_garden_pytorch_gemma_peft_finetuning_hf.ipynb\",\n",
" \"NOTEBOOK_ENVIRONMENT\": get_deploy_source(),\n",
" },\n",
" )\n",
" print(\"endpoint_name:\", endpoint.name)\n",
"\n",
" return model, endpoint\n",
"\n",
"\n",
"models[\"vllm_gpu\"], endpoints[\"vllm_gpu\"] = deploy_model_vllm(\n",
" model_name=common_util.get_job_name_with_datetime(prefix=\"gemma-vllm-serve\"),\n",
" base_model_id=f\"google/{base_model_id}\",\n",
" publisher=\"google\",\n",
" publisher_model_id=\"gemma\",\n",
" model_id=merged_model_output_dir,\n",
" service_account=SERVICE_ACCOUNT,\n",
" machine_type=serve_machine_type,\n",
" accelerator_type=serve_accelerator_type,\n",
" accelerator_count=serve_accelerator_count,\n",
" max_model_len=max_model_len,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "2UYUNn60G_4U"
},
"outputs": [],
"source": [
"# @title Predict\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint with text prompts. Sampling parameters supported by vLLM can be found [here](https://docs.vllm.ai/en/latest/dev/sampling_params.html).\n",
"\n",
"# @markdown Example:\n",
"\n",
"# @markdown ```\n",
"# @markdown Human: What is a car?\n",
"# @markdown Assistant: A car, or a motor car, is a road-connected human-transportation system used to move people or goods from one place to another. The term also encompasses a wide range of vehicles, including motorboats, trains, and aircrafts. Cars typically have four wheels, a cabin for passengers, and an engine or motor. They have been around since the early 19th century and are now one of the most popular forms of transportation, used for daily commuting, shopping, and other purposes.\n",
"# @markdown ```\n",
"# @markdown Additionally, you can moderate the generated text with Vertex AI. See [Moderate text documentation](https://cloud.google.com/natural-language/docs/moderating-text) for more details.\n",
"\n",
"# Loads an existing endpoint instance using the endpoint name:\n",
"# - Using `endpoint_name = endpoint.name` allows us to get the\n",
"# endpoint name of the endpoint `endpoint` created in the cell\n",
"# above.\n",
"# - Alternatively, you can set `endpoint_name = \"1234567890123456789\"` to load\n",
"# an existing endpoint with the ID 1234567890123456789.\n",
"# You may uncomment the code below to load an existing endpoint.\n",
"\n",
"# endpoint_name = \"\" # @param {type:\"string\"}\n",
"# aip_endpoint_name = (\n",
"# f\"projects/{PROJECT_ID}/locations/{REGION}/endpoints/{endpoint_name}\"\n",
"# )\n",
"# endpoint = aiplatform.Endpoint(aip_endpoint_name)\n",
"\n",
"prompt = \"What is a car?\" # @param {type: \"string\"}\n",
"# @markdown If you encounter an issue like `ServiceUnavailable: 503 Took too long to respond when processing`, you can reduce the maximum number of output tokens, by lowering `max_tokens`.\n",
"max_tokens = 50 # @param {type:\"integer\"}\n",
"temperature = 1.0 # @param {type:\"number\"}\n",
"top_p = 1.0 # @param {type:\"number\"}\n",
"top_k = 1 # @param {type:\"integer\"}\n",
"# @markdown Set `raw_response` to `True` to obtain the raw model output. Set `raw_response` to `False` to apply additional formatting in the structure of `\"Prompt:\\n{prompt.strip()}\\nOutput:\\n{output}\"`.\n",
"raw_response = False # @param {type:\"boolean\"}\n",
"\n",
"# Overrides parameters for inferences.\n",
"instances = [\n",
" {\n",
" \"prompt\": prompt,\n",
" \"max_tokens\": max_tokens,\n",
" \"temperature\": temperature,\n",
" \"top_p\": top_p,\n",
" \"top_k\": top_k,\n",
" \"raw_response\": raw_response,\n",
" },\n",
"]\n",
"response = endpoints[\"vllm_gpu\"].predict(\n",
" instances=instances, use_dedicated_endpoint=use_dedicated_endpoint\n",
")\n",
"\n",
"for prediction in response.predictions:\n",
" print(prediction)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "2T_cXYJhG_4U"
},
"outputs": [],
"source": [
"# @markdown Text moderation analyzes a document against a list of safety attributes, which include \"harmful categories\" and topics that may be considered sensitive.\n",
"\n",
"for generated_text in response.predictions:\n",
" # Send a request to the API.\n",
" response = moderate_text(generated_text)\n",
" # Show the results.\n",
" show_text_moderation(generated_text, response)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "af21a3cff1e0"
},
"source": [
"## Clean up resources"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "911406c1561e"
},
"outputs": [],
"source": [
"# Delete the train job.\n",
"train_job.delete()\n",
"\n",
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continuous charges that may incur.\n",
"\n",
"# Undeploy model and delete endpoint.\n",
"for endpoint in endpoints.values():\n",
" endpoint.delete(force=True)\n",
"\n",
"# Delete models.\n",
"for model in models.values():\n",
" model.delete()\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_NAME"
]
}
],
"metadata": {
"colab": {
"name": "model_garden_pytorch_gemma_peft_finetuning_hf.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
@@ -0,0 +1,557 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"id": "53d024d9",
"metadata": {
"cellView": "form",
"id": "483138c1a042"
},
"outputs": [],
"source": [
"# Copyright 2026 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",
"# https://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",
"id": "ed49cc75",
"metadata": {
"id": "595822edfb75"
},
"source": [
"# Vertex AI Model Garden - Kimi-K3 (Deployment)\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_pytorch_kimi_k3_deployment.ipynb\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\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%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_pytorch_kimi_k3_deployment.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_pytorch_kimi_k3_deployment.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"id": "ed230c19",
"metadata": {
"id": "39399f72d0ea"
},
"source": [
"<!-- @publisher_model_name moonshotai/kimi-k3@kimi-k3 -->\n",
"<!-- @model_name Kimi-K3 -->\n",
"\n",
"## Overview\n",
"\n",
"This notebook demonstrates serving [Kimi-K3](https://huggingface.co/moonshotai/Kimi-K3) with [SGLang](https://github.com/sgl-project/sglang) on `a4-highgpu-8g` machines with NVIDIA B200 GPUs on Vertex AI.\n",
"\n",
"Kimi-K3 is a state-of-the-art large language model from Moonshot AI, featuring advanced reasoning and tool-calling capabilities. This notebook deploys Kimi-K3 using multi-host GPU serving with DSPARK speculative decoding ([Kimi-K3-DSpark](https://huggingface.co/RadixArk/Kimi-K3-DSpark)).\n",
"\n",
"### Objective\n",
"\n",
"- Deploy Kimi-K3 with SGLang on GPU using multi-host serving and [Spot VMs](https://cloud.google.com/compute/docs/instances/spot) (Optional). Multi-host GPU serving is a preview feature.\n",
"\n",
"### File a bug\n",
"\n",
"File a bug on [GitHub](https://github.com/GoogleCloudPlatform/vertex-ai-samples/issues/new) if you encounter any issue with the notebook.\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
"cell_type": "markdown",
"id": "a303257c",
"metadata": {
"id": "69453bf7230e"
},
"source": [
"## Before you begin"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "730d5d48",
"metadata": {
"cellView": "form",
"id": "90053fcb22ce"
},
"outputs": [],
"source": [
"# @title Request for quota\n",
"\n",
"# @markdown To deploy with a4-highgpu-8g (8 x B200) machines, check that you have sufficient quota: [CustomModelServingB200GPUsPerProjectPerRegion](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_b200_gpus). Find the available region(s) [here](https://cloud.google.com/vertex-ai/docs/general/locations#region_considerations).\n",
"\n",
"# @markdown If you don't have sufficient quota, request for quota following the instructions at [\"Request a quota adjustment\"](https://cloud.google.com/docs/quotas/view-manage#requesting_higher_quota).\n",
"\n",
"# @markdown You can also use Compute Engine reservations with Vertex Prediction following the instructions [here](https://cloud.google.com/vertex-ai/docs/predictions/use-reservations). Note that the GCE quota for the shared reservation will be managed separately. Shared reservation is the only GCE consumption mode."
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "bf54d918",
"metadata": {
"cellView": "form",
"id": "8ab79f2d05a7"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform==1.103.0'\n",
"\n",
"# Import the necessary packages\n",
"import importlib\n",
"import os\n",
"import time\n",
"from typing import Tuple\n",
"\n",
"import requests\n",
"from google import auth\n",
"from google.cloud import aiplatform\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"if os.environ.get(\"VERTEX_PRODUCT\") != \"COLAB_ENTERPRISE\":\n",
" ! pip install --upgrade tensorflow\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.notebooks.community.model_garden.docker_source_codes.notebook_util.common_util\"\n",
")\n",
"\n",
"LABEL = \"sglang_gpu\"\n",
"models, endpoints = {}, {}\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"if not REGION:\n",
" REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION)\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"\n",
"import vertexai\n",
"\n",
"vertexai.init(\n",
" project=PROJECT_ID,\n",
" location=REGION,\n",
")"
]
},
{
"cell_type": "markdown",
"id": "de5d3133",
"metadata": {
"id": "b5bac0d360d8"
},
"source": [
"## Specify model artifacts\n",
"\n",
"For a more reliable deployment, it is recommended to upload the Kimi-K3 base model (`moonshotai/Kimi-K3`) and its speculative draft model (`RadixArk/Kimi-K3-DSpark`) to a personal Google Cloud Storage (GCS) bucket beforehand.\n",
"\n",
"In the cell below, you can optionally specify custom GCS paths to your uploaded model artifacts. If provided, the deployment will use your GCS paths; otherwise, it will default to the pre-staged model paths."
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "93a1c593",
"metadata": {
"cellView": "form",
"id": "27afc5ed43b0"
},
"outputs": [],
"source": [
"# @title Specify custom GCS model paths\n",
"\n",
"# @markdown **[Optional]** Specify custom GCS paths to your uploaded Kimi-K3 base model and speculative draft model artifacts. If left empty, default pre-staged paths will be used.\n",
"BASE_MODEL_GCS_URI = \"\" # @param {type:\"string\"}\n",
"SPECULATIVE_DRAFT_MODEL_GCS_URI = \"\" # @param {type:\"string\"}\n",
"\n",
"BASE_MODEL_ID = \"moonshotai/Kimi-K3\"\n",
"\n",
"if BASE_MODEL_GCS_URI and BASE_MODEL_GCS_URI.strip():\n",
" base_model_path = BASE_MODEL_GCS_URI.strip()\n",
" print(f\"Using custom base model GCS path: {base_model_path}\")\n",
"else:\n",
" base_model_path = \"gs://vertex-model-garden-restricted-us/moonshotai/Kimi-K3\"\n",
" print(f\"Using default base model path: {base_model_path}\")\n",
"\n",
"if SPECULATIVE_DRAFT_MODEL_GCS_URI and SPECULATIVE_DRAFT_MODEL_GCS_URI.strip():\n",
" speculative_draft_model_path = SPECULATIVE_DRAFT_MODEL_GCS_URI.strip()\n",
" print(\n",
" \"Using custom speculative draft model GCS path:\"\n",
" f\" {speculative_draft_model_path}\"\n",
" )\n",
"else:\n",
" speculative_draft_model_path = \"RadixArk/Kimi-K3-DSpark\"\n",
" print(\n",
" \"Using default speculative draft model path:\" f\" {speculative_draft_model_path}\"\n",
" )"
]
},
{
"cell_type": "markdown",
"id": "2f079624",
"metadata": {
"id": "8439d66690dd"
},
"source": [
"## Deploy Kimi-K3 with SGLang"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "40f70977",
"metadata": {
"cellView": "form",
"id": "7c2e6903ad6d"
},
"outputs": [],
"source": [
"# @title Deploy Kimi-K3 model on Vertex AI\n",
"\n",
"# @markdown This section deploys the Kimi-K3 model to a Vertex AI Prediction Endpoint. It takes ~30 minutes to finish.\n",
"\n",
"# @markdown The pre-built serving docker image for SGLang.\n",
"# @markdown The current B200 serving support in Model Garden is preliminary and will be continuously improved in the future.\n",
"SGLANG_DOCKER_URI = (\n",
" \"us-docker.pkg.dev/agent-platform-mg-public/containers/sglang-airlock:kimi-k3\"\n",
")\n",
"\n",
"# @markdown Choose whether to use a [Spot VM](https://cloud.google.com/compute/docs/instances/spot) for the deployment.\n",
"is_spot = False # @param {type:\"boolean\"}\n",
"\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint). Note that [dedicated endpoint does not support VPC Service Controls](https://cloud.google.com/vertex-ai/docs/predictions/choose-endpoint-type), uncheck the box if you are using VPC-SC.\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"# @markdown Find Vertex AI prediction supported accelerators and regions at https://cloud.google.com/vertex-ai/docs/predictions/configure-compute.\n",
"accelerator_type = \"NVIDIA_B200\" # @param [\"NVIDIA_B200\"] {isTemplate:true}\n",
"if accelerator_type == \"NVIDIA_B200\":\n",
" accelerator_count = 8\n",
" machine_type = \"a4-highgpu-8g\"\n",
"else:\n",
" raise ValueError(\"Sample deployment options are not available.\")\n",
"multihost_gpu_node_count = 2\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=int(accelerator_count * multihost_gpu_node_count),\n",
" is_for_training=False,\n",
" is_spot=is_spot,\n",
")\n",
"\n",
"\n",
"def poll_operation(op_name: str) -> bool:\n",
" creds, _ = auth.default()\n",
" auth_req = auth.transport.requests.Request()\n",
" creds.refresh(auth_req)\n",
" headers = {\n",
" \"Authorization\": f\"Bearer {creds.token}\",\n",
" }\n",
" get_resp = requests.get(\n",
" f\"https://{REGION}-aiplatform.googleapis.com/ui/{op_name}\",\n",
" headers=headers,\n",
" )\n",
" opjs = get_resp.json()\n",
" if \"error\" in opjs:\n",
" raise ValueError(f\"Operation failed: {opjs['error']}\")\n",
" return opjs.get(\"done\", False)\n",
"\n",
"\n",
"def poll_and_wait(op_name: str, total_wait: int, interval: int = 60):\n",
" waited = 0\n",
" while not poll_operation(op_name):\n",
" if waited > total_wait:\n",
" raise TimeoutError(\"Operation timed out\")\n",
" print(\n",
" f\"\\rStill waiting for operation... Waited time in second: {waited:<6}\",\n",
" end=\"\",\n",
" flush=True,\n",
" )\n",
" waited += interval\n",
" time.sleep(interval)\n",
"\n",
"\n",
"def deploy_model_kimi_k3_sglang(\n",
" model_name: str,\n",
" base_model_path: str,\n",
" speculative_draft_model_path: str,\n",
" machine_type: str,\n",
" accelerator_type: str,\n",
" accelerator_count: int,\n",
" multihost_gpu_node_count: int,\n",
" use_dedicated_endpoint: bool = False,\n",
" is_spot: bool = False,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys Kimi-K3 with SGLang into Vertex AI.\"\"\"\n",
" endpoint = aiplatform.Endpoint.create(\n",
" display_name=f\"{model_name}-endpoint\",\n",
" dedicated_endpoint_enabled=use_dedicated_endpoint,\n",
" )\n",
"\n",
" sglang_args = [\n",
" \"./entrypoint.sh\",\n",
" f\"--model={base_model_path}\",\n",
" \"--trust-remote-code\",\n",
" f\"--tp-size={int(accelerator_count * multihost_gpu_node_count)}\",\n",
" \"--mem-fraction-static=0.85\",\n",
" \"--disable-flashinfer-autotune\",\n",
" \"--enable-metrics\",\n",
" \"--watchdog-timeout=3600\",\n",
" \"--reasoning-parser=kimi_k3\",\n",
" \"--tool-call-parser=kimi_k3\",\n",
" '--model-loader-extra-config={\"enable_multithread_load\": true}',\n",
" \"--mamba-full-memory-ratio=0.43\",\n",
" \"--speculative-algorithm=DSPARK\",\n",
" f\"--speculative-draft-model-path={speculative_draft_model_path}\",\n",
" \"--speculative-dspark-block-size=7\",\n",
" \"--enable-linear-replayssm-spec\",\n",
" \"--enable-hierarchical-cache\",\n",
" ]\n",
"\n",
" env_vars = {\n",
" \"MODEL_ID\": BASE_MODEL_ID,\n",
" \"DEPLOY_SOURCE\": \"notebook\",\n",
" \"NCCL_DEBUG\": \"TRACE\",\n",
" }\n",
"\n",
" try:\n",
" if \"HF_TOKEN\" in os.environ and os.environ[\"HF_TOKEN\"]:\n",
" env_vars[\"HF_TOKEN\"] = os.environ[\"HF_TOKEN\"]\n",
" except Exception:\n",
" pass\n",
"\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=SGLANG_DOCKER_URI,\n",
" serving_container_command=[\"./gcs_download_launcher.sh\"],\n",
" serving_container_args=sglang_args,\n",
" serving_container_ports=[30000],\n",
" serving_container_predict_route=\"/vertex_generate\",\n",
" serving_container_health_route=\"/health\",\n",
" serving_container_environment_variables=env_vars,\n",
" serving_container_shared_memory_size_mb=(32 * 1024), # 32768 MB\n",
" serving_container_deployment_timeout=7200,\n",
" model_garden_source_model_name=\"publishers/moonshotai/models/kimi-k3@kimi-k3\",\n",
" )\n",
" print(\n",
" f\"Deploying {model_name} on {machine_type} with\"\n",
" f\" {int(accelerator_count * multihost_gpu_node_count)} {accelerator_type}\"\n",
" \" GPU(s).\"\n",
" )\n",
"\n",
" creds, _ = auth.default()\n",
" auth_req = auth.transport.requests.Request()\n",
" creds.refresh(auth_req)\n",
"\n",
" url = f\"https://{REGION}-aiplatform.googleapis.com/ui/projects/{PROJECT_ID}/locations/{REGION}/endpoints/{endpoint.name}:deployModel\"\n",
" headers = {\n",
" \"Content-Type\": \"application/json\",\n",
" \"Authorization\": f\"Bearer {creds.token}\",\n",
" }\n",
" data = {\n",
" \"deployedModel\": {\n",
" \"model\": model.resource_name,\n",
" \"displayName\": model_name,\n",
" \"dedicatedResources\": {\n",
" \"machineSpec\": {\n",
" \"machineType\": machine_type,\n",
" \"multihostGpuNodeCount\": multihost_gpu_node_count,\n",
" \"acceleratorType\": accelerator_type,\n",
" \"acceleratorCount\": accelerator_count,\n",
" },\n",
" \"minReplicaCount\": 1,\n",
" \"maxReplicaCount\": 1,\n",
" },\n",
" \"system_labels\": {\n",
" \"NOTEBOOK_NAME\": \"model_garden_pytorch_kimi_k3_deployment.ipynb\",\n",
" \"NOTEBOOK_ENVIRONMENT\": common_util.get_deploy_source(),\n",
" },\n",
" },\n",
" }\n",
" if is_spot:\n",
" data[\"deployedModel\"][\"dedicatedResources\"][\"spot\"] = True\n",
" response = requests.post(url, headers=headers, json=data)\n",
" print(f\"Deploy Model response: {response.json()}\")\n",
" if response.status_code != 200 or \"name\" not in response.json():\n",
" raise ValueError(f\"Failed to deploy model: {response.text}\")\n",
" poll_and_wait(response.json()[\"name\"], 7200)\n",
" print(\"endpoint_name:\", endpoint.name)\n",
"\n",
" return model, endpoint\n",
"\n",
"\n",
"models[\"sglang_gpu\"], endpoints[\"sglang_gpu\"] = deploy_model_kimi_k3_sglang(\n",
" model_name=common_util.get_job_name_with_datetime(prefix=\"kimi-k3-serve\"),\n",
" base_model_path=base_model_path,\n",
" speculative_draft_model_path=speculative_draft_model_path,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" multihost_gpu_node_count=multihost_gpu_node_count,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" is_spot=is_spot,\n",
")"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "28f2ac72",
"metadata": {
"cellView": "form",
"id": "111c3738fa0d"
},
"outputs": [],
"source": [
"# @title Raw predict\n",
"\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint with text prompts. Sampling parameters supported by SGLang can be found [here](https://docs.sglang.ai/backend/sampling_params.html).\n",
"\n",
"# @markdown Example:\n",
"\n",
"# @markdown ```\n",
"# @markdown User: What is the best way to diagnose and fix a flickering light in my house?\n",
"# @markdown Assistant: Okay, so I need to figure out how to diagnose and fix a flickering light in my house. Hmm, where do I start? Let's think. First, I remember that flickering lights can be caused by various issues. Maybe the bulb is loose? That's a common problem. Let me start with the simplest things first.\n",
"# @markdown ```\n",
"# @markdown Additionally, you can moderate the generated text with Vertex AI. See [Moderate text documentation](https://cloud.google.com/natural-language/docs/moderating-text) for more details.\n",
"\n",
"# Loads an existing endpoint instance using the endpoint name:\n",
"# - Using `endpoint_name = endpoint.name` allows us to get the\n",
"# endpoint name of the endpoint `endpoint` created in the cell\n",
"# above.\n",
"# - Alternatively, you can set `endpoint_name = \"1234567890123456789\"` to load\n",
"# an existing endpoint with the ID 1234567890123456789.\n",
"# You may uncomment the code below to load an existing endpoint.\n",
"\n",
"# endpoint_name = \"\" # @param {type:\"string\"}\n",
"# aip_endpoint_name = (\n",
"# f\"projects/{PROJECT_ID}/locations/{REGION}/endpoints/{endpoint_name}\"\n",
"# )\n",
"# endpoint = aiplatform.Endpoint(aip_endpoint_name)\n",
"\n",
"prompt = \"What is a car?\" # @param {type: \"string\"}\n",
"# @markdown If you encounter an issue like `ServiceUnavailable: 503 Took too long to respond when processing`, you can reduce the maximum number of output tokens, by lowering `max_tokens`.\n",
"max_new_tokens = 1024 # @param {type:\"integer\"}\n",
"temperature = 0.6 # @param {type:\"number\"}\n",
"top_p = 0.95 # @param {type:\"number\"}\n",
"\n",
"# Overrides parameters for inferences.\n",
"instances = [{\"text\": prompt}]\n",
"parameters = {\n",
" \"sampling_params\": {\n",
" \"max_new_tokens\": max_new_tokens,\n",
" \"temperature\": temperature,\n",
" \"top_p\": top_p,\n",
" }\n",
"}\n",
"\n",
"response = endpoints[\"sglang_gpu\"].predict(\n",
" instances=instances, use_dedicated_endpoint=use_dedicated_endpoint\n",
")\n",
"\n",
"for prediction in response.predictions:\n",
" print(prediction)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "markdown",
"id": "e2f22420",
"metadata": {
"id": "32884c1e7bcd"
},
"source": [
"## Clean up resources"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "20c1aaf8",
"metadata": {
"cellView": "form",
"id": "d57a7d442f24"
},
"outputs": [],
"source": [
"# @title Delete the models and endpoints\n",
"\n",
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continuous charges that may incur.\n",
"\n",
"# Undeploy model and delete endpoint.\n",
"for endpoint in endpoints.values():\n",
" endpoint.delete(force=True)\n",
"\n",
"# Delete models.\n",
"for model in models.values():\n",
" model.delete()\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_NAME"
]
}
],
"metadata": {
"colab": {
"name": "model_garden_pytorch_kimi_k3_deployment.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
File diff suppressed because it is too large Load Diff
@@ -1,682 +0,0 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "7d9bbf86da5e"
},
"outputs": [],
"source": [
"# Copyright 2024 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",
"# https://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": "99c1c3fc2ca5"
},
"source": [
"# Vertex AI Model Garden - LLaMA2 (PEFT Finetuning)\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_pytorch_llama2_peft_finetuning.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_pytorch_llama2_peft_finetuning.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "3de7470326a2"
},
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates downloading [LLaMA2 models](https://huggingface.co/meta-llama), finetuning with parameter efficient finetuning libraries ([PEFT](https://github.com/huggingface/peft)), and deploying the finetuned model on Vertex AI.\n",
"\n",
"### Objective\n",
"\n",
"- Download prebuilt LLaMA2 models.\n",
"- Finetune and deploy LLaMA2 models with Vertex AI SDK.\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing) and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "264c07757582"
},
"source": [
"## Run the notebook"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "QJgmw34Xwctp"
},
"outputs": [],
"source": [
"# @title (Optional) Finetune with Vertex AI Pipeline\n",
"\n",
"# @markdown Vertex Model Garden offers a pre-configured pipeline that can be launched from the UI, which will fine-tune, evaluate, upload, and deploy your desired LLaMA2 model.\n",
"# @markdown This pipeline currently supports [huggingface datasets](https://huggingface.co/datasets) for finetuning.\n",
"\n",
"# @markdown To launch a LLaMA 2 finetuning pipeline, open the [LLaMA 2 model card](https://console.cloud.google.com/vertex-ai/publishers/google/model-garden/139) and click the \"FINE-TUNE\" button.\n",
"# @markdown Then, click \"CREATE RUN\" button near the top of the pipeline details page, and follow the instructions to fill in pipeline parameters.\n",
"\n",
"# @markdown Learn about [Vertex AI Pipelines](https://cloud.google.com/vertex-ai/docs/pipelines/introduction)."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "IoPsYDwDdFBf"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. [Optional] [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"# @markdown 3. If you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus).\n",
"\n",
"# @markdown > | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
"# @markdown | a2-ultragpu-1g | 1 NVIDIA_A100_80GB | us-central1, us-east4, europe-west4, asia-southeast1, us-east4 |\n",
"# @markdown | a3-highgpu-2g | 2 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-4g | 4 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, us-east5, europe-west4, us-west1, asia-southeast1 |\n",
"\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform>=1.64.0'\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"# Import the necessary packages\n",
"import importlib\n",
"import os\n",
"import uuid\n",
"from datetime import datetime\n",
"from typing import Tuple\n",
"\n",
"from google.cloud import aiplatform\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.notebooks.community.model_garden.docker_source_codes.notebook_util.common_util\"\n",
")\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, change the value below.\n",
"now = datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}\"\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"EXPERIMENT_BUCKET = os.path.join(BUCKET_URI, \"peft\")\n",
"MODEL_BUCKET = os.path.join(BUCKET_URI, \"llama2\")\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default BUCKET_URI and SERVICE_ACCOUNT if they were not specified by the user.\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"! gcloud config set project $PROJECT_ID"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "zI30m3bqDtCj"
},
"outputs": [],
"source": [
"# @title Access LLaMA2 models on Vertex AI for GPU based serving\n",
"# @markdown The original models from Meta are converted into the Hugging Face format for serving in Vertex AI.\n",
"# @markdown Accept the model agreement to access the models:\n",
"# @markdown 1. Open the [LLaMA2 model card](https://console.cloud.google.com/vertex-ai/publishers/google/model-garden/139) from [Vertex AI Model Garden](https://cloud.google.com/model-garden).\n",
"# @markdown 2. Review and accept the agreement in the pop-up window on the model card page. If you have previously accepted the model agreement, there will not be a pop-up window on the model card page and this step is not needed.\n",
"# @markdown 3. A Cloud Storage bucket (starting with `gs://`) containing LLaMA 2 pretrained and finetuned models will be shared under the “Documentation” section and its “Get started” subsection.\n",
"\n",
"\n",
"VERTEX_AI_MODEL_GARDEN_LLAMA2 = \"\" # @param {type:\"string\", isTemplate:true}\n",
"assert (\n",
" VERTEX_AI_MODEL_GARDEN_LLAMA2\n",
"), \"Model artifact path is required. Click the agreement of LLaMA2 in Vertex AI Model Garden, and get the GCS path of LLaMA2 model artifacts.\"\n",
"print(\n",
" \"Copying LLaMA2 model artifacts from\",\n",
" VERTEX_AI_MODEL_GARDEN_LLAMA2,\n",
" \"to\",\n",
" MODEL_BUCKET,\n",
")\n",
"\n",
"! gsutil -m cp -R $VERTEX_AI_MODEL_GARDEN_LLAMA2/* $MODEL_BUCKET\n",
"\n",
"# The pre-built serving and training docker images.\n",
"VLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-vllm-serve:20240326_0916_RC00\"\n",
"TRAIN_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-peft-train:20240321_0936_RC00\"\n",
"\n",
"\n",
"def get_job_name_with_datetime(prefix: str) -> str:\n",
" \"\"\"Gets the job name with date time when triggering training or deployment\n",
" jobs in Vertex AI.\n",
" \"\"\"\n",
" return prefix + datetime.now().strftime(\"_%Y%m%d_%H%M%S\")\n",
"\n",
"\n",
"def deploy_model_vllm(\n",
" model_name: str,\n",
" model_id: str,\n",
" service_account: str,\n",
" machine_type: str = \"g2-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
" max_model_len: int = 4096,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys trained models with vLLM into Vertex AI.\"\"\"\n",
" common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=False,\n",
" )\n",
"\n",
" endpoint = aiplatform.Endpoint.create(display_name=f\"{model_name}-endpoint\")\n",
"\n",
" vllm_args = [\n",
" \"--host=0.0.0.0\",\n",
" \"--port=7080\",\n",
" f\"--tensor-parallel-size={accelerator_count}\",\n",
" \"--swap-space=16\",\n",
" \"--gpu-memory-utilization=0.8\",\n",
" f\"--max-model-len={max_model_len}\",\n",
" \"--max-num-batched-tokens=4096\",\n",
" \"--disable-log-stats\",\n",
" ]\n",
"\n",
" env_vars = {\"MODEL_ID\": model_id, \"DEPLOY_SOURCE\": \"notebook\"}\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=VLLM_DOCKER_URI,\n",
" serving_container_command=[\"python\", \"-m\", \"vllm.entrypoints.api_server\"],\n",
" serving_container_args=vllm_args,\n",
" serving_container_ports=[7080],\n",
" serving_container_predict_route=\"/generate\",\n",
" serving_container_health_route=\"/ping\",\n",
" serving_container_environment_variables=env_vars,\n",
" artifact_uri=model_id,\n",
" model_garden_source_model_name=\"publishers/meta/models/llama2\"\n",
" )\n",
" print(\n",
" f\"Deploying {model_name} on {machine_type} with {accelerator_count} {accelerator_type} GPU(s).\"\n",
" )\n",
" model.deploy(\n",
" endpoint=endpoint,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" deploy_request_timeout=1800,\n",
" service_account=service_account,\n",
" system_labels={\n",
" \"NOTEBOOK_NAME\": \"model_garden_pytorch_llama2_peft_finetuning.ipynb\"\n",
" },\n",
" )\n",
" print(\"endpoint_name:\", endpoint.name)\n",
"\n",
" print(\"To load this existing endpoint from a different session:\")\n",
" print(\"from google.cloud import aiplatform\")\n",
" print(\n",
" f'endpoint = aiplatform.Endpoint(\"projects/{PROJECT_ID}/locations/{REGION}/endpoints/{endpoint.name}\")'\n",
" )\n",
" return model, endpoint"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "Ax6GlzzMk-sc"
},
"outputs": [],
"source": [
"# @title Set training dataset\n",
"\n",
"# @markdown This notebook uses [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) dataset as an example.\n",
"# @markdown You can set `dataset_name` to any existing [Hugging Face dataset](https://huggingface.co/datasets) name, and set `instruct_column_in_dataset` to the name of the dataset column containing training data. The [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) has only one column `text`, and therefore we set `instruct_column_in_dataset` to `text` in this notebook.\n",
"\n",
"# @markdown #### (Optional) Prepare a custom JSONL dataset for finetuning\n",
"\n",
"# @markdown You can prepare a JSONL file where each line is a valid JSON string as your custom training dataset. For example, here is one line from the [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) dataset:\n",
"# @markdown ```\n",
"# @markdown {\"text\": \"### Human: Hola### Assistant: \\u00a1Hola! \\u00bfEn qu\\u00e9 puedo ayudarte hoy?\"}\n",
"# @markdown ```\n",
"\n",
"# @markdown The JSON object has a key `text`, which should match `instruct_column_in_dataset`; The value should be one training data point, i.e. a string. After you prepared your JSONL file, you can either upload it to [Hugging Face datasets](https://huggingface.co/datasets) or [Google Cloud Storage](https://cloud.google.com/storage).\n",
"\n",
"# @markdown - To upload a JSONL dataset to [Hugging Face datasets](https://huggingface.co/datasets), follow the instructions on [Uploading Datasets](https://huggingface.co/docs/hub/en/datasets-adding). Then, set `dataset_name` to the name of your newly created dataset on Hugging Face.\n",
"\n",
"# @markdown - To upload a JSONL dataset to [Google Cloud Storage](https://cloud.google.com/storage), follow the instructions on [Upload objects from a filesystem](https://cloud.google.com/storage/docs/uploading-objects). Then, set `dataset_name` to the `gs://` URI to your JSONL file. For example: `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`.\n",
"\n",
"# @markdown Optionally update the `instruct_column_in_dataset` field below if your JSON objects use a key other than the default `text`.\n",
"\n",
"# @markdown #### (Optional) Format your data with custom JSON template\n",
"\n",
"# @markdown Sometimes, your dataset might have multiple text columns and you want to construct the training data with a template. You can prepare a JSON template in the following format:\n",
"\n",
"# @markdown ```\n",
"# @markdown {\n",
"# @markdown \"description\": \"A short template for vertex sample dataset.\",\n",
"# @markdown \"prompt_input\": \"{input_text}{output_text}\",\n",
"# @markdown \"prompt_no_input\": \"{input_text}{output_text}\"\n",
"# @markdown }\n",
"# @markdown ```\n",
"\n",
"# @markdown As an example, the template above can be used to format the following training data (this line comes from `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`):\n",
"\n",
"# @markdown ```\n",
"# @markdown {\"input_text\":\"TRANSCRIPT: \\nREASON FOR EVALUATION:,\\n\\n LABEL:\",\"output_text\":\"Chiropractic\"}\n",
"# @markdown ```\n",
"\n",
"# @markdown This example template simply concatenates `input_text` with `output_text`. You can set `template` to `vertex_sample` to try out this built-in template with the dataset `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`, or build more complicated JSON templates such as [the alpaca example](https://github.com/tloen/alpaca-lora/blob/main/templates/alpaca.json). To use your own JSON template, [upload it to Google Cloud Storage](https://cloud.google.com/storage/docs/uploading-objects) and put the `gs://` URI in the `template` field below. Leave `instruct_column_in_dataset` as `text`.\n",
"\n",
"# Hugging Face dataset name or gs:// URI to a custom JSONL dataset.\n",
"dataset_name = \"timdettmers/openassistant-guanaco\" # @param {type:\"string\"}\n",
"\n",
"# Name of the dataset column containing training text input.\n",
"instruct_column_in_dataset = \"text\" # @param {type:\"string\"}\n",
"\n",
"# Optional. Template name or gs:// URI to a custom template.\n",
"template = \"\" # @param {type:\"string\"}"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "e1289e21a9d3"
},
"outputs": [],
"source": [
"# @title Finetune with PEFT\n",
"\n",
"# @markdown This section demonstrates how to finetune the LLaMA 2 models with PEFT LoRA.\n",
"\n",
"# @markdown By default, the model will be finetuned for 500 steps on a batch size of 1 to save GPU resources.\n",
"# @markdown Finetuning `llama2-7b` models is expected to take around 30 minutes.\n",
"# @markdown To customize finetuning settings and parameters, click \"Show code\" to see more details.\n",
"\n",
"# @markdown Set the base model id.\n",
"base_model_id = \"llama2-7b-hf\" # @param [\"llama2-7b-hf\", \"llama2-7b-chat-hf\", \"llama2-13b-hf\", \"llama2-13b-chat-hf\", \"llama2-70b-hf\", \"llama2-70b-chat-hf\"]\n",
"model_id = os.path.join(MODEL_BUCKET, base_model_id)\n",
"\n",
"# @markdown Set the accelerator type.\n",
"accelerator_type = \"NVIDIA_TESLA_V100\" # @param [\"NVIDIA_TESLA_V100\", \"NVIDIA_L4\", \"NVIDIA_TESLA_A100\"]\n",
"\n",
"# Worker pool spec.\n",
"# Find Vertex AI supported accelerators and regions in:\n",
"# https://cloud.google.com/vertex-ai/docs/training/configure-compute\n",
"machine_type = None\n",
"if \"7b\" in model_id:\n",
" if accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 1\n",
" elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-standard-16\"\n",
" accelerator_count = 2\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-8\"\n",
" accelerator_count = 1\n",
"elif \"13b\" in model_id:\n",
" if accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 1\n",
" elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-standard-32\"\n",
" accelerator_count = 4\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-24\"\n",
" accelerator_count = 2\n",
"elif \"70b\" in model_id:\n",
" if accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-4g\"\n",
" accelerator_count = 4\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-96\"\n",
" accelerator_count = 8\n",
"\n",
"if machine_type is None:\n",
" raise ValueError(\n",
" f\"Recommended machine settings not found for: {accelerator_type}. To use another another accelerator, edit this code block to set an appropriate `machine_type`, `accelerator_type`, and `accelerator_count` in worker_pool_specs.\"\n",
" )\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=True,\n",
")\n",
"\n",
"job_name = get_job_name_with_datetime(\"llama2-train\")\n",
"output_dir = os.path.join(EXPERIMENT_BUCKET, job_name)\n",
"merge_job_name = get_job_name_with_datetime(\"llama2-merge\")\n",
"merged_model_output_dir = os.path.join(EXPERIMENT_BUCKET, merge_job_name)\n",
"finetune_precision_mode = \"float16\"\n",
"\n",
"replica_count = 1\n",
"\n",
"# Runs 500 training steps.\n",
"max_steps = 500 # @param {type: \"integer\"}\n",
"per_device_train_batch_size = 1\n",
"# LoRA parameters.\n",
"lora_rank = 16 # @param {type: \"integer\"}\n",
"lora_alpha = 32\n",
"lora_dropout = 0.05\n",
"\n",
"flags = {\n",
" \"learning_rate\": 2e-4,\n",
" \"precision_mode\": finetune_precision_mode,\n",
" \"task\": \"instruct-lora\",\n",
" \"per_device_train_batch_size\": per_device_train_batch_size,\n",
" \"dataset_name\": dataset_name,\n",
" \"instruct_column_in_dataset\": instruct_column_in_dataset,\n",
" \"template\": template,\n",
" \"pretrained_model_id\": model_id,\n",
" \"output_dir\": output_dir,\n",
" \"merge_base_and_lora_output_dir\": merged_model_output_dir,\n",
" \"warmup_steps\": 10,\n",
" \"max_steps\": max_steps,\n",
" \"lora_rank\": lora_rank,\n",
" \"lora_alpha\": lora_alpha,\n",
" \"lora_dropout\": lora_dropout,\n",
"}\n",
"\n",
"\n",
"# Add labels for the finetuning job.\n",
"labels = {\n",
" \"mg-source\": \"notebook\",\n",
" \"mg-notebook-name\": \"model_garden_pytorch_llama2_peft_finetuning.ipynb\".split(\".\")[\n",
" 0\n",
" ],\n",
"}\n",
"\n",
"labels[\"mg-tune\"] = \"publishers-meta-models-llama2\"\n",
"versioned_model_id = base_model_id\n",
"labels[\"versioned-mg-tune\"] = f\"{labels['mg-tune']}-{versioned_model_id}\"\n",
"\n",
"train_job = aiplatform.CustomJob(\n",
" display_name=job_name,\n",
" worker_pool_specs=[\n",
" {\n",
" \"machine_spec\": {\n",
" \"machine_type\": machine_type,\n",
" \"accelerator_type\": accelerator_type,\n",
" \"accelerator_count\": accelerator_count,\n",
" },\n",
" \"replica_count\": replica_count,\n",
" \"container_spec\": {\n",
" \"image_uri\": TRAIN_DOCKER_URI,\n",
" \"args\": [\"--{}={}\".format(k, v) for k, v in flags.items()],\n",
" },\n",
" }\n",
" ],\n",
" staging_bucket=STAGING_BUCKET,\n",
" labels=labels,\n",
")\n",
"train_job.run()\n",
"\n",
"print(\"The finetuned models of different trials can be found at: \", output_dir)\n",
"print(\n",
" \"The finetuned model merged with the base model can be found at: \",\n",
" merged_model_output_dir,\n",
")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "h0hGj09CuRFQ"
},
"outputs": [],
"source": [
"# @title Deploy\n",
"# @markdown This section uploads the model to Model Registry and deploys it on the Endpoint. It takes 15 minutes to 1 hour to finish.\n",
"# @markdown Click \"Show code\" to see more details.\n",
"\n",
"print(\"Deploying models in: \", merged_model_output_dir)\n",
"\n",
"# The max_model_len must not exceed the model's context length.\n",
"# A larger max_model_len will require more GPU memory.\n",
"max_model_len = 2048\n",
"# Worker pool spec.\n",
"# Find Vertex AI supported accelerators and regions in:\n",
"# https://cloud.google.com/vertex-ai/docs/training/configure-compute\n",
"machine_type = None\n",
"accelerator_type = \"NVIDIA_L4\" # @param [\"NVIDIA_L4\", \"NVIDIA_TESLA_V100\", \"NVIDIA_TESLA_A100\", \"NVIDIA_H100_80GB\"]\n",
"\n",
"if \"7b\" in model_id:\n",
" if accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 1\n",
" elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-standard-8\"\n",
" accelerator_count = 1\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-8\"\n",
" accelerator_count = 1\n",
"elif \"13b\" in model_id:\n",
" if accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 1\n",
" elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-standard-16\"\n",
" accelerator_count = 2\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-24\"\n",
" accelerator_count = 2\n",
"elif \"70b\" in model_id:\n",
" if accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-4g\"\n",
" accelerator_count = 4\n",
" elif accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-96\"\n",
" accelerator_count = 8\n",
" elif accelerator_type == \"NVIDIA_H100_80GB\":\n",
" machine_type = \"a3-highgpu-4g\"\n",
" accelerator_count = 4\n",
"\n",
"if machine_type is None:\n",
" raise ValueError(\n",
" f\"Recommended machine settings not found for: {accelerator_type}. To use another another accelerator, edit this code block to pass in an appropriate `machine_type`, `accelerator_type`, and `accelerator_count` to the deploy_model_vllm function.\"\n",
" )\n",
"\n",
"model, endpoint = deploy_model_vllm(\n",
" model_name=get_job_name_with_datetime(prefix=\"llama-vllm-serve\"),\n",
" model_id=merged_model_output_dir,\n",
" service_account=SERVICE_ACCOUNT,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" max_model_len=max_model_len,\n",
")\n",
"print(\"endpoint_name:\", endpoint.name)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "vgU_qYHNuy3w"
},
"outputs": [],
"source": [
"# @title Predict\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint with text prompts.\n",
"\n",
"# @markdown Here we use an example from the [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) to show the finetuning outcome:\n",
"\n",
"# @markdown ```\n",
"# @markdown ### Human: How would the Future of AI in 10 Years look?### Assistant: Predicting the future is always a challenging task, but here are some possible ways that AI could evolve over the next 10 years: Continued advancements in deep learning: Deep learning has been one of the main drivers of recent AI breakthroughs, and we can expect continued advancements in this area. This may include improvements to existing algorithms, as well as the development of new architectures that are better suited to specific types of data and tasks. Increased use of AI in healthcare: AI has the potential to revolutionize healthcare, by improving the accuracy of diagnoses, developing new treatments, and personalizing patient care. We can expect to see continued investment in this area, with more healthcare providers and researchers using AI to improve patient outcomes. Greater automation in the workplace: Automation is already transforming many industries, and AI is likely to play an increasingly important role in this process. We can expect to see more jobs being automated, as well as the development of new types of jobs that require a combination of human and machine skills. More natural and intuitive interactions with technology: As AI becomes more advanced, we can expect to see more natural and intuitive ways of interacting with technology. This may include voice and gesture recognition, as well as more sophisticated chatbots and virtual assistants. Increased focus on ethical considerations: As AI becomes more powerful, there will be a growing need to consider its ethical implications. This may include issues such as bias in AI algorithms, the impact of automation on employment, and the use of AI in surveillance and policing. Overall, the future of AI in 10 years is likely to be shaped by a combination of technological advancements, societal changes, and ethical considerations. While there are many exciting possibilities for AI in the future, it will be important to carefully consider its potential impact on society and to work towards ensuring that its benefits are shared fairly and equitably.\n",
"# @markdown ```\n",
"\n",
"# @markdown Additionally, you can moderate the generated text with Vertex AI. See [Moderate text documentation](https://cloud.google.com/natural-language/docs/moderating-text) for more details.\n",
"\n",
"# Loads an existing endpoint instance using the endpoint name:\n",
"# - Using `endpoint_name = endpoint.name` allows us to get the\n",
"# endpoint name of the endpoint `endpoint` created in the cell\n",
"# above.\n",
"# - Alternatively, you can set `endpoint_name = \"1234567890123456789\"` to load\n",
"# an existing endpoint with the ID 1234567890123456789.\n",
"# You may uncomment the code below to load an existing endpoint.\n",
"\n",
"# endpoint_name = \"\" # @param {type:\"string\"}\n",
"# aip_endpoint_name = (\n",
"# f\"projects/{PROJECT_ID}/locations/{REGION}/endpoints/{endpoint_name}\"\n",
"# )\n",
"# endpoint = aiplatform.Endpoint(aip_endpoint_name)\n",
"\n",
"prompt = \"How would the Future of AI in 10 Years look?\" # @param {type: \"string\"}\n",
"max_tokens = 128 # @param {type:\"integer\"}\n",
"temperature = 1.0 # @param {type:\"number\"}\n",
"top_p = 0.9 # @param {type:\"number\"}\n",
"top_k = 1 # @param {type:\"integer\"}\n",
"\n",
"# Overrides max_tokens and top_k parameters during inferences.\n",
"# If you encounter the issue like `ServiceUnavailable: 503 Took too long to respond when processing`,\n",
"# you can reduce the max length, such as set max_tokens as 20.\n",
"instances = [\n",
" {\n",
" \"prompt\": f\"### Human: {prompt}### Assistant: \",\n",
" \"max_tokens\": max_tokens,\n",
" \"temperature\": temperature,\n",
" \"top_p\": top_p,\n",
" \"top_k\": top_k,\n",
" },\n",
"]\n",
"response = endpoint.predict(instances=instances)\n",
"\n",
"for prediction in response.predictions:\n",
" print(prediction)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "911406c1561e"
},
"outputs": [],
"source": [
"# @title Clean up resources\n",
"\n",
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continuous charges that may incur.\n",
"\n",
"if train_job._gca_resource.name:\n",
" # Training job is submitted.\n",
" train_job.delete()\n",
"\n",
"# Undeploy model and delete endpoint.\n",
"endpoint.delete(force=True)\n",
"\n",
"# Delete model.\n",
"model.delete()\n",
"\n",
"# Delete Cloud Storage objects that were created.\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_URI"
]
}
],
"metadata": {
"colab": {
"name": "model_garden_pytorch_llama2_peft_finetuning.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -1,857 +0,0 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "7d9bbf86da5e"
},
"outputs": [],
"source": [
"# Copyright 2026 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",
"# https://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": "99c1c3fc2ca5"
},
"source": [
"# Vertex AI Model Garden - Llama 3 Finetuning\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_pytorch_llama3_finetuning.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_pytorch_llama3_finetuning.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "3de7470326a2"
},
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates finetuning and deploying Llama 3 models with Vertex AI. All of the examples in this notebook use parameter efficient finetuning methods [PEFT (LoRA)](https://github.com/huggingface/peft) to reduce training and storage costs. LoRA (Low-Rank Adaptation) is one approach of Parameter Efficient FineTuning (PEFT), where pretrained model weights are frozen and rank decomposition matrices representing the change in model weights are trained during finetuning. Read more about LoRA in the following publication: [Hu, E.J., Shen, Y., Wallis, P., Allen-Zhu, Z., Li, Y., Wang, S., Wang, L. and Chen, W., 2021. Lora: Low-rank adaptation of large language models. *arXiv preprint arXiv:2106.09685*](https://arxiv.org/abs/2106.09685).\n",
"\n",
"After finetuning, we can deploy models on Vertex with GPU.\n",
"\n",
"\n",
"### Objective\n",
"\n",
"- Finetune Llama 3 models with Vertex AI Custom Training Jobs.\n",
"- Deploy finetuned Llama 3 models on Vertex AI Prediction.\n",
"- Send prediction requests to your finetuned Llama 3 models.\n",
"\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "264c07757582"
},
"source": [
"## Before you begin"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "855d6b96f291"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. **[Optional]** [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 3. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 4. If you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus). You can request for quota following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota).\n",
"\n",
"# @markdown | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
"# @markdown | a2-ultragpu-1g | 1 NVIDIA_A100_80GB | us-central1, us-east4, europe-west4, asia-southeast1, us-east4 |\n",
"# @markdown | a3-highgpu-2g | 2 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-4g | 4 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"\n",
"# @markdown 5. [Make sure that you have GPU quota for Vertex Training (finetuing) and Vertex Prediction (serving)](https://cloud.google.com/docs/quotas/view-manage). The quota name for Vertex Training is \"Custom model training your-gpu-type per region\" and the quota name for Vertex Prediction is \"Custom model serving your-gpu-type per region\" such as `Custom model training Nvidia L4 GPUs per region` and `Custom model serving Nvidia L4 GPUs per region` for L4 GPUs. [Submit a quota increase request](https://cloud.google.com/docs/quotas/view-manage#requesting_higher_quota) if additional quota is needed. At minimum, running this notebook requires 4 L4s for finetuning and 1 L4 for serving. More GPUs may be needed for larger models and different finetuning configurations. To secure GPUs for larger models, ask your customer engineer to get you allowlisted for a Shared Reservation or a Dynamic Workload Scheduler.\n",
"\n",
"# Import the necessary packages\n",
"# Upgrade Vertex AI SDK.\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform>=1.64.0'\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"import datetime\n",
"import importlib\n",
"import os\n",
"import uuid\n",
"from typing import Tuple\n",
"\n",
"from google.cloud import aiplatform\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.notebooks.community.model_garden.docker_source_codes.notebook_util.common_util\"\n",
")\n",
"\n",
"models, endpoints = {}, {}\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"if not REGION:\n",
" if not os.environ.get(\"GOOGLE_CLOUD_REGION\"):\n",
" raise ValueError(\n",
" \"REGION must be set. See\"\n",
" \" https://cloud.google.com/vertex-ai/docs/general/locations for\"\n",
" \" available cloud locations.\"\n",
" )\n",
" REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, change the value yourself below.\n",
"now = datetime.datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}\"\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"MODEL_BUCKET = os.path.join(BUCKET_URI, \"llama3\")\n",
"\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default SERVICE_ACCOUNT.\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/storage.admin\"\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/aiplatform.user\""
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "36c21f10355f"
},
"outputs": [],
"source": [
"# @title Access Llama 3 models\n",
"\n",
"# @markdown For GPU based finetuning and serving, choose between accessing Llama 3 models on [Hugging Face](https://huggingface.co/)\n",
"# @markdown or Vertex AI as described below.\n",
"\n",
"# @markdown If you already obtained access to Llama 3 models on [Hugging Face](https://huggingface.co/), you can load models from there.\n",
"# @markdown Alternatively, you can also load the original Llama 3 models for finetuning and serving from Vertex AI after accepting the agreement.\n",
"\n",
"# @markdown **Only select and fill one of the following sections.**\n",
"# fmt: off\n",
"LOAD_MODEL_FROM = \"Hugging Face\" # @param [\"Hugging Face\", \"Google Cloud\"] {isTemplate:true}\n",
"# fmt: on\n",
"\n",
"# @markdown ---\n",
"\n",
"# @markdown ### Access Llama 3 models on Hugging Face for GPU based finetuning and serving\n",
"# @markdown You must provide a Hugging Face User Access Token (with read access) to access the Llama 3 models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"\n",
"HF_TOKEN = \"\" # @param {type:\"string\", isTemplate:true}\n",
"if LOAD_MODEL_FROM == \"Hugging Face\":\n",
" assert (\n",
" HF_TOKEN\n",
" ), \"Provide a read HF_TOKEN to load models from Hugging Face, or select a different model source.\"\n",
"\n",
"# @markdown *--- Or ---*\n",
"# @markdown ### Access Llama 3 models on Vertex AI for GPU based serving\n",
"# @markdown The original models from Meta are converted into the Hugging Face format for serving in Vertex AI.\n",
"# @markdown Accept the model agreement to access the models:\n",
"# @markdown 1. Open the [Llama 3 model card](https://console.cloud.google.com/vertex-ai/publishers/meta/model-garden/llama3) from [Vertex AI Model Garden](https://cloud.google.com/model-garden).\n",
"# @markdown 2. Review and accept the agreement in the pop-up window on the model card page. If you have previously accepted the model agreement, there will not be a pop-up window on the model card page and this step is not needed.\n",
"# @markdown 3. After accepting the agreement of Llama 3, a `gs://` URI containing Llama 3 pretrained and finetuned models will be shared.\n",
"# @markdown 4. Paste the URI in the `VERTEX_AI_MODEL_GARDEN_LLAMA3` field below.\n",
"\n",
"VERTEX_AI_MODEL_GARDEN_LLAMA3 = \"\" # @param {type:\"string\", isTemplate:true}\n",
"\n",
"if LOAD_MODEL_FROM == \"Google Cloud\":\n",
" assert (\n",
" VERTEX_AI_MODEL_GARDEN_LLAMA3\n",
" ), \"Click the agreement of Llama 3 in Vertex AI Model Garden, and get the GCS path of Llama 3 model artifacts.\"\n",
" print(\n",
" \"Copying Llama 3 model artifacts from\",\n",
" VERTEX_AI_MODEL_GARDEN_LLAMA3,\n",
" \"to \",\n",
" MODEL_BUCKET,\n",
" )\n",
" HF_TOKEN = \"\"\n",
"\n",
" ! gsutil -m cp -R $VERTEX_AI_MODEL_GARDEN_LLAMA3/* $MODEL_BUCKET\n",
"\n",
"# @markdown ---"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "cb56d402e84a"
},
"source": [
"## Finetune with HuggingFace PEFT and deploy with vLLM on GPUs"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "KwAW99YZHTdy"
},
"outputs": [],
"source": [
"# @title Set dataset\n",
"\n",
"# @markdown Use the Vertex AI SDK to create and run the custom training jobs.\n",
"\n",
"# @markdown This notebook uses [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) dataset as an example.\n",
"# @markdown You can set `dataset_name` to any existing [Hugging Face dataset](https://huggingface.co/datasets) name, and set `instruct_column_in_dataset` to the name of the dataset column containing training data. The [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) has only one column `text`, and therefore we set `instruct_column_in_dataset` to `text` in this notebook.\n",
"\n",
"# @markdown ### (Optional) Prepare a custom JSONL dataset for finetuning\n",
"\n",
"# @markdown You can prepare a JSONL file where each line is a valid JSON string as your custom training dataset. For example, here is one line from the [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) dataset:\n",
"# @markdown ```\n",
"# @markdown {\"text\": \"### Human: Hola### Assistant: \\u00a1Hola! \\u00bfEn qu\\u00e9 puedo ayudarte hoy?\"}\n",
"# @markdown ```\n",
"\n",
"# @markdown The JSON object has a key `text`, which should match `instruct_column_in_dataset`; The value should be one training data point, i.e. a string. After you prepared your JSONL file, you can either upload it to [Hugging Face datasets](https://huggingface.co/datasets) or [Google Cloud Storage](https://cloud.google.com/storage).\n",
"\n",
"# @markdown - To upload a JSONL dataset to [Hugging Face datasets](https://huggingface.co/datasets), follow the instructions on [Uploading Datasets](https://huggingface.co/docs/hub/en/datasets-adding). Then, set `dataset_name` to the name of your newly created dataset on Hugging Face.\n",
"\n",
"# @markdown - To upload a JSONL dataset to [Google Cloud Storage](https://cloud.google.com/storage), follow the instructions on [Upload objects from a filesystem](https://cloud.google.com/storage/docs/uploading-objects). Then, set `dataset_name` to the `gs://` URI to your JSONL file. For example: `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`.\n",
"\n",
"# @markdown Optionally update the `instruct_column_in_dataset` field below if your JSON objects use a key other than the default `text`.\n",
"\n",
"# @markdown ### (Optional) Format your data with custom JSON template\n",
"\n",
"# @markdown Sometimes, your dataset might have multiple text columns and you want to construct the training data with a template. You can prepare a JSON template in the following format:\n",
"\n",
"# @markdown ```\n",
"# @markdown {\n",
"# @markdown \"description\": \"Template used by Llama 3, accepting text-bison format.\",\n",
"# @markdown \"source\": \"https://cloud.google.com/vertex-ai/generative-ai/docs/models/tune-text-models-supervised#dataset-format\",\n",
"# @markdown \"prompt_input\": \"<|start_header_id|>user<|end_header_id|>\\n\\n{input_text}<|eot_id|><|start_header_id|>assistant<|end_header_id|>\\n\\n{output_text}<|eot_id|>\",\n",
"# @markdown \"instruction_separator\": \"<|start_header_id|>user<|end_header_id|>\\n\\n\",\n",
"# @markdown \"response_separator\": \"<|start_header_id|>assistant<|end_header_id|>\\n\\n\"\n",
"# @markdown }\n",
"# @markdown ```\n",
"\n",
"# @markdown As an example, the template above can be used to format the following training data (this line comes from `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`):\n",
"\n",
"# @markdown ```\n",
"# @markdown {\"input_text\":\"TRANSCRIPT: \\nREASON FOR EVALUATION:,\\n\\n LABEL:\",\"output_text\":\"Chiropractic\"}\n",
"# @markdown ```\n",
"\n",
"# @markdown This example template simply concatenates `input_text` with `output_text` with some special tokens in between.\n",
"# @markdown\n",
"# @markdown To try such custom dataset, you can make the following changes:\n",
"# @markdown 1. Set `template` to `llama3-text-bison`\n",
"# @markdown 1. Set `train_dataset_name` to `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`\n",
"# @markdown 1. Set `train_split_name` to `train`\n",
"# @markdown 1. Set `eval_dataset_name` to `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_eval_sample.jsonl`\n",
"# @markdown 1. Set `eval_split_name` to `train` (**NOT** `test`)\n",
"# @markdown 1. Set `instruct_column_in_dataset` as `input_text`.\n",
"\n",
"# Template name or gs:// URI to a custom template.\n",
"template = \"openassistant-guanaco\" # @param {type:\"string\"}\n",
"\n",
"# Hugging Face dataset name or gs:// URI to a custom JSONL dataset.\n",
"train_dataset_name = \"timdettmers/openassistant-guanaco\" # @param {type:\"string\"}\n",
"train_split_name = \"train\" # @param {type:\"string\"}\n",
"eval_dataset_name = \"timdettmers/openassistant-guanaco\" # @param {type:\"string\"}\n",
"eval_split_name = \"test\" # @param {type:\"string\"}\n",
"\n",
"# Name of the dataset column containing training text input.\n",
"instruct_column_in_dataset = \"text\" # @param {type:\"string\"}"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "ivVGS9dHXPOz"
},
"outputs": [],
"source": [
"# @title Finetune\n",
"# @markdown Use the Vertex AI SDK to create and run the custom training jobs.\n",
"\n",
"# @markdown **Note**:\n",
"# @markdown 1. We recommend setting `finetuning_precision_mode` to `4bit` because it enables using fewer hardware resources for finetuning.\n",
"# @markdown 1. We recommend using NVIDIA_L4 for 8B models and NVIDIA_A100_80GB for 70B models.\n",
"# @markdown 1. If `max_steps>0`, it will precedence over `epochs`. One can set a small `max_steps` value to quickly check the pipeline.\n",
"# @markdown 1. With the default setting, training takes between 1.5 ~ 2 hours.\n",
"\n",
"TRAIN_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-peft-train:stable_20240909\"\n",
"\n",
"\n",
"# The Llama 3 base model.\n",
"MODEL_ID = \"meta-llama/Meta-Llama-3-8B-Instruct\" # @param [\"meta-llama/Meta-Llama-3-8B\", \"meta-llama/Meta-Llama-3-8B-Instruct\", \"meta-llama/Meta-Llama-3-70B\", \"meta-llama/Meta-Llama-3-70B-Instruct\"] {isTemplate:true}\n",
"if LOAD_MODEL_FROM == \"Google Cloud\":\n",
" if MODEL_ID == \"meta-llama/Meta-Llama-3-8B\":\n",
" base_model_id = \"llama3-8b-hf\"\n",
" elif MODEL_ID == \"meta-llama/Meta-Llama-3-8B-Instruct\":\n",
" base_model_id = \"llama3-8b-chat-hf\"\n",
" elif MODEL_ID == \"meta-llama/Meta-Llama-3-70B\":\n",
" base_model_id = \"llama3-70b-hf\"\n",
" elif MODEL_ID == \"meta-llama/Meta-Llama-3-70B-Instruct\":\n",
" base_model_id = \"llama3-70b-chat-hf\"\n",
" else:\n",
" raise ValueError(f\"Undefined model ID: {MODEL_ID}.\")\n",
" base_model_id = os.path.join(MODEL_BUCKET, base_model_id)\n",
"else:\n",
" base_model_id = MODEL_ID\n",
"\n",
"# The accelerator to use.\n",
"accelerator_type = \"NVIDIA_L4\" # @param [\"NVIDIA_L4\", \"NVIDIA_A100_80GB\"]\n",
"\n",
"# Batch size for finetuning.\n",
"per_device_train_batch_size = 1 # @param{type:\"integer\"}\n",
"gradient_accumulation_steps = 8 # @param{type:\"integer\"}\n",
"# Maximum sequence length.\n",
"max_seq_length = 4096 # @param{type:\"integer\"}\n",
"# Setting a positive `max_steps` here will override `num_epochs`\n",
"max_steps = -1 # @param{type:\"integer\"}\n",
"num_epochs = 1.0 # @param{type:\"number\"}\n",
"# Precision mode for finetuning.\n",
"finetuning_precision_mode = \"4bit\" # @param [\"4bit\", \"8bit\", \"float16\"]\n",
"# Learning rate.\n",
"learning_rate = 5e-5 # @param{type:\"number\"}\n",
"lr_scheduler_type = \"cosine\" # @param{type:\"string\"}\n",
"# LoRA parameters.\n",
"lora_rank = 16 # @param{type:\"integer\"}\n",
"lora_alpha = 32 # @param{type:\"integer\"}\n",
"lora_dropout = 0.05 # @param{type:\"number\"}\n",
"enable_gradient_checkpointing = True\n",
"attn_implementation = \"flash_attention_2\"\n",
"optimizer = \"paged_adamw_32bit\"\n",
"warmup_ratio = \"0.01\"\n",
"report_to = \"tensorboard\"\n",
"save_steps = 10\n",
"logging_steps = save_steps\n",
"\n",
"# Worker pool spec.\n",
"machine_type = None\n",
"if \"8b\" in MODEL_ID.lower():\n",
" if accelerator_type == \"NVIDIA_L4\":\n",
" accelerator_count = 4\n",
" machine_type = \"g2-standard-48\"\n",
" else:\n",
" raise ValueError(\n",
" f\"Recommended machine settings not found for: {accelerator_type}. To use another accelerator, edit this code block to pass in an appropriate `machine_type`, `accelerator_type`, and `accelerator_count` to the deploy_model_vllm function by clicking `Show Code` and then modifying the code.\"\n",
" )\n",
"elif \"70b\" in MODEL_ID.lower():\n",
" if accelerator_type == \"NVIDIA_A100_80GB\":\n",
" accelerator_count = 4\n",
" machine_type = \"a2-ultragpu-4g\"\n",
" else:\n",
" raise ValueError(\n",
" f\"Recommended machine settings not found for: {accelerator_type}. To use another accelerator, edit this code block to pass in an appropriate `machine_type`, `accelerator_type`, and `accelerator_count` to the deploy_model_vllm function by clicking `Show Code` and then modifying the code.\"\n",
" )\n",
"else:\n",
" raise ValueError(f\"Unsupported model ID or GCS path: {MODEL_ID}.\")\n",
"\n",
"replica_count = 1\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=True,\n",
")\n",
"\n",
"job_name = common_util.get_job_name_with_datetime(\"llama3-lora-train\")\n",
"\n",
"base_output_dir = os.path.join(STAGING_BUCKET, job_name)\n",
"# Create a GCS folder to store the LORA adapter.\n",
"lora_output_dir = os.path.join(base_output_dir, \"adapter\")\n",
"# Create a GCS folder to store the merged model with the base model and the\n",
"# finetuned LORA adapter.\n",
"merged_model_output_dir = os.path.join(base_output_dir, \"merged-model\")\n",
"\n",
"# Add labels for the finetuning job.\n",
"labels = {\n",
" \"mg-source\": \"notebook\",\n",
" \"mg-notebook-name\": \"model_garden_pytorch_llama3_finetuning.ipynb\".split(\".\")[0],\n",
"}\n",
"\n",
"labels[\"mg-tune\"] = \"publishers-meta-models-llama3\"\n",
"versioned_model_id = base_model_id.split(\"/\")[1].lower().replace(\".\", \"-\")\n",
"labels[\"versioned-mg-tune\"] = f\"{labels['mg-tune']}-{versioned_model_id}\"\n",
"\n",
"eval_args = [\n",
" f\"--eval_dataset_path={eval_dataset_name}\",\n",
" f\"--eval_column={instruct_column_in_dataset}\",\n",
" f\"--eval_template={template}\",\n",
" f\"--eval_split={eval_split_name}\",\n",
" f\"--eval_steps={save_steps}\",\n",
" \"--eval_tasks=builtin_eval\",\n",
" \"--eval_metric_name=loss\",\n",
"]\n",
"\n",
"train_job_args = [\n",
" \"--config_file=vertex_vision_model_garden_peft/deepspeed_zero2_4gpu.yaml\",\n",
" \"--task=instruct-lora\",\n",
" \"--completion_only=True\",\n",
" f\"--pretrained_model_id={base_model_id}\",\n",
" f\"--dataset_name={train_dataset_name}\",\n",
" f\"--train_split_name={train_split_name}\",\n",
" f\"--instruct_column_in_dataset={instruct_column_in_dataset}\",\n",
" f\"--output_dir={lora_output_dir}\",\n",
" f\"--merge_base_and_lora_output_dir={merged_model_output_dir}\",\n",
" f\"--per_device_train_batch_size={per_device_train_batch_size}\",\n",
" f\"--gradient_accumulation_steps={gradient_accumulation_steps}\",\n",
" f\"--lora_rank={lora_rank}\",\n",
" f\"--lora_alpha={lora_alpha}\",\n",
" f\"--lora_dropout={lora_dropout}\",\n",
" f\"--max_steps={max_steps}\",\n",
" f\"--max_seq_length={max_seq_length}\",\n",
" f\"--learning_rate={learning_rate}\",\n",
" f\"--lr_scheduler_type={lr_scheduler_type}\",\n",
" f\"--precision_mode={finetuning_precision_mode}\",\n",
" f\"--enable_gradient_checkpointing={enable_gradient_checkpointing}\",\n",
" f\"--num_epochs={num_epochs}\",\n",
" f\"--attn_implementation={attn_implementation}\",\n",
" f\"--optimizer={optimizer}\",\n",
" f\"--warmup_ratio={warmup_ratio}\",\n",
" f\"--report_to={report_to}\",\n",
" f\"--logging_output_dir={base_output_dir}\",\n",
" f\"--save_steps={save_steps}\",\n",
" f\"--logging_steps={logging_steps}\",\n",
" f\"--template={template}\",\n",
" f\"--huggingface_access_token={HF_TOKEN}\",\n",
"] + eval_args\n",
"\n",
"# Create TensorBoard\n",
"tensorboard = aiplatform.Tensorboard.create(job_name)\n",
"exp = aiplatform.TensorboardExperiment.create(\n",
" tensorboard_experiment_id=job_name, tensorboard_name=tensorboard.name\n",
")\n",
"\n",
"# Pass training arguments and launch job.\n",
"train_job = aiplatform.CustomContainerTrainingJob(\n",
" display_name=job_name,\n",
" container_uri=TRAIN_DOCKER_URI,\n",
" labels=labels,\n",
")\n",
"\n",
"train_job.run(\n",
" args=train_job_args,\n",
" environment_variables={\"WANDB_DISABLED\": True},\n",
" replica_count=replica_count,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" boot_disk_size_gb=500,\n",
" service_account=SERVICE_ACCOUNT,\n",
" tensorboard=tensorboard.resource_name,\n",
" base_output_dir=base_output_dir,\n",
")\n",
"\n",
"print(\"LoRA adapter was saved in: \", lora_output_dir)\n",
"print(\"Trained and merged models were saved in: \", merged_model_output_dir)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "qmHW6m8xG_4U"
},
"outputs": [],
"source": [
"# @title Deploy\n",
"# @markdown This section uploads the model to Model Registry and deploys it on the Endpoint. It takes 15 minutes to 1 hour to finish.\n",
"\n",
"print(\"Deploying models in: \", merged_model_output_dir)\n",
"\n",
"# The pre-built serving docker image for vLLM.\n",
"VLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-vllm-serve:20240721_0916_RC00\"\n",
"\n",
"accelerator_type = \"NVIDIA_H100_80GB\" # @param [\"NVIDIA_L4\", \"NVIDIA_H100_80GB\"]\n",
"machine_type = None\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"# Find Vertex AI prediction supported accelerators and regions in [here](https://cloud.google.com/vertex-ai/docs/predictions/configure-compute).\n",
"if \"8b\" in MODEL_ID.lower():\n",
" if accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-12\"\n",
" accelerator_count = 1\n",
" elif accelerator_type == \"NVIDIA_H100_80GB\":\n",
" machine_type = \"a3-highgpu-2g\"\n",
" accelerator_count = 2\n",
"else:\n",
" if accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-96\"\n",
" accelerator_count = 8\n",
" elif accelerator_type == \"NVIDIA_H100_80GB\":\n",
" machine_type = \"a3-highgpu-4g\"\n",
" accelerator_count = 4\n",
"\n",
"if machine_type is None:\n",
" raise ValueError(\n",
" f\"Recommended GPU setting not found for: {accelerator_type} and {MODEL_ID.lower()}.\"\n",
" )\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=False,\n",
")\n",
"\n",
"gpu_memory_utilization = 0.85\n",
"max_model_len = 8192 # Maximum context length.\n",
"\n",
"# Ensure max_model_len does not exceed the limit\n",
"if max_model_len > 8192:\n",
" raise ValueError(\"max_model_len cannot exceed 8192\")\n",
"\n",
"\n",
"def deploy_model_vllm(\n",
" model_name: str,\n",
" model_id: str,\n",
" publisher: str,\n",
" publisher_model_id: str,\n",
" service_account: str,\n",
" base_model_id: str = None,\n",
" machine_type: str = \"g2-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
" gpu_memory_utilization: float = 0.9,\n",
" max_model_len: int = 4096,\n",
" dtype: str = \"auto\",\n",
" enable_trust_remote_code: bool = False,\n",
" enforce_eager: bool = False,\n",
" enable_lora: bool = False,\n",
" enable_chunked_prefill: bool = False,\n",
" enable_prefix_cache: bool = False,\n",
" host_prefix_kv_cache_utilization_target: float = 0.0,\n",
" max_loras: int = 1,\n",
" max_cpu_loras: int = 8,\n",
" use_dedicated_endpoint: bool = False,\n",
" max_num_seqs: int = 256,\n",
" model_type: str = None,\n",
" enable_llama_tool_parser: bool = False,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys trained models with vLLM into Vertex AI.\"\"\"\n",
" endpoint = aiplatform.Endpoint.create(\n",
" display_name=f\"{model_name}-endpoint\",\n",
" dedicated_endpoint_enabled=use_dedicated_endpoint,\n",
" )\n",
"\n",
" if not base_model_id:\n",
" base_model_id = model_id\n",
"\n",
" # See https://docs.vllm.ai/en/latest/models/engine_args.html for a list of possible arguments with descriptions.\n",
" vllm_args = [\n",
" \"python\",\n",
" \"-m\",\n",
" \"vllm.entrypoints.api_server\",\n",
" \"--host=0.0.0.0\",\n",
" \"--port=8080\",\n",
" f\"--model={model_id}\",\n",
" f\"--tensor-parallel-size={accelerator_count}\",\n",
" \"--swap-space=16\",\n",
" f\"--max-model-len={max_model_len}\",\n",
" f\"--dtype={dtype}\",\n",
" f\"--max-loras={max_loras}\",\n",
" f\"--max-cpu-loras={max_cpu_loras}\",\n",
" f\"--max-num-seqs={max_num_seqs}\",\n",
" \"--disable-log-stats\",\n",
" ]\n",
"\n",
" if gpu_memory_utilization:\n",
" vllm_args.append(f\"--gpu-memory-utilization={gpu_memory_utilization}\")\n",
"\n",
" if enable_trust_remote_code:\n",
" vllm_args.append(\"--trust-remote-code\")\n",
"\n",
" if enforce_eager:\n",
" vllm_args.append(\"--enforce-eager\")\n",
"\n",
" if enable_lora:\n",
" vllm_args.append(\"--enable-lora\")\n",
"\n",
" if enable_chunked_prefill:\n",
" vllm_args.append(\"--enable-chunked-prefill\")\n",
"\n",
" if enable_prefix_cache:\n",
" vllm_args.append(\"--enable-prefix-caching\")\n",
"\n",
" if 0 < host_prefix_kv_cache_utilization_target < 1:\n",
" vllm_args.append(\n",
" f\"--host-prefix-kv-cache-utilization-target={host_prefix_kv_cache_utilization_target}\"\n",
" )\n",
"\n",
" if model_type:\n",
" vllm_args.append(f\"--model-type={model_type}\")\n",
"\n",
" if enable_llama_tool_parser:\n",
" vllm_args.append(\"--enable-auto-tool-choice\")\n",
" vllm_args.append(\"--tool-call-parser=vertex-llama-3\")\n",
"\n",
" env_vars = {\n",
" \"MODEL_ID\": base_model_id,\n",
" \"DEPLOY_SOURCE\": \"notebook\",\n",
" }\n",
"\n",
" # HF_TOKEN is not a compulsory field and may not be defined.\n",
" try:\n",
" if HF_TOKEN:\n",
" env_vars[\"HF_TOKEN\"] = HF_TOKEN\n",
" except NameError:\n",
" pass\n",
"\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=VLLM_DOCKER_URI,\n",
" serving_container_args=vllm_args,\n",
" serving_container_ports=[8080],\n",
" serving_container_predict_route=\"/generate\",\n",
" serving_container_health_route=\"/ping\",\n",
" serving_container_environment_variables=env_vars,\n",
" serving_container_shared_memory_size_mb=(16 * 1024), # 16 GB\n",
" serving_container_deployment_timeout=7200,\n",
" model_garden_source_model_name=(\n",
" f\"publishers/{publisher}/models/{publisher_model_id}\"\n",
" ),\n",
" )\n",
" print(\n",
" f\"Deploying {model_name} on {machine_type} with {accelerator_count} {accelerator_type} GPU(s).\"\n",
" )\n",
" model.deploy(\n",
" endpoint=endpoint,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" deploy_request_timeout=1800,\n",
" service_account=service_account,\n",
" system_labels={\n",
" \"NOTEBOOK_NAME\": \"model_garden_pytorch_llama3_finetuning.ipynb\",\n",
" \"NOTEBOOK_ENVIRONMENT\": common_util.get_deploy_source(),\n",
" },\n",
" )\n",
" print(\"endpoint_name:\", endpoint.name)\n",
"\n",
" return model, endpoint\n",
"\n",
"\n",
"models[\"vllm_gpu\"], endpoints[\"vllm_gpu\"] = deploy_model_vllm(\n",
" model_name=common_util.get_job_name_with_datetime(prefix=\"llama3-vllm-serve\"),\n",
" model_id=merged_model_output_dir,\n",
" publisher=\"meta\",\n",
" publisher_model_id=\"llama3\",\n",
" service_account=SERVICE_ACCOUNT,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" gpu_memory_utilization=gpu_memory_utilization,\n",
" max_model_len=max_model_len,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "2UYUNn60G_4U"
},
"outputs": [],
"source": [
"# @title Predict\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint with text prompts. Sampling parameters supported by vLLM can be found [here](https://docs.vllm.ai/en/latest/dev/sampling_params.html).\n",
"\n",
"# @markdown Example:\n",
"\n",
"# @markdown ```\n",
"# @markdown Human: What is a car?\n",
"# @markdown Assistant: A car, or a motor car, is a road-connected human-transportation system used to move people or goods from one place to another. The term also encompasses a wide range of vehicles, including motorboats, trains, and aircrafts. Cars typically have four wheels, a cabin for passengers, and an engine or motor. They have been around since the early 19th century and are now one of the most popular forms of transportation, used for daily commuting, shopping, and other purposes.\n",
"# @markdown ```\n",
"# @markdown Additionally, you can moderate the generated text with Vertex AI. See [Moderate text documentation](https://cloud.google.com/natural-language/docs/moderating-text) for more details.\n",
"\n",
"# Loads an existing endpoint instance using the endpoint name:\n",
"# - Using `endpoint_name = endpoint.name` allows us to get the\n",
"# endpoint name of the endpoint `endpoint` created in the cell\n",
"# above.\n",
"# - Alternatively, you can set `endpoint_name = \"1234567890123456789\"` to load\n",
"# an existing endpoint with the ID 1234567890123456789.\n",
"# You may uncomment the code below to load an existing endpoint.\n",
"\n",
"# endpoint_name = \"\" # @param {type:\"string\"}\n",
"# aip_endpoint_name = (\n",
"# f\"projects/{PROJECT_ID}/locations/{REGION}/endpoints/{endpoint_name}\"\n",
"# )\n",
"# endpoint = aiplatform.Endpoint(aip_endpoint_name)\n",
"\n",
"prompt = \"What is a car?\" # @param {type: \"string\"}\n",
"# @markdown If you encounter an issue like `ServiceUnavailable: 503 Took too long to respond when processing`, you can reduce the maximum number of output tokens, by lowering `max_tokens`.\n",
"max_tokens = 50 # @param {type:\"integer\"}\n",
"temperature = 1.0 # @param {type:\"number\"}\n",
"top_p = 1.0 # @param {type:\"number\"}\n",
"top_k = 1 # @param {type:\"integer\"}\n",
"# @markdown Set `raw_response` to `True` to obtain the raw model output. Set `raw_response` to `False` to apply additional formatting in the structure of `\"Prompt:\\n{prompt.strip()}\\nOutput:\\n{output}\"`.\n",
"raw_response = False # @param {type:\"boolean\"}\n",
"\n",
"# Overrides parameters for inferences.\n",
"instances = [\n",
" {\n",
" \"prompt\": prompt,\n",
" \"max_tokens\": max_tokens,\n",
" \"temperature\": temperature,\n",
" \"top_p\": top_p,\n",
" \"top_k\": top_k,\n",
" \"raw_response\": raw_response,\n",
" },\n",
"]\n",
"response = endpoints[\"vllm_gpu\"].predict(\n",
" instances=instances, use_dedicated_endpoint=use_dedicated_endpoint\n",
")\n",
"\n",
"for prediction in response.predictions:\n",
" print(prediction)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "af21a3cff1e0"
},
"source": [
"## Clean up resources"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "911406c1561e"
},
"outputs": [],
"source": [
"# @title Delete the model and endpoint\n",
"\n",
"train_job.delete()\n",
"\n",
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continuous charges that may incur.\n",
"\n",
"# Undeploy model and delete endpoint.\n",
"for endpoint in endpoints.values():\n",
" endpoint.delete(force=True)\n",
"\n",
"# Delete models.\n",
"for model in models.values():\n",
" model.delete()\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_NAME"
]
}
],
"metadata": {
"colab": {
"name": "model_garden_pytorch_llama3_finetuning.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
@@ -1,913 +0,0 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "iJc36RtD90jd"
},
"outputs": [],
"source": [
"# Copyright 2026 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",
"# https://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": "b9EezHSo90jf"
},
"source": [
"# Vertex AI Model Garden - Mistral-7B (PEFT)\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_pytorch_mistral_peft_tuning.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_pytorch_mistral_peft_tuning.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "ybMCVFh0_5R8"
},
"source": [
"## Overview\n",
"In this notebook you will learn how to fine tune Mistral with QLoRa and\n",
"deploy to Vertex AI endpoint.\n",
"\n",
"### Objective\n",
"\n",
"* Finetune and merge Mistral model using PEFT training docker image.\n",
"* Deploy the finetuned model with vLLM docker image on a Vertex AI Endpoint.\n",
"* Run inference on the deployed Vertex AI Endpoint.\n",
"\n",
"### File a bug\n",
"\n",
"File a bug on [GitHub](https://github.com/GoogleCloudPlatform/vertex-ai-samples/issues/new) if you encounter any issue with the notebook.\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "vzvFJU27a8si"
},
"source": [
"## Run the notebook"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "Q86d4aDSgGCu"
},
"outputs": [],
"source": [
"# @title Install Python Packages for Finetuning\n",
"\n",
"# @markdown 1. Install google-cloud-aiplatform package and restart the session if instructed.\n",
"! pip install --upgrade --quiet google-cloud-aiplatform==1.130.0\n",
"\n",
"# @markdown 2. Install packages to validate dataset with template.\n",
"! pip install --upgrade --quiet gcsfs==2024.3.1\n",
"! pip install --upgrade --quiet accelerate==0.31.0\n",
"! pip install --upgrade --quiet transformers==4.43.1\n",
"! pip install --upgrade --quiet datasets==2.19.2\n",
"\n",
"# Load local tensorboard.\n",
"%load_ext tensorboard"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "I-OjzhpyMHsu"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. For finetuning, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Frestricted_image_training_nvidia_a100_80gb_gpus)** to check if your project already has the required 8 Nvidia A100 80 GB GPUs in the us-central1 region. If yes, then run this notebook in the us-central1 region. If you do not have 8 Nvidia A100 80 GPUs or have more GPU requirements than this, then schedule your job with Nvidia H100 GPUs via Dynamic Workload Scheduler using [these instructions](https://cloud.google.com/vertex-ai/docs/training/schedule-jobs-dws). For Dynamic Workload Scheduler, check the [us-central1](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) or [europe-west4](https://console.cloud.google.com/iam-admin/quotas?location=europe-west4&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) quota for Nvidia H100 GPUs. If you do not have enough GPUs, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request quota.\n",
"\n",
"# @markdown 3. For serving, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_l4_gpus)** to check if your project already has the required 1 L4 GPU in the us-central1 region. If yes, then run this notebook in the us-central1 region. If you need more L4 GPUs for your project, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request more. Alternatively, if you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus).\n",
"\n",
"# @markdown > | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
"# @markdown | a2-ultragpu-1g | 1 NVIDIA_A100_80GB | us-central1, us-east4, europe-west4, asia-southeast1, us-east4 |\n",
"# @markdown | a3-highgpu-2g | 2 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-4g | 4 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"\n",
"# @markdown 4. **[Optional]** [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 5. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# Import the necessary packages\n",
"! rm -rf vertex-ai-samples && git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"! cd vertex-ai-samples && git reset --hard 0727e19520cf7957bceb701c248221bd3dbe4f1f\n",
"\n",
"import datetime\n",
"import importlib\n",
"import os\n",
"import uuid\n",
"from typing import Tuple\n",
"\n",
"from google.cloud import aiplatform\n",
"from google.cloud.aiplatform.compat.types import \\\n",
" custom_job as gca_custom_job_compat\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.notebooks.community.model_garden.docker_source_codes.notebook_util.common_util\"\n",
")\n",
"\n",
"models, endpoints = {}, {}\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"if not REGION:\n",
" if not os.environ.get(\"GOOGLE_CLOUD_REGION\"):\n",
" raise ValueError(\n",
" \"REGION must be set. See\"\n",
" \" https://cloud.google.com/vertex-ai/docs/general/locations for\"\n",
" \" available cloud locations.\"\n",
" )\n",
" REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, change the value yourself below.\n",
"now = datetime.datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}\"\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"MODEL_BUCKET = os.path.join(BUCKET_URI, \"mistral\")\n",
"\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default SERVICE_ACCOUNT.\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/storage.admin\"\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/aiplatform.user\""
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "5K169qf_udor"
},
"outputs": [],
"source": [
"# @title Set dataset\n",
"\n",
"# @markdown Use the Vertex AI SDK to create and run the custom training jobs.\n",
"\n",
"# @markdown This notebook uses [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) dataset as an example.\n",
"# @markdown You can set `dataset_name` to any existing [Hugging Face dataset](https://huggingface.co/datasets) name, and set `instruct_column_in_dataset` to the name of the dataset column containing training data. The [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) has only one column `text`, and therefore we set `instruct_column_in_dataset` to `text` in this notebook.\n",
"\n",
"# @markdown ### (Optional) Prepare a custom JSONL dataset for finetuning\n",
"\n",
"# @markdown You can prepare a JSONL file where each line is a valid JSON string as your custom training dataset. For example, here is one line from the [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) dataset:\n",
"# @markdown ```\n",
"# @markdown {\"text\": \"### Human: Hola### Assistant: \\u00a1Hola! \\u00bfEn qu\\u00e9 puedo ayudarte hoy?\"}\n",
"# @markdown ```\n",
"\n",
"# @markdown The JSON object has a key `text`, which should match `instruct_column_in_dataset`; The value should be one training data point, i.e. a string. After you prepared your JSONL file, you can either upload it to [Hugging Face datasets](https://huggingface.co/datasets) or [Google Cloud Storage](https://cloud.google.com/storage).\n",
"\n",
"# @markdown - To upload a JSONL dataset to [Hugging Face datasets](https://huggingface.co/datasets), follow the instructions on [Uploading Datasets](https://huggingface.co/docs/hub/en/datasets-adding). Then, set `dataset_name` to the name of your newly created dataset on Hugging Face.\n",
"\n",
"# @markdown - To upload a JSONL dataset to [Google Cloud Storage](https://cloud.google.com/storage), follow the instructions on [Upload objects from a filesystem](https://cloud.google.com/storage/docs/uploading-objects). Then, set `dataset_name` to the `gs://` URI to your JSONL file. For example: `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`.\n",
"\n",
"# @markdown Optionally update the `instruct_column_in_dataset` field below if your JSON objects use a key other than the default `text`.\n",
"\n",
"# @markdown ### (Optional) Format your data with custom JSON template\n",
"\n",
"# @markdown Sometimes, your dataset might have multiple text columns and you want to construct the training data with a template. You can prepare a JSON template in the following format:\n",
"\n",
"# @markdown ```\n",
"# @markdown {\n",
"# @markdown \"description\": \"Template that accepts text-bison format.\",\n",
"# @markdown \"source\": \"https://cloud.google.com/vertex-ai/generative-ai/docs/models/tune-text-models-supervised#dataset-format\",\n",
"# @markdown \"prompt_input\": \"\\n\\n<|start_header_id|>user<|end_header_id|>\\n\\n{input_text}<|eot_id|>\\n\\n<|start_header_id|>assistant<|end_header_id|>\\n\\n{output_text}<|eot_id|>\",\n",
"# @markdown \"instruction_separator\": \"<|start_header_id|>user<|end_header_id|>\\n\\n\",\n",
"# @markdown \"response_separator\": \"<|start_header_id|>assistant<|end_header_id|>\\n\\n\"\n",
"# @markdown }\n",
"# @markdown ```\n",
"\n",
"\n",
"# @markdown As an example, the template above can be used to format the following training data (this line comes from `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`):\n",
"\n",
"# @markdown ```\n",
"# @markdown {\"input_text\":\"TRANSCRIPT: \\nREASON FOR EVALUATION:,\\n\\n LABEL:\",\"output_text\":\"Chiropractic\"}\n",
"# @markdown ```\n",
"\n",
"# @markdown This example template simply concatenates `input_text` with `output_text` with some special tokens in between.\n",
"# @markdown\n",
"# @markdown To try such custom dataset, you can make the following changes:\n",
"# @markdown 1. Set `template` to `llama3-text-bison`\n",
"# @markdown 1. Set `train_dataset_name` to `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`\n",
"# @markdown 1. Set `train_split_name` to `train`\n",
"# @markdown 1. Set `eval_dataset_name` to `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_eval_sample.jsonl`\n",
"# @markdown 1. Set `eval_split_name` to `train` (**NOT** `test`)\n",
"# @markdown 1. Set `instruct_column_in_dataset` as `input_text`.\n",
"\n",
"# Template name or gs:// URI to a custom template.\n",
"template = \"openassistant-guanaco\" # @param {type:\"string\"}\n",
"\n",
"# Hugging Face dataset name or gs:// URI to a custom JSONL dataset.\n",
"train_dataset_name = \"timdettmers/openassistant-guanaco\" # @param {type:\"string\"}\n",
"train_split_name = \"train\" # @param {type:\"string\"}\n",
"eval_dataset_name = \"timdettmers/openassistant-guanaco\" # @param {type:\"string\"}\n",
"eval_split_name = \"test\" # @param {type:\"string\"}\n",
"\n",
"# Name of the dataset column containing training text input.\n",
"instruct_column_in_dataset = \"text\" # @param {type:\"string\"}"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "ncoBBZXq2qxf"
},
"outputs": [],
"source": [
"# @title Set model\n",
"\n",
"# @markdown Select a model variant of Mistral.\n",
"base_model_id = \"mistralai/Mistral-7B-v0.1\" # @param [\"mistralai/Mistral-7B-v0.1\"] {isTemplate: true}\n",
"pretrained_model_id = f\"gs://vertex-model-garden-public-us/{base_model_id}\""
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "8MTGQCZTxDbN"
},
"outputs": [],
"source": [
"# @title Validate Dataset with Template\n",
"\n",
"# @markdown This section validates the train and eval datasets with the template before starting the fine tuning process.\n",
"\n",
"import transformers\n",
"\n",
"dataset_validation_util = importlib.import_module(\n",
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.dataset_validation_util\"\n",
")\n",
"\n",
"if dataset_validation_util.is_gcs_path(pretrained_model_id):\n",
" # Download tokenizer.\n",
" ! mkdir tokenizer\n",
" ! gsutil cp {pretrained_model_id}/tokenizer.json ./tokenizer\n",
" ! gsutil cp {pretrained_model_id}/config.json ./tokenizer\n",
" tokenizer_path = \"./tokenizer\"\n",
" access_token = \"\"\n",
"else:\n",
" tokenizer_path = pretrained_model_id\n",
" access_token = HF_TOKEN\n",
"\n",
"tokenizer = transformers.AutoTokenizer.from_pretrained(\n",
" tokenizer_path,\n",
" trust_remote_code=False,\n",
" use_fast=True,\n",
" token=access_token,\n",
")\n",
"\n",
"# Validate the train dataset.\n",
"dataset_validation_util.validate_dataset_with_template(\n",
" dataset_name=train_dataset_name,\n",
" split=train_split_name,\n",
" input_column=instruct_column_in_dataset,\n",
" template=template,\n",
" use_multiprocessing=False,\n",
" tokenizer=tokenizer,\n",
")\n",
"\n",
"# Validate the eval dataset.\n",
"dataset_validation_util.validate_dataset_with_template(\n",
" dataset_name=eval_dataset_name,\n",
" split=eval_split_name,\n",
" input_column=instruct_column_in_dataset,\n",
" template=template,\n",
" use_multiprocessing=False,\n",
" tokenizer=tokenizer,\n",
")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "885Vf4o8hbbo"
},
"outputs": [],
"source": [
"# @title Finetune\n",
"\n",
"# @markdown This section demonstrates how to finetune the Mistral-7B model and merge the finetuned LoRA adapter with the base model on Vertex AI. It uses the Vertex AI SDK to create and run the custom training jobs.\n",
"\n",
"# @markdown The training job takes approximately between 10 to 20 mins to set-up. Once done, the training job is expected to take around 20 mins with the default configurations. To find the training time, throughput, and memory usage of your training job, you can go to the training logs and check the log line of the last training epoch.\n",
"\n",
"# @markdown **Note**:\n",
"# @markdown 1. We recommend setting `finetuning_precision_mode` to `4bit` because it enables using fewer hardware resources for finetuning.\n",
"# @markdown 1. If `max_steps > 0`, it takes precedence over `epochs`. One can set a small `max_steps` value to quickly check the pipeline.\n",
"\n",
"# @markdown Acceletor type to use for training.\n",
"accelerator_type = \"NVIDIA_A100_80GB\" # @param [\"NVIDIA_A100_80GB\", \"NVIDIA_H100_80GB\"]\n",
"\n",
"# The pre-built training docker image.\n",
"if accelerator_type == \"NVIDIA_A100_80GB\":\n",
" repo = \"us-docker.pkg.dev/vertex-ai-restricted\"\n",
" is_restricted_image = True\n",
" is_dynamic_workload_scheduler = False\n",
" dws_kwargs = {}\n",
"else:\n",
" repo = \"us-docker.pkg.dev/vertex-ai\"\n",
" is_restricted_image = False\n",
" is_dynamic_workload_scheduler = True\n",
" dws_kwargs = {\n",
" \"max_wait_duration\": 1800, # 30 minutes\n",
" \"scheduling_strategy\": gca_custom_job_compat.Scheduling.Strategy.FLEX_START,\n",
" }\n",
"\n",
"TRAIN_DOCKER_URI = (\n",
" f\"{repo}/vertex-vision-model-garden-dockers/pytorch-peft-train:stable_20240909\"\n",
")\n",
"\n",
"# Worker pool spec.\n",
"if accelerator_type == \"NVIDIA_A100_80GB\":\n",
" per_node_accelerator_count = 8\n",
" machine_type = \"a2-ultragpu-8g\"\n",
"elif accelerator_type == \"NVIDIA_H100_80GB\":\n",
" per_node_accelerator_count = 8\n",
" machine_type = \"a3-highgpu-8g\"\n",
"else:\n",
" raise ValueError(\n",
" f\"Recommended machine settings not found for: {accelerator_type}. To use another accelerator type, edit this code block to pass in an appropriate `machine_type`, `accelerator_type`, and `per_node_accelerator_count` to the deploy_model_vllm function by clicking `Show Code` and then modifying the code.\"\n",
" )\n",
"\n",
"# @markdown Batch size for finetuning.\n",
"per_device_train_batch_size = 1 # @param{type:\"integer\"}\n",
"# @markdown Number of updates steps to accumulate the gradients for, before performing a backward/update pass.\n",
"gradient_accumulation_steps = 4 # @param{type:\"integer\"}\n",
"# @markdown Maximum sequence length.\n",
"max_seq_length = 4096 # @param{type:\"integer\"}\n",
"# @markdown Setting a positive `max_steps` here will override `num_epochs`.\n",
"max_steps = -1 # @param{type:\"integer\"}\n",
"num_epochs = 1.0 # @param{type:\"number\"}\n",
"# @markdown Precision mode for finetuning.\n",
"finetuning_precision_mode = \"4bit\" # @param [\"4bit\", \"8bit\", \"float16\"]\n",
"# @markdown Learning rate.\n",
"learning_rate = 5e-5 # @param{type:\"number\"}\n",
"# @markdown The scheduler type to use.\n",
"lr_scheduler_type = \"cosine\" # @param{type:\"string\"}\n",
"# @markdown LoRA parameters.\n",
"lora_rank = 16 # @param{type:\"integer\"}\n",
"lora_alpha = 32 # @param{type:\"integer\"}\n",
"lora_dropout = 0.05 # @param{type:\"number\"}\n",
"# Activates gradient checkpointing for the current model (may be referred to as activation checkpointing or checkpoint activations in other frameworks).\n",
"enable_gradient_checkpointing = True\n",
"# Attention implementation to use in the model.\n",
"attn_implementation = \"flash_attention_2\"\n",
"# The optimizer for which to schedule the learning rate.\n",
"optimizer = \"paged_adamw_32bit\"\n",
"# Define the proportion of training to be dedicated to a linear warmup where learning rate gradually increases.\n",
"warmup_ratio = \"0.01\"\n",
"# The list or string of integrations to report the results and logs to.\n",
"report_to = \"tensorboard\"\n",
"# Number of updates steps before two checkpoint saves.\n",
"save_steps = 10\n",
"# Number of update steps between two logs.\n",
"logging_steps = save_steps\n",
"# Train precision of the model.\n",
"train_precision = \"float16\"\n",
"\n",
"replica_count = 1\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=per_node_accelerator_count * replica_count,\n",
" is_for_training=True,\n",
" is_restricted_image=is_restricted_image,\n",
" is_dynamic_workload_scheduler=is_dynamic_workload_scheduler,\n",
")\n",
"\n",
"# Setup training job.\n",
"job_name = common_util.get_job_name_with_datetime(\"mistral-lora-train\")\n",
"\n",
"base_output_dir = os.path.join(STAGING_BUCKET, job_name)\n",
"# Create a GCS folder to store the LORA adapter.\n",
"lora_output_dir = os.path.join(base_output_dir, \"adapter\")\n",
"# Create a GCS folder to store the merged model with the base model and the\n",
"# finetuned LORA adapter.\n",
"merged_model_output_dir = os.path.join(base_output_dir, \"merged-model\")\n",
"\n",
"# Add labels for the finetuning job.\n",
"labels = {\n",
" \"mg-source\": \"notebook\",\n",
" \"mg-notebook-name\": \"model_garden_pytorch_mistral_peft_tuning.ipynb\".split(\".\")[0],\n",
"}\n",
"\n",
"labels[\"mg-tune\"] = \"publishers-mistralai-models-mistral\"\n",
"versioned_model_id = base_model_id.split(\"/\")[1].lower().replace(\".\", \"-\")\n",
"labels[\"versioned-mg-tune\"] = f\"{labels['mg-tune']}-{versioned_model_id}\"\n",
"\n",
"eval_args = [\n",
" f\"--eval_dataset_path={eval_dataset_name}\",\n",
" f\"--eval_column={instruct_column_in_dataset}\",\n",
" f\"--eval_template={template}\",\n",
" f\"--eval_split={eval_split_name}\",\n",
" f\"--eval_steps={save_steps}\",\n",
" \"--eval_tasks=builtin_eval\",\n",
" \"--eval_metric_name=loss\",\n",
"]\n",
"\n",
"train_job_args = [\n",
" \"--config_file=vertex_vision_model_garden_peft/deepspeed_zero2_8gpu.yaml\",\n",
" \"--task=instruct-lora\",\n",
" \"--completion_only=False\",\n",
" f\"--pretrained_model_id={pretrained_model_id}\",\n",
" f\"--dataset_name={train_dataset_name}\",\n",
" f\"--train_split_name={train_split_name}\",\n",
" f\"--instruct_column_in_dataset={instruct_column_in_dataset}\",\n",
" f\"--output_dir={lora_output_dir}\",\n",
" f\"--merge_base_and_lora_output_dir={merged_model_output_dir}\",\n",
" f\"--per_device_train_batch_size={per_device_train_batch_size}\",\n",
" f\"--gradient_accumulation_steps={gradient_accumulation_steps}\",\n",
" f\"--lora_rank={lora_rank}\",\n",
" f\"--lora_alpha={lora_alpha}\",\n",
" f\"--lora_dropout={lora_dropout}\",\n",
" f\"--max_steps={max_steps}\",\n",
" f\"--max_seq_length={max_seq_length}\",\n",
" f\"--learning_rate={learning_rate}\",\n",
" f\"--lr_scheduler_type={lr_scheduler_type}\",\n",
" f\"--precision_mode={finetuning_precision_mode}\",\n",
" f\"--train_precision={train_precision}\",\n",
" f\"--enable_gradient_checkpointing={enable_gradient_checkpointing}\",\n",
" f\"--num_epochs={num_epochs}\",\n",
" f\"--attn_implementation={attn_implementation}\",\n",
" f\"--optimizer={optimizer}\",\n",
" f\"--warmup_ratio={warmup_ratio}\",\n",
" f\"--report_to={report_to}\",\n",
" f\"--logging_output_dir={base_output_dir}\",\n",
" f\"--save_steps={save_steps}\",\n",
" f\"--logging_steps={logging_steps}\",\n",
" f\"--template={template}\",\n",
"] + eval_args\n",
"\n",
"train_job = aiplatform.CustomContainerTrainingJob(\n",
" display_name=job_name,\n",
" container_uri=TRAIN_DOCKER_URI,\n",
" labels=labels,\n",
")\n",
"\n",
"print(\"Running training job with args:\")\n",
"print(\" \\\\\\n\".join(train_job_args))\n",
"# Pass training arguments and launch job.\n",
"train_job.run(\n",
" args=train_job_args,\n",
" replica_count=replica_count,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=per_node_accelerator_count,\n",
" boot_disk_size_gb=500,\n",
" service_account=SERVICE_ACCOUNT,\n",
" base_output_dir=base_output_dir,\n",
" sync=False, # Non-blocking call to run.\n",
" **dws_kwargs,\n",
")\n",
"\n",
"# Wait until resource has been created.\n",
"train_job.wait_for_resource_creation()\n",
"\n",
"print(\"LoRA adapter will be saved in:\", lora_output_dir)\n",
"print(\"Trained and merged models will be saved in:\", merged_model_output_dir)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "cLNQwLMGmzlR"
},
"outputs": [],
"source": [
"# @title Run TensorBoard\n",
"# @markdown This section shows how to launch TensorBoard in a [Cloud Shell](https://cloud.google.com/shell/docs).\n",
"# @markdown 1. Click the Cloud Shell icon(![terminal](https://github.com/google/material-design-icons/blob/master/png/action/terminal/materialicons/24dp/1x/baseline_terminal_black_24dp.png?raw=true)) on the top right to open the Cloud Shell.\n",
"# @markdown 2. Copy the `tensorboard` command shown below by running this cell.\n",
"# @markdown 3. Paste and run the command in the Cloud Shell to launch TensorBoard.\n",
"# @markdown 4. Once the command runs (You may have to click `Authorize` if prompted), click the link starting with `http://localhost`.\n",
"\n",
"# @markdown Note: You may need to wait around 10 minutes after the job starts in order for the TensorBoard logs to be written to the GCS bucket.\n",
"print(f\"Command to copy: tensorboard --logdir {base_output_dir}/logs\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "GyDWPdV1NjMT"
},
"outputs": [],
"source": [
"# @title Deploy\n",
"# @markdown This section uploads the model to Model Registry and deploys it on the Endpoint. It takes 15 minutes to 1 hour to finish depending on the size of model.\n",
"\n",
"if train_job.end_time is None:\n",
" print(\"Waiting for the training job to finish...\")\n",
" train_job.wait()\n",
" print(\"The training job has finished.\")\n",
"\n",
"print(\"Deploying models in:\", merged_model_output_dir)\n",
"\n",
"# The pre-built serving docker image for vLLM.\n",
"VLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-vllm-serve:20240721_0916_RC00\"\n",
"\n",
"# Find Vertex AI prediction supported accelerators and regions [here](https://cloud.google.com/vertex-ai/docs/predictions/configure-compute).\n",
"# @markdown Accelerator type to use for serving.\n",
"accelerator_type = \"NVIDIA_L4\" # @param [\"NVIDIA_L4\", \"NVIDIA_TESLA_V100\", \"NVIDIA_TESLA_T4\", \"NVIDIA_TESLA_A100\"]\n",
"\n",
"if accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-8\"\n",
" accelerator_count = 1\n",
"elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-standard-16\"\n",
" accelerator_count = 2\n",
"elif accelerator_type == \"NVIDIA_TESLA_T4\":\n",
" machine_type = \"n1-standard-16\"\n",
" accelerator_count = 2\n",
"elif accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 1\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=False,\n",
")\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"gpu_memory_utilization = 0.85\n",
"max_model_len = 8192 # Maximum context length.\n",
"\n",
"# Ensure max_model_len does not exceed the limit\n",
"if max_model_len > 8192:\n",
" raise ValueError(\"max_model_len cannot exceed 8192\")\n",
"\n",
"\n",
"def get_deploy_source() -> str:\n",
" \"\"\"Gets deploy_source string based on running environment.\"\"\"\n",
" vertex_product = os.environ.get(\"VERTEX_PRODUCT\", \"\")\n",
" if vertex_product == \"COLAB_ENTERPRISE\":\n",
" return \"notebook_colab_enterprise\"\n",
" elif vertex_product == \"WORKBENCH_INSTANCE\":\n",
" return \"notebook_workbench\"\n",
" else:\n",
" # Legacy workbench, legacy colab, or other custom environments.\n",
" return \"notebook_environment_unspecified\"\n",
"\n",
"\n",
"def deploy_model_vllm(\n",
" model_name: str,\n",
" model_id: str,\n",
" publisher: str,\n",
" publisher_model_id: str,\n",
" service_account: str,\n",
" base_model_id: str = None,\n",
" machine_type: str = \"g2-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
" gpu_memory_utilization: float = 0.9,\n",
" max_model_len: int = 4096,\n",
" dtype: str = \"auto\",\n",
" enable_trust_remote_code: bool = False,\n",
" enforce_eager: bool = False,\n",
" enable_lora: bool = False,\n",
" enable_chunked_prefill: bool = False,\n",
" enable_prefix_cache: bool = False,\n",
" host_prefix_kv_cache_utilization_target: float = 0.0,\n",
" max_loras: int = 1,\n",
" max_cpu_loras: int = 8,\n",
" use_dedicated_endpoint: bool = False,\n",
" max_num_seqs: int = 256,\n",
" model_type: str = None,\n",
" enable_llama_tool_parser: bool = False,\n",
" is_spot: bool = False,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys trained models with vLLM into Vertex AI.\"\"\"\n",
" endpoint = aiplatform.Endpoint.create(\n",
" display_name=f\"{model_name}-endpoint\",\n",
" dedicated_endpoint_enabled=use_dedicated_endpoint,\n",
" )\n",
"\n",
" if not base_model_id:\n",
" base_model_id = model_id\n",
"\n",
" # See https://docs.vllm.ai/en/latest/models/engine_args.html for a list of possible arguments with descriptions.\n",
" vllm_args = [\n",
" \"python\",\n",
" \"-m\",\n",
" \"vllm.entrypoints.api_server\",\n",
" \"--host=0.0.0.0\",\n",
" \"--port=8080\",\n",
" f\"--model={model_id}\",\n",
" f\"--tensor-parallel-size={accelerator_count}\",\n",
" \"--swap-space=16\",\n",
" f\"--max-model-len={max_model_len}\",\n",
" f\"--dtype={dtype}\",\n",
" f\"--max-loras={max_loras}\",\n",
" f\"--max-cpu-loras={max_cpu_loras}\",\n",
" f\"--max-num-seqs={max_num_seqs}\",\n",
" \"--disable-log-stats\",\n",
" ]\n",
"\n",
" if gpu_memory_utilization:\n",
" vllm_args.append(f\"--gpu-memory-utilization={gpu_memory_utilization}\")\n",
"\n",
" if enable_trust_remote_code:\n",
" vllm_args.append(\"--trust-remote-code\")\n",
"\n",
" if enforce_eager:\n",
" vllm_args.append(\"--enforce-eager\")\n",
"\n",
" if enable_lora:\n",
" vllm_args.append(\"--enable-lora\")\n",
"\n",
" if enable_chunked_prefill:\n",
" vllm_args.append(\"--enable-chunked-prefill\")\n",
"\n",
" if enable_prefix_cache:\n",
" vllm_args.append(\"--enable-prefix-caching\")\n",
"\n",
" if 0 < host_prefix_kv_cache_utilization_target < 1:\n",
" vllm_args.append(\n",
" f\"--host-prefix-kv-cache-utilization-target={host_prefix_kv_cache_utilization_target}\"\n",
" )\n",
"\n",
" if model_type:\n",
" vllm_args.append(f\"--model-type={model_type}\")\n",
"\n",
" if enable_llama_tool_parser:\n",
" vllm_args.append(\"--enable-auto-tool-choice\")\n",
" vllm_args.append(\"--tool-call-parser=vertex-llama-3\")\n",
"\n",
" env_vars = {\n",
" \"MODEL_ID\": base_model_id,\n",
" \"DEPLOY_SOURCE\": \"notebook\",\n",
" }\n",
"\n",
" # HF_TOKEN is not a compulsory field and may not be defined.\n",
" try:\n",
" if HF_TOKEN:\n",
" env_vars[\"HF_TOKEN\"] = HF_TOKEN\n",
" except NameError:\n",
" pass\n",
"\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=VLLM_DOCKER_URI,\n",
" serving_container_args=vllm_args,\n",
" serving_container_ports=[8080],\n",
" serving_container_predict_route=\"/generate\",\n",
" serving_container_health_route=\"/ping\",\n",
" serving_container_environment_variables=env_vars,\n",
" serving_container_shared_memory_size_mb=(16 * 1024), # 16 GB\n",
" serving_container_deployment_timeout=7200,\n",
" model_garden_source_model_name=(\n",
" f\"publishers/{publisher}/models/{publisher_model_id}\"\n",
" ),\n",
" )\n",
" print(\n",
" f\"Deploying {model_name} on {machine_type} with {accelerator_count} {accelerator_type} GPU(s).\"\n",
" )\n",
" model.deploy(\n",
" endpoint=endpoint,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" deploy_request_timeout=1800,\n",
" service_account=service_account,\n",
" spot=is_spot,\n",
" system_labels={\n",
" \"NOTEBOOK_NAME\": \"model_garden_pytorch_mistral_peft_tuning.ipynb\",\n",
" \"NOTEBOOK_ENVIRONMENT\": get_deploy_source(),\n",
" },\n",
" )\n",
" print(\"endpoint_name:\", endpoint.name)\n",
"\n",
" return model, endpoint\n",
"\n",
"\n",
"models[\"vllm_gpu\"], endpoints[\"vllm_gpu\"] = deploy_model_vllm(\n",
" model_name=common_util.get_job_name_with_datetime(prefix=\"mistral-vllm-serve\"),\n",
" model_id=merged_model_output_dir,\n",
" publisher=\"mistral-ai\",\n",
" publisher_model_id=\"mistral\",\n",
" service_account=SERVICE_ACCOUNT,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" gpu_memory_utilization=gpu_memory_utilization,\n",
" max_model_len=max_model_len,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"# @markdown Click \"Show code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "4v2Mnui4tH1X"
},
"outputs": [],
"source": [
"# @title Predict\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint with text prompts. Sampling parameters supported by vLLM can be found [here](https://docs.vllm.ai/en/latest/dev/sampling_params.html).\n",
"\n",
"# @markdown Example:\n",
"\n",
"# @markdown ```\n",
"# @markdown Human: What is a car?\n",
"# @markdown Assistant: A car, or a motor car, is a road-connected human-transportation system used to move people or goods from one place to another. The term also encompasses a wide range of vehicles, including motorboats, trains, and aircrafts. Cars typically have four wheels, a cabin for passengers, and an engine or motor. They have been around since the early 19th century and are now one of the most popular forms of transportation, used for daily commuting, shopping, and other purposes.\n",
"# @markdown ```\n",
"# @markdown Additionally, you can moderate the generated text with Vertex AI. See [Moderate text documentation](https://cloud.google.com/natural-language/docs/moderating-text) for more details.\n",
"\n",
"# Loads an existing endpoint instance using the endpoint name:\n",
"# - Using `endpoint_name = endpoint.name` allows us to get the\n",
"# endpoint name of the endpoint `endpoint` created in the cell\n",
"# above.\n",
"# - Alternatively, you can set `endpoint_name = \"1234567890123456789\"` to load\n",
"# an existing endpoint with the ID 1234567890123456789.\n",
"# You may uncomment the code below to load an existing endpoint.\n",
"\n",
"# endpoint_name = \"\" # @param {type:\"string\"}\n",
"# aip_endpoint_name = (\n",
"# f\"projects/{PROJECT_ID}/locations/{REGION}/endpoints/{endpoint_name}\"\n",
"# )\n",
"# endpoint = aiplatform.Endpoint(aip_endpoint_name)\n",
"\n",
"prompt = \"What is a car?\" # @param {type: \"string\"}\n",
"# @markdown If you encounter an issue like `ServiceUnavailable: 503 Took too long to respond when processing`, you can reduce the maximum number of output tokens, by lowering `max_tokens`.\n",
"max_tokens = 50 # @param {type:\"integer\"}\n",
"temperature = 1.0 # @param {type:\"number\"}\n",
"top_p = 1.0 # @param {type:\"number\"}\n",
"top_k = 1 # @param {type:\"integer\"}\n",
"# @markdown Set `raw_response` to `True` to obtain the raw model output. Set `raw_response` to `False` to apply additional formatting in the structure of `\"Prompt:\\n{prompt.strip()}\\nOutput:\\n{output}\"`.\n",
"raw_response = False # @param {type:\"boolean\"}\n",
"\n",
"# Overrides parameters for inferences.\n",
"instances = [\n",
" {\n",
" \"prompt\": prompt,\n",
" \"max_tokens\": max_tokens,\n",
" \"temperature\": temperature,\n",
" \"top_p\": top_p,\n",
" \"top_k\": top_k,\n",
" \"raw_response\": raw_response,\n",
" },\n",
"]\n",
"response = endpoints[\"vllm_gpu\"].predict(\n",
" instances=instances, use_dedicated_endpoint=use_dedicated_endpoint\n",
")\n",
"\n",
"for prediction in response.predictions:\n",
" print(prediction)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "x9EMCOUJ6-ji"
},
"outputs": [],
"source": [
"# @title Delete the model and endpoint\n",
"\n",
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continuous charges that may incur.\n",
"\n",
"# Undeploy model and delete endpoint.\n",
"for endpoint in endpoints.values():\n",
" endpoint.delete(force=True)\n",
"\n",
"# Delete models.\n",
"for model in models.values():\n",
" model.delete()\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_NAME"
]
}
],
"metadata": {
"accelerator": "GPU",
"colab": {
"name": "model_garden_pytorch_mistral_peft_tuning.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
@@ -1,919 +0,0 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "iJc36RtD90jd"
},
"outputs": [],
"source": [
"# Copyright 2026 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",
"# https://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": "b9EezHSo90jf"
},
"source": [
"# Vertex AI Model Garden - Mixtral-8x7B (PEFT)\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_pytorch_mixtral_peft_tuning.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_pytorch_mixtral_peft_tuning.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "ybMCVFh0_5R8"
},
"source": [
"## Overview\n",
"In this notebook you will learn how to fine tune Mixtral-8x7B with QLoRa and\n",
"deploy to Vertex AI endpoint.\n",
"\n",
"### Objective\n",
"\n",
"* Finetune and merge Mixtral-8x7B model with PEFT training docker image.\n",
"* Deploy the finetuned model with vLLM docker image on a Vertex AI Endpoint.\n",
"* Run inference on the deployed Vertex AI Endpoint.\n",
"\n",
"### File a bug\n",
"\n",
"File a bug on [GitHub](https://github.com/GoogleCloudPlatform/vertex-ai-samples/issues/new) if you encounter any issue with the notebook.\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "vzvFJU27a8si"
},
"source": [
"## Run the notebook"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "DzESmydvgME9"
},
"outputs": [],
"source": [
"# @title Install Python Packages for Finetuning\n",
"\n",
"# @markdown 1. Install google-cloud-aiplatform package and restart the session if instructed.\n",
"! pip install --upgrade --quiet google-cloud-aiplatform==1.130.0\n",
"\n",
"# @markdown 2. Install packages to validate dataset with template.\n",
"! pip install --upgrade --quiet gcsfs==2024.3.1\n",
"! pip install --upgrade --quiet accelerate==0.31.0\n",
"! pip install --upgrade --quiet transformers==4.43.1\n",
"! pip install --upgrade --quiet datasets==2.19.2\n",
"\n",
"# Load local tensorboard.\n",
"%load_ext tensorboard"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "I-OjzhpyMHsu"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. For finetuning, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Frestricted_image_training_nvidia_a100_80gb_gpus)** to check if your project already has the required 8 Nvidia A100 80 GB GPUs in the us-central1 region. If yes, then run this notebook in the us-central1 region. If you do not have 8 Nvidia A100 80 GPUs or have more GPU requirements than this, then schedule your job with Nvidia H100 GPUs via Dynamic Workload Scheduler using [these instructions](https://cloud.google.com/vertex-ai/docs/training/schedule-jobs-dws). For Dynamic Workload Scheduler, check the [us-central1](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) or [europe-west4](https://console.cloud.google.com/iam-admin/quotas?location=europe-west4&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) quota for Nvidia H100 GPUs. If you do not have enough GPUs, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request quota.\n",
"\n",
"# @markdown 3. For serving, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_l4_gpus)** to check if your project already has the required 1 L4 GPU in the us-central1 region. If yes, then run this notebook in the us-central1 region. If you need more L4 GPUs for your project, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request more. Alternatively, if you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus).\n",
"\n",
"# @markdown > | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
"# @markdown | a2-ultragpu-1g | 1 NVIDIA_A100_80GB | us-central1, us-east4, europe-west4, asia-southeast1, us-east4 |\n",
"# @markdown | a3-highgpu-2g | 2 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-4g | 4 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"\n",
"# @markdown 4. **[Optional]** [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 5. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# Import the necessary packages\n",
"! rm -rf vertex-ai-samples && git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"! cd vertex-ai-samples && git reset --hard 0727e19520cf7957bceb701c248221bd3dbe4f1f\n",
"\n",
"import datetime\n",
"import importlib\n",
"import os\n",
"import uuid\n",
"from typing import Tuple\n",
"\n",
"from google.cloud import aiplatform\n",
"from google.cloud.aiplatform.compat.types import \\\n",
" custom_job as gca_custom_job_compat\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.notebooks.community.model_garden.docker_source_codes.notebook_util.common_util\"\n",
")\n",
"\n",
"models, endpoints = {}, {}\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"if not REGION:\n",
" if not os.environ.get(\"GOOGLE_CLOUD_REGION\"):\n",
" raise ValueError(\n",
" \"REGION must be set. See\"\n",
" \" https://cloud.google.com/vertex-ai/docs/general/locations for\"\n",
" \" available cloud locations.\"\n",
" )\n",
" REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, change the value yourself below.\n",
"now = datetime.datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}\"\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"MODEL_BUCKET = os.path.join(BUCKET_URI, \"mixtral\")\n",
"\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default SERVICE_ACCOUNT.\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/storage.admin\"\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/aiplatform.user\""
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "rwP8nr8jnNdt"
},
"outputs": [],
"source": [
"# @title Set dataset\n",
"\n",
"# @markdown Use the Vertex AI SDK to create and run the custom training jobs.\n",
"\n",
"# @markdown This notebook uses [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) dataset as an example.\n",
"# @markdown You can set `dataset_name` to any existing [Hugging Face dataset](https://huggingface.co/datasets) name, and set `instruct_column_in_dataset` to the name of the dataset column containing training data. The [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) has only one column `text`, and therefore we set `instruct_column_in_dataset` to `text` in this notebook.\n",
"\n",
"# @markdown ### (Optional) Prepare a custom JSONL dataset for finetuning\n",
"\n",
"# @markdown You can prepare a JSONL file where each line is a valid JSON string as your custom training dataset. For example, here is one line from the [timdettmers/openassistant-guanaco](https://huggingface.co/datasets/timdettmers/openassistant-guanaco) dataset:\n",
"# @markdown ```\n",
"# @markdown {\"text\": \"### Human: Hola### Assistant: \\u00a1Hola! \\u00bfEn qu\\u00e9 puedo ayudarte hoy?\"}\n",
"# @markdown ```\n",
"\n",
"# @markdown The JSON object has a key `text`, which should match `instruct_column_in_dataset`; The value should be one training data point, i.e. a string. After you prepared your JSONL file, you can either upload it to [Hugging Face datasets](https://huggingface.co/datasets) or [Google Cloud Storage](https://cloud.google.com/storage).\n",
"\n",
"# @markdown - To upload a JSONL dataset to [Hugging Face datasets](https://huggingface.co/datasets), follow the instructions on [Uploading Datasets](https://huggingface.co/docs/hub/en/datasets-adding). Then, set `dataset_name` to the name of your newly created dataset on Hugging Face.\n",
"\n",
"# @markdown - To upload a JSONL dataset to [Google Cloud Storage](https://cloud.google.com/storage), follow the instructions on [Upload objects from a filesystem](https://cloud.google.com/storage/docs/uploading-objects). Then, set `dataset_name` to the `gs://` URI to your JSONL file. For example: `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`.\n",
"\n",
"# @markdown Optionally update the `instruct_column_in_dataset` field below if your JSON objects use a key other than the default `text`.\n",
"\n",
"# @markdown ### (Optional) Format your data with custom JSON template\n",
"\n",
"# @markdown Sometimes, your dataset might have multiple text columns and you want to construct the training data with a template. You can prepare a JSON template in the following format:\n",
"\n",
"# @markdown ```\n",
"# @markdown {\n",
"# @markdown \"description\": \"Template that accepts text-bison format.\",\n",
"# @markdown \"source\": \"https://cloud.google.com/vertex-ai/generative-ai/docs/models/tune-text-models-supervised#dataset-format\",\n",
"# @markdown \"prompt_input\": \"\\n\\n<|start_header_id|>user<|end_header_id|>\\n\\n{input_text}<|eot_id|>\\n\\n<|start_header_id|>assistant<|end_header_id|>\\n\\n{output_text}<|eot_id|>\",\n",
"# @markdown \"instruction_separator\": \"<|start_header_id|>user<|end_header_id|>\\n\\n\",\n",
"# @markdown \"response_separator\": \"<|start_header_id|>assistant<|end_header_id|>\\n\\n\"\n",
"# @markdown }\n",
"# @markdown ```\n",
"\n",
"\n",
"# @markdown As an example, the template above can be used to format the following training data (this line comes from `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`):\n",
"\n",
"# @markdown ```\n",
"# @markdown {\"input_text\":\"TRANSCRIPT: \\nREASON FOR EVALUATION:,\\n\\n LABEL:\",\"output_text\":\"Chiropractic\"}\n",
"# @markdown ```\n",
"\n",
"# @markdown This example template simply concatenates `input_text` with `output_text` with some special tokens in between.\n",
"# @markdown\n",
"# @markdown To try such custom dataset, you can make the following changes:\n",
"# @markdown 1. Set `template` to `llama3-text-bison`\n",
"# @markdown 1. Set `train_dataset_name` to `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_train_sample.jsonl`\n",
"# @markdown 1. Set `train_split_name` to `train`\n",
"# @markdown 1. Set `eval_dataset_name` to `gs://cloud-samples-data/vertex-ai/model-evaluation/peft_eval_sample.jsonl`\n",
"# @markdown 1. Set `eval_split_name` to `train` (**NOT** `test`)\n",
"# @markdown 1. Set `instruct_column_in_dataset` as `input_text`.\n",
"\n",
"# Template name or gs:// URI to a custom template.\n",
"template = \"openassistant-guanaco\" # @param {type:\"string\"}\n",
"\n",
"# Hugging Face dataset name or gs:// URI to a custom JSONL dataset.\n",
"train_dataset_name = \"timdettmers/openassistant-guanaco\" # @param {type:\"string\"}\n",
"train_split_name = \"train\" # @param {type:\"string\"}\n",
"eval_dataset_name = \"timdettmers/openassistant-guanaco\" # @param {type:\"string\"}\n",
"eval_split_name = \"test\" # @param {type:\"string\"}\n",
"\n",
"# Name of the dataset column containing training text input.\n",
"instruct_column_in_dataset = \"text\" # @param {type:\"string\"}"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "lkCVPgWl2vxv"
},
"outputs": [],
"source": [
"# @title Set model\n",
"\n",
"# @markdown Select a model variant of Mixtral.\n",
"base_model_id = \"mistralai/Mixtral-8x7B-v0.1\" # @param [\"mistralai/Mixtral-8x7B-v0.1\"] {isTemplate: true}\n",
"pretrained_model_id = f\"gs://vertex-model-garden-public-us/{base_model_id}\""
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "YZPqfZ-FvPXS"
},
"outputs": [],
"source": [
"# @title Validate Dataset with Template\n",
"\n",
"# @markdown This section validates the train and eval datasets with the template before starting the fine tuning process.\n",
"\n",
"import transformers\n",
"\n",
"dataset_validation_util = importlib.import_module(\n",
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.dataset_validation_util\"\n",
")\n",
"\n",
"if dataset_validation_util.is_gcs_path(pretrained_model_id):\n",
" # Download tokenizer.\n",
" ! mkdir tokenizer\n",
" ! gsutil cp {pretrained_model_id}/tokenizer.json ./tokenizer\n",
" ! gsutil cp {pretrained_model_id}/config.json ./tokenizer\n",
" tokenizer_path = \"./tokenizer\"\n",
" access_token = \"\"\n",
"else:\n",
" tokenizer_path = pretrained_model_id\n",
" access_token = HF_TOKEN\n",
"\n",
"tokenizer = transformers.AutoTokenizer.from_pretrained(\n",
" tokenizer_path,\n",
" trust_remote_code=False,\n",
" use_fast=True,\n",
" token=access_token,\n",
")\n",
"\n",
"# Validate the train dataset.\n",
"dataset_validation_util.validate_dataset_with_template(\n",
" dataset_name=train_dataset_name,\n",
" split=train_split_name,\n",
" input_column=instruct_column_in_dataset,\n",
" template=template,\n",
" use_multiprocessing=False,\n",
" tokenizer=tokenizer,\n",
")\n",
"\n",
"# Validate the eval dataset.\n",
"dataset_validation_util.validate_dataset_with_template(\n",
" dataset_name=eval_dataset_name,\n",
" split=eval_split_name,\n",
" input_column=instruct_column_in_dataset,\n",
" template=template,\n",
" use_multiprocessing=False,\n",
" tokenizer=tokenizer,\n",
")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "885Vf4o8hbbo"
},
"outputs": [],
"source": [
"# @title Finetune\n",
"\n",
"# @markdown This section demonstrates how to finetune the Mixtral-8x7B model and merge the finetuned LoRA adapter with the base model on Vertex AI. It uses the Vertex AI SDK to create and run the custom training jobs.\n",
"\n",
"# @markdown The training job takes approximately between 10 to 20 mins to set-up. Once done, the training job is expected to take around 90 mins with the default configurations. To find the training time, throughput, and memory usage of your training job, you can go to the training logs and check the log line of the last training epoch.\n",
"\n",
"# @markdown **Note**:\n",
"# @markdown 1. We recommend setting `finetuning_precision_mode` to `4bit` because it enables using fewer hardware resources for finetuning.\n",
"# @markdown 1. If `max_steps > 0`, it takes precedence over `epochs`. One can set a small `max_steps` value to quickly check the pipeline.\n",
"\n",
"# @markdown Acceletor type to use for training.\n",
"accelerator_type = \"NVIDIA_A100_80GB\" # @param [\"NVIDIA_A100_80GB\", \"NVIDIA_H100_80GB\"]\n",
"\n",
"# The pre-built training docker image.\n",
"if accelerator_type == \"NVIDIA_A100_80GB\":\n",
" repo = \"us-docker.pkg.dev/vertex-ai-restricted\"\n",
" is_restricted_image = True\n",
" is_dynamic_workload_scheduler = False\n",
" dws_kwargs = {}\n",
"else:\n",
" repo = \"us-docker.pkg.dev/vertex-ai\"\n",
" is_restricted_image = False\n",
" is_dynamic_workload_scheduler = True\n",
" dws_kwargs = {\n",
" \"max_wait_duration\": 1800, # 30 minutes\n",
" \"scheduling_strategy\": gca_custom_job_compat.Scheduling.Strategy.FLEX_START,\n",
" }\n",
"\n",
"TRAIN_DOCKER_URI = (\n",
" f\"{repo}/vertex-vision-model-garden-dockers/pytorch-peft-train:stable_20240909\"\n",
")\n",
"\n",
"# Worker pool spec.\n",
"if accelerator_type == \"NVIDIA_A100_80GB\":\n",
" per_node_accelerator_count = 8\n",
" machine_type = \"a2-ultragpu-8g\"\n",
"elif accelerator_type == \"NVIDIA_H100_80GB\":\n",
" per_node_accelerator_count = 8\n",
" machine_type = \"a3-highgpu-8g\"\n",
"else:\n",
" raise ValueError(\n",
" f\"Recommended machine settings not found for: {accelerator_type}. To use another accelerator type, edit this code block to pass in an appropriate `machine_type`, `accelerator_type`, and `per_node_accelerator_count` to the deploy_model_vllm function by clicking `Show Code` and then modifying the code.\"\n",
" )\n",
"\n",
"# @markdown Batch size for finetuning.\n",
"per_device_train_batch_size = 1 # @param{type:\"integer\"}\n",
"# @markdown Number of updates steps to accumulate the gradients for, before performing a backward/update pass.\n",
"gradient_accumulation_steps = 4 # @param{type:\"integer\"}\n",
"# @markdown Maximum sequence length.\n",
"max_seq_length = 4096 # @param{type:\"integer\"}\n",
"# @markdown Setting a positive `max_steps` here will override `num_epochs`.\n",
"max_steps = -1 # @param{type:\"integer\"}\n",
"num_epochs = 1.0 # @param{type:\"number\"}\n",
"# @markdown Precision mode for finetuning.\n",
"finetuning_precision_mode = \"4bit\" # @param [\"4bit\"]\n",
"# @markdown Learning rate.\n",
"learning_rate = 5e-5 # @param{type:\"number\"}\n",
"# @markdown The scheduler type to use.\n",
"lr_scheduler_type = \"cosine\" # @param{type:\"string\"}\n",
"# @markdown LoRA parameters.\n",
"lora_rank = 16 # @param{type:\"integer\"}\n",
"lora_alpha = 32 # @param{type:\"integer\"}\n",
"lora_dropout = 0.05 # @param{type:\"number\"}\n",
"# Activates gradient checkpointing for the current model (may be referred to as activation checkpointing or checkpoint activations in other frameworks).\n",
"enable_gradient_checkpointing = True\n",
"# Attention implementation to use in the model.\n",
"attn_implementation = \"flash_attention_2\"\n",
"# The optimizer for which to schedule the learning rate.\n",
"optimizer = \"paged_adamw_32bit\"\n",
"# Define the proportion of training to be dedicated to a linear warmup where learning rate gradually increases.\n",
"warmup_ratio = \"0.01\"\n",
"# The list or string of integrations to report the results and logs to.\n",
"report_to = \"tensorboard\"\n",
"# Number of updates steps before two checkpoint saves.\n",
"save_steps = 10\n",
"# Number of update steps between two logs.\n",
"logging_steps = save_steps\n",
"# Train precision of the model.\n",
"train_precision = \"float16\"\n",
"\n",
"replica_count = 1\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=per_node_accelerator_count * replica_count,\n",
" is_for_training=True,\n",
" is_restricted_image=is_restricted_image,\n",
" is_dynamic_workload_scheduler=is_dynamic_workload_scheduler,\n",
")\n",
"\n",
"# Setup training job.\n",
"job_name = common_util.get_job_name_with_datetime(\"mixtral-lora-train\")\n",
"\n",
"base_output_dir = os.path.join(STAGING_BUCKET, job_name)\n",
"# Create a GCS folder to store the LORA adapter.\n",
"lora_output_dir = os.path.join(base_output_dir, \"adapter\")\n",
"# Create a GCS folder to store the merged model with the base model and the\n",
"# finetuned LORA adapter.\n",
"merged_model_output_dir = os.path.join(base_output_dir, \"merged-model\")\n",
"\n",
"# Add labels for the finetuning job.\n",
"labels = {\n",
" \"mg-source\": \"notebook\",\n",
" \"mg-notebook-name\": \"model_garden_pytorch_mixtral_peft_tuning.ipynb\".split(\".\")[0],\n",
"}\n",
"\n",
"labels[\"mg-tune\"] = \"publishers-mistralai-models-mixtral\"\n",
"versioned_model_id = base_model_id.split(\"/\")[1].lower().replace(\".\", \"-\")\n",
"labels[\"versioned-mg-tune\"] = f\"{labels['mg-tune']}-{versioned_model_id}\"\n",
"\n",
"eval_args = [\n",
" f\"--eval_dataset_path={eval_dataset_name}\",\n",
" f\"--eval_column={instruct_column_in_dataset}\",\n",
" f\"--eval_template={template}\",\n",
" f\"--eval_split={eval_split_name}\",\n",
" f\"--eval_steps={save_steps}\",\n",
" \"--eval_tasks=builtin_eval\",\n",
" \"--eval_metric_name=loss\",\n",
"]\n",
"\n",
"train_job_args = [\n",
" \"--config_file=vertex_vision_model_garden_peft/deepspeed_zero2_8gpu.yaml\",\n",
" \"--task=instruct-lora\",\n",
" \"--completion_only=False\",\n",
" f\"--pretrained_model_id={pretrained_model_id}\",\n",
" f\"--dataset_name={train_dataset_name}\",\n",
" f\"--train_split_name={train_split_name}\",\n",
" f\"--instruct_column_in_dataset={instruct_column_in_dataset}\",\n",
" f\"--output_dir={lora_output_dir}\",\n",
" f\"--merge_base_and_lora_output_dir={merged_model_output_dir}\",\n",
" f\"--per_device_train_batch_size={per_device_train_batch_size}\",\n",
" f\"--gradient_accumulation_steps={gradient_accumulation_steps}\",\n",
" f\"--lora_rank={lora_rank}\",\n",
" f\"--lora_alpha={lora_alpha}\",\n",
" f\"--lora_dropout={lora_dropout}\",\n",
" f\"--max_steps={max_steps}\",\n",
" f\"--max_seq_length={max_seq_length}\",\n",
" f\"--learning_rate={learning_rate}\",\n",
" f\"--lr_scheduler_type={lr_scheduler_type}\",\n",
" f\"--precision_mode={finetuning_precision_mode}\",\n",
" f\"--train_precision={train_precision}\",\n",
" f\"--enable_gradient_checkpointing={enable_gradient_checkpointing}\",\n",
" f\"--num_epochs={num_epochs}\",\n",
" f\"--attn_implementation={attn_implementation}\",\n",
" f\"--optimizer={optimizer}\",\n",
" f\"--warmup_ratio={warmup_ratio}\",\n",
" f\"--report_to={report_to}\",\n",
" f\"--logging_output_dir={base_output_dir}\",\n",
" f\"--save_steps={save_steps}\",\n",
" f\"--logging_steps={logging_steps}\",\n",
" f\"--template={template}\",\n",
"] + eval_args\n",
"\n",
"train_job = aiplatform.CustomContainerTrainingJob(\n",
" display_name=job_name,\n",
" container_uri=TRAIN_DOCKER_URI,\n",
" labels=labels,\n",
")\n",
"\n",
"print(\"Running training job with args:\")\n",
"print(\" \\\\\\n\".join(train_job_args))\n",
"# Pass training arguments and launch job.\n",
"train_job.run(\n",
" args=train_job_args,\n",
" replica_count=replica_count,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=per_node_accelerator_count,\n",
" boot_disk_size_gb=500,\n",
" service_account=SERVICE_ACCOUNT,\n",
" base_output_dir=base_output_dir,\n",
" sync=False, # Non-blocking call to run.\n",
" **dws_kwargs,\n",
")\n",
"\n",
"# Wait until resource has been created.\n",
"train_job.wait_for_resource_creation()\n",
"\n",
"print(\"LoRA adapter will be saved in:\", lora_output_dir)\n",
"print(\"Trained and merged models will be saved in:\", merged_model_output_dir)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "gLzO6p0Gm2BM"
},
"outputs": [],
"source": [
"# @title Run TensorBoard\n",
"# @markdown This section shows how to launch TensorBoard in a [Cloud Shell](https://cloud.google.com/shell/docs).\n",
"# @markdown 1. Click the Cloud Shell icon(![terminal](https://github.com/google/material-design-icons/blob/master/png/action/terminal/materialicons/24dp/1x/baseline_terminal_black_24dp.png?raw=true)) on the top right to open the Cloud Shell.\n",
"# @markdown 2. Copy the `tensorboard` command shown below by running this cell.\n",
"# @markdown 3. Paste and run the command in the Cloud Shell to launch TensorBoard.\n",
"# @markdown 4. Once the command runs (You may have to click `Authorize` if prompted), click the link starting with `http://localhost`.\n",
"\n",
"# @markdown Note: You may need to wait around 10 minutes after the job starts in order for the TensorBoard logs to be written to the GCS bucket.\n",
"print(f\"Command to copy: tensorboard --logdir {base_output_dir}/logs\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "GyDWPdV1NjMT"
},
"outputs": [],
"source": [
"# @title Deploy\n",
"# @markdown This section uploads the model to Model Registry and deploys it on the Endpoint. It takes 15 minutes to 1 hour to finish depending on the size of model.\n",
"\n",
"if train_job.end_time is None:\n",
" print(\"Waiting for the training job to finish...\")\n",
" train_job.wait()\n",
" print(\"The training job has finished.\")\n",
"\n",
"print(\"Deploying models in:\", merged_model_output_dir)\n",
"\n",
"# The pre-built serving docker image for vLLM.\n",
"VLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-vllm-serve:20240721_0916_RC00\"\n",
"\n",
"dtype = \"auto\"\n",
"\n",
"# @markdown L4 GPUs are good serving solutions and are more cost effective than V100s for 8x7B models. The 8x22B models only works with A100/H100 GPUs now.\n",
"\n",
"# Find Vertex AI prediction supported accelerators and regions [here](https://cloud.google.com/vertex-ai/docs/predictions/configure-compute).\n",
"# @markdown Accelerator type to use for serving.\n",
"accelerator_type = \"NVIDIA_L4\" # @param [\"NVIDIA_L4\", \"NVIDIA_TESLA_V100\", \"NVIDIA_H100_80GB\"]\n",
"\n",
"if accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-96\"\n",
" accelerator_count = 8\n",
"elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" machine_type = \"n1-highmem-32\"\n",
" accelerator_count = 8\n",
" dtype = \"float16\"\n",
"elif accelerator_type == \"NVIDIA_H100_80GB\":\n",
" machine_type = \"a3-highgpu-8g\"\n",
" accelerator_count = 8\n",
"\n",
"if \"22B\" in base_model_id and accelerator_type != \"NVIDIA_H100_80GB\":\n",
" raise ValueError(\"8x22B model version only works with H100/A100 GPUs.\")\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=False,\n",
")\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"gpu_memory_utilization = 0.85\n",
"max_model_len = 8192 # Maximum context length.\n",
"\n",
"# Ensure max_model_len does not exceed the limit\n",
"if max_model_len > 8192:\n",
" raise ValueError(\"max_model_len cannot exceed 8192\")\n",
"\n",
"\n",
"def get_deploy_source() -> str:\n",
" \"\"\"Gets deploy_source string based on running environment.\"\"\"\n",
" vertex_product = os.environ.get(\"VERTEX_PRODUCT\", \"\")\n",
" if vertex_product == \"COLAB_ENTERPRISE\":\n",
" return \"notebook_colab_enterprise\"\n",
" elif vertex_product == \"WORKBENCH_INSTANCE\":\n",
" return \"notebook_workbench\"\n",
" else:\n",
" # Legacy workbench, legacy colab, or other custom environments.\n",
" return \"notebook_environment_unspecified\"\n",
"\n",
"\n",
"def deploy_model_vllm(\n",
" model_name: str,\n",
" model_id: str,\n",
" publisher: str,\n",
" publisher_model_id: str,\n",
" service_account: str,\n",
" base_model_id: str = None,\n",
" machine_type: str = \"g2-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
" gpu_memory_utilization: float = 0.9,\n",
" max_model_len: int = 4096,\n",
" dtype: str = \"auto\",\n",
" enable_trust_remote_code: bool = False,\n",
" enforce_eager: bool = False,\n",
" enable_lora: bool = False,\n",
" enable_chunked_prefill: bool = False,\n",
" enable_prefix_cache: bool = False,\n",
" host_prefix_kv_cache_utilization_target: float = 0.0,\n",
" max_loras: int = 1,\n",
" max_cpu_loras: int = 8,\n",
" use_dedicated_endpoint: bool = False,\n",
" max_num_seqs: int = 256,\n",
" model_type: str = None,\n",
" enable_llama_tool_parser: bool = False,\n",
" is_spot: bool = False,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys trained models with vLLM into Vertex AI.\"\"\"\n",
" endpoint = aiplatform.Endpoint.create(\n",
" display_name=f\"{model_name}-endpoint\",\n",
" dedicated_endpoint_enabled=use_dedicated_endpoint,\n",
" )\n",
"\n",
" if not base_model_id:\n",
" base_model_id = model_id\n",
"\n",
" # See https://docs.vllm.ai/en/latest/models/engine_args.html for a list of possible arguments with descriptions.\n",
" vllm_args = [\n",
" \"python\",\n",
" \"-m\",\n",
" \"vllm.entrypoints.api_server\",\n",
" \"--host=0.0.0.0\",\n",
" \"--port=8080\",\n",
" f\"--model={model_id}\",\n",
" f\"--tensor-parallel-size={accelerator_count}\",\n",
" \"--swap-space=16\",\n",
" f\"--max-model-len={max_model_len}\",\n",
" f\"--dtype={dtype}\",\n",
" f\"--max-loras={max_loras}\",\n",
" f\"--max-cpu-loras={max_cpu_loras}\",\n",
" f\"--max-num-seqs={max_num_seqs}\",\n",
" \"--disable-log-stats\",\n",
" ]\n",
"\n",
" if gpu_memory_utilization:\n",
" vllm_args.append(f\"--gpu-memory-utilization={gpu_memory_utilization}\")\n",
"\n",
" if enable_trust_remote_code:\n",
" vllm_args.append(\"--trust-remote-code\")\n",
"\n",
" if enforce_eager:\n",
" vllm_args.append(\"--enforce-eager\")\n",
"\n",
" if enable_lora:\n",
" vllm_args.append(\"--enable-lora\")\n",
"\n",
" if enable_chunked_prefill:\n",
" vllm_args.append(\"--enable-chunked-prefill\")\n",
"\n",
" if enable_prefix_cache:\n",
" vllm_args.append(\"--enable-prefix-caching\")\n",
"\n",
" if 0 < host_prefix_kv_cache_utilization_target < 1:\n",
" vllm_args.append(\n",
" f\"--host-prefix-kv-cache-utilization-target={host_prefix_kv_cache_utilization_target}\"\n",
" )\n",
"\n",
" if model_type:\n",
" vllm_args.append(f\"--model-type={model_type}\")\n",
"\n",
" if enable_llama_tool_parser:\n",
" vllm_args.append(\"--enable-auto-tool-choice\")\n",
" vllm_args.append(\"--tool-call-parser=vertex-llama-3\")\n",
"\n",
" env_vars = {\n",
" \"MODEL_ID\": base_model_id,\n",
" \"DEPLOY_SOURCE\": \"notebook\",\n",
" }\n",
"\n",
" # HF_TOKEN is not a compulsory field and may not be defined.\n",
" try:\n",
" if HF_TOKEN:\n",
" env_vars[\"HF_TOKEN\"] = HF_TOKEN\n",
" except NameError:\n",
" pass\n",
"\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=VLLM_DOCKER_URI,\n",
" serving_container_args=vllm_args,\n",
" serving_container_ports=[8080],\n",
" serving_container_predict_route=\"/generate\",\n",
" serving_container_health_route=\"/ping\",\n",
" serving_container_environment_variables=env_vars,\n",
" serving_container_shared_memory_size_mb=(16 * 1024), # 16 GB\n",
" serving_container_deployment_timeout=7200,\n",
" model_garden_source_model_name=(\n",
" f\"publishers/{publisher}/models/{publisher_model_id}\"\n",
" ),\n",
" )\n",
" print(\n",
" f\"Deploying {model_name} on {machine_type} with {accelerator_count} {accelerator_type} GPU(s).\"\n",
" )\n",
" model.deploy(\n",
" endpoint=endpoint,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" deploy_request_timeout=1800,\n",
" service_account=service_account,\n",
" spot=is_spot,\n",
" system_labels={\n",
" \"NOTEBOOK_NAME\": \"model_garden_pytorch_mixtral_peft_tuning.ipynb\",\n",
" \"NOTEBOOK_ENVIRONMENT\": get_deploy_source(),\n",
" },\n",
" )\n",
" print(\"endpoint_name:\", endpoint.name)\n",
"\n",
" return model, endpoint\n",
"\n",
"\n",
"models[\"vllm_gpu\"], endpoints[\"vllm_gpu\"] = deploy_model_vllm(\n",
" model_name=common_util.get_job_name_with_datetime(prefix=\"mixtral-vllm-serve\"),\n",
" model_id=merged_model_output_dir,\n",
" publisher=\"mistral-ai\",\n",
" publisher_model_id=\"mixtral\",\n",
" service_account=SERVICE_ACCOUNT,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" gpu_memory_utilization=gpu_memory_utilization,\n",
" max_model_len=max_model_len,\n",
" dtype=dtype,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"# @markdown Click \"Show code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "4v2Mnui4tH1X"
},
"outputs": [],
"source": [
"# @title Predict\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint with text prompts. Sampling parameters supported by vLLM can be found [here](https://docs.vllm.ai/en/latest/dev/sampling_params.html).\n",
"\n",
"# @markdown Example:\n",
"\n",
"# @markdown ```\n",
"# @markdown Human: What is a car?\n",
"# @markdown Assistant: A car, or a motor car, is a road-connected human-transportation system used to move people or goods from one place to another. The term also encompasses a wide range of vehicles, including motorboats, trains, and aircrafts. Cars typically have four wheels, a cabin for passengers, and an engine or motor. They have been around since the early 19th century and are now one of the most popular forms of transportation, used for daily commuting, shopping, and other purposes.\n",
"# @markdown ```\n",
"# @markdown Additionally, you can moderate the generated text with Vertex AI. See [Moderate text documentation](https://cloud.google.com/natural-language/docs/moderating-text) for more details.\n",
"\n",
"# Loads an existing endpoint instance using the endpoint name:\n",
"# - Using `endpoint_name = endpoint.name` allows us to get the\n",
"# endpoint name of the endpoint `endpoint` created in the cell\n",
"# above.\n",
"# - Alternatively, you can set `endpoint_name = \"1234567890123456789\"` to load\n",
"# an existing endpoint with the ID 1234567890123456789.\n",
"# You may uncomment the code below to load an existing endpoint.\n",
"\n",
"# endpoint_name = \"\" # @param {type:\"string\"}\n",
"# aip_endpoint_name = (\n",
"# f\"projects/{PROJECT_ID}/locations/{REGION}/endpoints/{endpoint_name}\"\n",
"# )\n",
"# endpoint = aiplatform.Endpoint(aip_endpoint_name)\n",
"\n",
"prompt = \"What is a car?\" # @param {type: \"string\"}\n",
"# @markdown If you encounter an issue like `ServiceUnavailable: 503 Took too long to respond when processing`, you can reduce the maximum number of output tokens, by lowering `max_tokens`.\n",
"max_tokens = 50 # @param {type:\"integer\"}\n",
"temperature = 1.0 # @param {type:\"number\"}\n",
"top_p = 1.0 # @param {type:\"number\"}\n",
"top_k = 1 # @param {type:\"integer\"}\n",
"# @markdown Set `raw_response` to `True` to obtain the raw model output. Set `raw_response` to `False` to apply additional formatting in the structure of `\"Prompt:\\n{prompt.strip()}\\nOutput:\\n{output}\"`.\n",
"raw_response = False # @param {type:\"boolean\"}\n",
"\n",
"# Overrides parameters for inferences.\n",
"instances = [\n",
" {\n",
" \"prompt\": prompt,\n",
" \"max_tokens\": max_tokens,\n",
" \"temperature\": temperature,\n",
" \"top_p\": top_p,\n",
" \"top_k\": top_k,\n",
" \"raw_response\": raw_response,\n",
" },\n",
"]\n",
"response = endpoints[\"vllm_gpu\"].predict(\n",
" instances=instances, use_dedicated_endpoint=use_dedicated_endpoint\n",
")\n",
"\n",
"for prediction in response.predictions:\n",
" print(prediction)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "x9EMCOUJ6-ji"
},
"outputs": [],
"source": [
"# @title Delete the model and endpoint\n",
"\n",
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continuous charges that may incur.\n",
"\n",
"# Undeploy model and delete endpoint.\n",
"for endpoint in endpoints.values():\n",
" endpoint.delete(force=True)\n",
"\n",
"# Delete models.\n",
"for model in models.values():\n",
" model.delete()\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_NAME"
]
}
],
"metadata": {
"accelerator": "GPU",
"colab": {
"name": "model_garden_pytorch_mixtral_peft_tuning.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -155,7 +155,7 @@
"source": [
"# @title Request for quota\n",
"\n",
"# @markdown By default, the quota for H100 deployment `Custom model serving per region` is 0. You need to request for H100 quota following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota)."
"# @markdown By default, the quota for H100 deployment `Custom model serving per region` is 0. You need to request for H100 quota following the instructions at [\"Request a quota adjustment\"](https://cloud.google.com/docs/quotas/view-manage#requesting_higher_quota)."
]
},
{
@@ -81,7 +81,8 @@
"\n",
"### Request For TPU Quota\n",
"\n",
"By default, the quota for TPU training [Custom model training TPU v5e cores per region](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_training_tpu_v5e) is 0. TPU quota is only available in `us-west1`, `us-west4`, `us-central1`. You can request for higher TPU quota following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota). It is suggested to request at least 4 v5e to run this notebook."
"By default, the quota for TPU training [Custom model training TPU v5e cores per region](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_training_tpu_v5e) is 0. TPU quota is only available in `us-west1`, `us-west4`, `us-central1`. You can request for higher TPU quota following the instructions at [\"Request a quota adjustment\"](https://cloud.google.com/docs/quotas/view-manage#requesting_higher_quota)\n",
". It is suggested to request at least 4 v5e to run this notebook."
]
},
{
@@ -0,0 +1,221 @@
```
The customer provides input Zarr files in a GCS bucket path via the
`--input_data_gcs_dir` flag.
#### How to Use
##### Flag: `--input_data_gcs_dir`
`--input_data_gcs_dir` flag for specifying custom input data:
```
--input_data_gcs_dir=gs://customer-bucket/path-to-input-data
```
##### GCS Bucket Setup
We recommend using the **same GCS bucket** for both input data and output
predictions.
The customer should place their input Zarr files under a path within their
existing output bucket, e.g.:
```
gs://customer-bucket/custom-inputs/ ← input Zarr files go here
gs://customer-bucket/outputs/ ← model predictions are written here
```
#### Input File Format
##### Zarr V3 Format Requirement
**All custom input Zarr files MUST be in Zarr V3 format.**
When creating custom input files, ensure they are saved in Zarr V3 format:
```python
import xarray as xa
# When creating input files, ensure they are saved in Zarr V3 format
dataset.to_zarr("path/to/output.zarr", zarr_format=3)
```
If V2 format files are provided as custom input, the job raises
an error when attempting to read them.
##### Example Dataset
Opening an example input Zarr file with xarray (init time 2026-03-18 00:00 UTC):
```python
>>> import xarray as xa
>>> ds = xa.open_zarr("2026_A1D03180000031800011.zarr")
>>> ds
<xarray.Dataset> Size: ...
Dimensions: (isobaricInhPa: 13, latitude: 721, longitude: 1440)
Coordinates:
* isobaricInhPa (isobaricInhPa) float64 13 ...
* latitude (latitude) float64 721 ...
* longitude (longitude) float64 1440 ...
number int64 ...
step int64 ...
surface float64 ...
time int64 ...
valid_time int64 ...
Data variables: (13 total)
msl (latitude, longitude) float32 ...
q (isobaricInhPa, latitude, longitude) float32 ...
sst (latitude, longitude) float32 ...
t (isobaricInhPa, latitude, longitude) float32 ...
t2m (latitude, longitude) float32 ...
u (isobaricInhPa, latitude, longitude) float32 ...
u10 (latitude, longitude) float32 ...
u100 (latitude, longitude) float32 ...
v (isobaricInhPa, latitude, longitude) float32 ...
v10 (latitude, longitude) float32 ...
v100 (latitude, longitude) float32 ...
w (isobaricInhPa, latitude, longitude) float32 ...
z (isobaricInhPa, latitude, longitude) float32 ...
```
##### Variables
All data variables are `float32`.
Only the variables and levels needed for inference are listed here. Your zarr
may contain additional variables and levels.
###### Pressure Level Variables (3D: `isobaricInhPa` × `latitude` × `longitude`)
Shape: `[13, 721, 1440]`
These variables are used for input at the following pressure levels (hPa):
`50, 100, 150, 200, 250, 300, 400, 500, 600, 700, 850, 925, 1000`
| Short Name | Description |
| :--- | :--- |
| `q` | specific humidity |
| `t` | temperature |
| `u` | u component of wind |
| `v` | v component of wind |
| `w` | vertical velocity |
| `z` | geopotential |
###### Surface Variables (2D: `latitude` × `longitude`)
Shape: `[721, 1440]`
| Short Name | Description |
| :--- | :--- |
| `msl` | mean sea level pressure |
| `sst` | sea surface temperature |
| `t2m` | 2m temperature |
| `u10` | 10m u component of wind |
| `u100` | 100m u component of wind |
| `v10` | 10m v component of wind |
| `v100` | 100m v component of wind |
###### Scalar Coordinates
| Name | dtype | Description |
|--------------|-----------|----------------------------------------------|
| `number` | `int64` | Ensemble member number |
| `step` | `int64` | Forecast step |
| `surface` | `float64` | Surface level indicator |
| `time` | `int64` | Timestamp (units: days since init time, calendar: proleptic_gregorian) |
| `valid_time` | `int64` | Validity time |
##### Dimension Coordinates
| Coordinate | dtype | Shape | Description |
|-----------------|-----------|----------|--------------------------------|
| `isobaricInhPa` | `float64` | `[13]` | 13 pressure levels in hPa |
| `latitude` | `float64` | `[721]` | 0.25° resolution, 721 points |
| `longitude` | `float64` | `[1440]` | 0.25° resolution, 1440 points |
##### Spatial Resolution
The data is at **0.25° resolution** globally:
- Latitude: 721 points (90°N to 90°S)
- Longitude: 1440 points (0° to 359.75°E)
#### File Naming and Structure
##### File Naming Convention
The inference binary expects input Zarr files to follow a specific file naming
convention. Each file corresponds to a specific forecast initialization time and
uses the following format:
```
<year>_<config><stream><MMDDHHMMMMDDHHMMEE>.zarr
```
Where:
- `<year>`: 4-digit year (e.g., `2025`)
- `<config>`: Data config name (default: `A1`)
- `<stream>`: `D` for 00/12 UTC init times, `S` for 06/18 UTC
- First `MMDDHHMM`: Month, day, hour, minute of the forecast init time
- Second `MMDDHHMM`: Month, day, hour, minute of the validity time
- `EE`: Experiment version (default: `1`)
The validity minute is hardcoded to `01`.
###### Examples
For a forecast initialized at **2026-03-18 00:00 UTC** using fc0:
```
2026_A1D03180000031800011.zarr
```
For a forecast initialized at **2026-03-17 06:00 UTC** using fc0:
```
2026_A1S03170600031706011.zarr
```
##### Success Sentinels
Each Zarr file directory must contain a `success` sentinel file to signal that
the data is complete and ready for reading:
```
gs://<bucket>/custom-inputs/
├── 2026_A1D03180000031800011.zarr/
│ ├── .zmetadata
│ ├── <array data>
│ └── success ← required sentinel file
├── 2026_A1D03171200031712011.zarr/
│ ├── .zmetadata
│ ├── <array data>
│ └── success
```
The sentinel is a zero-byte file named `success` placed inside each `.zarr`
directory. The job will not proceed with inference until all required input
sentinels exist.
##### Number of Input Files
The model typically requires **2 input timestamps** (the forecast init time and
6 hours prior). For example, for a forecast initialized at 2026-03-18 12:00 UTC,
the binary expects:
1. `2026_A1D03181200031812011.zarr` (init time: 12:00 UTC)
2. `2026_A1D03180600031806011.zarr` (6 hours prior: 06:00 UTC)
#### Caveats and Limitations
1. **No fine-tuning guarantee**: The model is trained on ECMWF HRES
data. Using custom inputs from a different source may degrade
forecast quality, especially for features sensitive to the initial
condition source.
2. **Variable completeness**: All variables listed above must be present in the
custom input files. Missing variables will cause the inference to fail.
3. **Temporal alignment**: Custom input timestamps must align with valid HRES
forecast hours (00, 06, 12, or 18 UTC).
4. **File format**: Only Zarr V3 format is accepted.
+3 -4
View File
@@ -4,10 +4,9 @@
## Overview
**Disclaimer:**
*Experimental*\
This product is subject to the "Pre-GA Offerings Terms" in the General Service Terms section of the [Service Specific Terms](https://cloud.google.com/terms/service-terms#1). Pre-GA products are available "as is" and might have limited support. For more information, see the [launch stage descriptions](https://cloud.google.com/products#product-launch-stages). <!-- disableFinding(LINE_OVER_80) -->
This product is subject to the General Service Terms section of the [Service Specific Terms](https://cloud.google.com/terms/service-terms#1). GA products follows standard support. For more information, see the [launch stage descriptions](https://cloud.google.com/products#product-launch-stages). <!-- disableFinding(LINE_OVER_80) -->
Access to the forecasting capabilities requires application and approval. Users must be added to an allowlist to generate forecasts using this service. Review pricing details at [Vertex AI Custom Training pricing,](https://cloud.google.com/vertex-ai/pricing?hl=en&e=48754805#custom-trained-models) [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) before running. <!-- disableFinding(LINE_OVER_80) -->
Please contact your Google Cloud sales person for using WeatherNext 2 on Agent Platform. In case you do not have a salesperson you work with, please submit the form using the Request access button in this model card and a Google Cloud representative will contact you. <!-- disableFinding(LINE_OVER_80) -->
**Overview**
@@ -75,4 +74,4 @@ To utilize these models for generating real-time forecasts via this service:
Running these models *will incur costs* for the GPUs and other Google Cloud resources used. Learn about [Vertex AI Custom Training pricing](https://cloud.google.com/vertex-ai/pricing?hl=en&e=48754805#custom-trained-models), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate an estimate.
## Quick start
The quickest way to get started with the WeatherNext in Google Cloud Platform is to run [our example notebook](weathernext_2_early_access_program.ipynb) in [Google Colab](https://colab.research.google.com/).
The quickest way to get started with the WeatherNext in Google Cloud Platform is to run [our example notebook](weathernext_2_ic_pc.ipynb) in [Google Colab](https://colab.research.google.com/).
@@ -0,0 +1,583 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "RirpY96_3zL0"
},
"outputs": [],
"source": [
"# Copyright 2026 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",
"# https://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": "52d0938f"
},
"source": [
"# WeatherNext 2 (Using DWS)\n",
"<table align=\"left\">\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://colab.research.google.com/github/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/weathernext/weathernext_2_dws.ipynb\">\n",
" <img 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%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%weathernext%2Fweathernext_2_dws.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/vertex-ai-samples/blob/main/notebooks/community/weathernext/weathernext_2_dws.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/vertex-ai-samples/blob/main/notebooks/community/weathernext/weathernext_2_dws.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>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "DGFzy9MTuYL_"
},
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates running [WeatherNext 2 inference on Google Cloud Vertex AI](https://developers.google.com/weathernext/guides/access-vmg). WeatherNext 2 is Google's latest medium-range probabilistic forecasting model, principally an operational version the FGN model ([published June 2025](https://arxiv.org/abs/2506.10772)). More information is available in the [WeatherNext documentation](https://developers.google.com/weathernext).\n",
"\n",
"### Objective\n",
"\n",
"- Configure the model inputs for distributed, multi-host inference on H100 or A100 GPUs.\n",
"- Run WeatherNext 2 model forecasts in parallel.\n",
"- Visualize forecast results.\n",
"\n",
"### Costs\n",
"\n",
"This uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage.\n",
"\n",
"\n",
"## Before you begin\n",
"\n",
"### Request For GPU Quota\n",
"\n",
"**WARNING:** Make sure you have sufficient GPU quota allocated for the inference configuration (i.e. `num_samples`) before running Vertex Jobs. Otherwise, some Vertex jobs may run while others will fail which would produce\n",
"incomplete results.\n",
"\n",
"\n",
"By default, the quota for GPUs is 0. You can request a higher quota by following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota).\n",
"\n",
"You will need to request quota for either **NVIDIA H100 80GB GPUs** or **NVIDIA A100 80GB GPUs** in your selected region. The total number of GPUs you request must be sufficient for your largest planned forecast (i.e., `num_samples`).\n",
"\n",
"You should request for the following quota:\n",
"\n",
"- Service: `Vertex AI API`\n",
"- Name: `Custom model training preemptible Nvidia A100 80GB GPUs per region` OR `Custom model training preemptible Nvidia H100 GPUs per region`\n",
"\n",
"### Custom Inputs Guide\n",
"\n",
"Please refer to the guide [here](CUSTOM_INPUTS_GUIDE.md)."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "yt4GxkKcDD7Y"
},
"outputs": [],
"source": [
"# @title Install python packages\n",
"\n",
"# Note that you may need to restart the kernel after this step.\n",
"# If so, continue to the next cell after restarting.\n",
"\n",
"print(\"Installing python packages.\")\n",
"\n",
"! pip3 install \\\n",
" google-cloud-aiplatform==1.129.0 \\\n",
" xarray[complete]"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "dO2mF4CPfpHW1ZKbIsut8r69"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. **[Optional]** [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched.\n",
"\n",
"\n",
"BUCKET_URI = \"gs://my-bucket\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 3. Select a region that has the required GPUs available.\n",
"\n",
"REGION = \"us-central1\" # @param {type:\"string\"}\n",
"\n",
"# Import the necessary packages\n",
"import datetime\n",
"import os\n",
"import re\n",
"\n",
"from google.cloud import aiplatform\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"now = datetime.datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" raise ValueError(\"GCS Bucket URI is invalid!\")\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" f\"Bucket region {bucket_region} is different from notebook region {REGION}\"\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"# Initialize Vertex AI API.\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Utility functions\n",
"\n",
"\n",
"def get_job_name_with_datetime(prefix: str) -> str:\n",
" return prefix + datetime.datetime.now().strftime(\"_%Y%m%d_%H%M%S\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "22BW43yjps8D"
},
"outputs": [],
"source": [
"# @title Configure Model Parameters\n",
"# @markdown Configure the hardware and input parameters for the WeatherNext 2 forecast.\n",
"\n",
"\n",
"# @markdown ### Hardware Configuration for Distributed Inference\n",
"# @markdown - **`machine_type`**: Select a valid machine type. `a3-highgpu` series use NVIDIA H100 80GB GPUs. `a2-ultragpu` series use NVIDIA A100 80GB GPUs.\n",
"# @markdown - **`num_samples`**: The total number of ensemble members to generate.\n",
"# @markdown The number of machine replicas will be calculated automatically (`num_samples` / GPUs per machine). **Therefore, `num_samples` must be a multiple of the number of GPUs in your selected `machine_type`.**\n",
"# @markdown - **`scheduling_strategy`**: The [strategy](https://cloud.google.com/vertex-ai/docs/reference/rest/v1beta1/CustomJobSpec#Strategy) used to acquire machines for the job. Defaults to [Dynamic Workload Scheduler](https://docs.cloud.google.com/vertex-ai/docs/training/schedule-jobs-dws) (FLEX_START).\n",
"machine_type = \"a3-highgpu-1g\" # @param [\"a3-highgpu-1g\", \"a3-highgpu-2g\", \"a3-highgpu-4g\", \"a3-highgpu-8g\", \"a2-ultragpu-1g\", \"a2-ultragpu-2g\", \"a2-ultragpu-4g\", \"a2-ultragpu-8g\"]\n",
"num_samples = 8 # @param {type:\"integer\"}\n",
"scheduling_strategy = \"FLEX_START\" # @param [\"FLEX_START\", \"SPOT\", \"STANDARD\"]\n",
"\n",
"# @markdown ### Forecast Configuration\n",
"# @markdown - **`forecast_init_time`**: The starting time for the forecast in ISO 8601 format (e.g., `2025-09-21T00:00:00Z`). Models are available for dates from 2024 onwards.\n",
"# @markdown - **`horizon_hrs`**: The desired length of the forecast in hours (e.g., 240 for a 10-day forecast).\n",
"# @markdown - **`model_seed`**: Choose a specific model seed (1-4) or select \"all\" to run inference with all four seeds in parallel for improved accuracy.\n",
"# @markdown - **`enable_hourly_prediction`**: If checked, the model will generate 1-hour predictions.\n",
"forecast_init_time = \"2025-11-20T00:00:00Z\" # @param {type:\"string\"}\n",
"horizon_hrs = 72 # @param {type:\"integer\"}\n",
"model_seed = \"all\" # @param [\"1\", \"2\", \"3\", \"4\", \"all\"]\n",
"enable_hourly_prediction = True # @param {type:\"boolean\"}\n",
"\n",
"# --- Parameter Validation and Configuration ---\n",
"\n",
"# Derive accelerator type and count from the chosen machine type\n",
"if machine_type.startswith(\"a3-highgpu\"):\n",
" accelerator_type = \"NVIDIA_H100_80GB\"\n",
"elif machine_type.startswith(\"a2-ultragpu\"):\n",
" accelerator_type = \"NVIDIA_A100_80GB\"\n",
"else:\n",
" raise ValueError(f\"Invalid machine type selected: {machine_type}.\")\n",
"\n",
"try:\n",
" # Extract the number of GPUs from the machine type string, e.g., 'a3-highgpu-4g' -> 4\n",
" accelerators_per_machine = int(re.search(r\"-(\\d+)g$\", machine_type).group(1))\n",
"except (AttributeError, ValueError):\n",
" raise ValueError(\n",
" f\"Could not determine accelerator count from machine type: {machine_type}\"\n",
" )\n",
"\n",
"seeds_to_run = [1, 2, 3, 4] if model_seed == \"all\" else [int(model_seed)]\n",
"num_seeds_to_run = len(seeds_to_run)\n",
"\n",
"num_samples_per_seed = num_samples\n",
"if len(seeds_to_run) > 1:\n",
" if num_samples % num_seeds_to_run != 0:\n",
" raise ValueError(\n",
" f\"`num_samples` ({num_samples}) is not divisible by the number of seeds to run ({num_seeds_to_run}.\"\n",
" )\n",
" num_samples_per_seed = num_samples // num_seeds_to_run\n",
"\n",
"# Validate that num_samples is a multiple of accelerators_per_machine\n",
"if num_samples_per_seed % accelerators_per_machine != 0:\n",
" raise ValueError(\n",
" f\"`num_samples_per_seed` ({num_samples_per_seed}) must be a multiple of the GPUs per machine ({accelerators_per_machine} for {machine_type}).\"\n",
" )\n",
"\n",
"# Calculate the number of replicas per seed\n",
"replica_count_per_seed = num_samples_per_seed // accelerators_per_machine\n",
"\n",
"# Ensure that there are enough samples\n",
"if num_samples_per_seed % accelerators_per_machine != 0:\n",
" raise ValueError(\n",
" f\"`num_samples_per_seed` ({num_samples_per_seed}) must be a multiple of the GPUs per machine ({accelerators_per_machine} for {machine_type}).\"\n",
" )\n",
"\n",
"# Calculate total GPUs needed for all jobs\n",
"total_gpus_needed = num_samples * (4 if model_seed == \"all\" else 1)\n",
"\n",
"# Set Docker URI\n",
"WEATHERNEXT2_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai-restricted/vertex-vision-model-garden-dockers/weather-next-2-inference.gpu.0-1:latest\"\n",
"\n",
"print(\"--- Job Configuration Summary ---\")\n",
"print(f\"Total Samples: {num_samples}\")\n",
"print(f\"Machine Type: {machine_type}\")\n",
"print(f\"Accelerator Type: {accelerator_type}\")\n",
"print(f\"GPUs per Machine: {accelerators_per_machine}\")\n",
"print(f\"Total number seeds to run: {num_seeds_to_run}\")\n",
"print(f\"Total number samples per seed: {num_samples_per_seed}\")\n",
"print(f\"Calculated Machine Replicas Per Seed: {replica_count_per_seed}\")\n",
"print(f\"Total GPUs per Job: {num_samples}\")\n",
"print(f\"Total GPUs across all Jobs (ensure sufficient quota): {total_gpus_needed}\")\n",
"print(f\"Docker Image: {WEATHERNEXT2_DOCKER_URI}\")\n",
"print(\"---------------------------------\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "VSuGx2jmpwP0"
},
"outputs": [],
"source": [
"# @title Run Forecasts\n",
"# @markdown This section creates and runs one or more Vertex AI Custom Training Jobs to generate the forecasts.\n",
"# @markdown **This operation is asynchronous.** The jobs will be submitted and this cell will complete quickly.\n",
"# @markdown You must monitor the job progress in the Google Cloud Console (https://console.cloud.google.com/vertex-ai/training/custom-jobs).\n",
"\n",
"from google.cloud.aiplatform.compat.types import \\\n",
" custom_job as gca_custom_job_compat\n",
"\n",
"print(f\"Submitting {len(seeds_to_run)} job(s) to run in parallel.\")\n",
"\n",
"launched_jobs = []\n",
"output_dirs = {}\n",
"\n",
"if scheduling_strategy == \"FLEX_START\":\n",
" SCHEDULLING_STRATEGY = gca_custom_job_compat.Scheduling.Strategy.FLEX_START\n",
"elif scheduling_strategy == \"SPOT\":\n",
" SCHEDULLING_STRATEGY = gca_custom_job_compat.Scheduling.Strategy.SPOT\n",
"else:\n",
" SCHEDULLING_STRATEGY = gca_custom_job_compat.Scheduling.Strategy.STANDARD\n",
"\n",
"for seed in seeds_to_run:\n",
" output_dir = os.path.join(BUCKET_URI, \"weathernext2_outputs\")\n",
" output_dirs[seed] = output_dir\n",
"\n",
" docker_args_list = [\n",
" f\"--pred_root_dir={output_dir}\",\n",
" f\"--num_samples={num_samples_per_seed}\",\n",
" f\"--horizon_hrs={horizon_hrs}\",\n",
" f\"--forecast_init_time={forecast_init_time}\",\n",
" f\"--model_seed={seed}\",\n",
" f\"--enable_hourly_prediction={enable_hourly_prediction}\",\n",
" # Uncomment the line below if you want to use custom inputs.\n",
" # f\"--input_data_gcs_dir={BUCKET_URI}/custom_inputs/\",\n",
" ]\n",
"\n",
" JOB_NAME = get_job_name_with_datetime(\n",
" prefix=f\"wn2-forecast-s{seed}-n{num_samples_per_seed}\"\n",
" )\n",
" print(f\"\\n--- Submitting Job for Seed {seed} ---\")\n",
" print(f\"JOB_NAME: {JOB_NAME}\")\n",
"\n",
" job = aiplatform.CustomContainerTrainingJob(\n",
" display_name=JOB_NAME,\n",
" container_uri=WEATHERNEXT2_DOCKER_URI,\n",
" )\n",
"\n",
" job.run(\n",
" args=docker_args_list,\n",
" replica_count=replica_count_per_seed,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerators_per_machine,\n",
" scheduling_strategy=SCHEDULLING_STRATEGY,\n",
" # Change this to True if you need to debug why the job hasn't started\n",
" sync=False,\n",
" )\n",
" launched_jobs.append(job)\n",
" print(\n",
" \"--> Job submitted successfully. Monitor it in the Google Cloud Console at https://console.cloud.google.com/vertex-ai/training/custom-jobs\"\n",
" )\n",
"\n",
"print(\"\\nAll forecast jobs have been submitted.\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "x6ZVqWRopWcI"
},
"outputs": [],
"source": [
"# @title Visualize Forecasts (Unified)\n",
"# @markdown Select which forecast output you want to visualize. This single component\n",
"# @markdown can handle both the standard 6-hourly predictions and the datasets\n",
"# @markdown with 1-hour model (which have a 'subtime' dimension).\n",
"# @markdown If you run into `Error loading Zarr store: unrecognized engine 'zarr'...` try restarting the runtime session and reruning this cell.\n",
"\n",
"# @markdown ---\n",
"# @markdown ### Visualization Settings\n",
"# @markdown - **`model_seed_to_visualize`**: Choose a specific model seed (1-4) to visualize. This should be one of the model seeds selected in the **Forecast Configuration** above.\n",
"# @markdown - **`time_steps_to_visualize`**: Choose to visualize 1-hourly or 6-hourly forecasts. If 1-hourly is selected, ensure `enable_hourly_prediction` was selected in the **Forecast Configuration** above.\n",
"# @markdown - **`variable_to_visualize`**: Choose the weather variable to visualize. See the [WeatherNext documentation](https://developers.google.com/weathernext/guides/model-specs-vmg) for variable names and descriptions.\n",
"# @markdown - **`sample_to_visualize`**: Choose the sample (ensemble member) to visualize.\n",
"# @markdown - **`plot_size`**: Choose the size of the plot to generate.\n",
"model_seed_to_visualize = \"4\" # @param [\"1\", \"2\", \"3\", \"4\"]\n",
"time_steps_to_visualize = \"6-Hourly\" # @param [\"6-Hourly\", \"1-Hourly\"]\n",
"variable_to_visualize = \"2m_temperature\" # @param {type:\"string\"}\n",
"sample_to_visualize = 0 # @param {type:\"integer\"}\n",
"plot_size = 8 # @param {type:\"number\"}\n",
"level_to_visualize = None\n",
"# @markdown ---\n",
"\n",
"\n",
"# --- Helper Functions ---\n",
"\n",
"\n",
"def init_time_to_folder_path(init_time: str) -> str:\n",
" \"\"\"\n",
" Convert init time to expected GCS folder path.\n",
" \"\"\"\n",
" init_date, init_time = init_time.split(\"T\")\n",
" return f\"{init_date.replace('-', '')}_{init_time[0:2]}hr\"\n",
"\n",
"\n",
"# override these if you'd like to visualize a different set of forecasts\n",
"visualize_bucket = BUCKET_URI\n",
"visualize_init_date = forecast_init_time\n",
"\n",
"# set paths based on chosen model seed, bucket, and init date\n",
"path_to_6hr_zarr = f\"{BUCKET_URI}/weathernext2_outputs/weathernext_2_seed_{model_seed_to_visualize}/{init_time_to_folder_path(visualize_init_date)}_01_preds/predictions.zarr/\"\n",
"path_to_1hr_zarr = f\"{BUCKET_URI}/weathernext2_outputs/weathernext_2_seed_{model_seed_to_visualize}_hourly/{init_time_to_folder_path(visualize_init_date)}_01_preds/predictions.zarr/\"\n",
"\n",
"\n",
"import datetime\n",
"from typing import Optional\n",
"\n",
"import matplotlib\n",
"import matplotlib.animation as animation\n",
"import matplotlib.pyplot as plt\n",
"import numpy as np\n",
"import xarray\n",
"from IPython.display import HTML\n",
"\n",
"matplotlib.rcParams[\"animation.embed_limit\"] = 500\n",
"\n",
"\n",
"# --- Helper Functions ---\n",
"\n",
"\n",
"def select_data(\n",
" data: xarray.Dataset,\n",
" variable: str,\n",
" level: Optional[int] = None,\n",
") -> xarray.Dataset:\n",
" \"\"\"Selects a variable from the dataset and optionally a level.\"\"\"\n",
" data = data[variable]\n",
" if \"batch\" in data.dims:\n",
" data = data.isel(batch=0)\n",
" if level is not None and \"level\" in data.coords:\n",
" data = data.sel(level=level)\n",
" return data\n",
"\n",
"\n",
"def scale_data(\n",
" data: xarray.Dataset,\n",
" center: Optional[float] = None,\n",
" robust: bool = False,\n",
") -> tuple[xarray.Dataset, matplotlib.colors.Normalize, str]:\n",
" \"\"\"Scales the data for visualization.\"\"\"\n",
" vmin = np.nanpercentile(data.values, (2 if robust else 0))\n",
" vmax = np.nanpercentile(data.values, (98 if robust else 100))\n",
" if center is not None:\n",
" diff = max(vmax - center, center - vmin)\n",
" vmin = center - diff\n",
" vmax = center + diff\n",
" return (\n",
" data,\n",
" matplotlib.colors.Normalize(vmin, vmax),\n",
" (\"RdBu_r\" if center is not None else \"viridis\"),\n",
" )\n",
"\n",
"\n",
"def create_forecast_animation(\n",
" dataset: xarray.Dataset,\n",
" fig_title: str,\n",
" plot_size: float = 5,\n",
" robust: bool = False,\n",
") -> HTML:\n",
" \"\"\"\n",
" Creates a forecast animation from an xarray Dataset.\n",
" It intelligently handles datasets with or without a 'subtime' dimension.\n",
" \"\"\"\n",
" # --- Data Preparation ---\n",
" # Check if the data still has 'subtime'). If so, stack dimensions.\n",
" # Otherwise, just rename the 'time' dimension for consistency.\n",
" if \"subtime\" in dataset.dims:\n",
" print(\"Detected 'subtime' dimension. Stacking for hourly animation.\")\n",
" # Stack 'time' and 'subtime' into a single animation dimension\n",
" plot_data = dataset.stack(animation_step=(\"time\", \"subtime\")).transpose(\n",
" \"animation_step\", \"lat\", \"lon\"\n",
" )\n",
" else:\n",
" print(\"No 'subtime' dimension found. Using 'time' for 6-hourly animation.\")\n",
" # Use 'time' as the animation dimension\n",
" plot_data = dataset.rename({\"time\": \"animation_step\"})\n",
"\n",
" # Now, the animation dimension is always called 'animation_step'\n",
" max_steps = plot_data.sizes[\"animation_step\"]\n",
" init_time = plot_data.coords[\"init_time\"].values\n",
"\n",
" # Scale the data for color mapping\n",
" scaled_data, norm, cmap = scale_data(plot_data, robust=robust)\n",
"\n",
" # --- Plotting Setup ---\n",
" figure = plt.figure(figsize=(plot_size * 2, plot_size))\n",
" ax = figure.add_subplot(1, 1, 1)\n",
" ax.set_xticks([])\n",
" ax.set_yticks([])\n",
" figure.suptitle(fig_title, fontsize=16)\n",
" figure.tight_layout(rect=[0, 0.03, 1, 0.95]) # Adjust for title\n",
"\n",
" im = ax.imshow(\n",
" scaled_data.isel(animation_step=0), norm=norm, origin=\"lower\", cmap=cmap\n",
" )\n",
"\n",
" plt.colorbar(\n",
" mappable=im,\n",
" ax=ax,\n",
" orientation=\"vertical\",\n",
" pad=0.02,\n",
" aspect=16,\n",
" shrink=0.75,\n",
" cmap=cmap,\n",
" extend=(\"both\" if robust else \"neither\"),\n",
" )\n",
"\n",
" # --- Animation Update Function ---\n",
" def update(frame):\n",
" # Get the coordinates for the current frame\n",
" step_coords = plot_data[\"animation_step\"][frame].coords\n",
"\n",
" # Calculate total offset and valid time based on available coordinates\n",
" if \"subtime\" in step_coords: # Hourly data\n",
" total_offset = step_coords[\"time\"].values + step_coords[\"subtime\"].values\n",
" else: # 6-hourly data\n",
" total_offset = step_coords[\"animation_step\"].values\n",
"\n",
" total_hours = total_offset / np.timedelta64(1, \"h\")\n",
" valid_time = init_time + total_offset\n",
" valid_time_str = np.datetime_as_string(valid_time, unit=\"m\").replace(\"T\", \" \")\n",
"\n",
" new_title = (\n",
" f\"{fig_title}\\n\"\n",
" f\"Valid: {valid_time_str} UTC (Forecast: +{total_hours:.1f}h)\"\n",
" )\n",
" figure.suptitle(new_title, fontsize=16)\n",
" im.set_data(scaled_data.isel(animation_step=frame))\n",
"\n",
" # --- Create and Display Animation ---\n",
" ani = animation.FuncAnimation(\n",
" fig=figure, func=update, frames=max_steps, interval=250\n",
" )\n",
" plt.close(figure.number)\n",
" return HTML(ani.to_html5_video())\n",
"\n",
"\n",
"# --- Main Visualization Logic ---\n",
"\n",
"# 1. Select the correct path based on the user's dropdown choice\n",
"if time_steps_to_visualize == \"6-Hourly\":\n",
" path_to_zarr = path_to_6hr_zarr\n",
"elif time_steps_to_visualize == \"1-Hourly\":\n",
" path_to_zarr = path_to_1hr_zarr\n",
"else:\n",
" raise ValueError(\"Invalid visualization target selected.\")\n",
"\n",
"print(f\"Loading data from: {path_to_zarr}\")\n",
"\n",
"# 2. Load the dataset\n",
"try:\n",
" full_dataset = xarray.open_zarr(path_to_zarr)\n",
"except Exception as e:\n",
" print(f\"Error loading Zarr store: {e}\")\n",
" # This is a common point of failure, so we exit gracefully.\n",
"else:\n",
" # 3. Select the specific data slice for visualization\n",
" data_for_vis = full_dataset.isel(sample=sample_to_visualize)\n",
" variable_data = select_data(data_for_vis, variable_to_visualize, level_to_visualize)\n",
"\n",
" # 4. Generate the title\n",
" title = f\"{variable_to_visualize} (Sample {sample_to_visualize})\"\n",
" if level_to_visualize:\n",
" title += f\" at {level_to_visualize} hPa\"\n",
"\n",
" # 5. Create and display the animation\n",
" display(create_forecast_animation(variable_data, title, plot_size, robust=True))"
]
}
],
"metadata": {
"colab": {
"name": "weathernext_2_dws.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
@@ -1,532 +1,11 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"id": "RirpY96_3zL0",
"metadata": {},
"outputs": [],
"source": [
"# Copyright 2026 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",
"# https://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",
"id": "52d0938f",
"metadata": {},
"source": [
"# WeatherNext 2 Early Access Program\n",
"<table align=\"left\">\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://colab.research.google.com/github/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/weathernext/weathernext_2_early_access_program.ipynb\">\n",
" <img 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%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%weathernext%2Fweathernext_2_early_access_program.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/vertex-ai-samples/blob/main/notebooks/community/weathernext/weathernext_2_early_access_program.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/vertex-ai-samples/blob/main/notebooks/community/weathernext/weathernext_2_early_access_program.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>"
]
},
{
"cell_type": "markdown",
"id": "DGFzy9MTuYL_",
"metadata": {},
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates running [WeatherNext 2 inference on Google Cloud Vertex AI](https://developers.google.com/weathernext/guides/access-vmg). WeatherNext 2 is Google's latest medium-range probabilistic forecasting model, principally an operational version the FGN model ([published June 2025](https://arxiv.org/abs/2506.10772)). More information is available in the [WeatherNext documentation](https://developers.google.com/weathernext).\n",
"\n",
"### Objective\n",
"\n",
"- Configure the model inputs for distributed, multi-host inference on H100 or A100 GPUs.\n",
"- Run WeatherNext 2 model forecasts in parallel.\n",
"- Visualize forecast results.\n",
"\n",
"### Costs\n",
"\n",
"This uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage.\n",
"\n",
"\n",
"## Before you begin\n",
"\n",
"### Request For GPU Quota\n",
"\n",
"**WARNING:** Make sure you have sufficient GPU quota allocated for the inference configuration (i.e. `num_samples`) before running Vertex Jobs. Otherwise, some Vertex jobs may run while others will fail which would produce\n",
"incomplete results.\n",
"\n",
"\n",
"By default, the quota for GPUs is 0. You can request a higher quota by following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota).\n",
"\n",
"You will need to request quota for either **NVIDIA H100 80GB GPUs** or **NVIDIA A100 80GB GPUs** in your selected region. The total number of GPUs you request must be sufficient for your largest planned forecast (i.e., `num_samples`).\n",
"\n",
"You should request for the following quota:\n",
"\n",
"- Service: `Vertex AI API`\n",
"- Name: `Custom model training preemptible Nvidia A100 80GB GPUs per region` OR `Custom model training preemptible Nvidia H100 GPUs per region`"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "yt4GxkKcDD7Y",
"metadata": {},
"outputs": [],
"source": [
"# @title Install python packages\n",
"\n",
"# Note that you may need to restart the kernel after this step.\n",
"# If so, continue to the next cell after restarting.\n",
"\n",
"print(\"Installing python packages.\")\n",
"\n",
"! pip3 install \\\n",
" google-cloud-aiplatform==1.129.0 \\\n",
" xarray[complete]"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "dO2mF4CPfpHW1ZKbIsut8r69",
"metadata": {
"tags": []
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. **[Optional]** [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched.\n",
"\n",
"\n",
"BUCKET_URI = \"gs://my-bucket\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 3. Select a region that has the required GPUs available.\n",
"\n",
"REGION = \"us-central1\" # @param {type:\"string\"}\n",
"\n",
"# Import the necessary packages\n",
"import datetime\n",
"import importlib\n",
"import os\n",
"import uuid\n",
"import glob\n",
"from google.cloud import aiplatform, storage\n",
"\n",
"import json\n",
"import math\n",
"import re\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"now = datetime.datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" raise ValueError(\"GCS Bucket URI is invalid!\")\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(f\"Bucket region {bucket_region} is different from notebook region {REGION}\")\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"# Initialize Vertex AI API.\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Utility functions\n",
"def get_job_name_with_datetime(prefix: str) -> str:\n",
" return prefix + datetime.datetime.now().strftime(\"_%Y%m%d_%H%M%S\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "22BW43yjps8D",
"metadata": {},
"outputs": [],
"source": [
"# @title Configure Model Parameters\n",
"# @markdown Configure the hardware and input parameters for the WeatherNext 2 forecast.\n",
"\n",
"# @markdown ### Hardware Configuration for Distributed Inference\n",
"# @markdown - **`machine_type`**: Select a valid machine type. `a3-highgpu` series use NVIDIA H100 80GB GPUs. `a2-ultragpu` series use NVIDIA A100 80GB GPUs.\n",
"# @markdown - **`num_samples`**: The total number of ensemble members to generate.\n",
"# @markdown The number of machine replicas will be calculated automatically (`num_samples` / GPUs per machine). **Therefore, `num_samples` must be a multiple of the number of GPUs in your selected `machine_type`.**\n",
"# @markdown - **`scheduling_strategy`**: The [strategy](https://cloud.google.com/vertex-ai/docs/reference/rest/v1beta1/CustomJobSpec#Strategy) used to acquire machines for the job. Defaults to [Dynamic Workload Scheduler](https://docs.cloud.google.com/vertex-ai/docs/training/schedule-jobs-dws) (FLEX_START).\n",
"machine_type = \"a3-highgpu-1g\" #@param [\"a3-highgpu-1g\", \"a3-highgpu-2g\", \"a3-highgpu-4g\", \"a3-highgpu-8g\", \"a2-ultragpu-1g\", \"a2-ultragpu-2g\", \"a2-ultragpu-4g\", \"a2-ultragpu-8g\"]\n",
"num_samples = 8 #@param {type:\"integer\"}\n",
"scheduling_strategy = \"FLEX_START\" #@param [\"FLEX_START\", \"SPOT\", \"STANDARD\"]\n",
"\n",
"# @markdown ### Forecast Configuration\n",
"# @markdown - **`forecast_init_time`**: The starting time for the forecast in ISO 8601 format (e.g., `2025-09-21T00:00:00Z`). Models are available for dates from 2024 onwards.\n",
"# @markdown - **`horizon_hrs`**: The desired length of the forecast in hours (e.g., 240 for a 10-day forecast).\n",
"# @markdown - **`model_seed`**: Choose a specific model seed (1-4) or select \"all\" to run inference with all four seeds in parallel for improved accuracy.\n",
"# @markdown - **`enable_hourly_prediction`**: If checked, the model will generate 1-hour predictions.\n",
"forecast_init_time = \"2025-11-20T00:00:00Z\" #@param {type:\"string\"}\n",
"horizon_hrs = 72 #@param {type:\"integer\"}\n",
"model_seed = \"all\" # @param [\"1\", \"2\", \"3\", \"4\", \"all\"]\n",
"enable_hourly_prediction = True # @param {type:\"boolean\"}\n",
"\n",
"# --- Parameter Validation and Configuration ---\n",
"\n",
"# Derive accelerator type and count from the chosen machine type\n",
"if machine_type.startswith('a3-highgpu'):\n",
" accelerator_type = 'NVIDIA_H100_80GB'\n",
"elif machine_type.startswith('a2-ultragpu'):\n",
" accelerator_type = 'NVIDIA_A100_80GB'\n",
"else:\n",
" raise ValueError(f\"Invalid machine type selected: {machine_type}.\")\n",
"\n",
"try:\n",
" # Extract the number of GPUs from the machine type string, e.g., 'a3-highgpu-4g' -> 4\n",
" accelerators_per_machine = int(re.search(r'-(\\d+)g$', machine_type).group(1))\n",
"except (AttributeError, ValueError):\n",
" raise ValueError(f\"Could not determine accelerator count from machine type: {machine_type}\")\n",
"\n",
"seeds_to_run = [1, 2, 3, 4] if model_seed == \"all\" else [int(model_seed)]\n",
"num_seeds_to_run = len(seeds_to_run)\n",
"\n",
"num_samples_per_seed = num_samples\n",
"if len(seeds_to_run) > 1:\n",
" if num_samples % num_seeds_to_run != 0:\n",
" raise ValueError(f\"`num_samples` ({num_samples}) is not divisible by the number of seeds to run ({num_seeds_to_run}.\")\n",
" num_samples_per_seed = num_samples // num_seeds_to_run\n",
"\n",
"# Validate that num_samples is a multiple of accelerators_per_machine\n",
"if num_samples_per_seed % accelerators_per_machine != 0:\n",
" raise ValueError(f\"`num_samples_per_seed` ({num_samples_per_seed}) must be a multiple of the GPUs per machine ({accelerators_per_machine} for {machine_type}).\")\n",
"\n",
"# Calculate the number of replicas per seed\n",
"replica_count_per_seed = num_samples_per_seed // accelerators_per_machine\n",
"\n",
"# Ensure that there are enough samples\n",
"if num_samples_per_seed % accelerators_per_machine != 0:\n",
" raise ValueError(f\"`num_samples_per_seed` ({num_samples_per_seed}) must be a multiple of the GPUs per machine ({accelerators_per_machine} for {machine_type}).\")\n",
"\n",
"# Calculate total GPUs needed for all jobs\n",
"total_gpus_needed = num_samples * (4 if model_seed == \"all\" else 1)\n",
"\n",
"# Set Docker URI\n",
"WEATHERNEXT2_DOCKER_URI = 'us-docker.pkg.dev/vertex-ai-restricted/vertex-vision-model-garden-dockers/weather-next-2-inference.gpu.0-1:latest'\n",
"\n",
"print(\"--- Job Configuration Summary ---\")\n",
"print(f\"Total Samples: {num_samples}\")\n",
"print(f\"Machine Type: {machine_type}\")\n",
"print(f\"Accelerator Type: {accelerator_type}\")\n",
"print(f\"GPUs per Machine: {accelerators_per_machine}\")\n",
"print(f\"Total number seeds to run: {num_seeds_to_run}\")\n",
"print(f\"Total number samples per seed: {num_samples_per_seed}\")\n",
"print(f\"Calculated Machine Replicas Per Seed: {replica_count_per_seed}\")\n",
"print(f\"Total GPUs per Job: {num_samples}\")\n",
"print(f\"Total GPUs across all Jobs (ensure sufficient quota): {total_gpus_needed}\")\n",
"print(f\"Docker Image: {WEATHERNEXT2_DOCKER_URI}\")\n",
"print(\"---------------------------------\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "VSuGx2jmpwP0",
"metadata": {},
"outputs": [],
"source": [
"# @title Run Forecasts\n",
"# @markdown This section creates and runs one or more Vertex AI Custom Training Jobs to generate the forecasts.\n",
"# @markdown **This operation is asynchronous.** The jobs will be submitted and this cell will complete quickly.\n",
"# @markdown You must monitor the job progress in the Google Cloud Console (https://console.cloud.google.com/vertex-ai/training/custom-jobs).\n",
"\n",
"from google.cloud.aiplatform.compat.types import \\\n",
" custom_job as gca_custom_job_compat\n",
"\n",
"print(f\"Submitting {len(seeds_to_run)} job(s) to run in parallel.\")\n",
"\n",
"launched_jobs = []\n",
"output_dirs = {}\n",
"\n",
"if scheduling_strategy == \"FLEX_START\":\n",
" SCHEDULLING_STRATEGY = gca_custom_job_compat.Scheduling.Strategy.FLEX_START\n",
"elif scheduling_strategy == \"SPOT\":\n",
" SCHEDULLING_STRATEGY = gca_custom_job_compat.Scheduling.Strategy.SPOT\n",
"else:\n",
" SCHEDULLING_STRATEGY = gca_custom_job_compat.Scheduling.Strategy.STANDARD\n",
"\n",
"for seed in seeds_to_run:\n",
" output_dir = os.path.join(BUCKET_URI, \"weathernext2_outputs\")\n",
" output_dirs[seed] = output_dir\n",
"\n",
" docker_args_list = [\n",
" f\"--pred_root_dir={output_dir}\",\n",
" f\"--num_samples={num_samples_per_seed}\",\n",
" f\"--horizon_hrs={horizon_hrs}\",\n",
" f\"--forecast_init_time={forecast_init_time}\",\n",
" f\"--model_seed={seed}\",\n",
" f\"--enable_hourly_prediction={enable_hourly_prediction}\",\n",
" ]\n",
"\n",
" JOB_NAME = get_job_name_with_datetime(prefix=f\"wn2-forecast-s{seed}-n{num_samples_per_seed}\")\n",
" print(f\"\\n--- Submitting Job for Seed {seed} ---\")\n",
" print(f\"JOB_NAME: {JOB_NAME}\")\n",
"\n",
" job = aiplatform.CustomContainerTrainingJob(\n",
" display_name=JOB_NAME,\n",
" container_uri=WEATHERNEXT2_DOCKER_URI,\n",
" )\n",
"\n",
" job.run(\n",
" args=docker_args_list,\n",
" replica_count=replica_count_per_seed,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerators_per_machine,\n",
" scheduling_strategy=SCHEDULLING_STRATEGY,\n",
" # Change this to True if you need to debug why the job hasn't started\n",
" sync=False\n",
" )\n",
" launched_jobs.append(job)\n",
" print(f\"--> Job submitted successfully. Monitor it in the Google Cloud Console at https://console.cloud.google.com/vertex-ai/training/custom-jobs\")\n",
"\n",
"print(\"\\nAll forecast jobs have been submitted.\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "x6ZVqWRopWcI",
"metadata": {
"cellView": "form"
},
"outputs": [],
"source": [
"# @title Visualize Forecasts (Unified)\n",
"# @markdown Select which forecast output you want to visualize. This single component\n",
"# @markdown can handle both the standard 6-hourly predictions and the datasets\n",
"# @markdown with 1-hour model (which have a 'subtime' dimension).\n",
"# @markdown If you run into `Error loading Zarr store: unrecognized engine 'zarr'...` try restarting the runtime session and reruning this cell.\n",
"\n",
"# @markdown ---\n",
"# @markdown ### Visualization Settings\n",
"# @markdown - **`model_seed_to_visualize`**: Choose a specific model seed (1-4) to visualize. This should be one of the model seeds selected in the **Forecast Configuration** above.\n",
"# @markdown - **`time_steps_to_visualize`**: Choose to visualize 1-hourly or 6-hourly forecasts. If 1-hourly is selected, ensure `enable_hourly_prediction` was selected in the **Forecast Configuration** above.\n",
"# @markdown - **`variable_to_visualize`**: Choose the weather variable to visualize. See the [WeatherNext documentation](https://developers.google.com/weathernext/guides/model-specs-vmg) for variable names and descriptions.\n",
"# @markdown - **`sample_to_visualize`**: Choose the sample (ensemble member) to visualize.\n",
"# @markdown - **`plot_size`**: Choose the size of the plot to generate.\n",
"model_seed_to_visualize = \"4\" # @param [\"1\", \"2\", \"3\", \"4\"]\n",
"time_steps_to_visualize = \"6-Hourly\" # @param [\"6-Hourly\", \"1-Hourly\"]\n",
"variable_to_visualize = \"2m_temperature\" # @param {type:\"string\"}\n",
"sample_to_visualize = 0 # @param {type:\"integer\"}\n",
"plot_size = 8 #@param {type:\"number\"}\n",
"level_to_visualize = None\n",
"# @markdown ---\n",
"\n",
"\n",
"# --- Helper Functions ---\n",
"\n",
"def init_time_to_folder_path(init_time: str) -> str:\n",
" \"\"\"\n",
" Convert init time to expected GCS folder path.\n",
" \"\"\"\n",
" init_date, init_time = init_time.split(\"T\")\n",
" return f\"{init_date.replace(\"-\",\"\")}_{init_time[0:2]}hr\"\n",
"\n",
"\n",
"# override these if you'd like to visualize a different set of forecasts\n",
"visualize_bucket = BUCKET_URI\n",
"visualize_init_date = forecast_init_time\n",
"\n",
"# set paths based on chosen model seed, bucket, and init date\n",
"path_to_6hr_zarr = f\"{BUCKET_URI}/weathernext2_outputs/weathernext_2_seed_{model_seed_to_visualize}/{init_time_to_folder_path(visualize_init_date)}_01_preds/predictions.zarr/\"\n",
"path_to_1hr_zarr = f\"{BUCKET_URI}/weathernext2_outputs/weathernext_2_seed_{model_seed_to_visualize}_hourly/{init_time_to_folder_path(visualize_init_date)}_01_preds/predictions.zarr/\"\n",
"\n",
"\n",
"import matplotlib\n",
"import xarray\n",
"from typing import Optional, Tuple\n",
"import matplotlib.pyplot as plt\n",
"import matplotlib.animation as animation\n",
"import math\n",
"from IPython.display import HTML\n",
"import numpy as np\n",
"import datetime\n",
"\n",
"matplotlib.rcParams['animation.embed_limit'] = 500\n",
"\n",
"\n",
"# --- Helper Functions ---\n",
"\n",
"def select_data(\n",
" data: xarray.Dataset,\n",
" variable: str,\n",
" level: Optional[int] = None,\n",
" ) -> xarray.Dataset:\n",
" \"\"\"Selects a variable from the dataset and optionally a level.\"\"\"\n",
" data = data[variable]\n",
" if \"batch\" in data.dims:\n",
" data = data.isel(batch=0)\n",
" if level is not None and \"level\" in data.coords:\n",
" data = data.sel(level=level)\n",
" return data\n",
"\n",
"def scale_data(\n",
" data: xarray.Dataset,\n",
" center: Optional[float] = None,\n",
" robust: bool = False,\n",
" ) -> tuple[xarray.Dataset, matplotlib.colors.Normalize, str]:\n",
" \"\"\"Scales the data for visualization.\"\"\"\n",
" vmin = np.nanpercentile(data.values, (2 if robust else 0))\n",
" vmax = np.nanpercentile(data.values, (98 if robust else 100))\n",
" if center is not None:\n",
" diff = max(vmax - center, center - vmin)\n",
" vmin = center - diff\n",
" vmax = center + diff\n",
" return (data, matplotlib.colors.Normalize(vmin, vmax),\n",
" (\"RdBu_r\" if center is not None else \"viridis\"))\n",
"\n",
"def create_forecast_animation(\n",
" dataset: xarray.Dataset,\n",
" fig_title: str,\n",
" plot_size: float = 5,\n",
" robust: bool = False,\n",
" ) -> HTML:\n",
" \"\"\"\n",
" Creates a forecast animation from an xarray Dataset.\n",
" It intelligently handles datasets with or without a 'subtime' dimension.\n",
" \"\"\"\n",
" # --- Data Preparation ---\n",
" # Check if the data still has 'subtime'). If so, stack dimensions.\n",
" # Otherwise, just rename the 'time' dimension for consistency.\n",
" if 'subtime' in dataset.dims:\n",
" print(\"Detected 'subtime' dimension. Stacking for hourly animation.\")\n",
" # Stack 'time' and 'subtime' into a single animation dimension\n",
" plot_data = dataset.stack(\n",
" animation_step=(\"time\", \"subtime\")\n",
" ).transpose(\"animation_step\", \"lat\", \"lon\")\n",
" else:\n",
" print(\"No 'subtime' dimension found. Using 'time' for 6-hourly animation.\")\n",
" # Use 'time' as the animation dimension\n",
" plot_data = dataset.rename({'time': 'animation_step'})\n",
"\n",
" # Now, the animation dimension is always called 'animation_step'\n",
" max_steps = plot_data.sizes[\"animation_step\"]\n",
" init_time = plot_data.coords['init_time'].values\n",
"\n",
" # Scale the data for color mapping\n",
" scaled_data, norm, cmap = scale_data(plot_data, robust=robust)\n",
"\n",
" # --- Plotting Setup ---\n",
" figure = plt.figure(figsize=(plot_size * 2, plot_size))\n",
" ax = figure.add_subplot(1, 1, 1)\n",
" ax.set_xticks([])\n",
" ax.set_yticks([])\n",
" figure.suptitle(fig_title, fontsize=16)\n",
" figure.tight_layout(rect=[0, 0.03, 1, 0.95]) # Adjust for title\n",
"\n",
" im = ax.imshow(\n",
" scaled_data.isel(animation_step=0), norm=norm, origin=\"lower\", cmap=cmap)\n",
"\n",
" plt.colorbar(\n",
" mappable=im, ax=ax, orientation=\"vertical\", pad=0.02,\n",
" aspect=16, shrink=0.75, cmap=cmap,\n",
" extend=(\"both\" if robust else \"neither\"))\n",
"\n",
" # --- Animation Update Function ---\n",
" def update(frame):\n",
" # Get the coordinates for the current frame\n",
" step_coords = plot_data['animation_step'][frame].coords\n",
"\n",
" # Calculate total offset and valid time based on available coordinates\n",
" if 'subtime' in step_coords: # Hourly data\n",
" total_offset = step_coords['time'].values + step_coords['subtime'].values\n",
" else: # 6-hourly data\n",
" total_offset = step_coords['animation_step'].values\n",
"\n",
" total_hours = total_offset / np.timedelta64(1, 'h')\n",
" valid_time = init_time + total_offset\n",
" valid_time_str = np.datetime_as_string(valid_time, unit='m').replace('T', ' ')\n",
"\n",
" new_title = (\n",
" f\"{fig_title}\\n\"\n",
" f\"Valid: {valid_time_str} UTC (Forecast: +{total_hours:.1f}h)\"\n",
" )\n",
" figure.suptitle(new_title, fontsize=16)\n",
" im.set_data(scaled_data.isel(animation_step=frame))\n",
"\n",
" # --- Create and Display Animation ---\n",
" ani = animation.FuncAnimation(\n",
" fig=figure, func=update, frames=max_steps, interval=250)\n",
" plt.close(figure.number)\n",
" return HTML(ani.to_html5_video())\n",
"\n",
"\n",
"# --- Main Visualization Logic ---\n",
"\n",
"# 1. Select the correct path based on the user's dropdown choice\n",
"if time_steps_to_visualize == \"6-Hourly\":\n",
" path_to_zarr = path_to_6hr_zarr\n",
"elif time_steps_to_visualize == \"1-Hourly\":\n",
" path_to_zarr = path_to_1hr_zarr\n",
"else:\n",
" raise ValueError(\"Invalid visualization target selected.\")\n",
"\n",
"print(f\"Loading data from: {path_to_zarr}\")\n",
"\n",
"# 2. Load the dataset\n",
"try:\n",
" full_dataset = xarray.open_zarr(path_to_zarr)\n",
"except Exception as e:\n",
" print(f\"Error loading Zarr store: {e}\")\n",
" # This is a common point of failure, so we exit gracefully.\n",
"else:\n",
" # 3. Select the specific data slice for visualization\n",
" data_for_vis = full_dataset.isel(sample=sample_to_visualize)\n",
" variable_data = select_data(data_for_vis, variable_to_visualize, level_to_visualize)\n",
"\n",
" # 4. Generate the title\n",
" title = f\"{variable_to_visualize} (Sample {sample_to_visualize})\"\n",
" if level_to_visualize:\n",
" title += f\" at {level_to_visualize} hPa\"\n",
"\n",
" # 5. Create and display the animation\n",
" display(create_forecast_animation(variable_data, title, plot_size, robust=True))"
"Latest notebook is [here.](weathernext_2_dws.ipynb)"
]
}
],
@@ -0,0 +1,26 @@
{
"cells": [
{
"cell_type": "markdown",
"id": "3c096b65",
"metadata": {
"id": "DGFzy9MTuYL_"
},
"source": [
"Latest notebook is [here.](weathernext_2_ic_pc.ipynb)\n"
]
}
],
"metadata": {
"colab": {
"name": "weathernext_2_ic_early_access_program.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
@@ -0,0 +1,751 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"id": "1a737002",
"metadata": {
"id": "c78774449a85"
},
"outputs": [],
"source": [
"# Copyright 2026 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",
"# https://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",
"id": "bf59250d",
"metadata": {
"id": "1ce8dc9a8d06"
},
"source": [
"# WeatherNext 2\n",
"<table align=\"left\">\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://colab.research.google.com/github/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/weathernext/weathernext_2_ic_pc.ipynb\">\n",
" <img 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%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fweathernext%2Fweathernext_2_ic_pc.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/vertex-ai-samples/main/notebooks/community/weathernext/weathernext_2_ic_pc.ipynb\">\n",
" <img src=\"https://www.gstatic.com/images/branding/gcpiconscolors/vertexai/v1/32px.svg\" alt=\"Enterprise Gemini Agent Platform logo\"><br> Open in Enterprise Gemini Agent Platform Workbench\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/weathernext/weathernext_2_ic_pc.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>"
]
},
{
"cell_type": "markdown",
"id": "3c096b65",
"metadata": {
"id": "9483477521e4"
},
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates running [WeatherNext 2 inference on Google Cloud Enterprise Gemini Agent Platform](https://developers.google.com/weathernext/guides/access-vmg) using **customer-provided initial conditions** (custom inputs). WeatherNext 2 is Google's latest medium-range probabilistic forecasting model, principally an operational version the FGN model ([published June 2025](https://arxiv.org/abs/2506.10772)). More information is available in the [WeatherNext documentation](https://developers.google.com/weathernext).\n",
"\n",
"WeatherNext 2 supports customer-provided initial conditions for inference. Instead of using the default ECMWF HRES real-time data, customers can supply their own input Zarr files (e.g., from GFS or their own analysis systems) to generate forecasts with WeatherNext models.\n",
"\n",
"> **Important**: The model is **not fine-tuned** on custom input data. Forecast performance when using custom inputs is **not guaranteed** to match the quality achieved with the default ECMWF HRES inputs. Customers should perform their own evaluation of output quality.\n",
"\n",
"### Objective\n",
"\n",
"- Configure the model inputs for distributed, multi-host inference on H100 or A100 GPUs.\n",
"- Provide custom initial conditions (Zarr files) for model inference.\n",
"- Run WeatherNext 2 model forecasts in parallel.\n",
"- Visualize forecast results.\n",
"\n",
"### Costs\n",
"\n",
"This uses billable components of Google Cloud:\n",
"\n",
"* [Gemini Enterprise Agent Platform]( https://docs.cloud.google.com/gemini-enterprise-agent-platform)\n",
"* [Cloud Storage](https://cloud.google.com/storage/docs)\n",
"* [Gemini Enterprise Agent Platform Persistent Resource](https://docs.cloud.google.com/gemini-enterprise-agent-platform/machine-learning/training/persistent-resource-create)\n",
"\n",
"Learn about [Gemini Enterprise Agent Platform pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), [Gemini Enterprise Platform Persistent Resource](https://cloud.google.com/products/gemini-enterprise-agent-platform/pricing#custom-trained-models) and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage.\n",
"\n",
"\n",
"## Before you begin\n",
"\n",
"### Request For GPU Quota\n",
"\n",
"**WARNING:** Make sure you have sufficient GPU quota allocated for the inference configuration (i.e. `num_samples`) before running Gemini Enterprise Agent Platform Jobs or provisioning a Persistent Resource. Otherwise, some Gemini Enterprise Agent Platform jobs may run while others will fail which would produce\n",
"incomplete results.\n",
"\n",
"\n",
"By default, the quota for GPUs is 0. You can request a higher quota by following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota).\n",
"\n",
"You will need to request quota for either **NVIDIA H100 80GB GPUs** or **NVIDIA A100 80GB GPUs** in your selected region. The total number of GPUs you request must be sufficient for your largest planned forecast (i.e., `num_samples`) or your provisioned Persistent Resource capacity.\n",
"\n",
"Depending on whether you use a **Persistent Resource** or run standard custom jobs, you should request the following quota under Service `Enterprise Gemini Agent Platform API`:\n",
"\n",
"#### Option 1: Persistent Resource\n",
"- Name: `Persistent resource Nvidia A100 80GB GPUs per region` OR `Persistent resource Nvidia H100 GPUs per region`\n",
"\n",
"#### Option 2: Standard Custom Model Training (Preemptible)\n",
"- Name: `Custom model training preemptible Nvidia A100 80GB GPUs per region` OR `Custom model training preemptible Nvidia H100 GPUs per region`\n",
"\n",
"### Custom Inputs Guide\n",
"\n",
"Please refer to the guide [here](CUSTOM_INPUTS_GUIDE.md)."
]
},
{
"cell_type": "markdown",
"id": "d656ed2b",
"metadata": {
"id": "9349a5bece68"
},
"source": [
"## Install packages - Restart the kernel after the installation"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "cbfab540",
"metadata": {
"id": "83ed640c5f9e"
},
"outputs": [],
"source": [
"# @title Install python packages\n",
"\n",
"# Note that you may need to restart the kernel after this step.\n",
"# If so, continue to the next cell after restarting.\n",
"\n",
"print(\"Installing python packages.\")\n",
"\n",
"! pip3 install \\\n",
" google-cloud-aiplatform==1.129.0 \\\n",
" xarray[complete]"
]
},
{
"cell_type": "markdown",
"id": "c123b19f",
"metadata": {
"id": "e0b91a0300f0"
},
"source": [
"## Authenticate to Google Cloud Platform"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "3b7e5c45",
"metadata": {
"id": "63c3ac57b2e3"
},
"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",
"id": "7abff806",
"metadata": {
"id": "fd0ba250117b"
},
"source": [
"## Set Google Cloud Project Pertinent Variables"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "b0805e58",
"metadata": {
"id": "3390c789e4cc"
},
"outputs": [],
"source": [
"# Replace my_gcp_project with your gcp project\n",
"PROJECT_ID = \"<my_gcp_project>\" # @param {type:\"string\"}\n",
"\n",
"# Alternatively collect the default cloud project id from the OS env variable\n",
"# PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"# @markdown 2. [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched.\n",
"\n",
"# Replace my_wn_bucket with your bucket\n",
"BUCKET_URI = \"gs://<my_wn_bucket>\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 3. Select a region that has the required GPUs available.\n",
"\n",
"# Select a ***US*** region region\n",
"REGION = \"us-central1\" # @param {type:\"string\"}\n",
"\n",
"\n",
"# @markdown 4. Provide the Persistence Resource ID\n",
"PERSISTENT_RESOURCE_ID = \"<my_persistent_resource_id>\" # @param {type:\"string\"}\n",
"\n",
"# Set Docker URI\n",
"# @markdown 5. Set the Image URI\n",
"WEATHERNEXT2_DOCKER_URI = \"us-central1-docker.pkg.dev/weathernext-1/wn25-private-preview-launch/weather-next-ic-2-inference.gpu.0-1:latest\""
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "9af9574a",
"metadata": {
"id": "5c37e09ee5f2"
},
"outputs": [],
"source": [
"!gcloud config set project $PROJECT_ID"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "d914b484",
"metadata": {
"id": "66cf7f2c6cfe"
},
"outputs": [],
"source": [
"# Import the necessary packages\n",
"import datetime\n",
"import os\n",
"import re\n",
"\n",
"from google.cloud import aiplatform\n",
"\n",
"# Enable the Enterprise Gemini Agent Platform API and Compute Engine API, if not already.\n",
"print(\"Enabling Enterprise Gemini Agent Platform API and Compute Engine API.\")\n",
"# ! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"now = datetime.datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" raise ValueError(\"GCS Bucket URI is invalid!\")\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" shell_output = !gcloud storage buckets describe {BUCKET_NAME} --format=\"value(location)\" | tr '[:upper:]' '[:lower:]'\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" f\"Bucket region {bucket_region} is different from notebook region {REGION}\"\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"# Initialize Enterprise Gemini Agent Platform API.\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Utility functions\n",
"\n",
"\n",
"def get_job_name_with_datetime(prefix: str) -> str:\n",
" return prefix + datetime.datetime.now().strftime(\"_%Y%m%d_%H%M%S\")"
]
},
{
"cell_type": "markdown",
"id": "95c31932",
"metadata": {
"id": "cb667a3392f7"
},
"source": [
"## Configure Model Parameters"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "f4854c78",
"metadata": {
"id": "65e9156c0070"
},
"outputs": [],
"source": [
"# @markdown ### Hardware Configuration for Distributed Inference\n",
"# @markdown - **`machine_type`**: Select a valid machine type. `a3-highgpu` series use NVIDIA H100 80GB GPUs. `a2-ultragpu` series use NVIDIA A100 80GB GPUs.\n",
"# @markdown - **`num_samples`**: The total number of ensemble members to generate.\n",
"# @markdown The number of machine replicas will be calculated automatically (`num_samples` / GPUs per machine). **Therefore, `num_samples` must be a multiple of the number of GPUs in your selected `machine_type`.**\n",
"machine_type = \"a2-ultragpu-1g\" # @param [\"a3-highgpu-1g\", \"a3-highgpu-2g\", \"a3-highgpu-4g\", \"a3-highgpu-8g\", \"a2-ultragpu-1g\", \"a2-ultragpu-2g\", \"a2-ultragpu-4g\", \"a2-ultragpu-8g\"]\n",
"num_samples = 4 # @param {type:\"integer\"}\n",
"\n",
"# @markdown ### Forecast Configuration\n",
"# @markdown - **`forecast_init_time`**: The starting time for the forecast in ISO 8601 format (e.g., `2025-09-21T00:00:00Z`). Models are available for dates from 2024 onwards.\n",
"# @markdown - **`horizon_hrs`**: The desired length of the forecast in hours (e.g., 240 for a 10-day forecast).\n",
"# @markdown - **`model_seed`**: Choose a specific model seed (1-4) or select \"all\" to run inference with all four seeds in parallel for improved accuracy.\n",
"# @markdown - **`enable_hourly_prediction`**: If checked, the model will generate 1-hour predictions.\n",
"forecast_init_time = \"2026-05-16T12:00:00Z\"\n",
"\n",
"horizon_hrs = 72 # @param {type:\"integer\"}\n",
"model_seed = \"all\" # @param [\"1\", \"2\", \"3\", \"4\", \"all\"]\n",
"enable_hourly_prediction = True # @param {type:\"boolean\"}\n",
"\n",
"# --- Parameter Validation and Configuration ---\n",
"\n",
"# Derive accelerator type and count from the chosen machine type\n",
"if machine_type.startswith(\"a3-highgpu\"):\n",
" accelerator_type = \"NVIDIA_H100_80GB\"\n",
"elif machine_type.startswith(\"a2-ultragpu\"):\n",
" accelerator_type = \"NVIDIA_A100_80GB\"\n",
"else:\n",
" raise ValueError(f\"Invalid machine type selected: {machine_type}.\")\n",
"\n",
"try:\n",
" # Extract the number of GPUs from the machine type string, e.g., 'a3-highgpu-4g' -> 4\n",
" accelerators_per_machine = int(re.search(r\"-(\\d+)g$\", machine_type).group(1))\n",
"except (AttributeError, ValueError):\n",
" raise ValueError(\n",
" f\"Could not determine accelerator count from machine type: {machine_type}\"\n",
" )\n",
"\n",
"seeds_to_run = [1, 2, 3, 4] if model_seed == \"all\" else [int(model_seed)]\n",
"num_seeds_to_run = len(seeds_to_run)\n",
"\n",
"num_samples_per_seed = num_samples\n",
"if len(seeds_to_run) > 1:\n",
" if num_samples % num_seeds_to_run != 0:\n",
" raise ValueError(\n",
" f\"`num_samples` ({num_samples}) is not divisible by the number of seeds to run ({num_seeds_to_run}.\"\n",
" )\n",
" num_samples_per_seed = num_samples // num_seeds_to_run\n",
"\n",
"# Validate that num_samples is a multiple of accelerators_per_machine\n",
"if num_samples_per_seed % accelerators_per_machine != 0:\n",
" raise ValueError(\n",
" f\"`num_samples_per_seed` ({num_samples_per_seed}) must be a multiple of the GPUs per machine ({accelerators_per_machine} for {machine_type}).\"\n",
" )\n",
"\n",
"# Calculate the number of replicas per seed\n",
"replica_count_per_seed = num_samples_per_seed // accelerators_per_machine\n",
"\n",
"# Ensure that there are enough samples\n",
"if num_samples_per_seed % accelerators_per_machine != 0:\n",
" raise ValueError(\n",
" f\"`num_samples_per_seed` ({num_samples_per_seed}) must be a multiple of the GPUs per machine ({accelerators_per_machine} for {machine_type}).\"\n",
" )\n",
"\n",
"# Calculate total GPUs needed for all jobs\n",
"total_gpus_needed = num_samples * (4 if model_seed == \"all\" else 1)\n",
"\n",
"print(\"--- Job Configuration Summary ---\")\n",
"print(f\"Total Samples: {num_samples}\")\n",
"print(f\"Machine Type: {machine_type}\")\n",
"print(f\"Accelerator Type: {accelerator_type}\")\n",
"print(f\"GPUs per Machine: {accelerators_per_machine}\")\n",
"print(f\"Total number seeds to run: {num_seeds_to_run}\")\n",
"print(f\"Total number samples per seed: {num_samples_per_seed}\")\n",
"print(f\"Calculated Machine Replicas Per Seed: {replica_count_per_seed}\")\n",
"print(f\"Total GPUs per Job: {num_samples}\")\n",
"print(f\"Total GPUs across all Jobs (ensure sufficient quota): {total_gpus_needed}\")\n",
"print(f\"Docker Image: {WEATHERNEXT2_DOCKER_URI}\")\n",
"print(\"---------------------------------\")"
]
},
{
"cell_type": "markdown",
"id": "0ba7d73c",
"metadata": {
"id": "fe8b458ce060"
},
"source": [
"## Run forecast jobs"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "42329e44",
"metadata": {
"id": "6d83132ca9fa"
},
"outputs": [],
"source": [
"# @markdown This section creates and runs one or more Enterprise Gemini Agent Platform Custom Training Jobs to generate the forecasts.\n",
"# @markdown **This operation is asynchronous.** The jobs will be submitted and this cell will complete quickly.\n",
"# @markdown You must monitor the job progress in the Google Cloud Console (https://console.cloud.google.com/vertex-ai/training/custom-jobs).\n",
"\n",
"import time\n",
"\n",
"print(\n",
" f\"Submitting {len(seeds_to_run)} job(s) to target persistent resource: {PERSISTENT_RESOURCE_ID}\"\n",
")\n",
"\n",
"launched_jobs = []\n",
"output_dirs = {}\n",
"\n",
"for seed in seeds_to_run:\n",
" output_dir = os.path.join(BUCKET_URI, \"weathernext2_outputs\")\n",
" output_dirs[seed] = output_dir\n",
"\n",
" docker_args_list = [\n",
" f\"--pred_root_dir={output_dir}\",\n",
" f\"--num_samples={num_samples_per_seed}\",\n",
" f\"--horizon_hrs={horizon_hrs}\",\n",
" f\"--forecast_init_time={forecast_init_time}\",\n",
" f\"--model_seed={seed}\",\n",
" f\"--enable_hourly_prediction={enable_hourly_prediction}\",\n",
" # Comment the line below if you don't want to use custom inputs.\n",
" f\"--input_data_gcs_dir={BUCKET_URI}/custom_inputs/\",\n",
" ]\n",
"\n",
" # Define the worker pool spec for the custom job matching the warm pool specs\n",
" worker_pool_specs = [\n",
" {\n",
" \"machine_spec\": {\n",
" \"machine_type\": machine_type,\n",
" \"accelerator_type\": accelerator_type,\n",
" \"accelerator_count\": accelerators_per_machine,\n",
" },\n",
" \"disk_spec\": {\n",
" \"boot_disk_type\": \"pd-standard\",\n",
" \"boot_disk_size_gb\": 100,\n",
" },\n",
" \"replica_count\": 1,\n",
" \"container_spec\": {\n",
" \"image_uri\": WEATHERNEXT2_DOCKER_URI,\n",
" \"args\": docker_args_list,\n",
" },\n",
" }\n",
" ]\n",
"\n",
" JOB_NAME = get_job_name_with_datetime(prefix=f\"wn2-persistent-run-s{seed}\")\n",
" print(f\"\\n--- Submitting Job for Seed {seed} ---\")\n",
" print(f\"JOB_NAME: {JOB_NAME}\")\n",
"\n",
" custom_job = aiplatform.CustomJob(\n",
" display_name=JOB_NAME,\n",
" worker_pool_specs=worker_pool_specs,\n",
" project=PROJECT_ID,\n",
" location=REGION,\n",
" staging_bucket=f\"{BUCKET_URI}/custom_job_staging\",\n",
" )\n",
"\n",
" # Run the job asynchronously targeting the warm persistent resource pool!\n",
" # Since the 4-GPU persistent resource was created without explicit custom SA runtime permissions,\n",
" # we run using the default Enterprise Gemini Agent Platform Service Agent (which already has GCS and registry reader permissions).\n",
" custom_job.run(\n",
" persistent_resource_id=PERSISTENT_RESOURCE_ID,\n",
" disable_retries=True,\n",
" sync=False,\n",
" )\n",
"\n",
" # Wait for GCA resource generation details\n",
" print(\"Waiting for CustomJob details...\")\n",
" for _ in range(15):\n",
" try:\n",
" if custom_job.resource_name:\n",
" break\n",
" except Exception:\n",
" pass\n",
" time.sleep(1)\n",
"\n",
" job_id = custom_job.name.split(\"/\")[-1]\n",
" print(f\"--> Job submitted successfully. ID: {job_id}\")\n",
" print(\n",
" f\"Console Link: https://console.cloud.google.com/vertex-ai/locations/{REGION}/training/custom-jobs/{job_id}?project={PROJECT_ID}\"\n",
" )\n",
" launched_jobs.append(custom_job)\n",
"\n",
"print(\"\\nAll forecast jobs have been submitted. Starting real-time status monitor...\")\n",
"\n",
"# Monitoring loop\n",
"completed_jobs = set()\n",
"while len(completed_jobs) < len(launched_jobs):\n",
" for job in launched_jobs:\n",
" if job.name in completed_jobs:\n",
" continue\n",
"\n",
" job._sync_gca_resource()\n",
" state = job.state.value if hasattr(job.state, \"value\") else int(job.state)\n",
" timestamp = datetime.datetime.now().strftime(\"%H:%M:%S\")\n",
" print(\n",
" f\"[{timestamp}] Job '{job.display_name}': {job.state.name} (Code: {state})\"\n",
" )\n",
"\n",
" # Final states: 4 is SUCCEEDED, 5 is FAILED, 7 is CANCELLED\n",
" if state in [4, 5, 7]:\n",
" completed_jobs.add(job.name)\n",
" if state == 4:\n",
" print(f\"\\n🎉 SUCCESS: Job '{job.display_name}' completed successfully!\")\n",
" else:\n",
" print(\n",
" f\"\\n⚠️ FAILURE/CANCEL: Job '{job.display_name}' exited with code {state}. Error: {getattr(job, 'error', 'N/A')}\"\n",
" )\n",
"\n",
" if len(completed_jobs) < len(launched_jobs):\n",
" time.sleep(20)\n",
"\n",
"print(\"\\nAll forecast runs are complete!\")"
]
},
{
"cell_type": "markdown",
"id": "d11c21c4",
"metadata": {
"id": "112398663144"
},
"source": [
"## Visualize Forecasts (Unified) "
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "b04bd730",
"metadata": {
"id": "b869b8d30e0a"
},
"outputs": [],
"source": [
"# @markdown Select which forecast output you want to visualize. This single component\n",
"# @markdown can handle both the standard 6-hourly predictions and the datasets\n",
"# @markdown with 1-hour model (which have a 'subtime' dimension).\n",
"# @markdown If you run into `Error loading Zarr store: unrecognized engine 'zarr'...` try restarting the runtime session and reruning this cell.\n",
"\n",
"# @markdown ---\n",
"# @markdown ### Visualization Settings\n",
"# @markdown - **`model_seed_to_visualize`**: Choose a specific model seed (1-4) to visualize. This should be one of the model seeds selected in the **Forecast Configuration** above.\n",
"# @markdown - **`time_steps_to_visualize`**: Choose to visualize 1-hourly or 6-hourly forecasts. If 1-hourly is selected, ensure `enable_hourly_prediction` was selected in the **Forecast Configuration** above.\n",
"# @markdown - **`variable_to_visualize`**: Choose the weather variable to visualize. See the [WeatherNext documentation](https://developers.google.com/weathernext/guides/model-specs-vmg) for variable names and descriptions.\n",
"# @markdown - **`sample_to_visualize`**: Choose the sample (ensemble member) to visualize.\n",
"# @markdown - **`plot_size`**: Choose the size of the plot to generate.\n",
"model_seed_to_visualize = \"4\" # @param [\"1\", \"2\", \"3\", \"4\"]\n",
"time_steps_to_visualize = \"6-Hourly\" # @param [\"6-Hourly\", \"1-Hourly\"]\n",
"variable_to_visualize = \"2m_temperature\" # @param {type:\"string\"}\n",
"sample_to_visualize = 0 # @param {type:\"integer\"}\n",
"plot_size = 8 # @param {type:\"number\"}\n",
"level_to_visualize = None\n",
"# @markdown ---\n",
"\n",
"\n",
"# --- Helper Functions ---\n",
"\n",
"\n",
"def init_time_to_folder_path(init_time: str) -> str:\n",
" \"\"\"\n",
" Convert init time to expected GCS folder path.\n",
" \"\"\"\n",
" init_date, init_time = init_time.split(\"T\")\n",
" return f\"{init_date.replace('-', '')}_{init_time[0:2]}hr\"\n",
"\n",
"\n",
"# override these if you'd like to visualize a different set of forecasts\n",
"visualize_bucket = BUCKET_URI\n",
"visualize_init_date = forecast_init_time\n",
"\n",
"# set paths based on chosen model seed, bucket, and init date\n",
"path_to_6hr_zarr = f\"{BUCKET_URI}/weathernext2_outputs/weathernext_2_seed_{model_seed_to_visualize}/{init_time_to_folder_path(visualize_init_date)}_01_preds/predictions.zarr/\"\n",
"path_to_1hr_zarr = f\"{BUCKET_URI}/weathernext2_outputs/weathernext_2_seed_{model_seed_to_visualize}_hourly/{init_time_to_folder_path(visualize_init_date)}_01_preds/predictions.zarr/\"\n",
"\n",
"\n",
"from typing import Optional\n",
"\n",
"import matplotlib\n",
"import matplotlib.animation as animation\n",
"import matplotlib.pyplot as plt\n",
"import numpy as np\n",
"import xarray\n",
"from IPython.display import HTML\n",
"\n",
"matplotlib.rcParams[\"animation.embed_limit\"] = 500\n",
"\n",
"\n",
"# --- Helper Functions ---\n",
"\n",
"\n",
"def select_data(\n",
" data: xarray.Dataset,\n",
" variable: str,\n",
" level: Optional[int] = None,\n",
") -> xarray.Dataset:\n",
" \"\"\"Selects a variable from the dataset and optionally a level.\"\"\"\n",
" data = data[variable]\n",
" if \"batch\" in data.dims:\n",
" data = data.isel(batch=0)\n",
" if level is not None and \"level\" in data.coords:\n",
" data = data.sel(level=level)\n",
" return data\n",
"\n",
"\n",
"def scale_data(\n",
" data: xarray.Dataset,\n",
" center: Optional[float] = None,\n",
" robust: bool = False,\n",
") -> tuple[xarray.Dataset, matplotlib.colors.Normalize, str]:\n",
" \"\"\"Scales the data for visualization.\"\"\"\n",
" vmin = np.nanpercentile(data.values, (2 if robust else 0))\n",
" vmax = np.nanpercentile(data.values, (98 if robust else 100))\n",
" if center is not None:\n",
" diff = max(vmax - center, center - vmin)\n",
" vmin = center - diff\n",
" vmax = center + diff\n",
" return (\n",
" data,\n",
" matplotlib.colors.Normalize(vmin, vmax),\n",
" (\"RdBu_r\" if center is not None else \"viridis\"),\n",
" )\n",
"\n",
"\n",
"def create_forecast_animation(\n",
" dataset: xarray.Dataset,\n",
" fig_title: str,\n",
" plot_size: float = 5,\n",
" robust: bool = False,\n",
") -> HTML:\n",
" \"\"\"\n",
" Creates a forecast animation from an xarray Dataset.\n",
" It intelligently handles datasets with or without a 'subtime' dimension.\n",
" \"\"\"\n",
" # --- Data Preparation ---\n",
" # Check if the data still has 'subtime'). If so, stack dimensions.\n",
" # Otherwise, just rename the 'time' dimension for consistency.\n",
" if \"subtime\" in dataset.dims:\n",
" print(\"Detected 'subtime' dimension. Stacking for hourly animation.\")\n",
" # Stack 'time' and 'subtime' into a single animation dimension\n",
" plot_data = dataset.stack(animation_step=(\"time\", \"subtime\")).transpose(\n",
" \"animation_step\", \"lat\", \"lon\"\n",
" )\n",
" else:\n",
" print(\"No 'subtime' dimension found. Using 'time' for 6-hourly animation.\")\n",
" # Use 'time' as the animation dimension\n",
" plot_data = dataset.rename({\"time\": \"animation_step\"})\n",
"\n",
" # Now, the animation dimension is always called 'animation_step'\n",
" max_steps = plot_data.sizes[\"animation_step\"]\n",
" init_time = plot_data.coords[\"init_time\"].values\n",
"\n",
" # Scale the data for color mapping\n",
" scaled_data, norm, cmap = scale_data(plot_data, robust=robust)\n",
"\n",
" # --- Plotting Setup ---\n",
" figure = plt.figure(figsize=(plot_size * 2, plot_size))\n",
" ax = figure.add_subplot(1, 1, 1)\n",
" ax.set_xticks([])\n",
" ax.set_yticks([])\n",
" figure.suptitle(fig_title, fontsize=16)\n",
" figure.tight_layout(rect=[0, 0.03, 1, 0.95]) # Adjust for title\n",
"\n",
" im = ax.imshow(\n",
" scaled_data.isel(animation_step=0), norm=norm, origin=\"lower\", cmap=cmap\n",
" )\n",
"\n",
" plt.colorbar(\n",
" mappable=im,\n",
" ax=ax,\n",
" orientation=\"vertical\",\n",
" pad=0.02,\n",
" aspect=16,\n",
" shrink=0.75,\n",
" cmap=cmap,\n",
" extend=(\"both\" if robust else \"neither\"),\n",
" )\n",
"\n",
" # --- Animation Update Function ---\n",
" def update(frame):\n",
" # Get the coordinates for the current frame\n",
" step_coords = plot_data[\"animation_step\"][frame].coords\n",
"\n",
" # Calculate total offset and valid time based on available coordinates\n",
" if \"subtime\" in step_coords: # Hourly data\n",
" total_offset = step_coords[\"time\"].values + step_coords[\"subtime\"].values\n",
" else: # 6-hourly data\n",
" total_offset = step_coords[\"animation_step\"].values\n",
"\n",
" total_hours = total_offset / np.timedelta64(1, \"h\")\n",
" valid_time = init_time + total_offset\n",
" valid_time_str = np.datetime_as_string(valid_time, unit=\"m\").replace(\"T\", \" \")\n",
"\n",
" new_title = (\n",
" f\"{fig_title}\\n\"\n",
" f\"Valid: {valid_time_str} UTC (Forecast: +{total_hours:.1f}h)\"\n",
" )\n",
" figure.suptitle(new_title, fontsize=16)\n",
" im.set_data(scaled_data.isel(animation_step=frame))\n",
"\n",
" # --- Create and Display Animation ---\n",
" ani = animation.FuncAnimation(\n",
" fig=figure, func=update, frames=max_steps, interval=250\n",
" )\n",
" plt.close(figure.number)\n",
" return HTML(ani.to_html5_video())\n",
"\n",
"\n",
"# --- Main Visualization Logic ---\n",
"\n",
"# 1. Select the correct path based on the user's dropdown choice\n",
"if time_steps_to_visualize == \"6-Hourly\":\n",
" path_to_zarr = path_to_6hr_zarr\n",
"elif time_steps_to_visualize == \"1-Hourly\":\n",
" path_to_zarr = path_to_1hr_zarr\n",
"else:\n",
" raise ValueError(\"Invalid visualization target selected.\")\n",
"\n",
"print(f\"Loading data from: {path_to_zarr}\")\n",
"\n",
"# 2. Load the dataset\n",
"try:\n",
" full_dataset = xarray.open_zarr(path_to_zarr)\n",
"except Exception as e:\n",
" print(f\"Error loading Zarr store: {e}\")\n",
" # This is a common point of failure, so we exit gracefully.\n",
"else:\n",
" # 3. Select the specific data slice for visualization\n",
" data_for_vis = full_dataset.isel(sample=sample_to_visualize)\n",
" variable_data = select_data(data_for_vis, variable_to_visualize, level_to_visualize)\n",
"\n",
" # 4. Generate the title\n",
" title = f\"{variable_to_visualize} (Sample {sample_to_visualize})\"\n",
" if level_to_visualize:\n",
" title += f\" at {level_to_visualize} hPa\"\n",
"\n",
" # 5. Create and display the animation\n",
" display(create_forecast_animation(variable_data, title, plot_size, robust=True))"
]
}
],
"metadata": {
"colab": {
"name": "weathernext_2_ic_pc.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
@@ -72,6 +72,21 @@
"\n",
"### Available Anthropic Claude models\n",
"\n",
"#### Claude Fable 5.1\n",
"Claude Fable 5.1 delivers frontier intelligence for ambitious tasks across coding, scientific discovery, and enterprise workflows.\n",
"\n",
"#### Claude Opus 5\n",
"Claude Opus 5 is Anthropic's most advanced Opus model, powering long-running agents while delivering improvements in coding and professional work.\n",
"\n",
"#### Claude Sonnet 5\n",
"Claude Sonnet 5 is our most capable Sonnet model yet, built for coding, agents, and professional work at scale. It brings near-Opus intelligence to the model teams run at scale every day, with the same balance of capability, cost, and speed teams already rely on Sonnet for.\n",
"\n",
"#### Claude Fable 5\n",
"Claude Fable 5 is our next generation of intelligence for the hardest knowledge work and coding problems. It works independently for longer than any prior generally available Claude model: run it in an agent harness and it can work for days at a time, planning across stages, delegating to sub-agents, and checking its own work.\n",
"\n",
"#### Claude Opus 4.8\n",
"Claude Opus 4.8 is our most intelligent Opus model and the best generally available model for coding and agents, with deeper reasoning for enterprise workflows.\n",
"\n",
"#### Claude Opus 4.7\n",
"Claude Opus 4.7 is our most capable production model yet, advancing performance across coding, enterprise workflows, and long-running agentic tasks.\n",
"\n",
@@ -144,26 +159,6 @@
"## Get Started\n"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "0660e339bf3f"
},
"source": [
"### Install required packages\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "754611260f53"
},
"outputs": [],
"source": [
"%pip install -U -q httpx"
]
},
{
"cell_type": "markdown",
"metadata": {
@@ -209,22 +204,17 @@
},
"outputs": [],
"source": [
"MODEL = \"claude-opus-4-7\" # @param [\"claude-opus-4-7\",\"claude-sonnet-4-6\",\"claude-opus-4-6\",\"claude-opus-4-5\",\"claude-haiku-4-5\",\"claude-sonnet-4-5\",\"claude-opus-4-1\",\"claude-sonnet-4\",\"claude-opus-4\",\"claude-3-7-sonnet\",\"claude-3-5-sonnet-v2\",\"claude-3-5-haiku\",\"claude-3-5-sonnet\",\"claude-3-opus\",\"claude-3-haiku\"]\n",
"if MODEL == \"claude-opus-4-7\":\n",
" available_regions = [\n",
" \"global\",\n",
" \"us\",\n",
" \"eu\",\n",
" ]\n",
"elif MODEL == \"claude-sonnet-4-6\":\n",
" available_regions = [\n",
" \"us-east5\",\n",
" \"europe-west1\",\n",
" \"asia-southeast1\",\n",
" \"global\",\n",
" ]\n",
"elif MODEL == \"claude-opus-4-6\":\n",
" available_regions = [\n",
"MODEL = \"claude-fable-5-1\" # @param [\"claude-fable-5-1\",\"claude-opus-5\",\"claude-sonnet-5\",\"claude-fable-5\",\"claude-opus-4-8\",\"claude-opus-4-7\",\"claude-sonnet-4-6\",\"claude-opus-4-6\",\"claude-opus-4-5\",\"claude-haiku-4-5\",\"claude-sonnet-4-5\",\"claude-opus-4-1\",\"claude-sonnet-4\",\"claude-opus-4\",\"claude-3-7-sonnet\",\"claude-3-5-sonnet-v2\",\"claude-3-5-haiku\",\"claude-3-5-sonnet\",\"claude-3-opus\",\"claude-3-haiku\"]\n",
"# Available regions per model.\n",
"MODEL_AVAILABLE_REGIONS = {\n",
" \"claude-fable-5-1\": [\"global\", \"us\", \"eu\"],\n",
" \"claude-opus-5\": [\"global\", \"us\", \"eu\"],\n",
" \"claude-sonnet-5\": [\"global\", \"us\", \"eu\"],\n",
" \"claude-fable-5\": [\"global\", \"us\", \"eu\"],\n",
" \"claude-opus-4-8\": [\"global\", \"us\", \"eu\"],\n",
" \"claude-opus-4-7\": [\"global\", \"us\", \"eu\"],\n",
" \"claude-sonnet-4-6\": [\"us-east5\", \"europe-west1\", \"asia-southeast1\", \"global\"],\n",
" \"claude-opus-4-6\": [\n",
" \"us-east5\",\n",
" \"us-west4\",\n",
" \"us-east1\",\n",
@@ -234,31 +224,22 @@
" \"europe-north1\",\n",
" \"asia-southeast1\",\n",
" \"global\",\n",
" ]\n",
"elif MODEL == \"claude-opus-4-5\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"asia-southeast1\", \"global\"]\n",
"elif MODEL == \"claude-haiku-4-5\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"global\"]\n",
"elif MODEL == \"claude-sonnet-4-5\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"asia-southeast1\", \"global\"]\n",
"elif MODEL == \"claude-opus-4-1\":\n",
" available_regions = [\"us-east5\", \"europe-west4\", \"global\"]\n",
"elif MODEL == \"claude-sonnet-4\":\n",
" available_regions = [\"us-east5\", \"europe-west4\", \"global\"]\n",
"elif MODEL == \"claude-opus-4\":\n",
" available_regions = [\"us-east5\", \"europe-west4\", \"global\"]\n",
"elif MODEL == \"claude-3-7-sonnet\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"europe-west4\", \"global\"]\n",
"elif MODEL == \"claude-3-5-sonnet-v2\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"global\"]\n",
"elif MODEL == \"claude-3-5-haiku\":\n",
" available_regions = [\"us-east5\"]\n",
"elif MODEL == \"claude-3-5-sonnet\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"asia-southeast1\"]\n",
"elif MODEL == \"claude-3-opus\":\n",
" available_regions = [\"us-east5\"]\n",
"elif MODEL == \"claude-3-haiku\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"asia-southeast1\"]"
" ],\n",
" \"claude-opus-4-5\": [\"us-east5\", \"europe-west1\", \"asia-southeast1\", \"global\"],\n",
" \"claude-haiku-4-5\": [\"us-east5\", \"europe-west1\", \"global\"],\n",
" \"claude-sonnet-4-5\": [\"us-east5\", \"europe-west1\", \"asia-southeast1\", \"global\"],\n",
" \"claude-opus-4-1\": [\"us-east5\", \"europe-west4\", \"global\"],\n",
" \"claude-sonnet-4\": [\"us-east5\", \"europe-west4\", \"global\"],\n",
" \"claude-opus-4\": [\"us-east5\", \"europe-west4\", \"global\"],\n",
" \"claude-3-7-sonnet\": [\"us-east5\", \"europe-west1\", \"europe-west4\", \"global\"],\n",
" \"claude-3-5-sonnet-v2\": [\"us-east5\", \"europe-west1\", \"global\"],\n",
" \"claude-3-5-haiku\": [\"us-east5\"],\n",
" \"claude-3-5-sonnet\": [\"us-east5\", \"europe-west1\", \"asia-southeast1\"],\n",
" \"claude-3-opus\": [\"us-east5\"],\n",
" \"claude-3-haiku\": [\"us-east5\", \"europe-west1\", \"asia-southeast1\"],\n",
"}\n",
"\n",
"available_regions = MODEL_AVAILABLE_REGIONS[MODEL]\n"
]
},
{
@@ -354,7 +335,6 @@
"import base64\n",
"import json\n",
"\n",
"import httpx\n",
"import requests\n",
"from IPython.display import Image"
]
@@ -443,9 +423,7 @@
"id": "3d627efff784"
},
"source": [
"#### Encode And Preview Image\n",
"\n",
"We fetch sample images from Wikipedia using the httpx library, but you can use whatever image sources work for you."
"#### Encode And Preview Image"
]
},
{
@@ -456,13 +434,22 @@
},
"outputs": [],
"source": [
"image_url = \"https://upload.wikimedia.org/wikipedia/commons/thumb/a/a7/Camponotus_flavomarginatus_ant.jpg/300px-Camponotus_flavomarginatus_ant.jpg\"\n",
"image_b64 = base64.b64encode(httpx.get(image_url).content).decode(\"utf-8\")\n",
"image_url = \"https://upload.wikimedia.org/wikipedia/commons/a/a7/Camponotus_flavomarginatus_ant.jpg\"\n",
"\n",
"response = requests.get(image_url)\n",
"image = Image(response.content, width=300, height=200)\n",
"# Wikimedia Foundation blocks requests from default HTTP client libraries that\n",
"# do not specify a custom `User-Agent` header.\n",
"headers = {\n",
" \"User-Agent\": (\n",
" \"Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML,\"\n",
" \" like Gecko) Chrome/115.0.0.0 Safari/537.36\"\n",
" )\n",
"}\n",
"\n",
"image"
"response = requests.get(url=image_url, headers=headers)\n",
"response.raise_for_status()\n",
"\n",
"image_b64 = base64.b64encode(response.content).decode(\"utf-8\")\n",
"display(Image(data=response.content, width=300))"
]
},
{
@@ -506,8 +493,16 @@
" \"stream\": False,\n",
"}\n",
"\n",
"request = json.dumps(PAYLOAD)\n",
"!curl -X POST -H \"Authorization: Bearer $(gcloud auth print-access-token)\" -H \"Content-Type: application/json\" {ENDPOINT}/v1/projects/{PROJECT_ID}/locations/{LOCATION}/publishers/anthropic/models/{MODEL}:rawPredict -d '{request}'"
"# Save payload to request.json\n",
"with open(\"request.json\", \"w\") as f:\n",
" json.dump(PAYLOAD, f)\n",
"\n",
"# Pass the file to curl using -d @request.json\n",
"!curl -X POST \\\n",
" -H \"Authorization: Bearer $(gcloud auth print-access-token)\" \\\n",
" -H \"Content-Type: application/json\" \\\n",
" \"{ENDPOINT}/v1/projects/{PROJECT_ID}/locations/{LOCATION}/publishers/anthropic/models/{MODEL}:rawPredict\" \\\n",
" -d @request.json"
]
},
{
@@ -551,8 +546,16 @@
" \"stream\": True,\n",
"}\n",
"\n",
"request = json.dumps(PAYLOAD)\n",
"!curl -X POST -H \"Authorization: Bearer $(gcloud auth print-access-token)\" -H \"Content-Type: application/json\" {ENDPOINT}/v1/projects/{PROJECT_ID}/locations/{LOCATION}/publishers/anthropic/models/{MODEL}:streamRawPredict -d '{request}'"
"# Save payload to request.json\n",
"with open(\"request.json\", \"w\") as f:\n",
" json.dump(PAYLOAD, f)\n",
"\n",
"# Pass the file to curl using -d @request.json\n",
"!curl -X POST \\\n",
" -H \"Authorization: Bearer $(gcloud auth print-access-token)\" \\\n",
" -H \"Content-Type: application/json\" \\\n",
" {ENDPOINT}/v1/projects/{PROJECT_ID}/locations/{LOCATION}/publishers/anthropic/models/{MODEL}:streamRawPredict \\\n",
" -d @request.json"
]
},
{
@@ -590,8 +593,7 @@
},
"outputs": [],
"source": [
"! pip3 install -U -q 'anthropic[vertex]'\n",
"! pip3 install -U -q httpx"
"! pip3 install -U -q 'anthropic[vertex]'"
]
},
{
@@ -679,22 +681,17 @@
},
"outputs": [],
"source": [
"MODEL = \"claude-opus-4-7\" # @param [\"claude-opus-4-7\",\"claude-sonnet-4-6\",\"claude-opus-4-6\",\"claude-opus-4-5\",\"claude-haiku-4-5\",\"claude-sonnet-4-5\",\"claude-opus-4-1\",\"claude-sonnet-4\",\"claude-opus-4\",\"claude-3-7-sonnet\",\"claude-3-5-sonnet-v2\",\"claude-3-5-haiku\",\"claude-3-5-sonnet\",\"claude-3-opus\",\"claude-3-haiku\"]\n",
"if MODEL == \"claude-opus-4-7\":\n",
" available_regions = [\n",
" \"global\",\n",
" \"us\",\n",
" \"eu\",\n",
" ]\n",
"elif MODEL == \"claude-sonnet-4-6\":\n",
" available_regions = [\n",
" \"us-east5\",\n",
" \"europe-west1\",\n",
" \"asia-southeast1\",\n",
" \"global\",\n",
" ]\n",
"elif MODEL == \"claude-opus-4-6\":\n",
" available_regions = [\n",
"MODEL = \"claude-fable-5-1\" # @param [\"claude-fable-5-1\",\"claude-opus-5\",\"claude-sonnet-5\",\"claude-fable-5\",\"claude-opus-4-8\",\"claude-opus-4-7\",\"claude-sonnet-4-6\",\"claude-opus-4-6\",\"claude-opus-4-5\",\"claude-haiku-4-5\",\"claude-sonnet-4-5\",\"claude-opus-4-1\",\"claude-sonnet-4\",\"claude-opus-4\",\"claude-3-7-sonnet\",\"claude-3-5-sonnet-v2\",\"claude-3-5-haiku\",\"claude-3-5-sonnet\",\"claude-3-opus\",\"claude-3-haiku\"]\n",
"# Available regions per model.\n",
"MODEL_AVAILABLE_REGIONS = {\n",
" \"claude-fable-5-1\": [\"global\", \"us\", \"eu\"],\n",
" \"claude-opus-5\": [\"global\", \"us\", \"eu\"],\n",
" \"claude-sonnet-5\": [\"global\", \"us\", \"eu\"],\n",
" \"claude-fable-5\": [\"global\", \"us\", \"eu\"],\n",
" \"claude-opus-4-8\": [\"global\", \"us\", \"eu\"],\n",
" \"claude-opus-4-7\": [\"global\", \"us\", \"eu\"],\n",
" \"claude-sonnet-4-6\": [\"us-east5\", \"europe-west1\", \"asia-southeast1\", \"global\"],\n",
" \"claude-opus-4-6\": [\n",
" \"us-east5\",\n",
" \"us-west4\",\n",
" \"us-east1\",\n",
@@ -704,31 +701,22 @@
" \"europe-north1\",\n",
" \"asia-southeast1\",\n",
" \"global\",\n",
" ]\n",
"elif MODEL == \"claude-opus-4-5\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"asia-southeast1\", \"global\"]\n",
"elif MODEL == \"claude-haiku-4-5\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"global\"]\n",
"elif MODEL == \"claude-sonnet-4-5\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"asia-southeast1\", \"global\"]\n",
"elif MODEL == \"claude-opus-4-1\":\n",
" available_regions = [\"us-east5\", \"europe-west4\", \"global\"]\n",
"elif MODEL == \"claude-sonnet-4\":\n",
" available_regions = [\"us-east5\", \"europe-west4\", \"global\"]\n",
"elif MODEL == \"claude-opus-4\":\n",
" available_regions = [\"us-east5\", \"europe-west4\", \"global\"]\n",
"elif MODEL == \"claude-3-7-sonnet\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"europe-west4\", \"global\"]\n",
"elif MODEL == \"claude-3-5-sonnet-v2\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"global\"]\n",
"elif MODEL == \"claude-3-5-haiku\":\n",
" available_regions = [\"us-east5\"]\n",
"elif MODEL == \"claude-3-5-sonnet\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"asia-southeast1\"]\n",
"elif MODEL == \"claude-3-opus\":\n",
" available_regions = [\"us-east5\"]\n",
"elif MODEL == \"claude-3-haiku\":\n",
" available_regions = [\"us-east5\", \"europe-west1\", \"asia-southeast1\"]"
" ],\n",
" \"claude-opus-4-5\": [\"us-east5\", \"europe-west1\", \"asia-southeast1\", \"global\"],\n",
" \"claude-haiku-4-5\": [\"us-east5\", \"europe-west1\", \"global\"],\n",
" \"claude-sonnet-4-5\": [\"us-east5\", \"europe-west1\", \"asia-southeast1\", \"global\"],\n",
" \"claude-opus-4-1\": [\"us-east5\", \"europe-west4\", \"global\"],\n",
" \"claude-sonnet-4\": [\"us-east5\", \"europe-west4\", \"global\"],\n",
" \"claude-opus-4\": [\"us-east5\", \"europe-west4\", \"global\"],\n",
" \"claude-3-7-sonnet\": [\"us-east5\", \"europe-west1\", \"europe-west4\", \"global\"],\n",
" \"claude-3-5-sonnet-v2\": [\"us-east5\", \"europe-west1\", \"global\"],\n",
" \"claude-3-5-haiku\": [\"us-east5\"],\n",
" \"claude-3-5-sonnet\": [\"us-east5\", \"europe-west1\", \"asia-southeast1\"],\n",
" \"claude-3-opus\": [\"us-east5\"],\n",
" \"claude-3-haiku\": [\"us-east5\", \"europe-west1\", \"asia-southeast1\"],\n",
"}\n",
"\n",
"available_regions = MODEL_AVAILABLE_REGIONS[MODEL]\n"
]
},
{
@@ -815,7 +803,6 @@
"source": [
"import base64\n",
"\n",
"import httpx\n",
"import requests\n",
"from IPython.display import Image"
]
@@ -916,9 +903,7 @@
"id": "2fe57432a56d"
},
"source": [
"#### Encode And Preview Image\n",
"\n",
"We fetch sample images from Wikipedia using the httpx library, but you can use whatever image sources work for you."
"#### Encode And Preview Image"
]
},
{
@@ -929,14 +914,23 @@
},
"outputs": [],
"source": [
"image_url = \"https://upload.wikimedia.org/wikipedia/commons/thumb/a/a7/Camponotus_flavomarginatus_ant.jpg/300px-Camponotus_flavomarginatus_ant.jpg\"\n",
"image_url = \"https://upload.wikimedia.org/wikipedia/commons/a/a7/Camponotus_flavomarginatus_ant.jpg\"\n",
"image_media_type = \"image/jpeg\"\n",
"image_b64 = base64.b64encode(httpx.get(image_url).content).decode(\"utf-8\")\n",
"\n",
"response = requests.get(image_url)\n",
"image = Image(response.content, width=300, height=200)\n",
"# Wikimedia Foundation blocks requests from default HTTP client libraries that\n",
"# do not specify a custom `User-Agent` header.\n",
"headers = {\n",
" \"User-Agent\": (\n",
" \"Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML,\"\n",
" \" like Gecko) Chrome/115.0.0.0 Safari/537.36\"\n",
" )\n",
"}\n",
"\n",
"image"
"response = requests.get(url=image_url, headers=headers)\n",
"response.raise_for_status()\n",
"\n",
"image_b64 = base64.b64encode(response.content).decode(\"utf-8\")\n",
"display(Image(data=response.content, width=300))"
]
},
{
@@ -348,7 +348,7 @@
" \"gemini-embedding-001\",\n",
" \"text-multilingual-embedding-002\",\n",
"]:\n",
" raise ValueError(f\"OUTPUT_DIMENTIONALITY cannot be specified for model '{MODEL}'.\")\n",
" raise ValueError(f\"OUTPUT_DIMENSIONALITY cannot be specified for model '{MODEL}'.\")\n",
"if TASK in [\"QUESTION_ANSWERING\", \"FACT_VERIFICATION\"] and MODEL not in [\n",
" \"text-embedding-005\",\n",
" \"text-embedding-004\",\n",
@@ -1,5 +1,5 @@
google-cloud-aiplatform==1.138.0
numpy==2.4.4
pandas==3.0.1
datasets==2.18.0
smart_open[gcs]==7.5.1
google-cloud-aiplatform==1.165.0
numpy==2.5.2
pandas==3.0.5
datasets==5.0.1
smart_open[gcs]==8.0.1
@@ -1,4 +1,4 @@
google-cloud-aiplatform==1.138.0
numpy==2.4.4
pandas==3.0.1
datasets==2.18.0
google-cloud-aiplatform==1.165.0
numpy==2.5.2
pandas==3.0.5
datasets==5.0.1