Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 26 additions & 9 deletions .ci/scripts/tests/test_cu134_dependencies.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,7 +61,7 @@ def test_all_install_steps_preserve_exact_cu134_selection(self):
"torch==2.14.0.dev20260810+cu134",
"torchvision==0.29.0.dev20260811+cu134",
"torchaudio==2.11.0.dev20260811+cu134",
f"torchao==0.19.0.dev20260907+{ao_variant}",
f"torchao=={self.installer.CU134_TORCHAO_NIGHTLY_VERSION}+{ao_variant}",
}
for index, command in enumerate(commands):
required = (
Expand Down Expand Up @@ -101,7 +101,9 @@ def test_other_cuda_trains_keep_existing_pins(self):
cuda, machine
)
self.assertIn("torch==2.14.0", core)
self.assertIn("torchao==0.19.0.dev20260907", core)
self.assertIn(
f"torchao=={self.installer.TORCHAO_NIGHTLY_VERSION}", core
)
self.assertIn("torchvision==0.29.0", domains)
self.assertIn("torchaudio==2.11.0", domains)
self.assertFalse(any("==" in arg for arg in local))
Expand All @@ -119,7 +121,7 @@ def test_source_pinned_torch_is_not_replaced(self):
def test_no_cuda_keeps_default_pins(self):
core, _, domains, _ = self.install_commands(None)
self.assertIn("torch==2.14.0", core)
self.assertIn("torchao==0.19.0.dev20260907", core)
self.assertIn(f"torchao=={self.installer.TORCHAO_NIGHTLY_VERSION}", core)
self.assertIn("torchvision==0.29.0", domains)
self.assertIn("https://download.pytorch.org/whl/test/cpu", core)

Expand Down Expand Up @@ -239,14 +241,29 @@ def test_cu134_keeps_explicit_torchao_source_build(self):
any(arg.startswith("torchao==") for arg in command)
)
self.assertIn("torch==2.14.0.dev20260810+cu134", commands[-1])
self.assertIn("0.19.0+gitb7ac3aa", metadata.specifier)
source_version = self.installer.TORCHAO_NIGHTLY_VERSION.partition(
".dev"
)[0]
source_commit = subprocess.run(
["git", "rev-parse", "HEAD:third-party/ao"],
cwd=ROOT,
capture_output=True,
check=True,
text=True,
).stdout.strip()
self.assertIn(
f"{source_version}+git{source_commit[:7]}", metadata.specifier
)

