diff --git a/.ci/scripts/tests/test_cu134_dependencies.py b/.ci/scripts/tests/test_cu134_dependencies.py index 689b6be3e05..ab424f41212 100644 --- a/.ci/scripts/tests/test_cu134_dependencies.py +++ b/.ci/scripts/tests/test_cu134_dependencies.py @@ -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 = ( @@ -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)) @@ -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) @@ -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( diff --git a/.github/scripts/update_pytorch_pin.py b/.github/scripts/update_pytorch_pin.py index dbc48552d9b..c6a6318e965 100644 --- a/.github/scripts/update_pytorch_pin.py +++ b/.github/scripts/update_pytorch_pin.py @@ -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): @@ -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. @@ -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) @@ -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(?:CU\d+_)?TORCHAO_NIGHTLY_VERSION\s*=\s*["\'])[^"\']+(?P["\'])$', + rf"\g{nightly_version}\g", + 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). @@ -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: diff --git a/.github/workflows/weekly-pytorch-pin-bump.yml b/.github/workflows/weekly-pytorch-pin-bump.yml index 5441ad8b836..0aecdaa7594 100644 --- a/.github/workflows/weekly-pytorch-pin-bump.yml +++ b/.github/workflows/weekly-pytorch-pin-bump.yml @@ -1,4 +1,4 @@ -name: Weekly PyTorch Pin Bump +name: Weekly PyTorch and TorchAO Pin Bump on: schedule: @@ -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') @@ -75,11 +77,12 @@ jobs: read -r -d '' PR_BODY <