Compare commits
1 commit
main
...
torch-211-
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ba2d6ac8d8 |
3 changed files with 95 additions and 0 deletions
24
.github/workflows/version-compat-ci.yml
vendored
24
.github/workflows/version-compat-ci.yml
vendored
|
|
@ -202,6 +202,30 @@ jobs:
|
|||
tests/version_compat/test_transformers_pinned_symbols.py \
|
||||
-v --tb=short
|
||||
|
||||
torch-211-readiness:
|
||||
name: torch 2.11 prebuilt-wheel readiness canary
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 6
|
||||
steps:
|
||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||
with:
|
||||
persist-credentials: false
|
||||
- uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6.2.0
|
||||
with:
|
||||
python-version: '3.12'
|
||||
cache: 'pip'
|
||||
- name: Install pytest only
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install 'pytest>=8'
|
||||
- name: Run torch 2.11 readiness canary
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: |
|
||||
PYTHONPATH=. python -m pytest \
|
||||
tests/version_compat/test_torch_211_prebuilt_readiness.py \
|
||||
-v --tb=short
|
||||
|
||||
# Optional second layer: actually `pip install` ONE representative
|
||||
# version of each package and verify unsloth + unsloth-zoo modules
|
||||
# import on it under the existing CUDA spoof. CPU-only, runs on
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import urllib.error
|
||||
|
|
@ -12,6 +13,23 @@ import urllib.request
|
|||
import pytest
|
||||
|
||||
|
||||
def fetch_json(url: str):
|
||||
"""GET a GitHub API URL and parse JSON. None on 404; skips on transient network errors."""
|
||||
req = urllib.request.Request(url, headers = {"Accept": "application/vnd.github+json"})
|
||||
token = os.environ.get("GITHUB_TOKEN") or os.environ.get("GH_TOKEN")
|
||||
if token:
|
||||
req.add_header("Authorization", f"Bearer {token}")
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout = 15) as r:
|
||||
return json.loads(r.read().decode("utf-8", errors = "replace"))
|
||||
except urllib.error.HTTPError as e:
|
||||
if e.code == 404:
|
||||
return None
|
||||
pytest.skip(f"GitHub API failed ({e.code}) for {url}")
|
||||
except (urllib.error.URLError, TimeoutError, json.JSONDecodeError) as e:
|
||||
pytest.skip(f"GitHub API failed ({e}) for {url}")
|
||||
|
||||
|
||||
def fetch_text(repo: str, ref: str, path: str) -> str | None:
|
||||
"""Fetch a file from GitHub raw. None on 404; skips on transient network errors."""
|
||||
url = f"https://raw.githubusercontent.com/{repo}/{ref}/{path}"
|
||||
|
|
|
|||
53
tests/version_compat/test_torch_211_prebuilt_readiness.py
Normal file
53
tests/version_compat/test_torch_211_prebuilt_readiness.py
Normal file
|
|
@ -0,0 +1,53 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team.
|
||||
"""Canary for raising the CUDA torch pin to 2.11.
|
||||
|
||||
The Studio installer selects torch-version-tagged prebuilt CUDA wheels for three
|
||||
accelerators: flash-attn (wheel_utils.flash_attn_package_version) and causal-conv1d
|
||||
/ mamba-ssm (worker._install_package_wheel_first). None of them publish torch 2.11
|
||||
wheels yet, so the CUDA torch pin stays < 2.11 (_CUDA_TORCH_PKG_SPEC in
|
||||
studio/install_python_stack.py); moving to 2.11 today would drop those wheels and
|
||||
force slow source builds. This test stays green while any of the three is missing
|
||||
torch 2.11 Linux wheels and fails once all three ship them, i.e. when it is finally
|
||||
safe to bump the pin.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
|
||||
from tests.version_compat._fetch import fetch_json
|
||||
|
||||
# human name -> GitHub repo
|
||||
_REPOS = {
|
||||
"flash-attn": "Dao-AILab/flash-attention",
|
||||
"causal-conv1d": "Dao-AILab/causal-conv1d",
|
||||
"mamba-ssm": "state-spaces/mamba",
|
||||
}
|
||||
_TORCH211_LINUX = re.compile(r"torch2\.11.*linux", re.IGNORECASE)
|
||||
|
||||
|
||||
def _has_torch211_linux_wheel(repo: str) -> bool:
|
||||
"""True if a recent release of ``repo`` publishes a torch 2.11 Linux wheel."""
|
||||
releases = fetch_json(f"https://api.github.com/repos/{repo}/releases?per_page=5")
|
||||
if not releases:
|
||||
return False
|
||||
return any(
|
||||
_TORCH211_LINUX.search(asset.get("name", ""))
|
||||
for rel in releases
|
||||
for asset in rel.get("assets", [])
|
||||
)
|
||||
|
||||
|
||||
def test_torch_211_prebuilt_wheels_not_all_ready():
|
||||
status = {name: _has_torch211_linux_wheel(repo) for name, repo in _REPOS.items()}
|
||||
ready = sorted(n for n, ok in status.items() if ok)
|
||||
pending = sorted(n for n, ok in status.items() if not ok)
|
||||
print(f"torch 2.11 wheel readiness -> ready: {ready or 'none'}; pending: {pending}")
|
||||
assert pending, (
|
||||
f"torch 2.11 Linux wheels are now published for all of {ready}. It is time "
|
||||
"to raise the CUDA torch pin: bump _CUDA_TORCH_PKG_SPEC to <2.12.0 (with the "
|
||||
"matching torchvision/torchaudio bounds) in studio/install_python_stack.py, "
|
||||
"add torch 2.11 to wheel_utils.flash_attn_package_version, and bump the "
|
||||
"causal-conv1d / mamba release tags in studio/backend/core/training/worker.py."
|
||||
)
|
||||
Loading…
Add table
Add a link
Reference in a new issue