Skip to content

Commit 7dc8dd9

Browse files
Automate TorchAO pin bumps (pytorch#23189)
## Summary Extend the weekly PyTorch pin-bump automation to update TorchAO at the same time. The implementation reuses the existing nightly-date parser and GitHub JSON request path, and reads `SUPPORTED_CUDA_VERSIONS` directly from `install_utils.py`. The TorchAO-specific code selects the newest nightly available for CPU and every supported CUDA channel, resolves the exact source SHA from the successful TorchAO wheel build, updates the nightly constants, and advances `third-party/ao`. Existing dependency tests now derive their expectations from those source-of-truth pins, so future bot PRs do not need to rewrite tests. The live lookup currently reproduces pytorch#23184: `0.19.0.dev20260907` at `b7ac3aacf9b6f1cb0c6bb29143329f601dfcfcb1`, with CUDA 12.6 limiting the common nightly date. Authored with assistance from Codex. ## Test plan - `python .ci/scripts/tests/test_cu134_dependencies.py` (9 tests) - `python -m py_compile .github/scripts/update_pytorch_pin.py` - `black --check .github/scripts/update_pytorch_pin.py .ci/scripts/tests/test_cu134_dependencies.py` - Parsed `.github/workflows/weekly-pytorch-pin-bump.yml` with Ruby YAML - `git diff --check` - Exercised the live TorchAO index and GitHub Actions lookup
1 parent c16dd0a commit 7dc8dd9

3 files changed

Lines changed: 157 additions & 21 deletions

File tree

‎.ci/scripts/tests/test_cu134_dependencies.py‎

Lines changed: 26 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -61,7 +61,7 @@ def test_all_install_steps_preserve_exact_cu134_selection(self):
6161
"torch==2.14.0.dev20260810+cu134",
6262
"torchvision==0.29.0.dev20260811+cu134",
6363
"torchaudio==2.11.0.dev20260811+cu134",
64-
f"torchao==0.19.0.dev20260907+{ao_variant}",
64+
f"torchao=={self.installer.CU134_TORCHAO_NIGHTLY_VERSION}+{ao_variant}",
6565
}
6666
for index, command in enumerate(commands):
6767
required = (
@@ -101,7 +101,9 @@ def test_other_cuda_trains_keep_existing_pins(self):
101101
cuda, machine
102102
)
103103
self.assertIn("torch==2.14.0", core)
104-
self.assertIn("torchao==0.19.0.dev20260907", core)
104+
self.assertIn(
105+
f"torchao=={self.installer.TORCHAO_NIGHTLY_VERSION}", core
106+
)
105107
self.assertIn("torchvision==0.29.0", domains)
106108
self.assertIn("torchaudio==2.11.0", domains)
107109
self.assertFalse(any("==" in arg for arg in local))
@@ -119,7 +121,7 @@ def test_source_pinned_torch_is_not_replaced(self):
119121
def test_no_cuda_keeps_default_pins(self):
120122
core, _, domains, _ = self.install_commands(None)
121123
self.assertIn("torch==2.14.0", core)
122-
self.assertIn("torchao==0.19.0.dev20260907", core)
124+
self.assertIn(f"torchao=={self.installer.TORCHAO_NIGHTLY_VERSION}", core)
123125
self.assertIn("torchvision==0.29.0", domains)
124126
self.assertIn("https://download.pytorch.org/whl/test/cpu", core)
125127

@@ -239,14 +241,29 @@ def test_cu134_keeps_explicit_torchao_source_build(self):
239241
any(arg.startswith("torchao==") for arg in command)
240242
)
241243
self.assertIn("torch==2.14.0.dev20260810+cu134", commands[-1])
242-
self.assertIn("0.19.0+gitb7ac3aa", metadata.specifier)
244+
source_version = self.installer.TORCHAO_NIGHTLY_VERSION.partition(
245+
".dev"
246+
)[0]
247+
source_commit = subprocess.run(
248+
["git", "rev-parse", "HEAD:third-party/ao"],
249+
cwd=ROOT,
250+
capture_output=True,
251+
check=True,
252+
text=True,
253+
).stdout.strip()
254+
self.assertIn(
255+
f"{source_version}+git{source_commit[:7]}", metadata.specifier
256+
)
243257

