unsloth/tests/python/test_torchcodec_torch_compat.py
2026-07-26 16:26:11 +00:00

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)}"
)