319 lines
11 KiB
Python
319 lines
11 KiB
Python
# SPDX-License-Identifier: AGPL-3.0-only
|
|
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
|
|
|
|
"""torch / torchcodec ABI guardrails (unslothai/unsloth#7225)."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import importlib.util
|
|
import re
|
|
import sys
|
|
import types
|
|
from pathlib import Path
|
|
|
|
|
|
REPO_ROOT = Path(__file__).resolve().parents[2]
|
|
PYPROJECT = REPO_ROOT / "pyproject.toml"
|
|
IMPORT_FIXES_PATH = REPO_ROOT / "unsloth" / "import_fixes.py"
|
|
EXTRAS_NO_DEPS_TXT = REPO_ROOT / "studio" / "backend" / "requirements" / "extras-no-deps.txt"
|
|
|
|
|
|
def _load_import_fixes_module():
|
|
spec = importlib.util.spec_from_file_location(
|
|
"unsloth_import_fixes_under_test",
|
|
IMPORT_FIXES_PATH,
|
|
)
|
|
assert spec and spec.loader
|
|
mod = importlib.util.module_from_spec(spec)
|
|
spec.loader.exec_module(mod)
|
|
return mod
|
|
|
|
|
|
def test_pyproject_declares_torch210_audio_extra_with_python_gate():
|
|
text = PYPROJECT.read_text(encoding = "utf-8")
|
|
assert "audio-torch210 = [" in text
|
|
assert "torchcodec>=0.10.0,<0.11.0" in text
|
|
assert "python_version >= '3.10'" in text
|
|
assert "audio-torch290 = [" in text
|
|
assert "audio-torch280 = [" in text
|
|
assert "\naudio = [" not in text
|
|
|
|
|
|
def _stub_torch(monkeypatch, version: str):
|
|
torch_mod = types.ModuleType("torch")
|
|
torch_mod.__version__ = version
|
|
monkeypatch.setitem(sys.modules, "torch", torch_mod)
|
|
|
|
|
|
def test_torch210_extras_bundle_audio_torch210():
|
|
text = PYPROJECT.read_text(encoding = "utf-8")
|
|
for extra in (
|
|
"cu128-torch2100",
|
|
"cu126-ampere-torch2100",
|
|
"rocm72-torch2100",
|
|
):
|
|
match = re.search(rf"^{extra} = \[(.*?)^\]", text, re.MULTILINE | re.DOTALL)
|
|
assert match is not None, extra
|
|
assert "unsloth[audio-torch210]" in match.group(1)
|
|
|
|
|
|
def test_torchcodec_matrix_matches_notebook_validator():
|
|
from scripts import notebook_validator as nv
|
|
fixes = _load_import_fixes_module()
|
|
assert fixes._TORCH_TORCHCODEC_MINORS == nv.TORCH_TORCHCODEC
|
|
|
|
|
|
def test_torchcodec_exclusive_upper_bound():
|
|
fixes = _load_import_fixes_module()
|
|
assert fixes._torchcodec_exclusive_upper("0.10") == "<0.11.0"
|
|
assert fixes._torchcodec_exclusive_upper("0.9") == "<0.10.0"
|
|
|
|
|
|
def test_torch290_rejects_torchcodec_07(monkeypatch):
|
|
import importlib.metadata
|
|
|
|
fixes = _load_import_fixes_module()
|
|
_stub_torch(monkeypatch, "2.9.0+cu128")
|
|
monkeypatch.setattr(importlib.metadata, "version", lambda _name: "0.7.0")
|
|
|
|
hint = fixes._torchcodec_version_mismatch_hint()
|
|
assert hint is not None
|
|
assert "audio-torch210" not in hint
|
|
|
|
|
|
def test_torch280_accepts_torchcodec_07(monkeypatch):
|
|
import importlib.metadata
|
|
|
|
fixes = _load_import_fixes_module()
|
|
_stub_torch(monkeypatch, "2.8.0+cu128")
|
|
monkeypatch.setattr(importlib.metadata, "version", lambda _name: "0.7.0")
|
|
|
|
assert fixes._torchcodec_version_mismatch_hint() is None
|
|
|
|
|
|
def test_torch210_rejects_torchcodec_011(monkeypatch):
|
|
import importlib.metadata
|
|
|
|
fixes = _load_import_fixes_module()
|
|
_stub_torch(monkeypatch, "2.10.0+cu128")
|
|
monkeypatch.setattr(
|
|
importlib.metadata,
|
|
"version",
|
|
lambda _name: "0.11.0",
|
|
)
|
|
|
|
hint = fixes._torchcodec_version_mismatch_hint()
|
|
assert hint is not None
|
|
assert "torchcodec 0.11.0" in hint
|
|
assert "audio-torch210" in hint
|
|
assert "<0.11.0" in hint
|
|
assert "<11.0" not in hint
|
|
|
|
|
|
def test_torch210_accepts_torchcodec_010(monkeypatch):
|
|
import importlib.metadata
|
|
|
|
fixes = _load_import_fixes_module()
|
|
_stub_torch(monkeypatch, "2.10.0+cu128")
|
|
monkeypatch.setattr(
|
|
importlib.metadata,
|
|
"version",
|
|
lambda _name: "0.10.0+cu128",
|
|
)
|
|
|
|
assert fixes._torchcodec_version_mismatch_hint() is None
|
|
|
|
|
|
def test_import_fixes_loads_on_python39_syntax():
|
|
"""Regression: module must import on 3.9 (postponed annotations for str | None)."""
|
|
fixes = _load_import_fixes_module()
|
|
assert callable(fixes._torchcodec_version_mismatch_hint)
|
|
|
|
|
|
# ── torch 2.11 (unslothai/unsloth#7225 follow-up) ──────────────────────
|
|
#
|
|
# install.sh's CUDA branch resolves torch 2.11.0 (the cu12x/cu13x indexes
|
|
# top out there), and torchcodec declares no `torch` dependency, so pip
|
|
# cannot catch a stale 0.10 pairing. These lock the 2.11 row in.
|
|
|
|
|
|
def _load_notebook_validator_module():
|
|
"""Load by path: `from scripts import ...` picks up whatever `scripts`
|
|
package happens to be on sys.path first, which is not always this repo's."""
|
|
spec = importlib.util.spec_from_file_location(
|
|
"unsloth_notebook_validator_under_test",
|
|
REPO_ROOT / "scripts" / "notebook_validator.py",
|
|
)
|
|
assert spec and spec.loader
|
|
mod = importlib.util.module_from_spec(spec)
|
|
# dataclasses resolves annotations through sys.modules, so the module has to
|
|
# be registered under its own name before it executes.
|
|
sys.modules[spec.name] = mod
|
|
try:
|
|
spec.loader.exec_module(mod)
|
|
except Exception:
|
|
sys.modules.pop(spec.name, None)
|
|
raise
|
|
return mod
|
|
|
|
|
|
def _load_install_python_stack():
|
|
studio_dir = REPO_ROOT / "studio"
|
|
if str(studio_dir) not in sys.path:
|
|
sys.path.insert(0, str(studio_dir))
|
|
import install_python_stack
|
|
|
|
return install_python_stack
|
|
|
|
|
|
def test_torch211_rejects_torchcodec_010(monkeypatch):
|
|
"""The guard must not be silent on the torch minor where the mismatch happens."""
|
|
import importlib.metadata
|
|
|
|
fixes = _load_import_fixes_module()
|
|
_stub_torch(monkeypatch, "2.11.0+cu128")
|
|
monkeypatch.setattr(
|
|
importlib.metadata,
|
|
"version",
|
|
lambda _name: "0.10.0+cu128",
|
|
)
|
|
|
|
hint = fixes._torchcodec_version_mismatch_hint()
|
|
assert hint is not None, "torch 2.11 + torchcodec 0.10 must not go unreported"
|
|
assert "torchcodec 0.10.0+cu128" in hint
|
|
assert "audio-torch211" in hint
|
|
assert ">=0.11" in hint
|
|
assert "<0.12.0" in hint
|
|
assert "audio-torch210" not in hint
|
|
|
|
|
|
def test_torch211_accepts_torchcodec_011(monkeypatch):
|
|
import importlib.metadata
|
|
|
|
fixes = _load_import_fixes_module()
|
|
_stub_torch(monkeypatch, "2.11.0+cu128")
|
|
monkeypatch.setattr(
|
|
importlib.metadata,
|
|
"version",
|
|
lambda _name: "0.11.1+cu128",
|
|
)
|
|
|
|
assert fixes._torchcodec_version_mismatch_hint() is None
|
|
|
|
|
|
def test_torch211_accepts_abi_stable_torchcodec(monkeypatch):
|
|
"""torchcodec 0.12+ targets torch >=2.11, so it is not locked to one minor."""
|
|
import importlib.metadata
|
|
|
|
fixes = _load_import_fixes_module()
|
|
for torch_version in ("2.11.0+cu128", "2.12.0", "2.13.0+cu130"):
|
|
for codec_version in ("0.12.0", "0.15.0+cu130"):
|
|
_stub_torch(monkeypatch, torch_version)
|
|
monkeypatch.setattr(importlib.metadata, "version", lambda _name, _v = codec_version: _v)
|
|
assert (
|
|
fixes._torchcodec_version_mismatch_hint() is None
|
|
), f"{torch_version} + torchcodec {codec_version} is supported upstream"
|
|
|
|
|
|
def test_torch210_still_rejects_abi_stable_torchcodec(monkeypatch):
|
|
"""The ABI-stable floor starts at torch 2.11: 2.10 keeps the exact pairing."""
|
|
import importlib.metadata
|
|
|
|
fixes = _load_import_fixes_module()
|
|
_stub_torch(monkeypatch, "2.10.0+cu128")
|
|
monkeypatch.setattr(importlib.metadata, "version", lambda _name: "0.15.0")
|
|
|
|
hint = fixes._torchcodec_version_mismatch_hint()
|
|
assert hint is not None
|
|
assert "audio-torch210" in hint
|
|
|
|
|
|
def test_notebook_validator_allows_abi_stable_pairing():
|
|
"""R-INST-004 is an error-severity rule: it must not fire on torch 2.11 + 0.12+."""
|
|
nv = _load_notebook_validator_module()
|
|
|
|
cell = '!pip install --no-deps "torch==2.11.0" "torchcodec==0.15.0"'
|
|
assert nv.rule_inst_004_torchcodec_torch(cell, {}, "nb.ipynb", 0) == []
|
|
|
|
stale = '!pip install --no-deps "torch==2.11.0" "torchcodec==0.10.0"'
|
|
findings = nv.rule_inst_004_torchcodec_torch(stale, {}, "nb.ipynb", 0)
|
|
assert len(findings) == 1
|
|
assert findings[0].rule == "R-INST-004"
|
|
|
|
|
|
def test_pyproject_declares_torch211_audio_extra_with_python_gate():
|
|
text = PYPROJECT.read_text(encoding = "utf-8")
|
|
match = re.search(r"^audio-torch211 = \[(.*?)^\]", text, re.MULTILINE | re.DOTALL)
|
|
assert match is not None, "pyproject must declare an audio-torch211 extra"
|
|
assert "torchcodec>=0.11.0,<0.12.0" in match.group(1)
|
|
assert "python_version >= '3.10'" in match.group(1)
|
|
|
|
|
|
def test_extras_no_deps_has_no_unconditional_torchcodec_pin():
|
|
"""A flat pin cannot serve both torch lines, so the installer picks the spec."""
|
|
lines = [
|
|
line.strip()
|
|
for line in EXTRAS_NO_DEPS_TXT.read_text(encoding = "utf-8").splitlines()
|
|
if line.strip() and not line.strip().startswith("#")
|
|
]
|
|
assert not any(line.lower().startswith("torchcodec") for line in lines), (
|
|
"extras-no-deps.txt must not pin torchcodec unconditionally; "
|
|
"install_python_stack._select_torchcodec_spec picks it per torch minor"
|
|
)
|
|
|
|
|
|
def test_select_torchcodec_spec_tracks_torch_minor():
|
|
ips = _load_install_python_stack()
|
|
assert ips._select_torchcodec_spec("2.11.0+cu128") == "torchcodec>=0.11.0,<0.12.0"
|
|
assert ips._select_torchcodec_spec("2.10.0+cu130") == "torchcodec>=0.10.0,<0.11.0"
|
|
assert ips._select_torchcodec_spec("2.9.1+cu128") == "torchcodec>=0.8.0,<0.10.0"
|
|
assert ips._select_torchcodec_spec("2.8.0+cu126") == "torchcodec>=0.6.0,<0.8.0"
|
|
|
|
|
|
def test_select_torchcodec_spec_never_caps_newer_torch_to_the_011_line():
|
|
"""0.11 is locked to torch 2.11 exactly and is absent from newer CUDA indexes,
|
|
so torch >2.11 must land on the open ABI-stable floor, not on <0.12.0."""
|
|
ips = _load_install_python_stack()
|
|
for version in ("2.12.0", "2.12.1+cu132", "2.13.0+cu130", "2.99.0"):
|
|
spec = ips._select_torchcodec_spec(version)
|
|
assert spec == ips._TORCHCODEC_ABI_STABLE_SPEC, version
|
|
assert "<" not in spec, version
|
|
|
|
|
|
def test_select_torchcodec_spec_falls_back_on_unknown_torch():
|
|
ips = _load_install_python_stack()
|
|
for value in (None, "", "not-a-version", "3.0.0", "2.rc1"):
|
|
assert ips._select_torchcodec_spec(value) == ips._TORCHCODEC_DEFAULT_SPEC
|
|
|
|
|
|
def test_select_torchcodec_spec_matches_pyproject_audio_extras():
|
|
"""The installer's specs and the pip extras must not drift apart."""
|
|
ips = _load_install_python_stack()
|
|
text = PYPROJECT.read_text(encoding = "utf-8")
|
|
for torch_version, extra in (
|
|
("2.11.0", "audio-torch211"),
|
|
("2.10.0", "audio-torch210"),
|
|
("2.9.0", "audio-torch290"),
|
|
("2.8.0", "audio-torch280"),
|
|
):
|
|
match = re.search(rf"^{extra} = \[(.*?)^\]", text, re.MULTILINE | re.DOTALL)
|
|
assert match is not None, extra
|
|
assert ips._select_torchcodec_spec(torch_version) in match.group(1), extra
|
|
|
|
|
|
def test_select_torchcodec_spec_matches_compat_matrix():
|
|
"""Installer specs must admit exactly the minors the compat matrix allows."""
|
|
from packaging.specifiers import SpecifierSet
|
|
|
|
fixes = _load_import_fixes_module()
|
|
ips = _load_install_python_stack()
|
|
probes = [f"0.{n}.0" for n in range(0, 16)]
|
|
for torch_minor, allowed in fixes._TORCH_TORCHCODEC_MINORS.items():
|
|
specifier = SpecifierSet(
|
|
ips._select_torchcodec_spec(f"{torch_minor}.0").split("torchcodec", 1)[1]
|
|
)
|
|
admitted = {p.rsplit(".", 1)[0] for p in probes if specifier.contains(p)}
|
|
assert admitted == allowed, (
|
|
f"torch {torch_minor}: installer admits {sorted(admitted)}, "
|
|
f"matrix allows {sorted(allowed)}"
|
|
)
|