unsloth/studio/backend/tests/test_torchao_select.py
Thomas Eric 🇧🇷 03cbe211a3
Studio: fix flash-attn and torchao install on Blackwell (sm_100+) GPUs (Closes #6961) (#6970)
* fix: Remove moot has_blackwell_gpu() function

Fixes unslothai/unsloth#6961. This function skipped flash-attn on Blackwell GPUs because no prebuilt wheel existed;
Dao-AILab now ships one and url_exists() already gates resolution.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* fix: use torchao 0.17.0 for Blackwell

Fixes #6961. Torchao 0.16.0's cpp extensions are built against CUDA 12, so on a CUDA-13
torch (cu130 / Blackwell) they fail to load with "libcudart.so.12: cannot
open shared object file". Select 0.17.0 there instead: its cpp targets torch
2.11, so it is skipped cleanly rather than crashing. CUDA-12 / ROCm / CPU
torch 2.10 keeps 0.16.0 and its working kernels.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>

* Condense torchao version-selection comments (no behavior change)

* Support torch 2.11 in the Studio installer via the torch2.10 prebuilt wheels

Map torch 2.11 to the torch2.10 prebuilt wheels for flash-attn, causal-conv1d,
and mamba through wheel_utils.prebuilt_wheel_torch_mm, applied in direct_wheel_url
(filename) and flash_attn_wheel_url (version). Those torch2.10 CUDA wheels load and
pass each project's own test suite on torch 2.11 (verified on B200), so a torch 2.11
environment gets the prebuilt accelerators instead of skipping or building from source.

Raise _CUDA_TORCH_PKG_SPEC to <2.12.0 (torchvision <0.27.0, torchaudio <2.12.0) so
the CUDA torch repair path can install torch 2.11, where torchao 0.17's cpp kernels
load cleanly. Add tests for the mapping.

* Keep has_blackwell_gpu as a False stub for future arch gating

* Restore has_blackwell_gpu as a return-False probe kept for future arch gating

Keep the nvidia-smi compute_cap detection and its two call sites, but short-circuit
with return False at the top so flash-attn is no longer skipped on Blackwell (sm_100+
now has prebuilt wheels and url_exists gates resolution). Drop the early return to
re-enable arch-based detection later.

---------

Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
Co-authored-by: Daniel Han <danielhanchen@gmail.com>
2026-07-08 06:38:10 -07:00

136 lines
5.6 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Tests for _select_torchao_spec in install_python_stack.py.
torchao's C++ extensions are built against one exact torch release, so the
installer must pick the torchao version matching the torch installed in the
venv (otherwise the cpp kernels are skipped). This pins that mapping.
"""
from __future__ import annotations
import sys
from pathlib import Path
from unittest.mock import MagicMock
import pytest
# install_python_stack.py lives at repo_root/studio/install_python_stack.py
_INSTALL_SCRIPT = Path(__file__).resolve().parents[2] / "install_python_stack.py"
def _load_module(monkeypatch):
"""(Re-)import install_python_stack and return it (mirrors test_pytorch_mirror)."""
sys.modules.pop("install_python_stack", None)
monkeypatch.syspath_prepend(str(_INSTALL_SCRIPT.parent))
import install_python_stack
return install_python_stack
@pytest.mark.parametrize(
"torch_version, expected",
[
# torch 2.10 on CUDA <= 12 -> 0.16.0 (its cpp is built for torch 2.10.0 and
# loads against the CUDA-12 PyPI wheel). Independent of patch level.
("2.10.0+cu128", "torchao==0.16.0"),
("2.10.0+cu126", "torchao==0.16.0"),
("2.10.0+rocm6.4", "torchao==0.16.0"),
("2.10.0+cpu", "torchao==0.16.0"),
("2.10.1", "torchao==0.16.0"),
("2.10.0", "torchao==0.16.0"),
# torch 2.10 on CUDA >= 13 (Blackwell / cu130): 0.16.0's CUDA-12 cpp can't
# load against a CUDA-13 torch (libcudart.so.12 error), so use 0.17.0.
("2.10.0+cu130", "torchao==0.17.0"),
("2.10.0+cu140", "torchao==0.17.0"),
# Pre-release / dev / rc builds: the minor is cleaned of non-digits; the
# CUDA tag still decides 0.16.0 vs 0.17.0.
("2.10.0rc1", "torchao==0.16.0"),
("2.10.0.dev20250804+cu130", "torchao==0.17.0"),
("2.10.0.dev20250804+cu128", "torchao==0.16.0"),
("2.10rc1", "torchao==0.16.0"),
# torch 2.11 (reachable via ROCm rocm7.2) and forward -> 0.17.0.
("2.11.0+cu130", "torchao==0.17.0"),
("2.11.0", "torchao==0.17.0"),
("2.12.0", "torchao==0.17.0"),
# torch <=2.9 keeps today's pin (already a correct match for 2.9.0).
("2.9.0+cu128", "torchao==0.14.0"),
("2.9.1", "torchao==0.14.0"),
("2.8.0", "torchao==0.14.0"),
("2.4.0", "torchao==0.14.0"),
# Unparseable / missing / non-2.x major -> conservative default.
(None, "torchao==0.14.0"),
("", "torchao==0.14.0"),
("garbage", "torchao==0.14.0"),
("2", "torchao==0.14.0"),
("3.0.0", "torchao==0.14.0"),
],
)
def test_select_torchao_spec(monkeypatch, torch_version, expected):
mod = _load_module(monkeypatch)
assert mod._select_torchao_spec(torch_version) == expected
def test_default_spec_matches_table(monkeypatch):
"""The default/floor stays the historical pin so older torch is unchanged."""
mod = _load_module(monkeypatch)
assert mod._TORCHAO_DEFAULT_SPEC == "torchao==0.14.0"
assert mod._select_torchao_spec("2.9.0") == mod._TORCHAO_DEFAULT_SPEC
@pytest.mark.parametrize(
("rocm_windows_torch_installed", "installed_torch_is_windows_rocm"),
[
(True, False),
(False, True),
],
)
def test_skips_torchao_on_windows_rocm(
monkeypatch, tmp_path, rocm_windows_torch_installed, installed_torch_is_windows_rocm
):
"""The overrides step must skip torchao on Windows ROCm: no working build exists
there (it imports an absent c10d backend and crashes transformers.quantizers),
so the installer skips it and relies on the runtime stub instead."""
mod = _load_module(monkeypatch)
installed_specs: list[str] = []
progress_labels: list[str] = []
def _record_pip_install(*args, **kwargs):
installed_specs.extend(str(arg) for arg in args)
return 0
unstructured_plugin = tmp_path / "unstructured"
github_plugin = tmp_path / "github"
unstructured_plugin.mkdir()
github_plugin.mkdir()
subprocess_result = MagicMock()
subprocess_result.returncode = 0
subprocess_result.stdout = ""
monkeypatch.setenv("SKIP_STUDIO_BASE", "1")
monkeypatch.setattr(mod, "IS_WINDOWS", True)
monkeypatch.setattr(mod, "IS_MACOS", False)
monkeypatch.setattr(mod, "IS_MAC_ARM", False)
monkeypatch.setattr(mod, "NO_TORCH", False)
monkeypatch.setattr(mod, "_rocm_windows_torch_installed", rocm_windows_torch_installed)
monkeypatch.setattr(
mod, "_installed_torch_is_windows_rocm", lambda: installed_torch_is_windows_rocm
)
monkeypatch.setattr(mod, "_bootstrap_uv", lambda: False)
monkeypatch.setattr(mod, "_repair_bad_anyio", lambda: None)
monkeypatch.setattr(mod, "_ensure_rocm_torch", lambda: None)
monkeypatch.setattr(mod, "_ensure_cuda_torch", lambda: None)
monkeypatch.setattr(mod, "_has_usable_nvidia_gpu", lambda: True)
monkeypatch.setattr(mod, "run", lambda *args, **kwargs: None)
monkeypatch.setattr(mod, "pip_install", _record_pip_install)
monkeypatch.setattr(mod, "_progress", lambda label: progress_labels.append(label))
monkeypatch.setattr(mod, "LOCAL_DD_UNSTRUCTURED_PLUGIN", unstructured_plugin)
monkeypatch.setattr(mod, "LOCAL_DD_GITHUB_PLUGIN", github_plugin)
monkeypatch.setattr(mod.subprocess, "run", lambda *args, **kwargs: subprocess_result)
assert mod.install_python_stack() == 0
assert not any(spec.startswith("torchao") for spec in installed_specs)
assert "dependency overrides (skipped, Windows ROCm)" in progress_labels