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
Rayan DasoriyaandCopybara-Service 24244351cd Add a new notebook for OSS distillation feasibility study.
PiperOrigin-RevId: 917879385
2026-05-19 09:37:47 -07:00
Vertex MG TeamandCopybara-Service bf0e1300a9 No public description
MG_DOCKER_CODES_PIPER_ORIGIN_REV_ID: 878451476
2026-05-13 12:58:56 -07:00
chnduandGitHub cf048b6fe4 Add live_api skills that help the user build their own liveapi service (#4511)
* Add live_api skills that help the user to build their own liveapi service.

Implementation are based on websocket. Support different coding languages.

* Update based on review

* Fix typos

* Update vertex to gemini enterprise.
2026-05-11 17:19:12 +00:00
Mend RenovateandGitHub 8c8820ecfa chore(deps): update dependency numpy to v2.4.4 (#4466) 2026-05-06 14:24:57 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
a1a52d8145 chore(deps): bump requests (#4488)
Bumps [requests](https://github.com/psf/requests) from 2.32.4 to 2.33.0.
- [Release notes](https://github.com/psf/requests/releases)
- [Changelog](https://github.com/psf/requests/blob/main/HISTORY.md)
- [Commits](https://github.com/psf/requests/compare/v2.32.4...v2.33.0)

---
updated-dependencies:
- dependency-name: requests
  dependency-version: 2.33.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-05-06 14:22:23 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
a06ce545e7 chore(deps): bump pillow (#4497)
Bumps [pillow](https://github.com/python-pillow/Pillow) from 12.1.1 to 12.2.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.2.0)

---
updated-dependencies:
- dependency-name: pillow
  dependency-version: 12.2.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-05-06 14:19:28 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
913780c4cb chore(deps): bump pillow (#4509)
Bumps [pillow](https://github.com/python-pillow/Pillow) from 10.3.0 to 12.2.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/10.3.0...12.2.0)

---
updated-dependencies:
- dependency-name: pillow
  dependency-version: 12.2.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-05-06 14:14:14 +00:00
gmaninatarajanandGitHub 71be46e7d8 feat: Updated whl file and package name as part of Vertex Model Garden setup (#4508) 2026-04-30 08:28:04 -04:00
Mayank SharanandGitHub 849e88a627 Vtc blog 2 (#4507)
* Adding reviewed version of VTC blog 2

* GCA suggested fixes

* Updating readme to have links
2026-04-29 19:10:45 +00:00
Jason DaiandGitHub daf56bcd0b Create Eval Quality Flywheel Skill for preview (#4505) 2026-04-23 16:37:08 +00:00
Mayank SharanandGitHub 563f423b93 Adding reviewed version of VTC blog 2 (#4502)
* Adding reviewed version of VTC blog 2

* GCA suggested fixes
2026-04-16 21:12:03 +00:00
ian1780andGitHub 292e540e96 Update anthropic_claude_intro.ipynb (#4501)
add opus 4.7 multi region endpoint support b/491171457
2026-04-16 17:59:12 +00:00
Sam-DecigaandGitHub 6c6a703c5a Ant nickel (#4500)
* Feat: Anthropic Opus-4-7 launch

* Feat: Anthropic Opus-4-7 launch
2026-04-16 12:18:47 -04:00
Sam-DecigaandGitHub 7ef83c6f73 MARS8 new Asian regions (#4498) 2026-04-14 19:59:45 +00:00
Vertex MG TeamandCopybara-Service aba6598109 use old docker hash for whisper model deployment.
PiperOrigin-RevId: 892120947
2026-03-30 23:08:58 -07:00
Sam-DecigaandGitHub 88a6b8037e Refactor: Anthropic NB (#4489) 2026-03-26 14:49:59 -04:00
Eric DongandGitHub 8845f7ab27 Update Gemini model references and availability details
Update Gemini versions.
2026-03-26 09:40:01 -04:00
Eric DongandGitHub 3b2e711a16 Update fine-tuning model reference in README
Updated the fine-tuning model reference from Gemini 1.5 Pro to Gemini 2.5 Pro in the README.
2026-03-25 16:44:52 -04:00
Eric DongandGitHub 5c0629cdc7 Revise README for Agent Skills in Vertex AI
Updated terminology and formatting for clarity.
2026-03-25 15:18:03 -04:00
Eric DongandGitHub e107d30807 chore: Add detailed installation instructions for skills (#4487) 2026-03-25 15:06:59 -04:00
Eric DongandGitHub b98ab36913 refactor: Add tool configuation in skills readme (#4486)
* chore: Update vertex-ai Skills readme

* refactor: Add tool configuation in  skills readme
2026-03-25 13:41:45 -04:00
Eric DongandGitHub f1d90b5a71 chore: Update vertex-ai Skills readme (#4485) 2026-03-25 11:36:41 -04:00
Eric DongandGitHub f848db6132 Revise README title and formatting for emphasis
Updated the title and emphasized 'Skills' in the README.
2026-03-25 10:16:42 -04:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
cc9fffd945 chore(deps): bump pillow (#4482)
Bumps [pillow](https://github.com/python-pillow/Pillow) from 10.3.0 to 12.1.1.
- [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/10.3.0...12.1.1)

---
updated-dependencies:
- dependency-name: pillow
  dependency-version: 12.1.1
  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-03-24 15:13:00 +00:00
Eric DongandGitHub 7606a1de03 chore: Update the skills readme with instructions (#4484) 2026-03-24 10:44:28 -04:00
Eric DongandGitHub cca59aa753 Update README.md
Remove icons
2026-03-24 10:12:04 -04:00
Eric DongandGitHub a1907da27a chore: Update readme (#4483) 2026-03-24 10:00:50 -04:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
e8cb7738d0 Bump pillow (#4443)
Bumps [pillow](https://github.com/python-pillow/Pillow) from 10.3.0 to 12.1.1.
- [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/10.3.0...12.1.1)

---
updated-dependencies:
- dependency-name: pillow
  dependency-version: 12.1.1
  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-03-24 13:52:57 +00:00
Mend RenovateandGitHub 18e8d603de chore(deps): update dependency black to v26.3.1 [security] (#4470) 2026-03-24 13:52:02 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
1bd5901fdb Bump black (#4468)
Bumps [black](https://github.com/psf/black) from 25.1.0 to 26.3.1.
- [Release notes](https://github.com/psf/black/releases)
- [Changelog](https://github.com/psf/black/blob/main/CHANGES.md)
- [Commits](https://github.com/psf/black/compare/25.1.0...26.3.1)

---
updated-dependencies:
- dependency-name: black
  dependency-version: 26.3.1
  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-03-24 13:51:38 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
5b245024cd Bump pyasn1 (#4474)
Bumps [pyasn1](https://github.com/pyasn1/pyasn1) from 0.6.2 to 0.6.3.
- [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.2...v0.6.3)

---
updated-dependencies:
- dependency-name: pyasn1
  dependency-version: 0.6.3
  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-03-24 13:50:31 +00:00
Eric DongandGitHub ba043c196c chore: Update skills readme with architecture (#4481) 2026-03-24 09:49:35 -04:00
Eric DongandGitHub 28ce8f6d7a chore: Update readme and template (#4480) 2026-03-24 09:42:45 -04:00
Eric DongandGitHub 3b5a8cad41 feat: Add Gen AI SDK skill for Vertex (#4479) 2026-03-24 09:29:01 -04:00
gmaninatarajanandGitHub c6d7971bc9 fix:Simplified authentication section and addressed timeout issues (#4478) 2026-03-23 18:50:35 -04:00
Eric DongandGitHub dbe28965cb feat: Add primary routing and readme for vertex ai skills (#4477) 2026-03-23 17:18:16 -04:00
Vertex MG TeamandCopybara-Service 0d34d6bbea update minimax m2 notebook.
PiperOrigin-RevId: 886885747
2026-03-20 11:14:55 -07:00
Sam-DecigaandGitHub a1d898f35e feat: Jina EmbV3 launch (#4476)
* feat: Jina EmbV3 launch

* feat: Jina EmbV3 launch
2026-03-20 08:19:38 -04:00
Eric DongandGitHub 772ee71bc3 feat: use Vertex AI MCP server (#4475)
* feat: use Vertex AI MCP server

* Address review comments
2026-03-19 11:02:28 -04:00
Sam-DecigaandGitHub 425851cedc feat: Nemotron3-Super model launch (#4473)
* feat: Nemotron3-Super model launch

* feat: Nemotron3-Super model launch

* feat: Nemotron3-Super model launch
2026-03-16 20:41:36 -04:00
Lav RaiandGitHub bcccbee164 Update distillation report. (#4472) 2026-03-16 18:40:10 +00:00
Lav RaiandGitHub 5ae325528a Add distillation report. (#4471) 2026-03-13 15:33:34 +00:00
vincentkt-googleandGitHub 86674effee Add and update existing vertex skills (#4467)
* Add and update existing vertex skills

- Add support for fine tuning for 1p gemini tuning
- Add support for deploying fine tuned model support
- Add support for running inference on MaaS models
- Add open model support for regions and cost estimating for 3p tuning

* fixing some of the commit errors

* updated scripts to use existing gemini 1.5 pro model

* swap gemini 1.5 pro to gemini 2.5 pro
2026-03-11 19:45:42 +00:00
vincentkt-googleandGitHub 8b4708c606 feat: add vertex ai skills to repo (#4454) 2026-03-05 17:50:41 +00:00
Yichen ZhouandCopybara-Service f3dd6cbca3 Update TimesFM-2.5 notebook for Model Garden.
PiperOrigin-RevId: 878726929
2026-03-04 16:42:55 -08:00
Rayan DasoriyaandCopybara-Service 1f9e93993c No public description
MG_DOCKER_CODES_PIPER_ORIGIN_REV_ID: 878215121
2026-03-03 18:26:58 -08:00
Vertex MG TeamandCopybara-Service 062835174e Updated the image default TAG to release
PiperOrigin-RevId: 877893209
2026-03-03 05:26:21 -08:00
Rayan DasoriyaandCopybara-Service cb4916f590 No public description
MG_DOCKER_CODES_PIPER_ORIGIN_REV_ID: 875506020
2026-02-25 21:44:41 -08:00
Damodar PanigrahiandGitHub 2933fe606b bug: remove A100 as recommended specs (#4450) 2026-02-25 14:05:36 +00:00
Damodar PanigrahiandGitHub b468809df7 bug: ahref update (#4449)
* bug: ahref update

* fix: linter

* fix: typo fix
2026-02-25 13:46:19 +00:00
Sam-DecigaandGitHub 0417d8b9c4 feat: Deprecate Claude 3 Haiku (#4448)
Deprecation start date: Feb. 23, 2026
End of Support date: Aug. 23, 2026
b/485993204
2026-02-23 20:57:43 -05:00
Damodar PanigrahiandGitHub 7750e83fbb fix: inference key change, finetuning jaxlib update (#4447) 2026-02-23 19:45:59 +00:00
Damodar PanigrahiandGitHub 649800e646 feat: Alphagenome finetuning notebook (#4445)
* feat: Alphagenome finetuning notebook

* Update cloudai_alphagenome_finetune.ipynb

Fixed the lint errors

* Update cloudai_alphagenome_finetune.ipynb

Fix lint errors

* feat: Add Alphagenome finetune

* feat: Include the Alphagenome finetuning notebook url in the readme. Update the codeowners

* feat: Add alphagegenome finetune notebook to the readme, add user to codeowners

* feat: fix spelling
2026-02-20 14:13:36 +00:00
Mend RenovateandGitHub 42b35056fa chore(deps): update dependency black to v26 (#4422) 2026-02-18 15:07:12 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
24974eda95 Bump protobuf (#4437)
Bumps [protobuf](https://github.com/protocolbuffers/protobuf) from 4.25.8 to 5.29.6.
- [Release notes](https://github.com/protocolbuffers/protobuf/releases)
- [Commits](https://github.com/protocolbuffers/protobuf/commits)

---
updated-dependencies:
- dependency-name: protobuf
  dependency-version: 5.29.6
  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-02-18 15:06:37 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
ad41377783 Bump protobuf (#4438)
Bumps [protobuf](https://github.com/protocolbuffers/protobuf) from 4.25.8 to 5.29.6.
- [Release notes](https://github.com/protocolbuffers/protobuf/releases)
- [Commits](https://github.com/protocolbuffers/protobuf/commits)

---
updated-dependencies:
- dependency-name: protobuf
  dependency-version: 5.29.6
  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-02-18 15:06:10 +00:00
Sam-DecigaandGitHub 36dea3ca01 feat: Anthropic Sonnet-4-6 launch (#4442)
Signed-off-by: Sam-Deciga <decigagarcia@google.com>
2026-02-17 14:02:30 -05:00
b75b2ea4d7 fix: updated with latest whl file version : alphagenome-0.4.2.6-py3-none-any.whl (#4439)
Co-authored-by: hyper-param <peeyusht@google.com>
2026-02-10 18:27:03 +00:00
Vertex MG TeamandCopybara-Service 28f7fc4445 Add SAM 3 notebook to Vertex AI Model Garden.
PiperOrigin-RevId: 866598579
2026-02-06 13:45:47 -08:00
201 changed files with 21556 additions and 24331 deletions
+2 -2
View File
@@ -2,9 +2,9 @@ git+https://github.com/tensorflow/docs
ipython
jupyter
nbconvert
black==25.12.0
black==26.5.1
pyupgrade==3.21.2
isort==7.0.0
isort==8.0.1
flake8==7.3.0
nbqa==1.9.1
+25 -9
View File
@@ -1,12 +1,12 @@
# ![Google Cloud](https://avatars.githubusercontent.com/u/2810941?s=60&v=4) Google Cloud Vertex AI Samples
This repository contains notebooks, code samples, sample apps, and other resources that demonstrate how to use, develop and manage machine learning and generative AI workflows using Google Cloud Vertex AI.
This repository contains notebooks, code samples, sample apps, skills, and other resources that demonstrate how to use, develop and manage machine learning and generative AI workflows using Google Cloud Vertex AI.
## Overview
[Vertex AI](https://cloud.google.com/vertex-ai) is a fully-managed, unified AI development platform for building and using generative AI. This repository is designed to help you get started with Vertex AI. Whether you're new to Vertex AI or an experienced ML practitioner, you'll find valuable resources here.
For more Vertex AI Generative AI notebook samples, please visit the Vertex AI [Generative AI](https://github.com/GoogleCloudPlatform/generative-ai) GitHub repository.
⚠️ For more Vertex AI Generative AI notebook samples, please visit the Vertex AI [Generative AI](https://github.com/GoogleCloudPlatform/generative-ai) GitHub repository.
## Explore, learn and contribute
@@ -16,11 +16,11 @@ You can explore, learn, and contribute to this repository to unleash the full po
Explore this repository, follow the links in the header section of each of the notebooks to -
![Colab](https://cloud.google.com/ml-engine/images/colab-logo-32px.png) Open and run the notebook in [Colab](https://colab.google/)\
![Colab Enterprise](https://cloud.google.com/ml-engine/images/colab-enterprise-logo-32px.png) Open and run the notebook in [Colab Enterprise](https://cloud.google.com/colab/docs/introduction)\
![Workbench](https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32) Open and run the notebook in [Vertex AI Workbench](https://cloud.google.com/vertex-ai/docs/workbench/introduction)\
![Github](https://cloud.google.com/ml-engine/images/github-logo-32px.png) View the notebook on Github
- Open and run the notebook in [Colab](https://colab.google/)
- Open and run the notebook in [Colab Enterprise](https://cloud.google.com/colab/docs/introduction)
- Open and run the notebook in [Vertex AI Workbench](https://cloud.google.com/vertex-ai/docs/workbench/introduction)
- View the notebook on Github
### Contribute
See the [Contributing Guide](https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/master/CONTRIBUTING.md).
@@ -35,7 +35,7 @@ To get started using Vertex AI, you must have a Google Cloud project.
## Repository structure
```bash
```text
├── notebooks
│ ├── official - Notebooks demonstrating use of each Vertex AI service
│ │ ├── automl
@@ -45,7 +45,23 @@ To get started using Vertex AI, you must have a Google Cloud project.
│ │ ├── model_garden
│ │ ├── ...
├── community-content - Sample code and tutorials contributed by the community
├── docs - Deep-dive documentation and advanced setup guides
└── skills - Suite of AI Agent "Skills" for Vertex AI
├── README.md # Developer guide for Vertex AI skills
├── vertex-ai/ # Primary router for Vertex AI tasks
│ └── SKILL.md # Entry point that routes across capabilities
├── genai-sdk/ # Gemini API usage with Gen AI SDK
│ └── SKILL.md # Guides for Python, JS/TS, Go, Java, C#
├── vertex-deploy/ # Deploying models to Endpoints
│ └── SKILL.md # Commands for open models & custom weights
├── vertex-inference/ # Inferencing with GenAI models
│ └── SKILL.md # Code samples for Gemini and OpenMaaS
└── vertex-tuning/ # Secondary router for model fine-tuning
├── SKILL.md # Router for tuning tasks
├── gemini/ # Fine-tuning first-party Gemini models
│ └── SKILL.md
└── open-model/ # Fine-tuning third-party open models
└── SKILL.md
```
## Examples
@@ -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==10.3.0
pillow==12.3.0
tf-agents==0.8.0
@@ -1,4 +1,4 @@
google-cloud-pubsub==2.5.0
pillow==10.3.0
pillow==12.3.0
tf-agents==0.8.0
tensorflow==2.12.1
@@ -1,5 +1,5 @@
dataclasses==0.6
google-cloud-aiplatform==1.8.1
tensorflow==2.12.1
pillow==10.3.0
pillow==12.3.0
tf-agents==0.8.0
@@ -2,9 +2,9 @@ dllogger@git+https://github.com/NVIDIA/dllogger@v1.0.0
# Fixing these libraries versions to avoid conflicting or broken packages.
immutabledict==4.2.1
protobuf==4.25.8
protobuf==5.29.6
opencv-python-headless==4.11.0.86
docutils==0.16
urllib3==2.6.3
urllib3==2.7.0
google-cloud-storage==3.0.0
retrying
@@ -1,7 +1,7 @@
absl-py==2.2.2
annotated-types==0.7.0
anyio==4.9.0
black==25.1.0
black==26.3.1
cachetools==5.5.2
certifi==2025.4.26
charset-normalizer==3.4.2
@@ -9,7 +9,7 @@ click==8.1.8
docstring_parser==0.16
google-api-core==2.24.2
google-auth==2.40.1
google-cloud-aiplatform==1.92.0
google-cloud-aiplatform==1.133.0
google-cloud-bigquery==3.31.0
google-cloud-core==2.4.3
google-cloud-resource-manager==1.14.2
@@ -24,26 +24,26 @@ grpcio-status==1.71.0
h11==0.16.0
httpcore==1.0.9
httpx==0.28.1
idna==3.10
idna==3.15
mypy_extensions==1.1.0
numpy==2.2.5
packaging==25.0
pathspec==0.12.1
platformdirs==4.3.8
proto-plus==1.26.1
protobuf==5.29.4
pyasn1==0.6.2
protobuf==5.29.6
pyasn1==0.6.4
pyasn1_modules==0.4.2
pydantic==2.11.4
pydantic_core==2.33.2
python-dateutil==2.9.0.post0
pytz==2025.2
requests==2.32.4
requests==2.33.0
rsa==4.9.1
shapely==2.1.0
six==1.17.0
sniffio==1.3.1
typing-inspection==0.4.0
typing_extensions==4.13.2
urllib3==2.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
+11
View File
@@ -0,0 +1,11 @@
# Agent Platform Training Clusters Blog Series
This directory contains deep-dive documentation, extended guides, and architectural references for Google Cloud Agent Platform Training Clusters.
## Contents
- **`vertex-training-cluster/`**: Documentation and setup guides for configuring and managing Agent Platform Training Clusters.
## Blog Posts
- [Model Distillation Best Practices](https://googlecloudplatform.github.io/vertex-ai-samples/vertex-training-cluster/model_distillation_best_practices): Explores off-policy model distillation, dataset curation, and hyperparameter scaling laws for training student models on Vertex AI.
- [Forgetting Mitigation via Data Mixing](https://googlecloudplatform.github.io/vertex-ai-samples/vertex-training-cluster/forgetting_mitigation_data_mixing): Discusses catastrophic forgetting in model fine-tuning and how to mitigate it using multi-domain data mixing on Vertex AI.
- [Multi-Turn Reinforcement Learning for τ²-bench](https://googlecloudplatform.github.io/vertex-ai-samples/vertex-training-cluster/multi_turn_reinforcement_learning_for_tau2_bench): Explores multi-turn RL training for tool-calling agents using GRPO on the τ²-bench customer service benchmark with NeMo RL.
File diff suppressed because one or more lines are too long
@@ -0,0 +1,448 @@
<script type="text/javascript" async
src="https://cdn.mathjax.org/mathjax/latest/MathJax.js?config=TeX-MML-AM_CHTML">
</script><br><br>
# VTC Multi-Domain Dataset: Mitigating Catastrophic Forgetting with Data Mixing
**Author:** [Mayank Sharan](mailto:mayanksharan@google.com)
## Table of Contents
* [Intro](#intro)
* [Background](#background)
* [Dataset Curation](#dataset-selection)
* [Forgetting Mitigation Best Practices](#forgetting-mitigation-best-practices)
* [Experimental Setup](#experimental-setup)
* [Mitigating Forgetting](#mitigating-forgetting)
* [Mixing Ratios](#mixing-ratios)
* [Different Starting Models](#different-starting-models)
* [Acknowledgements](#acknowledgements)
* [References](#references)
## Intro
In this entry of our blog series on model training best practices for Vertex AI Training Cluster (VTC) customers, we talk about catastrophic forgetting and how to mitigate it. We focus on tuning public models using supervised fine tuning (SFT) with a specialized domain dataset. With both open and closed source models performing well on general tasks the primary goal of training one's own models is to improve the performance on specialized tasks. This typically comes at the cost of the model forgetting general capabilities which can severely limit the utility of the trained model.
There are many possible interventions to limit forgetting, the most effective is mixing the target dataset with the actual dataset used in the model’s training. Since this is not available even for the most open source models, we have curated a multi-domain dataset that delivers the same benefits. This allows Vertex AI Training Cluster (VTC) customers to maintain and surpass frontier level model capabilities while training to further performance on specialized tasks.
<figure align="center" id="fig-teaser">
<table align="center" width="80%">
<tr>
<td align="center" width="100%">
<img src="images_data_mixing/teaser_forgetting.png" width="100%"><br>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 1: Impact of mixing VTC Post Training dataset on Forgetting (8B model). </b> <i>Comparing SFT runs using only a specialized target dataset (MedMCQA) vs a mix of the target dataset and the VTC Post training dataset. Forgetting across all non-target domains is significantly mitigated with no performance loss on the target metric. Qwen3 Public here is the instruction tuned public Qwen3 8B model and the other two models are trained starting from the base Qwen3 8B model using only the target dataset and a mix of target dataset with the VTC dataset.</i></sub>
</figcaption>
</figure>
We provide a thorough set of experiments to serve as a guide for reducing forgetting while post training the Qwen3 open-weight thinking model family, beginning from their base pre-trained checkpoints. Furthermore, we demonstrate the value our datasets provide across model sizes often surpassing the performance of the official Qwen3 models while preserving performance on the specialized task (See [Figure 1](#fig-teaser)). The Qwen3 family was specifically chosen for this study because its diverse range of parameter counts and the availability of both pre-trained and post-trained checkpoints provide an ideal environment for high-fidelity scaling analysis.
To ensure our findings can be applied to a broad set of applications we validate our findings across five model sizes: 0.6B, 1.7B, 4B, 8B and 14B parameters. To support our VTC community in accelerating their own development, all code, datasets, and experiment configurations used in this blog are being made available for use in your training workloads.
## Background
Loss landscapes for neural networks have always been a complex multidimensional manifold rather than the simple convex ones that gradient descent is built for. Forgetting is a well known phenomenon in model customization, the first academically recorded instance being (McCloskey and Cohen, 1989) [<a href="#ref1">1</a>]. These manifolds have become even more complex with the introduction of Large Language Models where the number of parameters being optimized are typically in the billions. This makes it hard to mathematically grasp issues like forgetting. [Figure 2](#fig-loss-landscape) demonstrates a geometric understanding of why forgetting happens and how data mixing can mitigate it.
<figure align="center" id="fig-loss-landscape">
<table align="center" width="80%">
<tr>
<td align="center" width="100%">
<img src="images_data_mixing/background_loss_landscape.png" width="100%"><br>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 2: Geometric Interpretation of Data Mixing to Mitigate Forgetting. </b> <i>Fine tuning objectives being meaningfully out of distribution from the pre-trained model often drives forgetting. Mixing in a dataset similar to the model distribution adjusts the objective enough to learn the new task without as much forgetting.</i></sub>
</figcaption>
</figure>
## Dataset Selection
Our primary requirements for a target dataset to run experiments to validate this were:
1. It should be out of distribution to cause forgetting
2. It should have an evaluation metric that it directly improves
3. It should be able to train the model to perform better than the counterpart generalist model
A good heuristic to determine where the data lies with respect to the model distribution is by calculating perplexity on samples from the dataset. Assuming
- <span>$$X={x_1, x_2, \dots, x_N}$$</span> is a dataset sample represented as sequence of tokens
- <span>$$P(x_i \mid x_{<i})$$</span> is the model likelihood of the i-th token given the sample till that token
Then the perplexity for this sample can be calculated as follows:
$$\begin{align*}
& ppl(X) = \exp \left( -\frac{1}{N} \sum_{i=1}^{N} \log P(x_i \mid x_{<i}) \right) \\
& = \exp \left( -\frac{1}{N} \log (\prod_{i=1}^{N} P(x_i \mid x_{<i})) \right)
\end{align*} $$
The product form of the equation shows that this is a direct measure of the joint probability of this sequence of tokens according to the model. Since this computation has a balancing negative sign to account for the negative log value a lower joint probability results in a higher perplexity value and vice versa. We evaluated the following datasets as out-of-distribution candidates:
- [MedMCQA](https://huggingface.co/datasets/syz-ml2025/medmcqa) : Multiple Choice Questions (MCQ) dataset focusing on the medical domain
- [BirdSQL](https://huggingface.co/datasets/birdsql/bird23-train-filtered) : Text to SQL generation dataset
- [HardGen](https://huggingface.co/datasets/Bingguang/HardGen) : Function calling dataset
We also calculate perplexity on [OpenR1-Math-220k](https://huggingface.co/datasets/open-r1/OpenR1-Math-220k) to provide a reference as we expect this to be in distribution for the model given the Qwen3 models are particularly strong in the math domain.
<table id="tab-perplexity" style="margin-left:auto; margin-right:auto;">
<thead>
<tr>
<th>Dataset \ Model</th>
<th>Qwen3-0.6B</th>
<th>Qwen3-8B</th>
<th>Ours-0.6B</th>
<th>Ours-8B</th>
</tr>
</thead>
<tbody>
<tr>
<td>MedMCQA</td>
<td>63.00</td>
<td>66.00</td>
<td>42.00</td>
<td>20.75</td>
</tr>
<tr>
<td>BirdSQL</td>
<td>38.50</td>
<td>55.50</td>
<td>45.50</td>
<td>17.00</td>
</tr>
<tr>
<td>HardGen</td>
<td>2.23</td>
<td>2.03</td>
<td>2.28</td>
<td>1.79</td>
</tr>
<tr>
<td>OpenR1-Math</td>
<td>8.63</td>
<td>9.75</td>
<td>6.44</td>
<td>5.34</td>
</tr>
</tbody>
<caption style="text-align: left;"><b>Table 1:</b> Perplexity score analysis with the public instruction tuned Qwen3 models and Qwen3 base models trained using the VTC dataset (Ours) to identify a suitable target dataset.</caption>
</table>
We see from [Table 1](#tab-perplexity) that OpenR1-Math-220k as we expected has low perplexity scores and HardGen shows an even lower perplexity score eliminating it from consideration. MedMCQA samples have high perplexity scores across all considered models. This dataset also has the advantage of a straightforward evaluation metric as we can use the validation split in the form of an MCQ verified evaluation.
Based on this analysis we choose MedMCQA as our target dataset for these experiments. Additionally, since we are training a thinking model and the dataset does not have thinking traces we use the Qwen3-235B model to inject thinking traces into the training samples.
## Forgetting Mitigation Best Practices
### Experimental Setup
#### Dataset Mixing
We tested the impact of how forgetting responds to mixing the base dataset in different ratios with the target dataset. The base dataset here refers to the multi domain SFT dataset we have developed (see our [distillation blog post](https://googlecloudplatform.github.io/vertex-ai-samples/vertex-training-cluster/model_distillation_best_practices) [<a href="#ref2">2</a>] for details of the generation process) that can replicate and on certain metrics beat the public Qwen3 models. The target dataset here refers to the MedMCQA dataset. It is important to understand that in all mixing scenarios where the target dataset is present we will use the complete target dataset as that is the reasonable course of action we expect any customer to take. This leads to the total number of samples used in training varying based on the mixing ratio.
We run 2 baseline experiments for each model size: using only the base dataset and only the target dataset. The mixing experiments are the base dataset being mixed in ratios of 0.9:0.1, 0.75:0.25 and 0.5:0.5. (0.9:0.1 means 90% of samples are from the base in-distribution dataset, while 10% are from the target out-of-distribution dataset.)
The base dataset is randomly subsampled for each of these experiments. For simpler reference and analysis let’s define a mixing ratio <span>$$0 \le \alpha < 1$$</span>, such that the final dataset mixture includes <span>$$ N'_{B} = \frac{\alpha}{1 - \alpha} N_T$$</span> samples from the base dataset where <span>$$N_T$$</span> is the number of samples in the target dataset. In each of these mixtures the complete target dataset is used, contributing <span>$$N_T$$</span> samples for a total training dataset size of <span>$$\frac{N_T}{1 - \alpha}$$</span>.
Since, our target dataset has 182,712 samples, this means that:
- 0.9:0.1 ratio (<span>$$\alpha = 0.9$$</span>) : Uses a total of 1,827,120 training samples
- 0.75:0.25 ratio (<span>$$\alpha = 0.75$$</span>) : Uses a total of 730,849 training samples
- 0.5:0.5 ratio (<span>$$\alpha = 0.5$$</span>) : Uses a total of 365,425 training samples
#### Evaluation
<table id="tab-eval-setup" style="margin-left:auto; margin-right:auto;">
<thead>
<tr>
<th>Capabilities</th>
<th>Benchmarks</th>
<th># Test Samples</th>
<th>Eval Metrics</th>
</tr>
</thead>
<tbody>
<tr>
<td rowspan="7">Math</td>
<td>AIME 24</td>
<td>30</td>
<td>pass@1 (average of 10)</td>
</tr>
<tr>
<td>AIME 25</td>
<td>30</td>
<td>pass@1 (average of 10)</td>
</tr>
<tr>
<td>BeyondAIME</td>
<td>100</td>
<td>pass@1 (average of 5)</td>
</tr>
<tr>
<td>Math 500</td>
<td>500</td>
<td>pass@1</td>
</tr>
<tr>
<td>HMMT 25</td>
<td>30</td>
<td>pass@1 (average of 10)</td>
</tr>
<tr>
<td>BRUMO 25</td>
<td>30</td>
<td>pass@1 (average of 10)</td>
</tr>
<tr>
<td>CMIMC 25</td>
<td>40</td>
<td>pass@1 (average of 10)</td>
</tr>
<tr>
<td rowspan="3">Science</td>
<td>GPQA</td>
<td>448</td>
<td>pass@1 (average of 5)</td>
</tr>
<tr>
<td>MMLU</td>
<td>14042</td>
<td>pass@1</td>
</tr>
<tr>
<td>MMLU Pro</td>
<td>12032</td>
<td>pass@1</td>
</tr>
<tr>
<td rowspan="2">Coding</td>
<td>HumanEval</td>
<td>164</td>
<td>pass@1 (average of 5)</td>
</tr>
<tr>
<td>LiveCodeBench v6</td>
<td>175</td>
<td>pass@1 (average of 5)</td>
</tr>
<tr>
<td>Instruction Following</td>
<td>IFEval</td>
<td>541</td>
<td>pass@1 (Strict Accuracy)</td>
</tr>
<tr>
<td>Reasoning</td>
<td>ARC-AGI 1</td>
<td>400</td>
<td>pass@1 (average of 5)</td>
</tr>
<tr>
<td>Medical (Target domain)</td>
<td>MedMCQA</td>
<td>4183</td>
<td>pass@1</td>
</tr>
</tbody>
<caption style="text-align: left;"><b>Table 2:</b> Comprehensive overview of task domains, evaluation benchmarks, and associated performance metrics.</caption>
</table>
Our evaluation benchmarks and metrics are detailed in [Table 2](#tab-eval-setup). To ensure statistical reliability on smaller datasets, we report metrics averaged over multiple independent runs to mitigate variance. For each domain with multiple evaluations, we utilize the average score across the core benchmarks as our primary performance indicator. To maintain a consistent comparison, both our trained models and the official Qwen3 thinking models were evaluated using standardized sampling parameters — `Temperature=0.6`, `Top-P=0.95`, `Top-K=20` and `Max-tokens=32768` — aligning with the recommended [best practices](https://huggingface.co/Qwen/Qwen3-14B#best-practices) from the official Qwen3 model card.
Note that we have separated MedMCQA as a target metric instead of including it in the Science domain. This is to ensure clear outcomes from our experiments and to demonstrate impacts on model performance without any interference.
#### Training
##### Vertex AI Training Cluster
All experiments and results presented were orchestrated using the [Vertex AI Training Cluster (VTC)](https://docs.cloud.google.com/vertex-ai/docs/training/training-clusters/overview). VTC is a managed Google Cloud service designed to simplify and accelerate large-scale AI workloads. It provides a simple managed user experience that enables optimized GPU scheduling, automated fault tolerance, high hardware resiliency, quick start recipes and science tooling which drastically reduces the time from cluster setup to production training and speeds up experimentation.
##### Training Framework and Hyperparameters
We utilize NVIDIA [NeMo RL](https://github.com/NVIDIA-NeMo/RL), an open library from the [NVIDIA NeMo framework](https://github.com/NVIDIA-NeMo/) as the primary training library, leveraging the Megatron backend for distributed scaling. Models are initialized from a Qwen3 Base checkpoint and fine-tuned with a 32,768 context window on curated datasets. Optimization is handled via AdamW (<span>$$\beta_1=0.9$$</span>, <span>$$\beta_2=0.95$$</span>, weight decay=0.1) using a linear warmup and cosine decay schedule. All training is conducted using BF16 mixed precision. There are many model sizes and dataset mixes used in the experimentation so the maximum learning rate is guided by learning rate scaling laws (see [distillation blog post](https://googlecloudplatform.github.io/vertex-ai-samples/vertex-training-cluster/model_distillation_best_practices#hyperparameter-scaling) [<a href="#ref2">2</a>] for more) available as a part of VTC. The value is validated by testing slight adjustments from the recommended value for each dataset mixture.
### Mitigating Forgetting
All models in this experiment are trained starting from the Qwen3 base checkpoint. We explore the impact of dataset mixing by comparing the public Qwen3 instruction tuned model performance with our two baselines — model trained with only the target dataset and model trained only with the base dataset — and with a model trained using a 0.9 ratio mix.
<figure align="center" id="fig3_data_mixing">
<table align="center" width="100%">
<tr>
<td align="center" width="50%">
<img src="images_data_mixing/fig3_math.png" width="100%"><br>
<sub><b>(a)</b> Math</sub>
</td>
<td align="center" width="50%">
<img src="images_data_mixing/fig3_science.png" width="100%"><br>
<sub><b>(b)</b> Science</sub>
</td>
</tr>
<tr>
<td align="center" width="50%">
<img src="images_data_mixing/fig3_coding.png" width="100%"><br>
<sub><b>(c)</b> Coding</sub>
</td>
<td align="center" width="50%">
<img src="images_data_mixing/fig3_ifeval.png" width="100%"><br>
<sub><b>(d)</b> IFEval</sub>
</td>
</tr>
<tr>
<td align="center" width="50%">
<img src="images_data_mixing/fig3_arc_agi.png" width="100%"><br>
<sub><b>(e)</b> ARC-AGI</sub>
</td>
<td align="center" width="50%">
<img src="images_data_mixing/fig3_medmcqa.png" width="100%"><br>
<sub><b>(f)</b> MedMCQA</sub>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 3: Performance with and without Data Mixing.</b> <i>A comparison across (a) Math, (b) Science, (c) Coding, (d) IFEval, (e) ARC-AGI and (f) MedMCQA benchmarks showing how data mixing impacts forgetting and performance on the target metric.</i></sub>
</figcaption>
</figure>
[Figure 3](#fig3_data_mixing) shows that for all non-target metrics other than Science using just the target dataset shows significant forgetting. Math and ARC-AGI are almost completely forgotten for all model sizes up to 8B parameters. The mixed dataset recovers the performance to similar levels as the base dataset. The base dataset delivers performance comparable to the public model in all domains and significantly better on ARC-AGI.
The Science domain evaluations do not suffer severe forgetting likely because MedMCQA is very close to this domain. In fact, for the 8B and 14B sizes due to these transfer learning dynamics the <span>$$\alpha = 0.9$$</span> model outperforms both the public instruction-tuned and the base dataset (<span>$$\alpha = 1$$</span>) models.
Performance on the target metric of MedMCQA follows expected behavior with best results achieved by the model when trained only with the target dataset. It is important to note that the <span>$$\alpha = 0.9$$</span> model for all sizes is still significantly better than the public instruction-tuned and base dataset (<span>$$\alpha = 1$$</span>) model and for all sizes other than the 0.6B mostly maintains the performance gains of the target dataset (<span>$$\alpha = 0$$</span>) model.
#### Key Observations
Combining these conclusions we can see that mixing with our base dataset:
- Matches and outperforms the public instruction tuned model on general tasks.
- Preserves the gains beyond the public model on target tasks.
- Provides additional gains on tasks from a similar domain.
### Mixing Ratios
Now that we know that mixing the base dataset almost eliminates forgetting it is important to understand how performance changes for different mixing configurations. This is also important to examine as it determines training length and hence the cost. We will compare models trained only with the target dataset to models trained using dataset mixes with <span>$$\alpha = 0.5, 0.75, 0.9$$</span>. The ratio mentioned here refers to the proportion of the dataset from the base dataset.
<figure align="center" id="fig4_mixing_ratios">
<table align="center" width="100%">
<tr>
<td align="center" width="50%">
<img src="images_data_mixing/fig4_math.png" width="100%"><br>
<sub><b>(a)</b> Math</sub>
</td>
<td align="center" width="50%">
<img src="images_data_mixing/fig4_science.png" width="100%"><br>
<sub><b>(b)</b> Science</sub>
</td>
</tr>
<tr>
<td align="center" width="50%">
<img src="images_data_mixing/fig4_coding.png" width="100%"><br>
<sub><b>(c)</b> Coding</sub>
</td>
<td align="center" width="50%">
<img src="images_data_mixing/fig4_ifeval.png" width="100%"><br>
<sub><b>(d)</b> IFEval</sub>
</td>
</tr>
<tr>
<td align="center" width="50%">
<img src="images_data_mixing/fig4_arc_agi.png" width="100%"><br>
<sub><b>(e)</b> ARC-AGI</sub>
</td>
<td align="center" width="50%">
<img src="images_data_mixing/fig4_medmcqa.png" width="100%"><br>
<sub><b>(f)</b> MedMCQA</sub>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 4: Performance across Mixing Ratios.</b> <i>A comparison across (a) Math, (b) Science, (c) Coding, (d) IFEval, (e) ARC-AGI and (f) MedMCQA benchmarks showing how dataset mixing ratios impact forgetting and performance on the target metric.</i></sub>
</figcaption>
</figure>
[Figure 4](#fig4_mixing_ratios) shows that for all non-target metrics mixing helps achieve better performance than just using the target dataset even with a <span>$$\alpha = 0.5$$</span> mix. As expected the performance on non target metrics worsens as we lower the ratio of the base dataset. This effect is more pronounced in the smaller size models and for datasets like ARC-AGI where the mixed training provides a lot more gain. These patterns confirm that the gains on non target metrics are directly correlated to the base dataset.
The effect while present for Science domain metrics is much less pronounced due to the cross domain characteristics. Even with lower ratios the performance for models 4B and larger holds, confirming that our target dataset of MedMCQA here contributes to limiting forgetting for this domain.
The performance on the target metric, MedMCQA, stays mostly consistent with dips mostly when going from <span>$$\alpha = 0.75$$</span> mix to <span>$$\alpha = 0.5$$</span> mix. This aligns well as in all cases we are doing a complete epoch on the target dataset. The performance mostly holding at mixing ratios indicates that the tradeoff on the target metrics is relatively low even at an aggressive mixing ratio like 0.5.
#### Key Observations
The mixing ratio comparison shows us that:
- A mixing ratio of 0.9 is the best for achieving gains on target tasks and limiting forgetting.
- A mixing ratio of even 0.5 limits forgetting well while only doubling the token budget compared to training without any mixing.
### Different Starting Models
We have trained all our models starting from Qwen3 base checkpoints. A natural question here might be: What happens if we train starting from the instruction tuned public Qwen3 checkpoints for our target task? In this section we examine this question and compare the instruction-tuned model tuned with the target dataset and an <span>$$\alpha = 0.9$$</span> mix to the instruction-tuned model itself and the base model tuned with an <span>$$\alpha = 0.9$$</span> mix.
<figure align="center" id="fig5_starting_models">
<table align="center" width="100%">
<tr>
<td align="center" width="50%">
<img src="images_data_mixing/fig5_math.png" width="100%"><br>
<sub><b>(a)</b> Math</sub>
</td>
<td align="center" width="50%">
<img src="images_data_mixing/fig5_science.png" width="100%"><br>
<sub><b>(b)</b> Science</sub>
</td>
</tr>
<tr>
<td align="center" width="50%">
<img src="images_data_mixing/fig5_coding.png" width="100%"><br>
<sub><b>(c)</b> Coding</sub>
</td>
<td align="center" width="50%">
<img src="images_data_mixing/fig5_ifeval.png" width="100%"><br>
<sub><b>(d)</b> IFEval</sub>
</td>
</tr>
<tr>
<td align="center" width="50%">
<img src="images_data_mixing/fig5_arc_agi.png" width="100%"><br>
<sub><b>(e)</b> ARC-AGI</sub>
</td>
<td align="center" width="50%">
<img src="images_data_mixing/fig5_medmcqa.png" width="100%"><br>
<sub><b>(f)</b> MedMCQA</sub>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 5: Performance across Starting Models.</b> <i>A comparison across (a) Math, (b) Science, (c) Coding, (d) IFEval, (e) ARC-AGI and (f) MedMCQA benchmarks showing how different starting models impact forgetting and performance on the target metric. Qwen 3 Public is the public instruction tuned Qwen3 model, &alpha;=0 (IT) and &alpha;=0.9 (IT) are the public instruction-tuned Qwen3 model trained only with the target dataset and the &alpha;=0.9 mixed dataset. &alpha;=0.9 (Base) is the base Qwen3 model trained on a 90% VTC dataset and 10% target dataset mix.</i></sub>
</figcaption>
</figure>
In [Figure 5](#fig5_starting_models), among the non-target metrics other than science we see a common trend that starting with the IT model and using only the target dataset (<span>$$\alpha = 0$$</span>) shows severe forgetting. The base model and the instruction-tuned model trained using the <span>$$\alpha = 0.9$$</span> mix match or surpass the performance of the public model. This shows that starting with an instruction-tuned model while better than starting with the base model is still not a solution to forgetting. This also shows the high quality of our dataset that it can provide further gains on the public instruction-tuned model.
Science domain metrics show different trends based on the model size. The advantage of data mixing is much more apparent in 0.6B and 1.7B models. Overall though there are no disadvantages to mixing across all model sizes. The IT model demonstrating significant forgetting is a clear indication that cross domain characteristics of our target dataset are not enough to mitigate forgetting on its own.
The performance of the target metric, MedMCQA, shows no additional gain when we train using only the target dataset except for the 0.6B model, whether the starting model is a base model or the IT model. For all model sizes other than the 0.6B model we also see that the <span>$$\alpha = 0.9$$</span> mix trained model does not lose any meaningful performance compared to the target dataset only trained models. All the models trained using the target dataset clearly improve on the public model.
#### Key Observations
The comparison of different starting models shows us:
- Using the instruction-tuned model as the starting model is better than the Base model.
- The IT model also shows catastrophic forgetting and loses performance on non target metrics.
- The <span>$$\alpha = 0.9$$</span> mix avoids forgetting even with the instruction-tuned starting model showing its robustness.
## Acknowledgements
We would like to express our sincere gratitude to the NVIDIA NeMo RL team–specifically Terry Kong– for their invaluable support throughout this project.
We would also like to express our gratitude to our VTC teammates: Mohammadreza Mohseni, Weiran Zhao, Fei Xia, Youbao Tang, Xuehan Xiong, Joseph Pagadora, Jiuqiang Tang, Bo Wu, Lav Rai, and Minwoo Park for developing the underlying datasets, providing infrastructure support, feedback, and insightful discussions throughout the project. We also thank Ting Yu, Shengyang Dai, Peng Xu, and Saurabh Tiwary for their leadership and support.
## References
<a id="ref1"></a>[1] McCloskey, Michael, and Neal J. Cohen. "Catastrophic interference in connectionist networks: The sequential learning problem." Psychology of learning and motivation. Vol. 24. Academic Press, 1989. 109-165.
<a id="ref2"></a>[2] Google Cloud. "Model Distillation Best Practices." Vertex AI Training Cluster Samples. Google, 2026. https://googlecloudplatform.github.io/vertex-ai-samples/vertex-training-cluster/model_distillation_best_practices.
Binary file not shown.

After

Width:  |  Height:  |  Size: 9.0 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 5.6 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 48 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 259 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 295 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 40 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 8.1 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 190 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 189 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 154 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 195 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 195 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 191 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 184 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 152 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 167 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 158 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 177 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 180 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 193 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 229 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 230 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 232 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 232 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 234 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 241 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 231 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 219 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 202 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 234 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 270 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 230 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 262 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 211 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 209 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 189 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 215 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 210 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 212 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 206 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 190 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 184 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 194 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 213 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 214 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 219 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 112 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 184 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 108 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 108 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 115 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 111 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 89 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 92 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 86 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 93 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 93 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 94 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 75 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 76 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 87 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 91 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 100 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 100 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 82 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 80 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 49 KiB

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,541 @@
<script type="text/javascript" async
src="https://cdn.mathjax.org/mathjax/latest/MathJax.js?config=TeX-MML-AM_CHTML">
</script><br><br>
# Model Distillation Best Practices
**Authors:** [Xuehan Xiong](mailto:xxman@google.com), [Youbao Tang](mailto:tangyoubao@google.com), [Fei Xia](mailto:feixia@google.com), [Bao Thach](mailto:baothach@google.com), [Joseph Pagadora](mailto:jcpagadora@google.com)
## Table of Contents
* [Intro](#intro)
* [Background](#background)
* [Dataset Curation](#dataset-curation)
* [Non-Agentic Tasks](#non-agentic-tasks)
* [Agentic Task (Tool Utilization)](#agentic-task-tool-utilization)
* [Model Distillation Experiments](#model-distillation-experiments)
* [Experimental Setup](#experimental-setup)
* [Choosing Teacher Models](#choosing-teacher-models)
* [Number of Rollouts per Prompt](#number-of-rollouts-per-prompt)
* [Rejection Sampling](#rejection-sampling)
* [Hyperparameter Scaling](#hyperparameter-scaling)
* [Scaling learning rates based on global batch size](#scaling-learning-rates-based-on-global-batch-size)
* [Scaling learning rates based on model parameters](#scaling-learning-rates-based-on-model-parameters)
* [Scaling learning rates based on token budget](#scaling-learning-rates-based-on-token-budget)
* [Key Takeaways](#key-takeaways)
* [Acknowledgements](#acknowledgements)
* [Appendix](./appendix.md#appendix)
## Intro
Welcome to the inaugural installment of our blog series dedicated to model training best practices for Vertex AI Training Cluster (VTC) customers. In this article, we examine model distillation—a popular cost-effective methodology for optimizing student models by leveraging the intelligence of high-capacity teacher models. Two primary distillation schemes are typically employed: **on-policy** and **off-policy distillation**. In an on-policy setting, the student model generates its own reasoning traces during training, which are then evaluated or corrected by a teacher model in real-time. While effective, this approach is computationally intensive and requires constant active inference from the teacher.
This blog focuses on **off-policy distillation**, a resource-efficient methodology where the student model is trained on a static, "gold-standard" dataset of reasoning traces previously curated by a teacher.
While online distillation often requires complex orchestration—like offloading the student to CPU while the teacher scores the trajectories to save VRAM—the off-policy approach simplifies the workflow by **completely decoupling generation from training.** By leveraging frontier-level models like [Qwen3-235B](https://huggingface.co/Qwen/Qwen3-235B-A22B-Thinking-2507) or [GLM-4.7 355B](https://huggingface.co/zai-org/GLM-4.7-FP8) to generate high-quality trajectories upfront, developers can:
* **Max out GPU Utilization:** Dedicate 100% of available VRAM and compute to the student's training phase without the overhead of model swapping.
* **Scale Independently:** Generate datasets once and reuse them for multiple student architectures or hyperparameter sweeps.
* **Simplify Orchestration:** Eliminate the need for multi-model memory management, allowing for a standard, high-throughput Supervised Fine-Tuning (SFT) pipeline.
This allows developers to achieve "big model" reasoning logic in smaller, deployable students without the logistical headache of maintaining a live teacher-student link.
<figure align="center" id="fig-teaser">
<table align="center" width="90%">
<tr>
<td align="center" width="50%">
<img src="images/teaser_arc_agi.png" width="100%"><br>
<sub><b>(a)</b> ARC-AGI 1</sub>
</td>
<td align="center" width="50%">
<img src="images/teaser_tau2.png" width="100%"><br>
<sub><b>(b)</b> &tau;<sup>2</sup>-bench</sub>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 1: Distilled Student Model Performance on Novel Domains.</b> <i>Comparison of student models fine-tuned via off-policy distillation versus official Qwen3 post-trained models of equivalent scale, showing significant performance gains on ARC-AGI 1 and &tau;<sup>2</sup>-bench.</i></sub>
</figcaption>
</figure>
We provide a rigorous, step-by-step framework for reproducing the advanced reasoning capabilities of the Qwen3 open-weight thinking model family, beginning from their base pre-trained checkpoints. Furthermore, we demonstrate how this same distillation pipeline can be applied to novel domains, such as [ARC-AGI 1](https://arcprize.org/arc-agi/1/) and [<span>$$\tau^2$$</span>-bench](https://github.com/sierra-research/tau2-bench), to develop student models that surpass the performance of official Qwen3 variants of equivalent scale (See [Figure 1](#fig-teaser)). The Qwen3 family was specifically chosen for this study because its diverse range of parameter counts and the availability of both pre-trained and post-trained checkpoints provide an ideal environment for high-fidelity scaling analysis.
To ensure the broad applicability of our findings, we conducted rigorous evaluations across four distinct task domains: competitive mathematics, instruction following, complex puzzle-solving, and tool utilization. We further validated these results across four model scales—1.7B, 4B, 8B, and 14B parameters—to demonstrate that our methodology remains consistent as model complexity increases.
To support our VTC community in accelerating their own development, all code, datasets, and experiment configurations used in this blog are being made available for use in your training workloads.
## Background
To establish a mathematical foundation for our experiments, we first delineate the key differences between the two predominant strategies in model distillation: off-policy distillation and on-policy distillation.
In off-policy distillation for an LLM, we assume:
* <span>$$z \sim P(Z)$$</span>: prompts drawn from some prompt distribution
* <span>$$x \sim P(X|z)$$</span>: responses generated by the teacher distribution conditioned on prompt <span>$$z$$</span>
* <span>$$Q(X|z)$$</span>: student distribution we want to train
The standard off-policy distillation objective minimizes the KL divergence from the student to the teacher:
$$\begin{align*}
& \min_Q \mathbb{E}_{Z} \left[\mathbb{E}_{X|Z}\left[\log \frac{P(X|Z)}{Q(X|Z)}\right]\right] \\
&= \min_Q \sum_{z} P(z) \sum_{x}P(x|z) \log\frac{P(x|z)}{Q(x|z)}
\end{align*} $$
Since the teacher distribution (<span>$$P$$</span>) is fixed, minimizing KL is equivalent to:
$$\begin{align*}
& \max_Q \mathbb{E}_{Z}\left[\mathbb{E}_{X|Z}\left[\log Q(X|Z)\right]\right] \\
&= \max_Q \sum_{z} P(z) \sum_{x}P(x|z) \log Q(x|z)
\end{align*} $$
Given a dataset of prompts and teacher-generated responses:
$$\begin{align*}
\{z_i\}_{i=1}^M \quad & \text{where} \quad z_i \sim P(Z) \\
\{x_{ij}\}_{i,j=1}^{M,N} \quad & \text{where} \quad x_{ij} \sim P(X | z_i)
\end{align*} $$
the empirical loss by Monte Carlo sampling becomes:
$$L_{\text{distill}}(Q) \approx -\frac{1}{MN}\sum_{i=1}^M\sum_{j=1}^N \log Q(x_{ij} | z_i). $$
This is simply maximum likelihood estimation for <span>$$Q$$</span> on teacher responses, which shares the same objective as Supervised Fine-tuning (SFT).
In on-policy distillation, responses are sampled from the student:
$$z \sim P(Z), \quad x \sim Q(X | z) $$
The teacher is only used to evaluate those student samples, so expectations are taken under <span>$$Q$$</span>, not <span>$$P$$</span>.
The natural objective is:
$$\min_Q \mathbb{E}_{Z} \left[\mathrm{KL}\big(Q(X|Z)||P(X|Z)\big) \right] $$
This objective defines the **reverse KL divergence**, which is characterized by its mode-seeking behavior. In this regime, the student model tends to concentrate its probability mass on the primary modes of the teacher distribution. This stands in contrast to the forward KL divergence used in off-policy distillation, which exhibits mean-seeking or mass-covering behavior, forcing the student to cover the entire support of the teacher’s distribution. Forward KL forces the student to allocate probability mass to *all* teacher modes, even those it cannot represent well. Under capacity constraints, this mass-covering approach produces a compromise distribution that can underperform a smaller, sharper target. This provides the intuition for our empirical results on Capacity Matching (Section [Choosing Teacher Models](#choosing-teacher-models)), where a same-sized teacher model proved most effective in some tasks.
## Dataset Curation
To facilitate the distillation of frontier-level reasoning, we established a high-fidelity data curation pipeline tailored to our four primary task domains: competitive mathematics, instruction following, reasoning/pattern recognition, and agentic tool utilization.
### Non-Agentic Tasks
#### Math
We selected [OpenR1-Math](https://huggingface.co/datasets/open-r1/OpenR1-Math-220k) (default subset) as our primary prompt source and implemented a multi-stage filtering pipeline to ensure the highest data fidelity:
1. **Verification filtering**: To ensure objective evaluation, we retained only "math-word-problem" types, discarding Multiple Choice Questions (MCQ) and proofs. MCQs were specifically excluded to mitigate the risk of the model arriving at a correct answer (25% baseline probability) through flawed reasoning chains.
2. **Near-duplicate removal**: We employed Locality-Sensitive Hashing (LSH) to identify and prune near-identical prompts within the training set and across our evaluation benchmarks, preventing data contamination and overfitting.
3. **Instructional sanitization**: We identified and removed hundreds of prompts containing extraneous translation instructions. This step ensures the student model remains focused on the mathematical reasoning task rather than defaulting to secondary objectives.
4. **Solution leakage prevention**: To enforce authentic problem-solving, we stripped prompts containing pre-existing solutions, which would otherwise provide the teacher model with an "open-book" advantage and degrade the quality of the distilled reasoning traces.
This rigorous curation process successfully refined the initial pool of 93,733 candidates into a high-quality dataset of 75,726 prompts.
#### Instruction Following
We selected the default partition of the [ifeval-like-data](https://huggingface.co/datasets/argilla/ifeval-like-data) dataset, comprising 550,000 unfiltered synthetic rows. To ensure data integrity, we applied a multi-stage refinement pipeline:
1. **Invalid sample pruning**: We discarded rows with missing language codes or malformed JSON within the "kwargs" field to maintain structural consistency.
2. **Conflict resolution**: We identified and removed pairs of mutually exclusive instructions that cannot be reliably evaluated together, utilizing a predefined mapping of instruction conflicts (`IFEVAL_INSTRUCTION_CONFLICTS`).
3. **Adherence verification**: Each teacher-generated response was rigorously assessed using the [lm_eval](https://github.com/EleutherAI/lm-evaluation-harness) library. Prompts that failed to meet their defined constraints were excluded.
4. **Strict accuracy filtering**: As a final quality gate, we retained only those samples where the response achieved "strict accuracy" at the prompt level, ensuring the student model learns from perfect examples of instruction following.
5. **LSH-based deduplication**: We utilized LSH to prune near-duplicate prompts within the training set and across the official IFEval benchmark to prevent contamination.
This pipeline successfully distilled the initial pool into 70,373 high-fidelity samples for our instruction-following training set.
#### Reasoning/Pattern Recognition
The [re-arc](https://github.com/michaelhodel/re-arc?tab=readme-ov-file) repository provides a way to programmatically synthesize [ARC-AGI-1](https://arcprize.org/arc-agi/1/) data. For each of the 400 training examples in the official ARC-AGI-1 dataset, re-arc provides a generator function to create similar puzzles following the same pattern (See [Figure 2](#fig-rearc-example) for one example). In total, we have generated 7926 puzzles for our experiments where we reserve 256 samples for validation and the rest for training.
<figure align="center" id="fig-rearc-example">
<table align="center" width="80%">
<tr>
<td align="center" width="100%">
<img src="images/arc_agi_original.png" width="100%"><br>
<sub><b>(a)</b> ARC-AGI original puzzles</sub>
</td>
</tr>
<tr>
<td align="center" width="510%">
<img src="images/arc_agi_generated.png" width="100%"><br>
<sub><b>(b)</b> Generated puzzles using re-arc</sub>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 2:</b> <i>Example of ARC-AGI 1 puzzle synthesis using the re-arc repository, showing an original training example and a generated similar puzzle following the same pattern.</i></sub>
</figcaption>
</figure>
#### Response Generation
Following prompt collection, we utilize a "thinking" teacher model to generate reasoning traces. To maintain a lean data pipeline, we store only the sampled tokens; given our ~150K vocabulary size, persisting full logits or log-probabilities would create prohibitive storage overhead. This approach is statistically grounded: as the number of samples increases, it provides an unbiased estimate of the KL divergence from the student model to the teacher model.
To maximize response diversity, we set both `Temperature` and `Top-P` to 1.0 during sampling. Finally, we prune any responses truncated by the maximum sequence length, as these instances frequently exhibit repetitive patterns that could degrade the student model’s performance.
### Agentic Task (Tool Utilization)
We utilized a specialized two-stage generation framework to synthesize high-complexity tool-use data for distillation training, leveraging sandboxed execution environments.
1. **Task Generation**: This initial phase analyzes a target agent's specific tool list to propose a variety of diverse, high-level topics. For each identified topic, the system synthesizes a comprehensive user scenario that includes the initial environment status, the necessary database state, and precise evaluation criteria required for verification.
2. **Trajectory Generation and Verification**: A verified <span>$$\tau^2$$</span>-bench sandbox is employed to execute each generated task several times in parallel, capturing a wide variety of trajectories. These execution outputs—comprising model responses, tool invocations, and subsequent state modifications—undergo a rigorous verification process. By applying deterministic checks such as action matching, database state differentials, and natural language assertions, the pipeline calculates the reward for every trajectory produced.
Once trajectories are verified, they are carefully remapped into final training configurations to maximize learning efficiency. For distillation, the complete reasoning trace is captured and enclosed within required `<think>` tags, ensuring architectural consistency for the thinking model. A sample task and trajectory are provided in the [Appendix: <span>$$\tau^2$$</span>-bench Synthetic Example](./appendix.md#-bench-synthetic-example).
## Model Distillation Experiments
### Experimental Setup
#### Evaluation
Our evaluation benchmarks and metrics are detailed in [Table 1](#tab-eval-setup) and the prompts can be found in the [Appendix: Prompts Used in Evaluation](./appendix.md#prompts-used-in-evaluation). To ensure statistical reliability on smaller datasets, we report metrics averaged over multiple independent runs to mitigate variance. For the Mathematics domain, we utilize the average score across six core benchmarks as our primary performance indicator, while granular results for individual benchmarks are provided in the [Appendix: Individual Math Benchmark Results](./appendix.md#individual-math-benchmark-results). For <span>$$\tau^2$$</span>-bench, we use GLM-4.7-FP8 as the user LLM and report the average score across three domains, “Telecom”, “Retail”, and “Airline”. To maintain a consistent comparison, both our distilled student models and the official Qwen3 thinking models were evaluated using standardized sampling parameters—`Temperature=0.6`, `Top-P=0.95`, and `Top-K=20`—aligning with the recommended [best practice](https://huggingface.co/Qwen/Qwen3-14B#best-practices) from the official Qwen3 model card.
<table id="tab-eval-setup" style="margin-left:auto; margin-right:auto;">
<thead>
<tr>
<th>Capabilities</th>
<th>Benchmarks</th>
<th># Test Samples</th>
<th>Eval Metrics</th>
</tr>
</thead>
<tbody>
<tr>
<td rowspan="6">Math</td>
<td>AIME 24</td>
<td>30</td>
<td>pass@1 (average of 16)</td>
</tr>
<tr>
<td>AIME 25</td>
<td>30</td>
<td>pass@1 (average of 16)</td>
</tr>
<tr>
<td>BeyondAIME</td>
<td>100</td>
<td>pass@1 (average of 5)</td>
</tr>
<tr>
<td>HMMT 25</td>
<td>30</td>
<td>pass@1 (average of 16)</td>
</tr>
<tr>
<td>BRUMO 25</td>
<td>30</td>
<td>pass@1 (average of 16)</td>
</tr>
<tr>
<td>CMIMC 25</td>
<td>40</td>
<td>pass@1 (average of 16)</td>
</tr>
<tr>
<td>Instruction Following</td>
<td>IFEval</td>
<td>541</td>
<td>pass@1 (Strict Accuracy)</td>
</tr>
<tr>
<td>Reasoning</td>
<td>ARC-AGI 1</td>
<td>400</td>
<td>pass@1 (average of 5)</td>
</tr>
<tr>
<td>Tool use</td>
<td>&tau;<sup>2</sup>-bench</td>
<td>278</td>
<td>pass@1 (average of 4)</td>
</tr>
</tbody>
<caption style="text-align: left;"><b>Table 1:</b> Comprehensive overview of task domains, evaluation benchmarks, and associated performance metrics.</caption>
</table>
#### Training
**Vertex AI Training Cluster**
To orchestrate the computational demands of our experiments, we utilized the [Vertex AI Training Cluster (VTC)](https://docs.cloud.google.com/vertex-ai/docs/training/training-clusters/overview). VTC is a managed Google Cloud service designed to simplify and accelerate large-scale AI workloads. It provides a familiar, open-source Slurm user experience that enables optimized GPU scheduling, automated fault tolerance, and high hardware resiliency, which drastically reduces the time from cluster setup to production training.
Our training infrastructure leverages VTC's high-performance [compute resources](https://docs.cloud.google.com/vertex-ai/docs/training/training-clusters/compute-resources), specifically utilizing A3-Mega (NVIDIA H100 GPUs), A3-Ultra (NVIDIA H200 GPUs), and A4 GPU (NVIDIA HGX B200) platforms powered by NVIDIA. To handle the communication overhead of distributed training, node connectivity is highly optimized for each hardware generation:
* A3-Mega clusters utilize [GPUDirect-TCPXO](https://docs.cloud.google.com/compute/docs/gpus/gpudirect) for low-latency host-bypass networking
* A3-Ultra and A4 clusters leverage [RoCE v2 (RDMA over Converged Ethernet)](https://cloud.google.com/blog/products/networking/rdma-rocev2-for-ai-workloads-on-google-cloud).
By leveraging network topologies specifically optimized for training on large clusters of GPUs, this environment provides the high throughput and scaling efficiency necessary to reliably train and finetune frontier-level models.
**Training Framework and Hyperparameters**
We utilize NVIDIA [NeMo RL](https://github.com/NVIDIA-NeMo/RL) as the primary training framework, leveraging a Megatron backend for distributed scaling. We implement the <span>$$\tau^2$$</span>-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 seamlessly integrated with the NeMo RL library for RL training runs.
Models are initialized from a Qwen3 Base checkpoint and fine-tuned with a 32,768 context window on curated datasets. Optimization is handled via AdamW (<span>$$\beta_1=0.9$$</span>, <span>$$\beta_2=0.95$$</span>, weight decay=0.1) using a linear warmup and cosine decay schedule. To manage computational load, we employ tensor parallelism (2-way for 1.7B–8B models; 4-way for 14B) alongside sequence parallelism, activation checkpointing, and ZeRO-2. For further reading on parallelization strategies, see this [ultrascale playbook](https://huggingface.co/spaces/nanotron/ultrascale-playbook). All training is conducted using BF16 mixed precision.
### Choosing Teacher Models
A critical decision in the distillation pipeline is the selection of an appropriate teacher model for a given task domain. While conventional wisdom often suggests that "bigger is better," our empirical results across four benchmarks (illustrated in [Figure 3](#fig-teacher-student-matrix)) reveal a more nuanced landscape. Specifically, on well-defined reasoning tasks like Mathematics and IFEval, capacity-matched (same-sized) teachers frequently outperform their larger counterparts. Conversely, on novel or highly complex domains like ARC-AGI and <span>$$\tau^2$$</span>-bench, massive teacher models remain the superior choice. Below, we provide a formal derivation to explain this capacity-matching phenomenon and the trade-offs between approximation bias and teacher error.
<figure align="center" id="fig-teacher-student-matrix">
<table align="center" width="100%">
<tr>
<td align="center" width="50%">
<img src="images/teacher_student_matrix_ifeval.png" width="100%"><br>
<sub><b>(a)</b> IFEval</sub>
</td>
<td align="center" width="50%">
<img src="images/teacher_student_matrix_math.png" width="100%"><br>
<sub><b>(b)</b> Math Average</sub>
</td>
</tr>
<tr>
<td align="center" width="50%">
<img src="images/teacher_student_matrix_arc_agi1.png" width="100%"><br>
<sub><b>(c)</b> ARC-AGI 1</sub>
</td>
<td align="center" width="50%">
<img src="images/teacher_student_matrix_tau2.png" width="100%"><br>
<sub><b>(d)</b> &tau;<sup>2</sup>-bench</sub>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 3: Teacher-Student Distillation Performance Matrices.</b> <i>A comparison of distillation outcomes across (a) IFEval, (b) Math, (c) ARC-AGI 1, and (d) &tau;<sup>2</sup>-bench benchmarks. The color intensity represents accuracy (%).</i></sub>
</figcaption>
</figure>
Let <span>$$Q$$</span> and <span>$$P$$</span> denote the student and teacher distribution, and the student model class (e.g., all 8B models following the same architecture) is <span>$$\mathbf{Q}_{8B}$$</span>.
**Optimal teacher model in theory**
Suppose the student is an 8B model. The best teacher model is the optimal 8B model <span>$$Q^*$$</span> where <span>$$Q^* \in \mathbf{Q}_{8B}$$</span> for this task because if we have infinite data and perfect optimization, distillation can recover this optimal model exactly.
**Why a larger teacher can still help**
Now suppose the teacher is a larger model (<span>$$P_L$$</span>). The student solves
$$Q^* = \text{argmin}_{Q\in\mathbf{Q}_{8B}} \mathrm{KL}(P_L|Q)$$
This is the best 8B approximation of the larger teacher. Two competing effects appear:
1. **Approximation bias**
Because the student is capacity-limited,
$$\inf_{Q \in \mathbf{Q}_{8B}} \mathrm{KL}(P_L|Q) > 0$$
So distillation from a very rich teacher may force the student to approximate a distribution it cannot represent well. This is the "mean-seeking / mass-covering" mentioned in the [Background](#background) section.
2. **Teacher suboptimality**
In practice we rarely have <span>$$Q^*$$</span>, the true optimal 8B model. Instead we have a trained 8B model <span>$$\hat{Q}$$</span>, which contains optimization error and data error. A larger teacher (<span>$$P_L$$</span>) may actually be closer to the true distribution (<span>$$P^*$$</span>).
If
$$\mathrm{KL}(P^*|P_L) < \mathrm{KL}(P^*|\hat{Q})$$
then projecting (<span>$$P_L$$</span>) onto the 8B class can produce a better 8B model than the original 8B model.
When an 8B student uses an 8B teacher, the projection error is inherently small due to matched capacity. However, if that 8B teacher is poorly optimized (e.g., on benchmarks like ARC-AGI and <span>$$\tau^2$$</span>-bench), its high **Teacher Error** dominates. Conversely, a massive model like Qwen3-235B or GLM 4.7, even if it has a higher **Projection Error** due to the size difference, can significantly lower the **Teacher Error** because it holds a more accurate approximation of the true distribution.
### Number of Rollouts per Prompt
We employ Monte Carlo sampling to approximate KL divergence, where expanding either the prompt set or the number of rollouts per prompt serves to reduce the variance of the estimate. However, high-fidelity prompts are often a finite resource—particularly in the Mathematics domain, which is constrained by the historical volume of competitive math problems. To compensate, we increase the number of teacher rollouts per prompt. This approach captures the teacher model’s **inherent uncertainty and multi-modal behavior** (e.g., discovering multiple valid reasoning paths to the same solution), enabling the student to map the full probability landscape rather than converging on a single, isolated trajectory.
<figure align="center" id="performance_vs_rollouts">
<table align="center" width="100%">
<tr>
<td align="center" width="50%">
<img src="images/response_scaling_ifeval.png" width="100%"><br>
<sub><b>(a)</b> IFEval</sub>
</td>
<td align="center" width="50%">
<img src="images/response_scaling_math.png" width="100%"><br>
<sub><b>(b)</b> Math Average</sub>
</td>
</tr>
<tr>
<td align="center" width="50%">
<img src="images/response_scaling_arc_agi.png" width="100%"><br>
<sub><b>(c)</b> ARC-AGI 1</sub>
</td>
<td align="center" width="50%">
<img src="images/response_scaling_tau2.png" width="100%"><br>
<sub><b>(d)</b> &tau;<sup>2</sup>-bench</sub>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 4: Performance vs. Number of Rollouts.</b> <i>A comparison across (a) IFEval, (b) Math, (c) ARC-AGI, and (d) &tau;<sup>2</sup>-bench benchmarks showing how performance scales as the number of teacher rollouts per prompt increases from 1 to 16.</i></sub>
</figcaption>
</figure>
[Figure 4](#performance_vs_rollouts) demonstrates that performance on the Math, ARC-AGI 1, and &tau;<sup>2</sup>-bench domains monotonically increases as the number of rollouts increases, while IFEval performance saturates at 8 rollouts.
### Rejection Sampling
This section explores whether rejection sampling on teacher responses—specifically, pruning trajectories that yield incorrect answers—enhances distillation performance. Formally, this approach **minimizes the KL divergence** against a **reweighted teacher distribution**, where incorrect paths are zero-weighted and valid paths are renormalized. [Figure 5](#fig_rejection_sampling) evaluates three strategies: utilizing the full response set, randomly subsampling to match the count of correct responses, and isolating correct responses only. Detailed acceptance rates for these task-teacher pairings are cataloged in [Table 2](./appendix.md#acceptance-rates-for-rejection-sampling) within the Appendix.
<figure align="center" id="fig_rejection_sampling">
<table align="center" width="100%">
<tr>
<td align="center" width="50%">
<img src="images/rejection_sampling_comparison_ifeval.png" width="100%"><br>
<sub><b>(a)</b> IFEval</sub>
</td>
<td align="center" width="50%">
<img src="images/rejection_sampling_comparison_math.png" width="100%"><br>
<sub><b>(b)</b> Math Average</sub>
</td>
</tr>
<tr>
<td align="center" width="50%">
<img src="images/rejection_sampling_comparison_arcagi.png" width="100%"><br>
<sub><b>(c)</b> ARC-AGI 1</sub>
</td>
<td align="center" width="50%">
<img src="images/rejection_sampling_comparison_tau2.png" width="100%"><br>
<sub><b>(d)</b> &tau;<sup>2</sup>-bench</sub>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 5: Performance with and without Rejection Sampling.</b> <i>A comparison across (a) IFEval, (b) Math, (c) ARC-AGI, and (d) &tau;<sup>2</sup>-bench benchmarks showing how distillation performance is impacted by the use of rejection sampling during data curation.</i></sub>
</figcaption>
</figure>
For instruction following (IFEval), rejection sampling provides a significant performance uplift within the same-sample-count regime. Notably, at the 14B model scale, this technique improves the student model by 2 percentage points compared to training on the full response set, despite utilizing a smaller volume of data. Conversely, for the Mathematics, ARC-AGI 1, <span>$$\tau^2$$</span>-bench domains, rejection sampling does not yield performance gains, with the highest accuracy achieved by utilizing all available teacher responses. This suggests that while rejection sampling refines the training distribution, it may also inadvertently prune "near-miss" cases or highly challenging problems that are essential for developing robust reasoning capabilities in those specific domains. For synthetic datasets like <span>$$\tau^2$$</span>-bench, imperfections in automated evaluation criteria may also cause the rejection of trajectories that contain high-quality reasoning traces despite an incorrect final answer.
## Hyperparameter Scaling
In practice, development-phase experimentation rarely mirrors the scale of final model training. To accelerate iteration, developers often conduct ablation studies using reduced token budgets and smaller model architectures. To help bridge this gap, we have compiled several **key rules of thumb** for translating hyperparameters from these 'proxy' settings to your final, full-scale production runs. All scaling studies below are conducted using the [IFEval-like](#instruction-following) dataset.
### Scaling learning rates based on global batch size
To accelerate training throughput, the most direct lever is scaling GPU resources and increasing the **Global Batch Size (B)**. However, effective scaling requires more than just hardware; the learning rate (<span>$$\eta$$</span>) must be precisely adjusted in tandem with the batch size to maintain an optimal convergence trajectory.
<figure align="center" id="fig-lr-scale-gbs">
<table align="center" width="80%">
<tr>
<td align="center" width="100%">
<img src="images/bs_scaling_ifeval.png" width="100%"><br>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 6: Scaling of Optimal Learning Rate (&eta;) with Global Batch Size (B).</b> <i>For these experiments, the global batch size is parameterized by the total number of tokens processed per optimization step. Blue points represent empirical findings from Qwen3 training runs. The red dashed line shows a linear regression fit in log-log space, indicating a power-law relationship.</i></sub>
</figcaption>
</figure>
To quantify this scaling relationship, we conducted a systematic hyperparameter sweep to identify the optimal learning rate across a broad spectrum of batch sizes. Analysis of these optimal pairings (illustrated in [Figure 6](#fig-lr-scale-gbs)) reveals a consistent logarithmic trend. Utilizing a least-squares fit, we derived a practical scaling law for production environments:
$$\log(\eta) = a \log(B) + b$$
For our specific configuration, we found <span>$$a = 0.578$$</span> and <span>$$b = -15.684$$</span>. By exponentiating both sides, we can express the learning rate as a power function of the batch size:
$$\eta = B^a \cdot e^b$$
**The Scaling Factor:**
This relationship allows us to predict how the learning rate should change when the global batch size is scaled by a factor of <span>$$C$$</span>. If we define <span>$$\eta'$$</span> as the new learning rate for a scaled batch size <span>$$(C \cdot B)$$</span>, the ratio of the new learning rate to the original is:
$$\frac{\eta'}{\eta} = \frac{(C \cdot B)^a \cdot e^b}{B^a \cdot e^b} = \left(\frac{C \cdot B}{B}\right)^a = C^a$$
**Practical Takeaway:**
This derivation provides a reliable heuristic for scaling your training runs on VTC. Essentially, when you scale your global batch size by <span>$$C$$</span>, you should scale your learning rate by <span>$$C^{0.578}$$</span>.
**Example:** If you double your global batch size (<span>$$C = 2$$</span>), your learning rate should increase by a factor of <span>$$2^{0.578} \approx 1.49$$</span>. This "1.5x rule" ensures that your model remains on the optimal convergence path even as you significantly increase compute throughput.
### Scaling learning rates based on model parameters
While scaling batch size helps with throughput, another fundamental factor influencing your learning rate is the scale of the model itself. As we transition from compact edge models to large-scale dense architectures, the optimal learning rate (<span>$$\eta$$</span>) shifts predictably.
<figure align="center" id="fig-lr-vs-val-loss">
<table align="center" width="80%">
<tr>
<td align="center" width="100%">
<img src="images/lr_scaling_model_size.png" width="100%"><br>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 7: Empirical Learning Rate Sweep across Model Scales.</b> <i>Validation loss is plotted against learning rate for five model sizes ranging from 0.6B to 14B parameters (embedding parameters are removed from the model size calculation). Stars indicate the observed minima for each configuration.</i></sub>
</figcaption>
</figure>
To map this shift for the Qwen3 dense model family, we evaluated five distinct model scales: 0.6B, 1.7B, 4B, 8B, and 14B parameters. For each architecture, we performed a log-scale grid search to pinpoint the optimal learning rate (illustrated in [Figure 7](#fig-lr-vs-val-loss)). By applying a least-squares fit to these empirical data points, we established a power-law relationship between model parameters (<span>$$N$$</span>) and the learning rate ([Figure 8](#fig-lr-scale-model-size)):
$$\log(\eta) = -0.646 \log(N) + 4.133$$
**Scaling by Model Size:**
Following a similar derivation to our batch size analysis, this formula allows us to predict the necessary adjustment when increasing model capacity. If the number of model parameters (<span>$$N$$</span>) increases by a factor of <span>$$C$$</span>, the optimal learning rate should be scaled by <span>$$C^{-0.646}$$</span>.
**Practical Takeaway:**
This inverse relationship means that as your model grows larger, your learning rate must become more conservative to maintain stability.
**Example:** If you decide to double your model size (<span>$$C = 2$$</span>), the optimal learning rate should be multiplied by <span>$$2^{-0.646} \approx 0.639$$</span>.
<figure align="center" id="fig-lr-scale-model-size">
<table align="center" width="80%">
<tr>
<td align="center" width="100%">
<img src="images/model_scaling_ifeval.png" width="100%"><br>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 8: Relationship between Optimal Learning Rate (&eta;) and Model Parameters (N).</b> <i>Empirical data points (blue) represent the best-performing learning rates across a parameter range of approximately **0.6B to 14B**. The red dashed line depicts the power-law scaling trend. The negative slope indicates that as parameter count increases, the learning rate must be scaled down according to a fixed ratio to maintain training efficiency.</i></sub>
</figcaption>
</figure>
### Scaling learning rates based on token budget
This section examines the relationship between optimal learning rates and the total training token budget. In contrast to the power-law relationships observed in pre-training literature (e.g., the [Chinchilla scaling laws](https://arxiv.org/abs/2203.15556)), our empirical findings indicate that the optimal learning rate remains stable as the token budget increases (see [Figure 9](#fig-lr-vs-token-budget)). This divergence suggests that the hyperparameter dynamics of Supervised Fine-Tuning (SFT) differ from those of initial pre-training phases.
<figure align="center" id="fig-lr-vs-token-budget">
<table align="center" width="80%">
<tr>
<td align="center" width="100%">
<img src="images/lr_scaling_token_budget.png" width="100%"><br>
</td>
</tr>
</table>
<figcaption align="left">
<sub><b>Figure 9: Scaling of Optimal Learning Rate (&eta;) with Total Token Budget (T).</b> <i>Empirical data showing the relationship between the learning rate and the number of training tokens.</i></sub>
</figcaption>
</figure>
**Practical Takeaway:**
Within an SFT framework, the optimal learning rate is largely invariant to changes in the training token budget, allowing for consistent hyperparameter application across varying dataset scales.
## Key Takeaways
Thanks for reading. We hope this distillation framework and these scaling insights help you achieve frontier-level performance for your own reasoning models on Vertex AI Training Cluster.
### Distillation Methodology & Teacher Selection
* **The "Capacity Matching" Nuance:** Bigger is not always better. For well-defined reasoning tasks (Mathematics and IFEval), **capacity-matched** (same-sized) teachers often outperform larger models. However, for novel or highly complex domains like ARC-AGI, massive teacher models remain superior.
* **Rollout Volume Matters:** Increasing the number of teacher rollouts per prompt captures the teacher’s inherent uncertainty and multiple valid reasoning paths. Performance generally increases with rollout count, though gains may saturate depending on the domain (e.g., IFEval saturates at 8 rollouts).
* **Rejection Sampling is Domain-Specific:** While rejection sampling (training only on correct answers) provides a significant uplift for **Instruction Following**, it does not yield gains in Mathematics, ARC-AGI, or <span>$$\tau^2$$</span>-bench. In complex reasoning domains, "near-miss" cases appear essential for building robustness.
### Hyperparameter Scaling
The blog establishes three critical "rules of thumb" for scaling Supervised Fine-Tuning (SFT) hyperparameters:
* **Batch Size Scaling:** When scaling the Global Batch Size (<span>$$B$$</span>) by a factor of <span>$$C$$</span>, the learning rate (<span>$$\eta$$</span>) should be scaled by <span>$$C^{0.578}$$</span>. For example, doubling the batch size suggests a **1.5x increase** in the learning rate.
* **Model Size Scaling:** As model parameters (<span>$$N$$</span>) increase, the learning rate must become more conservative. The optimal learning rate scales by <span>$$C^{-0.646}$$</span> when the model size is increased by factor <span>$$C$$</span>.
* **Token Budget Stability:** Unlike initial pre-training, the optimal learning rate for SFT is **largely invariant** to the total training token budget, allowing for consistent application across different dataset scales.
## Acknowledgements
We would like to express our sincere gratitude to the NVIDIA NeMo RL team–specifically Terry Kong– as well as the NVIDIA NeMo Gym team–specifically Brian Yu and Chris Wing– for their invaluable support throughout this project.
We would also like to express our gratitude to our VTC teammates: Mohammadreza Mohseni, Mayank Sharan, Weiran Zhao, Jiuqiang Tang, Bo Wu, Lav Rai, and Minwoo Park for their infrastructure support, feedback, and insightful discussions throughout the project. We also thank Ting Yu, Shengyang Dai, Peng Xu, and Saurabh Tiwary for their leadership and support.
@@ -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.
+5
View File
@@ -27,7 +27,12 @@
/vertex_endpoints/nvidia-triton/nvidia-triton-custom-container-prediction.ipynb @RajeshThallam
/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
+6 -2
View File
@@ -1,7 +1,8 @@
![AlphaGenome header image](https://raw.githubusercontent.com/google-deepmind/alphagenome/refs/heads/main/docs/source/_static/header.png)
# AlphaGenome
[**Overview**](#overview) | [**Use Cases**](#use-cases) | [**Documentation**](#documentation) | [**Pricing**](#pricing) | [**Quick start**](#quick-start)
[**Overview**](#overview) | [**Use Cases**](#use-cases) | [**Documentation**](#documentation) | [**Pricing**](#pricing) | [**Quick start inference**](#quick-start-inference) |
[**Quick start finetune**](#quick-start-finetune)
## Overview
**Disclaimer:** *Experimental*.
@@ -89,5 +90,8 @@ To utilize these models via this service:
* **Pricing information** will be shared directly with users upon approval
and placement on the allowlist.
## Quick start
## Quick start inference
The quickest way to get started with the AlphaGenome in Google Cloud Platform is to run [our example notebook](cloudai_alphagenome_vai_quickstart.ipynb) in [Google Colab](https://colab.research.google.com/).
## Quick start finetune
The quickest way to get started with the AlphaGenome fineutning in Google Cloud Platform is to run [our finetuning notebook](cloudai_alphagenome_finetune.ipynb) in [Google Cloud Platform Enterprise Colab](https://docs.cloud.google.com/colab/docs/introduction).
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
@@ -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",
@@ -387,6 +387,57 @@ def load_tokenizer(
return tokenizer
def _get_indices_for_valid_length(
dataset: Any,
input_column: str,
max_sequence_length: int,
tokenizer: transformers.PreTrainedTokenizer,
context_name: str = "the dataset",
) -> tuple[list[int], int, int]:
"""Gets indices of examples shorter than or equal to max_seq_length.
Args:
dataset: The dataset to check.
input_column: The input column in the dataset.
max_sequence_length: The maximum sequence length.
tokenizer: The tokenizer.
context_name: A name for the dataset used in log messages.
Returns:
A tuple of (indices_to_keep, original_length, dropped_samples).
"""
if not dataset:
return [], 0, 0
original_length = len(dataset)
indices_to_keep = [
i
for i, entry in enumerate(dataset)
if len(tokenizer(entry[input_column])["input_ids"]) <= max_sequence_length
]
dropped_samples = original_length - len(indices_to_keep)
if dropped_samples > 0:
examples_removed_percent = (dropped_samples * 100) / original_length
logging.info(
"(%.2f%%) of examples token length is <= max-seq-length(%d); (%.2f%%) >"
" max-seq-length in %s. %d example(s) were longer than max-seq-length.",
100 - examples_removed_percent,
max_sequence_length,
examples_removed_percent,
context_name,
dropped_samples,
)
else:
logging.info(
"No samples were dropped from %s because all samples are"
" shorter than max_sequence_length=%d.",
context_name,
max_sequence_length,
)
return indices_to_keep, original_length, dropped_samples
def get_filtered_dataset(
dataset: Any,
input_column: str,
@@ -411,33 +462,25 @@ def get_filtered_dataset(
ValueError: If more than `example_removed_threshold` of the dataset is
filtered out.
"""
actual_dataset_length = len(dataset)
filtered_dataset = dataset.filter(
lambda x: len(tokenizer(x[input_column])["input_ids"]) <= max_seq_length
)
filtered_dataset_length = len(filtered_dataset)
if actual_dataset_length != filtered_dataset_length:
examples_removed_percent = (
(actual_dataset_length - filtered_dataset_length)
* 100
/ actual_dataset_length
)
logging.info(
"(%.2f%%) of examples token length is <= max-seq-length(%d); (%.2f%%) >"
" max-seq-length. Filtering out %d example(s) which are longer than"
" max-seq-length.",
100 - examples_removed_percent,
max_seq_length,
examples_removed_percent,
actual_dataset_length - filtered_dataset_length,
)
if examples_removed_percent > example_removed_threshold:
raise ValueError(
"More than %.2f%% of the dataset is filtered out. This may be due to"
" small value of max-seq-length(%d) or incorrect template. Please"
" increase the max-seq-length or check the template."
% (examples_removed_percent, max_seq_length)
indices_to_keep, original_length, dropped_samples = (
_get_indices_for_valid_length(
dataset, input_column, max_seq_length, tokenizer, "the dataset"
)
)
if (
original_length > 0
and dropped_samples / original_length * 100 > example_removed_threshold
):
examples_removed_percent = (dropped_samples * 100) / original_length
raise ValueError(
f"More than {examples_removed_percent:.2f}% of the dataset is filtered"
" out. This may be due to small value of"
f" max-seq-length({max_seq_length}) or incorrect template. Please"
" increase the max-seq-length or check the template."
)
filtered_dataset = dataset.select(indices_to_keep)
print(f"Some formatted examples from the dataset are: {filtered_dataset[:5]}")
return filtered_dataset
@@ -502,6 +545,43 @@ def load_dataset_with_template(
return raw, templated
def drop_long_sequences(
dataset: Any,
dataset_with_template: Any,
input_column: str,
max_sequence_length: int,
tokenizer: transformers.PreTrainedTokenizer,
is_train: bool,
) -> tuple[Any, Any, int]:
"""Drops examples longer than max_seq_length from the dataset.
Args:
dataset: The dataset to filter.
dataset_with_template: The dataset with template to filter.
input_column: The input column in the dataset to be used.
max_sequence_length: The maximum sequence length.
tokenizer: The tokenizer.
is_train: Whether the dataset is for training.
Returns:
A tuple of (filtered_dataset, filtered_dataset_with_template,
dropped_samples).
"""
context_name = f"the {'train' if is_train else 'eval'} dataset"
indices_to_keep, _, dropped_samples = _get_indices_for_valid_length(
dataset_with_template,
input_column,
max_sequence_length,
tokenizer,
context_name,
)
filtered_dataset = dataset.select(indices_to_keep)
filtered_dataset_with_template = dataset_with_template.select(indices_to_keep)
return filtered_dataset, filtered_dataset_with_template, dropped_samples
def validate_dataset_with_template(
dataset_name: str,
split: str,
@@ -1,58 +0,0 @@
FROM nvidia/cuda:12.3.2-devel-ubuntu22.04
# Install basic libs
RUN apt-get update && apt-get upgrade -y && apt-get install -y --no-install-recommends \
cmake \
curl \
wget \
sudo \
gnupg \
libsm6 \
libxext6 \
libxrender-dev \
lsb-release \
ca-certificates \
build-essential \
git \
software-properties-common \
cuda-toolkit \
libcudnn8 \
apt-transport-https
RUN apt install -y --no-install-recommends python3.10 \
python3.10-venv \
python3.10-dev \
python3-pip
Run apt-get autoremove -y
RUN pip install --upgrade pip
RUN pip install --upgrade --ignore-installed \
"jax[cuda12]==0.4.26" -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html \
numpy==1.26.4 \
paxml==1.4.0 \
praxis==1.4.0 \
jaxlib==0.4.26 \
pandas==2.1.4 \
einshape==1.0.0 \
utilsforecast==0.1.10 \
huggingface_hub[cli]==0.23.0 \
google-cloud-aiplatform[prediction]==1.51.0 \
fastapi==0.109.1 \
flask==3.0.3 \
smart_open[gcs]==7.0.4 \
protobuf==3.19.6 \
scikit-learn==1.0.2 \
timesfm==1.0.1
# Download license.
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
# Move scaffold.
COPY model_oss/timesfm/main.py /app/main.py
COPY model_oss/timesfm/predictor.py /app/predictor.py
WORKDIR ..
# Spin off inference server.
CMD ["python3", "/app/main.py"]
@@ -1,71 +0,0 @@
"""Predict server for TimesFM."""
import json
import os
import flask
import predictor
from predictor import PredictionError
# Create the flask app.
app = flask.Flask(__name__)
_OK_STATUS = 200
_INTERNAL_ERROR_STATUS = 500
_BAD_REQUEST_STATUS = 400
_HOST = '0.0.0.0'
# Define the predictor and load the checkpoints.
predictor = predictor.TimesFMPredictor()
predictor.load(os.environ['AIP_STORAGE_URI'])
@app.route(os.environ['AIP_HEALTH_ROUTE'], methods=['GET'])
def health() -> flask.Response:
return flask.Response(status=_OK_STATUS)
@app.route(os.environ['AIP_PREDICT_ROUTE'], methods=['GET', 'POST'])
def predict() -> flask.Response:
"""Calls TimesFM for prediction.
Returns:
A `flask.Response` containing the prediction result in JSON.
"""
try:
body = flask.request.get_json(silent=True, force=True)
preprocessed_inputs = predictor.preprocess(body)
outputs = predictor.predict(preprocessed_inputs)
conf_level = preprocessed_inputs.get('conf_level')
if conf_level is not None:
postprocessed_outputs = predictor.postprocess_with_conf_level(
outputs, preprocessed_inputs['conf_level']
)
else:
postprocessed_outputs = predictor.postprocess(outputs)
return flask.Response(
json.dumps(postprocessed_outputs),
status=_OK_STATUS,
mimetype='application/json',
)
except PredictionError as e:
return flask.Response(
json.dumps({'error': str(e)}),
status=e.status_code,
mimetype='application/json',
)
except ValueError as e:
return flask.Response(
json.dumps({'error': str(e)}),
status=_BAD_REQUEST_STATUS,
mimetype='application/json',
)
except Exception as e: # pylint: disable=broad-exception-caught
return flask.Response(
json.dumps({'error': str(e)}),
status=_INTERNAL_ERROR_STATUS,
mimetype='application/json',
)
if __name__ == '__main__':
app.run(host=_HOST, port=os.environ['AIP_HTTP_PORT'])
@@ -1,609 +0,0 @@
"""Adapts a pretrained TimesFM to the CPR framework.
Documentation for the model is here:
https://github.com/google-research/timesfm
Model checkpoints can be found here:
https://www.huggingface.co/google/timesfm-1.0-200m
"""
from collections.abc import Sequence
import datetime
import os
from typing import Any
import fastapi
from google.cloud.aiplatform.utils import prediction_utils
from jax._src import config
import numpy as np
import scipy.stats as st
import timesfm
HTTPException = fastapi.HTTPException
_BACKEND = os.getenv("TIMESFM_BACKEND", default="cpu")
config.update(
"jax_platforms", {"cpu": "cpu", "gpu": "cuda", "tpu": ""}[_BACKEND]
)
TsArray = None | float | int | str | list["TsArray"]
_BAD_REQUEST_STATUS = 400
_EXPECTED_FORMAT = """
[NOTICE] TimesFM inference server expects input format:
{
"instances": [
{
"input": [0.0, 0.1, 0.2, ...],
"freq": 0, # optional, 0/1/2
"horizon": 12, # optional
"timestamp": ["2024-01-01", "2024-01-02", ...], # optional
"timestamp_format": "%Y-%m-%d", # optional
"dynamic_numerical_covariates": {
"dncov1": [1.0, 2.0, 1.5, ...],
"dncov2": [3.0, 1.1, 2.4, ...],
}, # optional
"dynamic_categorical_covariates": {
"dccov1": ["a", "b", "a", ...],
"dccov2": [0, 1, 0, ...],
}, # optional
"static_numerical_covariates": {
"sncov1": 1.0,
"sncov2": 2.0,
}, # optional
"static_categorical_covariates": {
"sccov1": "a",
"sccov2": "b",
}, # optional
"xreg_kwargs": {...}, # optional
},
{"input": [113.2, 15.0, 65.4, ...], ...},
{"input": [ 0.0, 10.0, 20.0, ...], ...},
...
]
}
"""
class PredictionError(Exception):
"""Custom exception for prediction errors."""
def __init__(self, message: str, status_code: int = _BAD_REQUEST_STATUS):
super().__init__(message)
self.status_code = status_code
self.message = message
def _raise_bad_request(message: str):
message = message + "\n" + _EXPECTED_FORMAT
raise PredictionError(
message=message,
status_code=_BAD_REQUEST_STATUS,
)
def _datetime_to_freq(dt1: datetime.datetime, dt2: datetime.datetime) -> int:
delta = dt2 - dt1
if delta.days <= 1:
return 0
elif delta.days <= 31:
return 1
else:
return 2
def _add_cov_to_dict(
index: int,
cov_input: dict[str, TsArray],
cov_dict: dict[str, list[TsArray]],
):
"""Adds covariates to the dictionary of covariates.
Args:
index: Index of the instance.
cov_input: Dictionary of covariates for the current instance.
cov_dict: Dictionary of covariates for all instances.
"""
if index == 0:
cov_dict.update({k: [v] for k, v in cov_input.items()})
else:
if set(cov_input.keys()) != set(cov_dict.keys()):
_raise_bad_request(
f"Instance {index}:"
" All instances must have the same set of covariates if any."
)
for k, v in cov_input.items():
cov_dict[k].append(v)
def _linear_interpolate_missing_timepoints(
timestamp: list[datetime.datetime],
value: list[float],
) -> tuple[list[datetime.datetime], list[TsArray]]:
"""Linearly interpolates missing timepoints in a timeseries."""
def _gcd_timelapse(t1, t2):
if (w := t2 % t1) == datetime.timedelta(0):
return t1
if t1 > t2:
return _gcd_timelapse(t2, t1)
return _gcd_timelapse(w, t1)
if len(timestamp) < 3:
return timestamp, value, False
no_missing = True
delta = timestamp[1] - timestamp[0]
if delta <= datetime.timedelta(0):
_raise_bad_request(
f"Timestamps must be in ascending order. Got {timestamp}"
)
for i in range(2, len(timestamp)):
delta_next = timestamp[i] - timestamp[i - 1]
if delta_next <= datetime.timedelta(0):
_raise_bad_request(
f"Timestamps must be in ascending order. Got {timestamp}"
)
delta_new = _gcd_timelapse(delta, delta_next)
if delta_new != delta:
no_missing = False
delta = delta_new
if no_missing:
return timestamp, value, False
new_timestamp = []
new_value = []
for i in range(len(timestamp) - 1):
new_timestamp.append(timestamp[i])
new_value.append(value[i])
if (num_deltas := int((timestamp[i + 1] - timestamp[i]) / delta + 0.5)) > 1:
value_delta = (value[i + 1] - value[i]) / num_deltas
for j in range(1, num_deltas):
new_timestamp.append(timestamp[i] + j * delta)
new_value.append(value[i] + j * value_delta)
new_timestamp.append(timestamp[-1])
new_value.append(value[-1])
return new_timestamp, new_value, True
class TimesFMPredictor:
"""Predictor class for time-series foundation model TimesFM."""
TIMESFM_MODEL_NAME = os.getenv(
"TIMESFM_MODEL_NAME", default="timesfm-1.0-200m"
)
CONTEXT_LEN = 512
INPUT_PATCH_LEN = 32
OUTPUT_PATCH_LEN = 128
NUM_LAYERS = 20
MODEL_DIMS = 1280
BACKEND = os.getenv("TIMESFM_BACKEND", default="cpu")
MAX_HORIZON = int(os.getenv("TIMESFM_HORIZON", default="128"))
def load(self, artifacts_uri: str = ""):
"""Initializes the model and preprocessing transforms.
Args:
artifacts_uri: Directory where state dict is stored. Can be a GCS URI or
local path.
"""
if not (os.path.isdir(artifacts_uri) or artifacts_uri.startswith("gs://")):
raise ValueError(
f"Provided artifact_uri is not a directory: {artifacts_uri}"
)
print(f"Downloading checkpoints from {artifacts_uri}")
prediction_utils.download_model_artifacts(artifacts_uri)
artifact_path = os.getcwd()
print(f"Loading checkpoints from {artifact_path}")
self._model = timesfm.TimesFm(
context_len=self.CONTEXT_LEN,
horizon_len=(
((self.MAX_HORIZON - 1) // self.OUTPUT_PATCH_LEN + 1)
* self.OUTPUT_PATCH_LEN
),
input_patch_len=self.INPUT_PATCH_LEN,
output_patch_len=self.OUTPUT_PATCH_LEN,
num_layers=self.NUM_LAYERS,
model_dims=self.MODEL_DIMS,
backend=self.BACKEND,
)
self._model.load_from_checkpoint(artifact_path)
print(f"Loaded TimesFM model from {artifact_path}")
def preprocess(
self, request_dict: dict[str, Sequence[dict[str, TsArray]]]
) -> dict[str, TsArray]:
"""Performs preprocessing.
By default, the server expects a request body consisting of a valid JSON
object. This will be parsed by the handler before it's evaluated by the
preprocess method.
Args:
request_dict: Parsed request body. We expect that the input consists of a
list of time-series forecast contexts. Each context should be in a
format convertible to JTensor by `jnp.array`.
Returns:
Time-series forecast contexts are passed as is from the input as a list.
"""
if "instances" not in request_dict:
_raise_bad_request('Request must contain "instances" as a top-level key.')
input_instances = request_dict["instances"]
if not input_instances or not isinstance(input_instances, list):
_raise_bad_request(
f"Received `instances` not a list. Got {type(input_instances)}"
)
inputs, freqs, timestamps, timestamp_formats = [], [], [], []
horizon_lens = []
conf_level = None
static_numerical_covariates, static_categorical_covariates = {}, {}
dynamic_numerical_covariates, dynamic_categorical_covariates = {}, {}
xreg_kwargs = {}
exists_missing = False
for index, each_input in enumerate(input_instances):
# 1. Add input time-series context.
if (
(not isinstance(each_input, dict))
or ("input" not in each_input)
or (len(each_input["input"]) < 2)
):
_raise_bad_request(
f"Instance {index}:"
" Invalid datatype. Each input example must have `input` key"
" mapped to a list of time-series forecast context with length > 1."
)
new_input = each_input["input"]
# 2. Process timestamps.
if "timestamp" not in each_input:
timestamps.append(None)
else:
if len(each_input["timestamp"]) != len(each_input["input"]):
_raise_bad_request(
f"Instance {index}:"
" Invalid datatype. `timestamp` if given must have same length as"
"`input`."
)
new_timestamp = [
datetime.datetime.fromisoformat(s) for s in each_input["timestamp"]
]
# Linearly interpolate missing timepoints and values.
new_timestamp, new_input, new_exists_missing = (
_linear_interpolate_missing_timepoints(new_timestamp, new_input)
)
exists_missing = exists_missing or new_exists_missing
timestamps.append(new_timestamp)
if "timestamp_format" in each_input:
timestamp_formats.append(each_input["timestamp_format"])
else:
timestamp_formats.append(None)
inputs.append(new_input)
# 3. Process frequency.
if "freq" in each_input:
freqs.append(each_input["freq"])
elif timestamps[index]:
freqs.append(
_datetime_to_freq(timestamps[index][0], timestamps[index][1])
)
else:
freqs.append(0)
# 4. Process covariate data.
for cov_category, cov_dict in [
("static_numerical_covariates", static_numerical_covariates),
("static_categorical_covariates", static_categorical_covariates),
("dynamic_numerical_covariates", dynamic_numerical_covariates),
("dynamic_categorical_covariates", dynamic_categorical_covariates),
]:
if cov_category in each_input:
_add_cov_to_dict(index, each_input[cov_category], cov_dict)
# 5. Process xreg config. Power user option. If nothing set we apply
# TimesFM default.
if "xreg_kwargs" in each_input:
if not xreg_kwargs:
xreg_kwargs = each_input["xreg_kwargs"]
elif xreg_kwargs != each_input["xreg_kwargs"]:
_raise_bad_request(
f"Instance {index}:"
" All instances must have the same xreg_kwargs if any."
)
# 6. Process horizon length.
if "horizon" in each_input:
if (w := each_input["horizon"]) > self.MAX_HORIZON:
_raise_bad_request(
f"Instance {index}: `horizon` must be <= maximum horizon"
f" {self.MAX_HORIZON}. Got {w}. To increase the maximum horizon,"
" recreate the endpoint with a higher `TIMESFM_HORIZON` env"
" value."
)
horizon_lens.append(w)
else:
horizon_lens.append(self.MAX_HORIZON)
# 7. Process conf level.
all_conf_levels = [
each_input.get("conf_level", None) for each_input in input_instances
]
defined_conf_levels = [cl for cl in all_conf_levels if cl is not None]
undefined_conf_levels = [cl for cl in all_conf_levels if cl is None]
if defined_conf_levels and undefined_conf_levels:
_raise_bad_request(
"Either all or none of the instances must define `conf_level`."
)
if defined_conf_levels:
unique_conf_levels = set(defined_conf_levels)
if len(unique_conf_levels) > 1:
_raise_bad_request("All instances must have the same `conf_level`.")
conf_level = unique_conf_levels.pop()
if not 0 <= conf_level <= 1:
_raise_bad_request(
f"`conf_level` must be between 0 and 1. Got {conf_level}."
)
else:
conf_level = None
return {
"inputs": inputs,
"freqs": freqs,
"timestamps": timestamps,
"timestamp_formats": timestamp_formats,
"exists_missing": exists_missing,
"static_numerical_covariates": static_numerical_covariates,
"static_categorical_covariates": static_categorical_covariates,
"dynamic_numerical_covariates": dynamic_numerical_covariates,
"dynamic_categorical_covariates": dynamic_categorical_covariates,
"xreg_kwargs": xreg_kwargs,
"horizon_lens": horizon_lens,
"conf_level": conf_level,
}
def predict(self, instances: dict[str, Any]) -> Any:
"""Performs prediction.
Args:
instances: A dictionary with two keys - `inputs` and `freq` where `inputs`
is list of time series forecast contexts. Each context time series
should be in a format convertible to JTensor by `jnp.array`. `freq` is
frequencies of each forecast context with values as 0 (high), 1 (medium)
and 2 (low). If not provided, all contexts are assumed to be high
frequency.
Returns:
A tuple of List:
- the mean forecast of size (# inputs, # forecast horizon),
- the full forecast (mean + quantiles) of size
(# inputs, # forecast horizon, 1 + # quantiles).
"""
(
inputs,
freqs,
timestamps,
timestamp_formats,
exists_missing,
static_numerical_covariates,
static_categorical_covariates,
dynamic_numerical_covariates,
dynamic_categorical_covariates,
xreg_kwargs,
horizon_lens,
) = (
instances["inputs"],
instances["freqs"],
instances["timestamps"],
instances["timestamp_formats"],
instances["exists_missing"],
instances["static_numerical_covariates"],
instances["static_categorical_covariates"],
instances["dynamic_numerical_covariates"],
instances["dynamic_categorical_covariates"],
instances["xreg_kwargs"],
instances["horizon_lens"],
)
if (
static_numerical_covariates
or static_categorical_covariates
or dynamic_numerical_covariates
or dynamic_categorical_covariates
):
if (
dynamic_categorical_covariates or dynamic_numerical_covariates
) and exists_missing:
_raise_bad_request(
"Dynamic covariates are not supported when input has missing"
" timestamps."
)
print("Detected covariates. Callng model.forecast_with_covariates.")
try:
point_forecast, _ = self._model.forecast_with_covariates(
inputs=inputs,
dynamic_numerical_covariates=dynamic_numerical_covariates,
dynamic_categorical_covariates=dynamic_categorical_covariates,
static_numerical_covariates=static_numerical_covariates,
static_categorical_covariates=static_categorical_covariates,
freq=freqs,
**xreg_kwargs,
)
# point_forecast is a list of np.ndarrays.
point_forecast = [p.tolist() for p in point_forecast]
quantile_forecast = None
except ValueError as e:
_raise_bad_request(f"model.forecast_with_covariates failed from {e}.")
return
else:
print("Calling model.forecast.")
point_forecast, quantile_forecast = self._model.forecast(
inputs=inputs, freq=freqs
)
# point_forecast and quantile_forecast are JTensors (np.ndarrays).
point_forecast = point_forecast.tolist()
quantile_forecast = quantile_forecast.tolist()
return (
point_forecast,
quantile_forecast,
timestamps,
timestamp_formats,
horizon_lens,
)
def postprocess(
self, forecasts: tuple[TsArray, TsArray, TsArray, TsArray]
) -> dict[str, list[dict[str, TsArray]]]:
"""Translates the model output.
Args:
forecasts: A tuple of List - the mean forecast of size (# inputs, #
forecast horizon), - the full forecast (mean + quantiles) of size (#
inputs, # forecast horizon, 1 + # quantiles).
Returns:
Dictionary containing the list of point forecasts and quantile forecasts
for each of the input time-series context.
"""
(
point_forecasts,
quantile_forecasts,
timestamps,
timestamp_formats,
horizon_lens,
) = forecasts
predictions = []
quantile_names = ["mean"] + [
f"p{int(quantile * 100)}" for quantile in self._model.model_p.quantiles
]
for i, point_forecast in enumerate(point_forecasts):
response = {"point_forecast": point_forecast[: horizon_lens[i]]}
if quantile_forecasts:
for j, quantile_name in enumerate(quantile_names):
response[quantile_name] = [x[j] for x in quantile_forecasts[i]][
: horizon_lens[i]
]
if timestamps[i]:
last_timestamp = timestamps[i][-1]
timestamp_delta = timestamps[i][-1] - timestamps[i][-2]
response["timestamp"] = []
for _ in range(len(point_forecast)):
last_timestamp = last_timestamp + timestamp_delta
response["timestamp"].append(
datetime.datetime.strftime(last_timestamp, timestamp_formats[i])
if timestamp_formats[i]
else last_timestamp.isoformat()
)
response["timestamp"] = response["timestamp"][: horizon_lens[i]]
predictions.append(response)
return {"predictions": predictions}
def postprocess_with_conf_level(
self,
forecasts: tuple[TsArray, TsArray, TsArray, TsArray, TsArray],
conf_level: float | None,
) -> dict[str, list[dict[str, TsArray]]]:
"""Translates the model output."""
lower_quantile = (1 - conf_level) / 2
higher_quantile = (1 + conf_level) / 2
_, quantile_forecast, _, _, horizon_lens = forecasts
response = self.postprocess(forecasts)
if quantile_forecast is None:
return response
# Note: The raw quantile forecast from TimesFM has the mean as the 0-th
# element. We strip it before passing to extend_quantiles.
quantile_forecast_np_array = np.array(quantile_forecast)
extended_forecasts = extend_quantiles(
quantile_forecast_np_array[..., 1:],
lower_quantile,
higher_quantile,
model_quantiles=self._model.model_p.quantiles,
)
lower_bounds = extended_forecasts["lower_bound"]
upper_bounds = extended_forecasts["upper_bound"]
for i, prediction in enumerate(response["predictions"]):
horizon = horizon_lens[i]
prediction["lower_bound"] = lower_bounds[i][:horizon].tolist()
prediction["upper_bound"] = upper_bounds[i][:horizon].tolist()
return response
def extend_quantiles(
quantile_forecast: np.ndarray,
lower_quantile: float,
higher_quantile: float,
model_quantiles: list[float],
) -> dict[str, np.ndarray]:
"""Extends the quantile forecast to the lower and upper bounds.
Args:
quantile_forecast: The quantile forecast from TimesFM.
lower_quantile: The lower quantile to extend to.
higher_quantile: The higher quantile to extend to.
model_quantiles: The quantiles used by the model.
Returns:
A dictionary containing the lower and upper bounds.
"""
if quantile_forecast.shape[2] != len(model_quantiles):
raise ValueError(
"Number of model quantiles should match the last dimension of the"
" quantile forecast. If you are using the raw TimesFM quantile forecast"
"output, you likely need to strip the 0-index which is the mean."
)
idx_median = model_quantiles.index(0.5)
idx_low_q = np.argmin(model_quantiles)
low_q = model_quantiles[idx_low_q]
if not (low_q < 0.5):
raise ValueError(
f"The lowest quantile {low_q=} provided in the forecast must be less"
" than 0.5."
)
idx_high_q = np.argmax(model_quantiles)
high_q = model_quantiles[idx_high_q]
if not (high_q > 0.5):
raise ValueError(
f"The highest quantile {high_q=} provided in the forecast must be"
" greater than 0.5."
)
positive_sigma = np.maximum(
0, quantile_forecast[..., idx_high_q] - quantile_forecast[..., idx_median]
) / st.norm.ppf(high_q)
negative_sigma = np.minimum(
0, quantile_forecast[..., idx_low_q] - quantile_forecast[..., idx_median]
) / st.norm.ppf(low_q)
lower_bound = quantile_forecast[
..., idx_median
] + negative_sigma * st.norm.ppf(lower_quantile)
upper_bound = quantile_forecast[
..., idx_median
] + positive_sigma * st.norm.ppf(higher_quantile)
return {"lower_bound": lower_bound, "upper_bound": upper_bound}
@@ -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",
@@ -387,6 +387,57 @@ def load_tokenizer(
return tokenizer
def _get_indices_for_valid_length(
dataset: Any,
input_column: str,
max_sequence_length: int,
tokenizer: transformers.PreTrainedTokenizer,
context_name: str = "the dataset",
) -> tuple[list[int], int, int]:
"""Gets indices of examples shorter than or equal to max_seq_length.
Args:
dataset: The dataset to check.
input_column: The input column in the dataset.
max_sequence_length: The maximum sequence length.
tokenizer: The tokenizer.
context_name: A name for the dataset used in log messages.
Returns:
A tuple of (indices_to_keep, original_length, dropped_samples).
"""
if not dataset:
return [], 0, 0
original_length = len(dataset)
indices_to_keep = [
i
for i, entry in enumerate(dataset)
if len(tokenizer(entry[input_column])["input_ids"]) <= max_sequence_length
]
dropped_samples = original_length - len(indices_to_keep)
if dropped_samples > 0:
examples_removed_percent = (dropped_samples * 100) / original_length
logging.info(
"(%.2f%%) of examples token length is <= max-seq-length(%d); (%.2f%%) >"
" max-seq-length in %s. %d example(s) were longer than max-seq-length.",
100 - examples_removed_percent,
max_sequence_length,
examples_removed_percent,
context_name,
dropped_samples,
)
else:
logging.info(
"No samples were dropped from %s because all samples are"
" shorter than max_sequence_length=%d.",
context_name,
max_sequence_length,
)
return indices_to_keep, original_length, dropped_samples
def get_filtered_dataset(
dataset: Any,
input_column: str,
@@ -411,33 +462,25 @@ def get_filtered_dataset(
ValueError: If more than `example_removed_threshold` of the dataset is
filtered out.
"""
actual_dataset_length = len(dataset)
filtered_dataset = dataset.filter(
lambda x: len(tokenizer(x[input_column])["input_ids"]) <= max_seq_length
)
filtered_dataset_length = len(filtered_dataset)
if actual_dataset_length != filtered_dataset_length:
examples_removed_percent = (
(actual_dataset_length - filtered_dataset_length)
* 100
/ actual_dataset_length
)
logging.info(
"(%.2f%%) of examples token length is <= max-seq-length(%d); (%.2f%%) >"
" max-seq-length. Filtering out %d example(s) which are longer than"
" max-seq-length.",
100 - examples_removed_percent,
max_seq_length,
examples_removed_percent,
actual_dataset_length - filtered_dataset_length,
)
if examples_removed_percent > example_removed_threshold:
raise ValueError(
"More than %.2f%% of the dataset is filtered out. This may be due to"
" small value of max-seq-length(%d) or incorrect template. Please"
" increase the max-seq-length or check the template."
% (examples_removed_percent, max_seq_length)
indices_to_keep, original_length, dropped_samples = (
_get_indices_for_valid_length(
dataset, input_column, max_seq_length, tokenizer, "the dataset"
)
)
if (
original_length > 0
and dropped_samples / original_length * 100 > example_removed_threshold
):
examples_removed_percent = (dropped_samples * 100) / original_length
raise ValueError(
f"More than {examples_removed_percent:.2f}% of the dataset is filtered"
" out. This may be due to small value of"
f" max-seq-length({max_seq_length}) or incorrect template. Please"
" increase the max-seq-length or check the template."
)
filtered_dataset = dataset.select(indices_to_keep)
print(f"Some formatted examples from the dataset are: {filtered_dataset[:5]}")
return filtered_dataset
@@ -502,6 +545,43 @@ def load_dataset_with_template(
return raw, templated
def drop_long_sequences(
dataset: Any,
dataset_with_template: Any,
input_column: str,
max_sequence_length: int,
tokenizer: transformers.PreTrainedTokenizer,
is_train: bool,
) -> tuple[Any, Any, int]:
"""Drops examples longer than max_seq_length from the dataset.
Args:
dataset: The dataset to filter.
dataset_with_template: The dataset with template to filter.
input_column: The input column in the dataset to be used.
max_sequence_length: The maximum sequence length.
tokenizer: The tokenizer.
is_train: Whether the dataset is for training.
Returns:
A tuple of (filtered_dataset, filtered_dataset_with_template,
dropped_samples).
"""
context_name = f"the {'train' if is_train else 'eval'} dataset"
indices_to_keep, _, dropped_samples = _get_indices_for_valid_length(
dataset_with_template,
input_column,
max_sequence_length,
tokenizer,
context_name,
)
filtered_dataset = dataset.select(indices_to_keep)
filtered_dataset_with_template = dataset_with_template.select(indices_to_keep)
return filtered_dataset, filtered_dataset_with_template, dropped_samples
def validate_dataset_with_template(
dataset_name: str,
split: str,

Some files were not shown because too many files have changed in this diff Show More