def test_wheel_torchao_bound_matches_selected_train(self):
for cuda, expected in (
((13, 4), "torchao>=0.19.0.dev20260907,<0.20"),
((13, 2), "torchao>=0.19.0.dev20260907,<0.20"),
(None, "torchao>=0.19.0.dev20260907,<0.20"),
):
for cuda in ((13, 4), (13, 2), None):
version = (
self.installer.CU134_TORCHAO_NIGHTLY_VERSION
if cuda == (13, 4)
else self.installer.TORCHAO_NIGHTLY_VERSION
)
major, minor = (int(part) for part in version.split(".")[:2])
expected = f"torchao>={version},<{major}.{minor + 1}"
self.utils.determine_torch_url.cache_clear()
with (
patch.object(
Expand Down
132 changes: 124 additions & 8 deletions .github/scripts/update_pytorch_pin.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,17 @@
#!/usr/bin/env python3

import base64
import hashlib
import json
import re
import runpy
import subprocess
import sys
import urllib.request
from pathlib import Path
from urllib.parse import unquote


TORCHAO_INDEX_URL = "https://download.pytorch.org/whl/nightly"


def parse_nightly_version(nightly_version):
Expand Down Expand Up @@ -44,6 +49,14 @@ def get_torch_nightly_version():
return match.group(1)


def get_json(url):
req = urllib.request.Request(url)
req.add_header("Accept", "application/vnd.github.v3+json")
req.add_header("User-Agent", "ExecuTorch-Bot")
with urllib.request.urlopen(req) as response:
return json.loads(response.read().decode())


def get_commit_hash_for_nightly(date_str):
"""
Fetch commit hash from PyTorch nightly branch for a given date.
Expand All @@ -58,13 +71,8 @@ def get_commit_hash_for_nightly(date_str):
params = f"?sha=nightly&per_page=50"
url = api_url + params

req = urllib.request.Request(url)
req.add_header("Accept", "application/vnd.github.v3+json")
req.add_header("User-Agent", "ExecuTorch-Bot")

try:
with urllib.request.urlopen(req) as response:
commits = json.loads(response.read().decode())
commits = get_json(url)
except Exception as e:
print(f"Error fetching commits: {e}", file=sys.stderr)
sys.exit(1)
Expand Down Expand Up @@ -104,6 +112,105 @@ def update_pytorch_pin(commit_hash):
print(f"Updated {pin_file} with commit hash: {commit_hash}")


def get_supported_torchao_channels():
cuda_versions = runpy.run_path("install_utils.py")["SUPPORTED_CUDA_VERSIONS"]
return ["cpu", *(f"cu{major}{minor}" for major, minor in cuda_versions)]


def get_torchao_versions(channel):
url = f"{TORCHAO_INDEX_URL}/{channel}/torchao/"
req = urllib.request.Request(url, headers={"User-Agent": "ExecuTorch-Bot"})
with urllib.request.urlopen(req) as response:
index_html = unquote(response.read().decode())

wheel_tags = {}
pattern = re.compile(
rf"^torchao-(\d+\.\d+\.\d+\.dev\d{{8}})\+{re.escape(channel)}-(.+)\.whl$"
)
for filename in re.findall(r'href="[^"]*/(torchao-[^"]+\.whl)"', index_html):
match = pattern.match(filename)
if match:
wheel_tags.setdefault(match.group(1), set()).add(match.group(2))

required_tags = ("py3-none-any", "aarch64") if channel == "cpu" else ("x86_64",)
return {
version
for version, tags in wheel_tags.items()
if all(any(required in tag for tag in tags) for required in required_tags)
}


def get_latest_torchao_nightly(max_date):
common_versions = None
channels = get_supported_torchao_channels()
for channel in channels:
versions = get_torchao_versions(channel)
common_versions = (
versions if common_versions is None else common_versions & versions
)

candidates = [
version
for version in common_versions or []
if version.rsplit(".dev", 1)[-1] <= max_date
]
if not candidates:
raise ValueError(
f"Could not find a TorchAO nightly on or before {max_date} for "
f"all supported channels: {', '.join(channels)}"
)
return max(candidates, key=lambda version: (version.rsplit(".dev", 1)[-1], version))


def get_torchao_commit_hash(nightly_version):
date = nightly_version.rsplit(".dev", 1)[-1]
formatted_date = parse_nightly_version(f"dev{date}")
url = (
"https://api.github.com/repos/pytorch/ao/actions/workflows/" # @lint-ignore
"build_wheels_linux_x86.yml/runs?event=schedule&status=success&"
f"created={formatted_date}&per_page=100"
)
runs = get_json(url).get("workflow_runs", [])
if not runs:
raise ValueError(
f"Could not find the successful TorchAO wheel build for {nightly_version}"
)
return runs[0]["head_sha"]


def update_torchao_pins(nightly_version, commit_hash):
requirements_path = Path("install_requirements.py")
content = requirements_path.read_text()
content, replacements = re.subn(
r'^(?P<prefix>(?:CU\d+_)?TORCHAO_NIGHTLY_VERSION\s*=\s*["\'])[^"\']+(?P<suffix>["\'])$',
rf"\g<prefix>{nightly_version}\g<suffix>",
content,
flags=re.MULTILINE,
)
if not replacements:
raise ValueError(f"Could not find TorchAO nightly pins in {requirements_path}")
requirements_path.write_text(content)

for command in (
["git", "submodule", "update", "--init", "third-party/ao"],
[
"git",
"-C",
"third-party/ao",
"fetch",
"--depth=1",
"origin",
commit_hash,
],
["git", "-C", "third-party/ao", "checkout", "--detach", commit_hash],
):
subprocess.run(command, check=True)
print(
f"Updated TorchAO nightly pins to {nightly_version} and third-party/ao "
f"to {commit_hash}"
)


def should_skip_file(filename):
"""
Check if a file should be skipped during sync (build files).
Expand Down Expand Up @@ -262,8 +369,17 @@ def main():
# Sync c10 directories from PyTorch
sync_c10_directories(commit_hash)

# Select the newest TorchAO nightly available for every supported CUDA
# channel and align the source submodule with the commit that built it.
max_torchao_date = date_str.replace("-", "")
torchao_version = get_latest_torchao_nightly(max_torchao_date)
print(f"Found TorchAO nightly version: {torchao_version}")
torchao_commit_hash = get_torchao_commit_hash(torchao_version)
print(f"Found TorchAO commit hash: {torchao_commit_hash}")
update_torchao_pins(torchao_version, torchao_commit_hash)

print(
"\n✅ Successfully updated PyTorch commit pin and synced c10 directories!"
"\n✅ Successfully updated PyTorch and TorchAO pins and synced c10 directories!"
)

except Exception as e:
Expand Down
11 changes: 7 additions & 4 deletions .github/workflows/weekly-pytorch-pin-bump.yml
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
name: Weekly PyTorch Pin Bump
name: Weekly PyTorch and TorchAO Pin Bump

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

nice


on:
schedule:
Expand Down Expand Up @@ -54,15 +54,17 @@ jobs:
git config user.email "pytorchbot@users.noreply.github.com"
git checkout -b "${BRANCH}"
git add torch_pin.py
git add install_requirements.py
git add .ci/docker/ci_commit_pins/pytorch.txt
git add runtime/core/portable_type/c10/
git add third-party/ao

if git diff --cached --quiet; then
echo "No changes to commit. Pin is already up to date."
exit 0
fi

git commit -m "Bump PyTorch pin to nightly ${{ steps.nightly.outputs.version }}"
git commit -m "Bump PyTorch and TorchAO pins for ${{ steps.nightly.outputs.version }}"
git push -u origin "${BRANCH}"

EXISTING=$(gh pr list --label "ci/pytorch-pin-bump" --state open --json number --jq '.[0].number')
Expand All @@ -75,11 +77,12 @@ jobs:
read -r -d '' PR_BODY <<EOF || true
## Summary

Automated weekly PyTorch pin bump.
Automated weekly PyTorch and TorchAO pin bump.

- Updates \`NIGHTLY_VERSION\` in \`torch_pin.py\` to \`${NIGHTLY}\`
- Updates \`.ci/docker/ci_commit_pins/pytorch.txt\` to the corresponding nightly commit hash
- Syncs c10 headers from PyTorch into \`runtime/core/portable_type/c10/\`
- Updates the TorchAO nightly pins and \`third-party/ao\` to the newest build available for every supported CUDA channel

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.

Expand All @@ -88,7 +91,7 @@ jobs:
PR_BODY=$(echo "${PR_BODY}" | sed 's/^ //')

gh pr create \
--title "Bump PyTorch pin to nightly ${NIGHTLY}" \
--title "Bump PyTorch and TorchAO pins for ${NIGHTLY}" \
--body "${PR_BODY}" \
--label "ci/pytorch-pin-bump" \
--label "ciflow/cuda"
Loading