244258
def test_wheel_torchao_bound_matches_selected_train(self):
245-
for cuda, expected in (
246-
((13, 4), "torchao>=0.19.0.dev20260907,<0.20"),
247-
((13, 2), "torchao>=0.19.0.dev20260907,<0.20"),
248-
(None, "torchao>=0.19.0.dev20260907,<0.20"),
249-
):
259+
for cuda in ((13, 4), (13, 2), None):
260+
version = (
261+
self.installer.CU134_TORCHAO_NIGHTLY_VERSION
262+
if cuda == (13, 4)
263+
else self.installer.TORCHAO_NIGHTLY_VERSION
264+
)
265+
major, minor = (int(part) for part in version.split(".")[:2])
266+
expected = f"torchao>={version},<{major}.{minor + 1}"
250267
self.utils.determine_torch_url.cache_clear()
251268
with (
252269
patch.object(

‎.github/scripts/update_pytorch_pin.py‎

Lines changed: 124 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,17 @@
11
#!/usr/bin/env python3
22

33
import base64
4-
import hashlib
54
import json
65
import re
6+
import runpy
7+
import subprocess
78
import sys
89
import urllib.request
910
from pathlib import Path
11+
from urllib.parse import unquote
12+
13+
14+
TORCHAO_INDEX_URL = "https://download.pytorch.org/whl/nightly"
1015

1116

1217
def parse_nightly_version(nightly_version):
@@ -44,6 +49,14 @@ def get_torch_nightly_version():
4449
return match.group(1)
4550

4651

52+
def get_json(url):
53+
req = urllib.request.Request(url)
54+
req.add_header("Accept", "application/vnd.github.v3+json")
55+
req.add_header("User-Agent", "ExecuTorch-Bot")
56+
with urllib.request.urlopen(req) as response:
57+
return json.loads(response.read().decode())
58+
59+
4760
def get_commit_hash_for_nightly(date_str):
4861
"""
4962
Fetch commit hash from PyTorch nightly branch for a given date.
@@ -58,13 +71,8 @@ def get_commit_hash_for_nightly(date_str):
5871
params = f"?sha=nightly&per_page=50"
5972
url = api_url + params
6073

61-
req = urllib.request.Request(url)
62-
req.add_header("Accept", "application/vnd.github.v3+json")
63-
req.add_header("User-Agent", "ExecuTorch-Bot")
64-
6574
try:
66-
with urllib.request.urlopen(req) as response:
67-
commits = json.loads(response.read().decode())
75+
commits = get_json(url)
6876
except Exception as e:
6977
print(f"Error fetching commits: {e}", file=sys.stderr)
7078
sys.exit(1)
@@ -104,6 +112,105 @@ def update_pytorch_pin(commit_hash):
104112
print(f"Updated {pin_file} with commit hash: {commit_hash}")
105113

106114

115+
def get_supported_torchao_channels():
116+
cuda_versions = runpy.run_path("install_utils.py")["SUPPORTED_CUDA_VERSIONS"]
117+
return ["cpu", *(f"cu{major}{minor}" for major, minor in cuda_versions)]
118+
119+
120+
def get_torchao_versions(channel):
121+
url = f"{TORCHAO_INDEX_URL}/{channel}/torchao/"
122+
req = urllib.request.Request(url, headers={"User-Agent": "ExecuTorch-Bot"})
123+
with urllib.request.urlopen(req) as response:
124+
index_html = unquote(response.read().decode())
125+
126+
wheel_tags = {}
127+
pattern = re.compile(
128+
rf"^torchao-(\d+\.\d+\.\d+\.dev\d{{8}})\+{re.escape(channel)}-(.+)\.whl$"
129+
)
130+
for filename in re.findall(r'href="[^"]*/(torchao-[^"]+\.whl)"', index_html):
131+
match = pattern.match(filename)
132+
if match:
133+
wheel_tags.setdefault(match.group(1), set()).add(match.group(2))
134+
135+
required_tags = ("py3-none-any", "aarch64") if channel == "cpu" else ("x86_64",)
136+
return {
137+
version
138+
for version, tags in wheel_tags.items()
139+
if all(any(required in tag for tag in tags) for required in required_tags)
140+
}
141+
142+
143+
def get_latest_torchao_nightly(max_date):
144+
common_versions = None
145+
channels = get_supported_torchao_channels()
146+
for channel in channels:
147+
versions = get_torchao_versions(channel)
148+
common_versions = (
149+
versions if common_versions is None else common_versions & versions
150+
)
151+
152+
candidates = [
153+
version
154+
for version in common_versions or []
155+
if version.rsplit(".dev", 1)[-1] <= max_date
156+
]
157+
if not candidates:
158+
raise ValueError(
159+
f"Could not find a TorchAO nightly on or before {max_date} for "
160+
f"all supported channels: {', '.join(channels)}"
161+
)
162+
return max(candidates, key=lambda version: (version.rsplit(".dev", 1)[-1], version))
163+
164+
165+
def get_torchao_commit_hash(nightly_version):
166+
date = nightly_version.rsplit(".dev", 1)[-1]
167+
formatted_date = parse_nightly_version(f"dev{date}")
168+
url = (
169+
"https://api.github.com/repos/pytorch/ao/actions/workflows/" # @lint-ignore
170+
"build_wheels_linux_x86.yml/runs?event=schedule&status=success&"
171+
f"created={formatted_date}&per_page=100"
172+
)
173+
runs = get_json(url).get("workflow_runs", [])
174+
if not runs:
175+
raise ValueError(
176+
f"Could not find the successful TorchAO wheel build for {nightly_version}"
177+
)
178+
return runs[0]["head_sha"]
179+
180+
181+
def update_torchao_pins(nightly_version, commit_hash):
182+
requirements_path = Path("install_requirements.py")
183+
content = requirements_path.read_text()
184+
content, replacements = re.subn(
185+
r'^(?P<prefix>(?:CU\d+_)?TORCHAO_NIGHTLY_VERSION\s*=\s*["\'])[^"\']+(?P<suffix>["\'])$',
186+
rf"\g<prefix>{nightly_version}\g<suffix>",
187+
content,
188+
flags=re.MULTILINE,
189+
)
190+
if not replacements:
191+
raise ValueError(f"Could not find TorchAO nightly pins in {requirements_path}")
192+
requirements_path.write_text(content)
193+
194+
for command in (
195+
["git", "submodule", "update", "--init", "third-party/ao"],
196+
[
197+
"git",
198+
"-C",
199+
"third-party/ao",
200+
"fetch",
201+
"--depth=1",
202+
"origin",
203+
commit_hash,
204+
],
205+
["git", "-C", "third-party/ao", "checkout", "--detach", commit_hash],
206+
):
207+
subprocess.run(command, check=True)
208+
print(
209+
f"Updated TorchAO nightly pins to {nightly_version} and third-party/ao "
210+
f"to {commit_hash}"
211+
)
212+
213+
107214
def should_skip_file(filename):
108215
"""
109216
Check if a file should be skipped during sync (build files).
@@ -262,8 +369,17 @@ def main():
262369
# Sync c10 directories from PyTorch
263370
sync_c10_directories(commit_hash)
264371

372+
# Select the newest TorchAO nightly available for every supported CUDA
373+
# channel and align the source submodule with the commit that built it.
374+
max_torchao_date = date_str.replace("-", "")
375+
torchao_version = get_latest_torchao_nightly(max_torchao_date)
376+
print(f"Found TorchAO nightly version: {torchao_version}")
377+
torchao_commit_hash = get_torchao_commit_hash(torchao_version)
378+
print(f"Found TorchAO commit hash: {torchao_commit_hash}")
379+
update_torchao_pins(torchao_version, torchao_commit_hash)
380+
265381
print(
266-
"\n✅ Successfully updated PyTorch commit pin and synced c10 directories!"
382+
"\n✅ Successfully updated PyTorch and TorchAO pins and synced c10 directories!"
267383
)
268384

269385
except Exception as e:

‎.github/workflows/weekly-pytorch-pin-bump.yml‎

Lines changed: 7 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
name: Weekly PyTorch Pin Bump
1+
name: Weekly PyTorch and TorchAO Pin Bump
22

33
on:
44
schedule:
@@ -54,15 +54,17 @@ jobs:
5454
git config user.email "pytorchbot@users.noreply.github.com"
5555
git checkout -b "${BRANCH}"
5656
git add torch_pin.py
57+
git add install_requirements.py
5758
git add .ci/docker/ci_commit_pins/pytorch.txt
5859
git add runtime/core/portable_type/c10/
60+
git add third-party/ao
5961
6062
if git diff --cached --quiet; then
6163
echo "No changes to commit. Pin is already up to date."
6264
exit 0
6365
fi
6466
65-
git commit -m "Bump PyTorch pin to nightly ${{ steps.nightly.outputs.version }}"
67+
git commit -m "Bump PyTorch and TorchAO pins for ${{ steps.nightly.outputs.version }}"
6668
git push -u origin "${BRANCH}"
6769
6870
EXISTING=$(gh pr list --label "ci/pytorch-pin-bump" --state open --json number --jq '.[0].number')
@@ -75,11 +77,12 @@ jobs:
7577
read -r -d '' PR_BODY <<EOF || true
7678
## Summary
7779
78-
Automated weekly PyTorch pin bump.
80+
Automated weekly PyTorch and TorchAO pin bump.
7981
8082
- Updates \`NIGHTLY_VERSION\` in \`torch_pin.py\` to \`${NIGHTLY}\`
8183
- Updates \`.ci/docker/ci_commit_pins/pytorch.txt\` to the corresponding nightly commit hash
8284
- Syncs c10 headers from PyTorch into \`runtime/core/portable_type/c10/\`
85+
- Updates the TorchAO nightly pins and \`third-party/ao\` to the newest build available for every supported CUDA channel
8386
8487
This PR was created automatically. If CI fails, Claude will attempt to fix issues (up to 3 attempts). If CI still fails, human review will be requested.
8588
@@ -88,7 +91,7 @@ jobs:
8891
PR_BODY=$(echo "${PR_BODY}" | sed 's/^ //')
8992
9093
gh pr create \
91-
--title "Bump PyTorch pin to nightly ${NIGHTLY}" \
94+
--title "Bump PyTorch and TorchAO pins for ${NIGHTLY}" \
9295
--body "${PR_BODY}" \
9396
--label "ci/pytorch-pin-bump" \
9497
--label "ciflow/cuda"

0 commit comments

Comments
 (0)