unsloth/tests/studio/install/test_rocm_support.py
Daniel Han 827d25931b
Stop 19 test files racing on one PowerShell startup cache (#9371)
* Stop 19 test files racing on one PowerShell startup cache

Backend CI run 32341628757 on `1c3dde199` finished `284 failed, 8498 passed`. Every
one of the 284 was a pwsh subprocess ending `died with <Signals.SIGABRT: 6>`, across
19 files that all read as Windows-installer regressions. None of them were. 222 of
the aborts land inside a two-second window, 88 at 07:09:30 and 133 at 07:09:31,
which is a mass kill of every live pwsh rather than independent per-test flakiness.

The cause
------------------------------------------------------------------------
Every `-NonInteractive` startup reads and rewrites an ~83 KB
`$XDG_CACHE_HOME/powershell/StartupProfileData-NonInteractive`, and XDG_CACHE_HOME
defaults to `$HOME/.cache`. Under `-n 4` all four xdist workers share one HOME, so
the whole job's pwsh processes race on one file and a startup that deserialises a
half-written one dies before it reaches our script. `Stack overflow.` is .NET's
failfast, which cannot unwind a blown stack, so it prints one line and calls
abort(); that is the SIGABRT (PowerShell/PowerShell#24461).

Measured twice, independently, 4000 startups per arm:

  run 1  shared cache dir     7/4000 died  {-11: 3, -6: 4}
         private cache dirs   0/4000
  run 2  shared cache dir    11/4000 died  {-11: 10, -6: 1}
         private cache dirs   0/4000

Three distinct crash shapes appeared, and each names the torn file rather than our
scripts: `Stack overflow.`, `System.IO.FileLoadException: The given assembly name`,
and `System.ArgumentException: String cannot have zero length.` 18 deaths in 8000
shared startups, 0 in 8000 private.

CI agrees from the other direction. Of the pwsh-heavy files in that run, exactly one
had zero failures, tests/test_windows_amd_gpu_scan_fallback.py, and it is the only
one that hands its child a private HOME, across roughly 80 startups where the run's
own rate predicts about 16 failures.

What is NOT established
------------------------------------------------------------------------
Neither experiment reproduces CI's rate. Roughly 20% of pwsh startups died there
against 0.2 to 0.3% here, and at CI's actual `-n 4` on this box I measured 0/1200 in
both arms: the race needed 48-way concurrency before it appeared at all. The likely
reason is that four workers on a 4-core runner are in real contention while four
threads on a 192-core box almost never overlap in the critical section, but that is
reasoning and not a measurement, so treat the mechanism as proven and the magnitude
as unexplained. That is also why this does not stop at removing the shared file.

Three layers, in order
------------------------------------------------------------------------
1. Remove the contended resource. One cache directory per xdist worker, fresh per
   session. Workers run their tests one at a time, so within a worker the startups
   stay sequential and the cache still does its job warm; across workers the
   directories are disjoint and there is nothing left to race on. Fresh rather than a
   stable path, because a cache torn by an earlier run would otherwise poison every
   later session on the same box.
2. Retry a run that produced no verdict. Three attempts, unslept, because the trigger
   is process startup rather than a resource that frees up.
3. Attribute what is left. A crash raises PwshInterpreterCrash naming the interpreter.

Layer 1 is the fix; 2 and 3 exist because of the unexplained magnitude above.

Deliberately NOT done: bounding pwsh concurrency with a lock, or giving up `-n 4`.
The workflow records 806.1s to 219.7s from that flag, and the contended resource can
be removed rather than rationed.

The rule that keeps this honest
------------------------------------------------------------------------
A signal is not a verdict, so retrying it papers over nothing: the script never ran
to its end. A normal exit is returned untouched on the first attempt whatever its
code, so a pwsh that runs and gives the WRONG answer still fails with its own
message. Getting that second half wrong would turn this into a way to retry real
regressions into green, which is worse than the bug it fixes, so both directions are
executed in tests/studio/test_pwsh_interpreter_crash_attribution.py against a real
SIGABRT rather than reviewed.

Mutation-tested: relaxing the crash test from `returncode < 0` to `returncode != 0`
fails test_a_clean_run_with_the_wrong_answer_still_fails_with_its_own_message and
test_a_clean_run_is_not_retried, which are exactly the two that guard that direction.

This also generalises `_run_pwsh` from tests/studio/test_install_phase_timing.py,
added earlier today for a second, signal-free shape: pwsh printing its "The
PowerShell process will exit" banner and exiting normally with empty stdout. That one
cannot be seen in the exit status, so it stays a text match.

Verified
------------------------------------------------------------------------
tests/python/test_windows_xformers_installer.py, tests/studio/test_install_phase_timing.py,
tests/studio/install/ and the new guard: 2635 passed, 3 skipped.
Guard alone: 5 passed.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Drop the subprocess import the pwsh conversion left behind

Source lint's import-hoist check is right: every subprocess.run in
test_windows_xformers_installer.py became run_pwsh, so `import subprocess` has no
references left except the one inside a comment explaining why run_pwsh is used
instead. Its wording names the shape exactly -- "was used before, now unused
(references re-pointed)" -- which is what a mechanical call-site rewrite leaves
behind.

Swept the other 18 converted files the same way with an AST pass rather than by
eye. This was the only real one: the remaining hits are `from __future__ import
annotations`, which every such scan reports, and a PropertyMock in
test_rocm_support.py that is present on main unchanged.

42 passed.

* Suppress the core dump on the forged SIGABRT

tests/test_deliberate_crashes_suppress_cores.py caught this: the abort child had no
PR_SET_DUMPABLE=0, so each of these aborts piped a multi-MB core to apport before the
child could be reaped. The guard is right and its message names the fix.

The child still exits -6 and PR_GET_DUMPABLE reads 0, so all five verdicts are
unchanged. Linux-only and non-fatal elsewhere: Windows has no CDLL(None) and pipes no
core, so arming it there would trade a no-op for a lost test.

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: danielhanchen <unslothai@gmail.com>
2026-08-20 04:22:41 -07:00

6674 lines
331 KiB
Python

"""AMD ROCm support tests across install pathways (all mocked, no AMD HW)."""
import contextlib
import importlib.util
import io
import json
import os
import re
import shutil
import subprocess
import sys
import tempfile
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import MagicMock, mock_open, patch, PropertyMock
import pytest
from unsloth_pwsh_runner import run_pwsh
# ── Load modules under test ──────────────────────────────────────────────────
PACKAGE_ROOT = Path(__file__).resolve().parents[3]
# install_llama_prebuilt.py
_PREBUILT_PATH = PACKAGE_ROOT / "studio" / "install_llama_prebuilt.py"
_PREBUILT_SPEC = importlib.util.spec_from_file_location(
"studio_install_llama_prebuilt", _PREBUILT_PATH
)
assert _PREBUILT_SPEC is not None and _PREBUILT_SPEC.loader is not None
prebuilt_mod = importlib.util.module_from_spec(_PREBUILT_SPEC)
sys.modules[_PREBUILT_SPEC.name] = prebuilt_mod
_PREBUILT_SPEC.loader.exec_module(prebuilt_mod)
HostInfo = prebuilt_mod.HostInfo
AssetChoice = prebuilt_mod.AssetChoice
PrebuiltFallback = prebuilt_mod.PrebuiltFallback
resolve_upstream_asset_choice = prebuilt_mod.resolve_upstream_asset_choice
runtime_patterns_for_choice = prebuilt_mod.runtime_patterns_for_choice
_apply_host_overrides = prebuilt_mod._apply_host_overrides
_normalize_forwarded_gfx = prebuilt_mod._normalize_forwarded_gfx
# install_python_stack.py
_STACK_PATH = PACKAGE_ROOT / "studio" / "install_python_stack.py"
_STACK_SPEC = importlib.util.spec_from_file_location("studio_install_python_stack", _STACK_PATH)
assert _STACK_SPEC is not None and _STACK_SPEC.loader is not None
stack_mod = importlib.util.module_from_spec(_STACK_SPEC)
sys.modules[_STACK_SPEC.name] = stack_mod
_STACK_SPEC.loader.exec_module(stack_mod)
# The probe prints its answer behind this marker, so chatter on either side of it cannot
# be mistaken for the answer. Mocked stdout has to carry it too.
_MARK = stack_mod._TORCH_PROBE_MARKER
_detect_rocm_version = stack_mod._detect_rocm_version
_ensure_rocm_torch = stack_mod._ensure_rocm_torch
_has_rocm_gpu = stack_mod._has_rocm_gpu
_has_usable_nvidia_gpu = stack_mod._has_usable_nvidia_gpu
_ROCM_TORCH_INDEX = stack_mod._ROCM_TORCH_INDEX
_windows_rocm_index_url = stack_mod._windows_rocm_index_url
_detect_windows_gfx_arch = stack_mod._detect_windows_gfx_arch
_install_bnb_windows_rocm = stack_mod._install_bnb_windows_rocm
@pytest.fixture(autouse = True)
def _reset_torch_runtime_probe():
"""The torch classification is memoized for the life of an install run, so one
test's mocked probe must not leak into the next."""
stack_mod._invalidate_torch_runtime_probe()
yield
stack_mod._invalidate_torch_runtime_probe()
def _extract_sh_function_body(source: str, name: str) -> str:
"""Return a shell function body from `source` by brace matching."""
needle = f"{name}() {{"
start = source.find(needle)
if start < 0:
return ""
depth = 0
i = start + len(needle) - 1 # land on the opening brace
n = len(source)
while i < n:
ch = source[i]
if ch == "{":
depth += 1
elif ch == "}":
depth -= 1
if depth == 0:
return source[start : i + 1]
i += 1
return source[start:]
# A dpkg-query -W stand-in that renders whichever showformat string it is handed,
# so this tests how install.sh ASKS for the version, not only how it parses the
# answer. It answers for a package in ANY state because the real tool does: only
# purged ones are left out, so a rocm-core removed with `apt remove` and never
# purged keeps reporting the version it had.
_DPKG_QUERY_STUB = r"""#!/bin/sh
_status='__STATUS__'
_ver='__VERSION__'
# dpkg's Status field is "<want> <error-flag> <status>".
case "$_status" in installed) _want=install ;; *) _want=deinstall ;; esac
_fmt=''
_found=''
while [ $# -gt 0 ]; do
case "$1" in
-f=*) _fmt=${1#-f=} ;;
--showformat=*) _fmt=${1#--showformat=} ;;
-f|--showformat) shift; _fmt=$1 ;;
-*) : ;;
rocm-core) _found=1 ;;
esac
shift
done
[ -n "$_found" ] || exit 1
[ -n "$_fmt" ] || _fmt='${Package}\t${Version}\n'
# Unrecognised fields render empty, like the real dpkg-query.
_out=$(printf '%s' "$_fmt" | sed \
-e "s|\${Package}|rocm-core|g" \
-e "s|\${Status}|$_want ok $_status|g" \
-e "s|\${db:Status-Status}|$_status|g" \
-e "s|\${db:Status-Want}|$_want|g" \
-e "s|\${db:Status-Eflag}|ok|g" \
-e "s|\${Version}|$_ver|g" \
-e "s|\${[^}]*}||g")
printf "$_out"
"""
def _write_dpkg_query_stub(
path: str,
version: str,
status: str = "installed",
) -> None:
with open(path, "w", encoding = "utf-8") as f:
f.write(_DPKG_QUERY_STUB.replace("__STATUS__", status).replace("__VERSION__", version))
os.chmod(path, 0o755)
# ── Helper: build HostInfo for different scenarios ──────────────────────────
def nvidia_host(**overrides) -> HostInfo:
"""NVIDIA Linux x86_64 host."""
defaults = dict(
system = "Linux",
machine = "x86_64",
is_windows = False,
is_linux = True,
is_macos = False,
is_x86_64 = True,
is_arm64 = False,
nvidia_smi = "/usr/bin/nvidia-smi",
driver_cuda_version = (12, 6),
compute_caps = ["89"],
visible_cuda_devices = None,
has_physical_nvidia = True,
has_usable_nvidia = True,
has_rocm = False,
)
defaults.update(overrides)
return HostInfo(**defaults)
def rocm_host(**overrides) -> HostInfo:
"""AMD ROCm Linux x86_64 host (no NVIDIA)."""
defaults = dict(
system = "Linux",
machine = "x86_64",
is_windows = False,
is_linux = True,
is_macos = False,
is_x86_64 = True,
is_arm64 = False,
nvidia_smi = None,
driver_cuda_version = None,
compute_caps = [],
visible_cuda_devices = None,
has_physical_nvidia = False,
has_usable_nvidia = False,
has_rocm = True,
)
defaults.update(overrides)
return HostInfo(**defaults)
def cpu_host(**overrides) -> HostInfo:
"""CPU-only Linux x86_64 host."""
defaults = dict(
system = "Linux",
machine = "x86_64",
is_windows = False,
is_linux = True,
is_macos = False,
is_x86_64 = True,
is_arm64 = False,
nvidia_smi = None,
driver_cuda_version = None,
compute_caps = [],
visible_cuda_devices = None,
has_physical_nvidia = False,
has_usable_nvidia = False,
has_rocm = False,
)
defaults.update(overrides)
return HostInfo(**defaults)
def macos_host(**overrides) -> HostInfo:
"""macOS arm64 host."""
defaults = dict(
system = "Darwin",
machine = "arm64",
is_windows = False,
is_linux = False,
is_macos = True,
is_x86_64 = False,
is_arm64 = True,
nvidia_smi = None,
driver_cuda_version = None,
compute_caps = [],
visible_cuda_devices = None,
has_physical_nvidia = False,
has_usable_nvidia = False,
has_rocm = False,
)
defaults.update(overrides)
return HostInfo(**defaults)
def windows_host(**overrides) -> HostInfo:
"""Windows x86_64 host."""
defaults = dict(
system = "Windows",
machine = "amd64",
is_windows = True,
is_linux = False,
is_macos = False,
is_x86_64 = True,
is_arm64 = False,
nvidia_smi = None,
driver_cuda_version = None,
compute_caps = [],
visible_cuda_devices = None,
has_physical_nvidia = False,
has_usable_nvidia = False,
has_rocm = False,
)
defaults.update(overrides)
return HostInfo(**defaults)
def windows_rocm_host(**overrides) -> HostInfo:
"""Windows x86_64 host with ROCm."""
defaults = dict(
system = "Windows",
machine = "amd64",
is_windows = True,
is_linux = False,
is_macos = False,
is_x86_64 = True,
is_arm64 = False,
nvidia_smi = None,
driver_cuda_version = None,
compute_caps = [],
visible_cuda_devices = None,
has_physical_nvidia = False,
has_usable_nvidia = False,
has_rocm = True,
)
defaults.update(overrides)
return HostInfo(**defaults)
# ── Upstream asset fixture ───────────────────────────────────────────────────
LLAMA_TAG = "b8508"
UPSTREAM_ASSETS = {
f"llama-{LLAMA_TAG}-bin-ubuntu-x64.tar.gz": f"https://example.com/{LLAMA_TAG}-linux-cpu.tar.gz",
f"llama-{LLAMA_TAG}-bin-ubuntu-rocm-7.2-x64.tar.gz": f"https://example.com/{LLAMA_TAG}-linux-rocm.tar.gz",
f"llama-{LLAMA_TAG}-bin-win-cpu-x64.zip": f"https://example.com/{LLAMA_TAG}-win-cpu.zip",
f"llama-{LLAMA_TAG}-bin-win-cuda-12.4-x64.zip": f"https://example.com/{LLAMA_TAG}-win-cuda.zip",
f"llama-{LLAMA_TAG}-bin-win-hip-radeon-x64.zip": f"https://example.com/{LLAMA_TAG}-win-hip.zip",
f"llama-{LLAMA_TAG}-bin-macos-arm64.tar.gz": f"https://example.com/{LLAMA_TAG}-macos-arm64.tar.gz",
f"llama-{LLAMA_TAG}-bin-macos-x64.tar.gz": f"https://example.com/{LLAMA_TAG}-macos-x64.tar.gz",
}
# TEST: install_llama_prebuilt.py -- resolve_upstream_asset_choice
class TestResolveUpstreamAssetChoice:
"""Verify that the asset selection logic picks the right binary for each platform."""
# The plain cpu-linux / windows-cpu / macos-arm64 routing cases live in
# test_selection_logic.py::TestResolveUpstreamAssetChoice (exact-name pins);
# this class keeps the ROCm/NVIDIA-precedence dialect only.
@patch.object(prebuilt_mod, "github_release_assets", return_value = UPSTREAM_ASSETS)
def test_nvidia_linux_gets_cpu_asset(self, mock_assets):
"""NVIDIA host should NOT hit the ROCm path -- gets CPU asset (CUDA handled elsewhere)."""
host = nvidia_host()
choice = resolve_upstream_asset_choice(host, LLAMA_TAG)
assert choice.install_kind == "linux-cpu"
assert "ubuntu-x64" in choice.name
assert "rocm" not in choice.name
@patch.object(prebuilt_mod, "github_release_assets", return_value = UPSTREAM_ASSETS)
def test_rocm_linux_gets_rocm_prebuilt(self, mock_assets):
"""AMD ROCm Linux host should get the ROCm prebuilt."""
host = rocm_host()
choice = resolve_upstream_asset_choice(host, LLAMA_TAG)
assert choice.install_kind == "linux-rocm"
assert "rocm" in choice.name
@patch.object(prebuilt_mod, "github_release_assets", return_value = UPSTREAM_ASSETS)
def test_windows_rocm_gets_hip_asset(self, mock_assets):
"""Windows ROCm host should get Windows HIP asset."""
host = windows_rocm_host()
choice = resolve_upstream_asset_choice(host, LLAMA_TAG)
assert choice.install_kind == "windows-hip"
assert "hip" in choice.name
@patch.object(prebuilt_mod, "github_release_assets", return_value = UPSTREAM_ASSETS)
def test_mixed_nvidia_rocm_prefers_nvidia(self, mock_assets):
"""Host with both NVIDIA and ROCm should use NVIDIA (CPU path here, CUDA elsewhere)."""
host = nvidia_host(has_rocm = True)
choice = resolve_upstream_asset_choice(host, LLAMA_TAG)
assert choice.install_kind == "linux-cpu"
assert "rocm" not in choice.name
@patch.object(prebuilt_mod, "github_release_assets")
def test_rocm_linux_no_prebuilt_falls_back(self, mock_assets):
"""AMD ROCm host should fall back to source build when no ROCm prebuilt exists."""
assets_without_rocm = {k: v for k, v in UPSTREAM_ASSETS.items() if "rocm" not in k}
mock_assets.return_value = assets_without_rocm
host = rocm_host()
with pytest.raises(PrebuiltFallback, match = "ROCm detected"):
resolve_upstream_asset_choice(host, LLAMA_TAG)
@patch.object(prebuilt_mod, "github_release_assets")
def test_windows_rocm_no_hip_falls_to_cpu(self, mock_assets):
"""Windows+ROCm with HIP prebuilt missing should fall through to CPU."""
assets_no_hip = {k: v for k, v in UPSTREAM_ASSETS.items() if "hip" not in k}
mock_assets.return_value = assets_no_hip
host = windows_rocm_host()
choice = resolve_upstream_asset_choice(host, LLAMA_TAG)
assert choice.install_kind == "windows-cpu"
@patch.object(prebuilt_mod, "github_release_assets", return_value = UPSTREAM_ASSETS)
def test_macos_rocm_impossible_has_rocm_false(self, mock_assets):
"""macOS host should never have has_rocm=True in practice; verify it gets macOS asset."""
host = macos_host(has_rocm = True)
choice = resolve_upstream_asset_choice(host, LLAMA_TAG)
assert choice.install_kind == "macos-arm64"
@patch.object(prebuilt_mod, "github_release_assets", return_value = UPSTREAM_ASSETS)
def test_linux_aarch64_rocm_gets_prebuilt_fallback(self, mock_assets):
"""Linux aarch64 with ROCm -- no x86_64 match, should raise PrebuiltFallback."""
host = rocm_host(machine = "aarch64", is_x86_64 = False, is_arm64 = True)
with pytest.raises(PrebuiltFallback):
resolve_upstream_asset_choice(host, LLAMA_TAG)
# TEST: install_llama_prebuilt.py -- runtime_patterns_for_choice
class TestRuntimePatterns:
"""Verify runtime file patterns for all install kinds."""
def test_linux_cpu_patterns(self):
choice = AssetChoice(
repo = "", tag = "", name = "", url = "", source_label = "", install_kind = "linux-cpu"
)
patterns = runtime_patterns_for_choice(choice)
assert "llama-server" in patterns
assert "llama-quantize" in patterns
# lib*.so* covers libllama/libggml/libmtmd plus the libllama-*-impl.so
# split from ggml-org/llama.cpp #23462 (between b9279 and b9283).
assert "lib*.so*" in patterns
def test_linux_cuda_patterns(self):
choice = AssetChoice(
repo = "", tag = "", name = "", url = "", source_label = "", install_kind = "linux-cuda"
)
patterns = runtime_patterns_for_choice(choice)
assert "lib*.so*" in patterns
def test_linux_rocm_patterns(self):
choice = AssetChoice(
repo = "", tag = "", name = "", url = "", source_label = "", install_kind = "linux-rocm"
)
patterns = runtime_patterns_for_choice(choice)
assert "lib*.so*" in patterns
assert "llama-server" in patterns
def test_windows_hip_patterns(self):
choice = AssetChoice(
repo = "",
tag = "",
name = "",
url = "",
source_label = "",
install_kind = "windows-hip",
)
patterns = runtime_patterns_for_choice(choice)
# Narrowed from "*.exe" to the two binaries Unsloth actually invokes.
assert "llama-server.exe" in patterns
assert "llama-quantize.exe" in patterns
assert "*.dll" in patterns
def test_macos_patterns(self):
choice = AssetChoice(
repo = "",
tag = "",
name = "",
url = "",
source_label = "",
install_kind = "macos-arm64",
)
patterns = runtime_patterns_for_choice(choice)
assert "lib*.dylib" in patterns
def test_diffusion_visual_server_kept(self):
# The DiffusionGemma visual-server must survive the prune so Unsloth can
# serve DiffusionGemma GGUFs natively.
for kind, name in (
("linux-cuda", "llama-diffusion-gemma-visual-server"),
("macos-arm64", "llama-diffusion-gemma-visual-server"),
("windows-cuda", "llama-diffusion-gemma-visual-server.exe"),
):
choice = AssetChoice(
repo = "", tag = "", name = "", url = "", source_label = "", install_kind = kind
)
assert name in runtime_patterns_for_choice(choice)
# TEST: install_llama_prebuilt.py -- HostInfo.has_rocm field
class TestHostInfoRocm:
"""Verify has_rocm field does not affect other HostInfo behavior."""
def test_has_rocm_default_false(self):
host = HostInfo(
system = "Linux",
machine = "x86_64",
is_windows = False,
is_linux = True,
is_macos = False,
is_x86_64 = True,
is_arm64 = False,
nvidia_smi = None,
driver_cuda_version = None,
compute_caps = [],
visible_cuda_devices = None,
has_physical_nvidia = False,
has_usable_nvidia = False,
)
assert host.has_rocm is False
def test_has_rocm_explicit_true(self):
host = rocm_host()
assert host.has_rocm is True
def test_nvidia_host_no_rocm(self):
host = nvidia_host()
assert host.has_rocm is False
assert host.has_usable_nvidia is True
def test_detect_host_has_rocm_detection_logic(self):
"""detect_host() should have ROCm GPU detection logic."""
import inspect
source = inspect.getsource(prebuilt_mod.detect_host)
# Must probe for actual GPU, not just tool presence.
assert "rocminfo" in source or "amd-smi" in source
def test_detect_host_windows_rocm_detection(self):
"""detect_host() source should have Windows-specific ROCm GPU detection."""
import inspect
source = inspect.getsource(prebuilt_mod.detect_host)
assert "hipinfo" in source or "amd-smi" in source
# TEST: install_python_stack.py -- _detect_rocm_version
class TestDetectRocmVersion:
"""Verify ROCm version detection from various sources."""
def test_no_rocm_returns_none(self, tmp_path):
"""No ROCm installed should return None."""
with patch.dict(os.environ, {"ROCM_PATH": str(tmp_path / "nonexistent")}):
with patch("shutil.which", return_value = None):
result = _detect_rocm_version()
assert result is None
def test_version_from_file(self, tmp_path):
"""Reads version from /opt/rocm/.info/version."""
info_dir = tmp_path / ".info"
info_dir.mkdir()
(info_dir / "version").write_text("7.1.0-12345\n")
with patch.dict(os.environ, {"ROCM_PATH": str(tmp_path)}):
result = _detect_rocm_version()
assert result == (7, 1)
def test_version_62(self, tmp_path):
"""Reads ROCm 6.2 version."""
info_dir = tmp_path / ".info"
info_dir.mkdir()
(info_dir / "version").write_text("6.2.0\n")
with patch.dict(os.environ, {"ROCM_PATH": str(tmp_path)}):
result = _detect_rocm_version()
assert result == (6, 2)
def test_hipconfig_fallback(self, tmp_path):
"""Falls back to hipconfig --version when file not found."""
with patch.dict(os.environ, {"ROCM_PATH": str(tmp_path / "nonexistent")}):
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"6.3.21234.2\n"
with patch("shutil.which", return_value = "/usr/bin/hipconfig"):
with patch("subprocess.run", return_value = mock_result):
result = _detect_rocm_version()
assert result == (6, 3)
def test_dpkg_fallback_without_hipconfig(self, tmp_path):
"""dpkg rocm-core fallback works when amd-smi and hipconfig are absent
(regression: a shadowing local re import raised UnboundLocalError)."""
def which(cmd):
return "/usr/bin/dpkg-query" if cmd == "dpkg-query" else None
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = "1:6.3.0-1\n"
with patch.dict(os.environ, {"ROCM_PATH": str(tmp_path / "nonexistent")}):
with patch("shutil.which", side_effect = which):
with patch("subprocess.run", return_value = mock_result):
assert _detect_rocm_version() == (6, 3)
def test_empty_version_file(self, tmp_path):
"""Empty version file should return None."""
info_dir = tmp_path / ".info"
info_dir.mkdir()
(info_dir / "version").write_text("")
with patch.dict(os.environ, {"ROCM_PATH": str(tmp_path)}):
with patch("shutil.which", return_value = None):
result = _detect_rocm_version()
assert result is None
def test_version_with_epoch_prefix(self, tmp_path):
"""Debian epoch prefix (2:6.2.0) -- version file has no epoch, so should parse."""
info_dir = tmp_path / ".info"
info_dir.mkdir()
(info_dir / "version").write_text("6.2.0\n")
with patch.dict(os.environ, {"ROCM_PATH": str(tmp_path)}):
result = _detect_rocm_version()
assert result == (6, 2)
def test_multiple_version_sources_first_wins(self, tmp_path):
"""When both .info/version and lib/rocm_version exist, first found wins."""
info_dir = tmp_path / ".info"
info_dir.mkdir()
(info_dir / "version").write_text("7.1.0\n")
lib_dir = tmp_path / "lib"
lib_dir.mkdir()
(lib_dir / "rocm_version").write_text("6.3.0\n")
with patch.dict(os.environ, {"ROCM_PATH": str(tmp_path)}):
result = _detect_rocm_version()
assert result == (7, 1) # .info/version checked first
def test_hipconfig_multiline_output(self, tmp_path):
"""hipconfig with multi-line output -- should use first line."""
with patch.dict(os.environ, {"ROCM_PATH": str(tmp_path / "nonexistent")}):
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"6.3.21234.2\nSome extra info\n"
with patch("shutil.which", return_value = "/usr/bin/hipconfig"):
with patch("subprocess.run", return_value = mock_result):
result = _detect_rocm_version()
assert result == (6, 3)
def test_hipconfig_timeout(self, tmp_path):
"""hipconfig that times out should return None."""
with patch.dict(os.environ, {"ROCM_PATH": str(tmp_path / "nonexistent")}):
with patch("shutil.which", return_value = "/usr/bin/hipconfig"):
with patch(
"subprocess.run",
side_effect = subprocess.TimeoutExpired("hipconfig", 5),
):
result = _detect_rocm_version()
assert result is None
# TEST: install_python_stack.py -- _ensure_rocm_torch
class TestEnsureRocmTorch:
"""Verify ROCm torch reinstall logic."""
# _infer_linux_amd_gfx_arch mocked to None: on a real Strix host the live
# /proc/cpuinfo would otherwise take the inferred-install path and break
# these "must not install" hosts (environment leak, not the code under test).
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_infer_linux_amd_gfx_arch", return_value = None)
def test_no_rocm_skips(self, mock_infer, mock_nvidia, mock_pip):
"""No ROCm toolchain should skip entirely."""
# Pin _detect_windows_gfx_arch to None so a real AMD test host's WMI
# fallback can't defeat the "no ROCm anywhere" premise.
with patch.object(stack_mod, "_detect_windows_gfx_arch", return_value = None):
with patch("os.path.isdir", return_value = False):
with patch("shutil.which", return_value = None):
_ensure_rocm_torch()
mock_pip.assert_not_called()
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = False)
@patch.object(stack_mod, "_infer_linux_amd_gfx_arch", return_value = "gfx1151")
@patch.object(stack_mod, "_detect_rocm_version", return_value = None)
def test_inferred_gfx_without_rocm_runtime_installs_amd_index(
self, mock_ver, mock_infer, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""Strix Halo without /dev/kfd must still get AMD gfx1151 wheels (unslothai#7301)."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.10.0+cpu||\n"
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
torch_call = str(mock_pip.call_args_list[0])
assert "gfx1151" in torch_call
assert "torch>=2.11.0,<2.12.0" in torch_call
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = False)
@patch.object(stack_mod, "_infer_linux_amd_gfx_arch", return_value = "gfx1151")
@patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = [])
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 1))
def test_inferred_gfx_not_overwritten_when_rocm_userland_readable(
self, mock_ver, mock_gfx, mock_infer, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""Codex P1 #7305: after an inferred per-arch install, do not fall through to the
generic pytorch.org/rocmX.Y reinstall just because has_hip_torch is still False.
Readable ROCm userland without /dev/kfd is exactly the case that used to overwrite
the AMD gfx wheels."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.10.0+cpu||\n"
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
assert mock_pip.call_count == 1, mock_pip.call_args_list
torch_call = str(mock_pip.call_args_list[0])
assert "gfx1151" in torch_call
assert "rocm7.1" not in torch_call
assert "download.pytorch.org" not in torch_call
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_infer_linux_amd_gfx_arch", return_value = "gfx1151")
@patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx1100"])
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 1))
def test_inference_yields_to_runtime_visible_gpu(
self, mock_ver, mock_gfx, mock_infer, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""When the runtime CAN enumerate a GPU, the cpuinfo inference must not
install wheels: a mixed Strix APU + dGPU box with the dGPU selected would
otherwise get gfx1151 wheels for a gfx1100 GPU. The runtime-visible arch
(Strix override / generic branch) decides instead."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.10.0+cpu||\n"
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
all_calls = str(mock_pip.call_args_list) + str(mock_pip_try.call_args_list)
assert "gfx1151" not in all_calls, all_calls
assert "rocm7.1" in all_calls, all_calls
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = [])
@patch.object(stack_mod, "_detect_rocm_version", return_value = None)
def test_gfx_override_installs_despite_visible_rocm(
self, mock_ver, mock_gfx, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""#7305 review: an explicit UNSLOTH_ROCM_GFX_ARCH is exempt from the
not-_has_rocm_gpu() gate (mirrors install.sh). A visible GPU with an
unreadable ROCm version must not silently discard the user's named arch
and leave CPU torch in place -- the per-arch install runs."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.10.0+cpu||\n"
with patch.dict(os.environ, {"UNSLOTH_ROCM_GFX_ARCH": "gfx1151"}):
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
assert mock_pip.call_count == 1, mock_pip.call_args_list
torch_call = str(mock_pip.call_args_list[0])
assert "gfx1151" in torch_call
assert "download.pytorch.org" not in torch_call
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 1))
def test_cuda_torch_on_amd_host_reinstalls(
self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""A CUDA-only torch build is unusable on an AMD-only host, so it must be
reinstalled to ROCm (has_hip_torch is driven by the empty HIP marker, not
by treating the CUDA version string as a HIP marker)."""
mock_probe = MagicMock()
mock_probe.returncode = 0
# Single-line probe: empty HIP marker before "|" for a CUDA build.
mock_probe.stdout = _MARK + "2.10.0+cu126||\n"
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
assert mock_pip.call_count == 1
assert "rocm7.1" in str(mock_pip.call_args_list[0])
@staticmethod
def _windows_repair(installed_family, gfx = "gfx1200"):
"""Run the Windows ROCm repair with an already-ROCm torch on disk.
Returns the pip_install_try mock so callers can assert on the reinstall.
"""
probe = MagicMock(returncode = 0, stdout = _MARK + "2.10.0+rocm7.1|7.1|\n")
pip_try = MagicMock(return_value = True)
with patch.dict(os.environ, {}, clear = False):
os.environ.pop("UNSLOTH_ROCM_TORCH_INSTALLED", None)
with (
patch.object(stack_mod, "IS_WINDOWS", True),
patch.object(stack_mod, "IS_MACOS", False),
patch.object(stack_mod, "_TORCH_BACKEND", ""),
patch.object(stack_mod, "_explicit_rocm_torch_index_url", return_value = None),
patch.object(
stack_mod, "_explicit_unknown_family_torch_index_url", return_value = None
),
patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False),
patch.object(stack_mod, "_detect_windows_gfx_arch", return_value = gfx),
patch.object(
stack_mod, "_installed_rocm_wheel_family", return_value = installed_family
),
patch.object(stack_mod, "_install_bnb_windows_rocm", return_value = True),
patch.object(stack_mod, "pip_install_try", pip_try),
patch("subprocess.run", return_value = probe),
):
_ensure_rocm_torch()
return pip_try
def test_wheel_family_change_forces_a_reinstall(self):
# A host with gfx103X wheels for its APU keeps them forever once the #7776 repick
# moves it to the dGPU: torch.version.hip only says "ROCm", so the repair path saw
# a ROCm build and stopped. setup.ps1 force-reinstalls, so this is `studio update`.
pip_try = self._windows_repair("gfx103x-all")
assert pip_try.call_count == 1
assert "gfx120X-all" in str(pip_try.call_args)
def test_matching_wheel_family_is_left_alone(self):
# Negative control: the right family must not be re-downloaded on every update.
assert self._windows_repair("gfx120x-all").call_count == 0
def test_unknown_wheel_family_is_left_alone(self):
# Older wheels predate the split runtime, so the family is unreadable and
# guessing would force a multi-GB reinstall on every update.
assert self._windows_repair(None).call_count == 0
@staticmethod
def _dists(*names):
out = []
for n in names:
d = MagicMock()
d.metadata = {"Name": n}
out.append(d)
return out
def test_installed_wheel_family_reads_the_rocm_metapackage(self):
# AMD's torch requires rocm[libraries], and that extra names the arch-specific
# runtime, so the installed `rocm` meta-package is the authoritative record.
reqs = [
"rocm-sdk-core==7.13.0",
'rocm-sdk-libraries-gfx120X-all==7.13.0; extra == "libraries"',
]
with patch("importlib.metadata.requires", return_value = reqs):
assert stack_mod._installed_rocm_wheel_family() == "gfx120x-all"
def test_orphaned_runtime_does_not_masquerade_as_the_active_family(self):
# pip never uninstalls the old rocm-sdk-libraries-<family> on a switch (different
# distribution name), so reading the first one found would report the orphan and
# redownload the stack on every update.
reqs = ['rocm-sdk-libraries-gfx120X-all==7.13.0; extra == "libraries"']
dists = self._dists("rocm_sdk_libraries_gfx103X-all", "rocm_sdk_libraries_gfx120X-all")
with patch("importlib.metadata.requires", return_value = reqs):
with patch("importlib.metadata.distributions", return_value = dists):
assert stack_mod._installed_rocm_wheel_family() == "gfx120x-all"
def test_two_runtimes_without_a_metapackage_are_unknowable(self):
# No `rocm` to arbitrate between two runtimes on disk, so leave the install alone.
dists = self._dists("rocm_sdk_libraries_gfx103X-all", "rocm_sdk_libraries_gfx120X-all")
with patch("importlib.metadata.requires", return_value = None):
with patch("importlib.metadata.distributions", return_value = dists):
assert stack_mod._installed_rocm_wheel_family() is None
def test_single_runtime_without_a_metapackage_still_reads(self):
dists = self._dists("rocm_sdk_libraries_gfx1151")
with patch("importlib.metadata.requires", return_value = None):
with patch("importlib.metadata.distributions", return_value = dists):
assert stack_mod._installed_rocm_wheel_family() == "gfx1151"
with patch("importlib.metadata.requires", return_value = None):
with patch("importlib.metadata.distributions", return_value = []):
assert stack_mod._installed_rocm_wheel_family() is None
def test_switched_host_does_not_reinstall_on_every_update(self):
# End to end: after the gfx103X -> gfx120X switch the orphan is still installed
# and the next `studio update` must do nothing.
reqs = ['rocm-sdk-libraries-gfx120X-all==7.13.0; extra == "libraries"']
dists = self._dists("rocm_sdk_libraries_gfx103X-all", "rocm_sdk_libraries_gfx120X-all")
with patch("importlib.metadata.requires", return_value = reqs):
with patch("importlib.metadata.distributions", return_value = dists):
family = stack_mod._installed_rocm_wheel_family()
assert self._windows_repair(family).call_count == 0
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 1))
def test_torch_already_has_hip_skips(self, mock_ver, mock_gpu, mock_nvidia, mock_pip):
"""If torch already has HIP, should skip ROCm reinstall."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.10.0+rocm7.1|7.1.12345|\n" # HIP marker + version
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
mock_pip.assert_not_called()
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 1))
def test_cpu_torch_probe_line_not_read_as_hip(self, mock_ver, mock_gpu, mock_nvidia, mock_pip):
"""A CPU build's probe line ("|2.10.0+cpu") must not read as HIP: the version
after the "|" separator is data, not a HIP marker, so has_hip_torch stays False
and the reinstall fires."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.10.0+cpu||\n"
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
with patch.object(stack_mod, "pip_install_try", return_value = True):
_ensure_rocm_torch()
assert mock_pip.call_count == 1
assert "rocm7.1" in str(mock_pip.call_args_list[0])
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 1))
def test_cpu_torch_gets_rocm_reinstall(
self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""CPU-only torch on ROCm host should trigger reinstall."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = "\n" # empty = no GPU backend
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
assert mock_pip.call_count == 1
assert "rocm7.1" in str(mock_pip.call_args_list[0])
assert mock_pip_try.call_count >= 1
assert "bitsandbytes" in str(mock_pip_try.call_args_list[0])
assert mock_pip_try.call_args.kwargs["force_pip"] is True
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (6, 3))
def test_rocm_63_selects_correct_tag(self, mock_ver, mock_gpu, mock_nvidia, mock_pip):
"""ROCm 6.3 should select rocm6.3 tag."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = "\n"
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
torch_call = mock_pip.call_args_list[0]
assert "rocm6.3" in str(torch_call)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (5, 0))
def test_old_rocm_skips(self, mock_ver, mock_gpu, mock_nvidia, mock_pip):
"""ROCm version too old (below 6.0) should skip."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = "\n"
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
mock_pip.assert_not_called()
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_infer_linux_amd_gfx_arch", return_value = None)
@patch.object(stack_mod, "_detect_rocm_version", return_value = None)
def test_version_unreadable_prints_warning(
self, mock_ver, mock_infer, mock_gpu, mock_nvidia, mock_pip, capsys
):
"""ROCm detected but version unreadable should print warning and skip."""
with patch("os.path.isdir", return_value = True):
_ensure_rocm_torch()
mock_pip.assert_not_called()
captured = capsys.readouterr()
assert "unreadable" in captured.out
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
def test_rocm_72_selects_72_tag(self, mock_ver, mock_gpu, mock_nvidia, mock_pip):
"""ROCm 7.2 should select rocm7.2 tag (now in mapping with torch 2.11.0)."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = "\n"
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
torch_call = mock_pip.call_args_list[0]
assert "rocm7.2" in str(torch_call)
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 14))
@patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx1150"])
def test_rocm_714_strix_routes_to_amd_arch_index(
self, mock_gfx, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""ROCm 7.14 caps to rocm7.2 on pytorch.org; Strix must use AMD gfx index."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.11.0+rocm7.2|7.14.60850|\n"
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
torch_call = str(mock_pip.call_args_list[0])
assert "gfx1150" in torch_call
assert "torch>=2.11.0,<2.12.0" in torch_call
@patch.object(stack_mod, "IS_MACOS", False)
@patch("platform.machine", return_value = "x86_64")
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 14))
@patch.object(
stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx1100", "gfx1100", "gfx1151"]
)
def test_mask_indexes_devices_not_deduplicated_arches(
self, mock_gfx, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try, mock_machine
):
# The mask names a device ordinal: on two gfx1100 plus a gfx1151, device 1 is the
# second gfx1100, but a deduplicated ['gfx1100','gfx1151'] reads index 1 as the
# Strix and routes the install to the Strix-only AMD index.
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.11.0+rocm7.2|7.14.60850|\n"
buf = io.StringIO()
with patch.dict(os.environ, {"CUDA_VISIBLE_DEVICES": "1"}, clear = False):
for _v in ("HIP_VISIBLE_DEVICES", "ROCR_VISIBLE_DEVICES"):
os.environ.pop(_v, None)
with contextlib.redirect_stdout(buf):
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
calls = str(mock_pip.call_args_list) + str(mock_pip_try.call_args_list)
assert "gfx1151" not in calls
assert "non-Strix runtime target (gfx1100)" in buf.getvalue()
@patch.object(stack_mod, "IS_MACOS", False)
@patch("platform.machine", return_value = "x86_64")
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 14))
@patch.object(
stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx1100", "gfx1100", "gfx1151"]
)
def test_mask_naming_the_strix_device_still_reroutes(
self, mock_gfx, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try, mock_machine
):
# Negative control: device 2 really is the Strix, so the reroute must still fire.
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.11.0+rocm7.2|7.14.60850|\n"
with patch.dict(os.environ, {"CUDA_VISIBLE_DEVICES": "2"}, clear = False):
for _v in ("HIP_VISIBLE_DEVICES", "ROCR_VISIBLE_DEVICES"):
os.environ.pop(_v, None)
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
calls = str(mock_pip.call_args_list) + str(mock_pip_try.call_args_list)
assert "gfx1151" in calls
# rocminfo repeats the gfx token per agent (Name, ISA, marketing name), so a flat
# findall counts one GPU several times and every later ordinal is wrong.
_ROCMINFO = """
*******
Agent 1
*******
Name: AMD Ryzen 9 7950X
Marketing Name: AMD Ryzen 9 7950X
*******
Agent 2
*******
Name: gfx1100
Marketing Name: AMD Radeon RX 7900 XTX
ISA Info:
ISA 1
Name: amdgcn-amd-amdhsa--gfx1100
*******
Agent 3
*******
Name: gfx1100
Marketing Name: AMD Radeon RX 7900 XTX
ISA Info:
ISA 1
Name: amdgcn-amd-amdhsa--gfx1100
*******
Agent 4
*******
Name: gfx1151
Marketing Name: AMD Radeon 8060S
ISA Info:
ISA 1
Name: amdgcn-amd-amdhsa--gfx1151
"""
@staticmethod
def _probe_gfx(out, dedup):
which = lambda n: "/usr/bin/rocminfo" if n == "rocminfo" else None # noqa: E731
with patch("shutil.which", side_effect = which):
with patch("subprocess.run", return_value = MagicMock(returncode = 0, stdout = out)):
return stack_mod._detect_amd_gfx_codes(dedup = dedup)
def test_gfx_probe_returns_one_entry_per_agent_not_per_token(self):
assert self._probe_gfx(self._ROCMINFO, False) == ["gfx1100", "gfx1100", "gfx1151"]
assert self._probe_gfx(self._ROCMINFO, True) == ["gfx1100", "gfx1151"]
def test_gfx_probe_falls_back_for_flat_output(self):
# amd-smi and test stubs emit no agent headers; those must still work.
flat = "Name: gfx1100\nName: gfx1100\nName: gfx1151\n"
assert self._probe_gfx(flat, False) == ["gfx1100", "gfx1151"]
assert self._probe_gfx(flat, True) == ["gfx1100", "gfx1151"]
def test_gfx_probe_records_which_tool_answered(self):
# Only rocminfo is mask-filtered, and only by ROCR, so the reroute needs this.
self._probe_gfx(self._ROCMINFO, False)
assert stack_mod._LAST_AMD_GFX_PROBE == "rocminfo"
def test_first_set_visible_mask_is_first_set_wins(self):
with patch.dict(os.environ, {}, clear = False):
for _v in ("HIP_VISIBLE_DEVICES", "ROCR_VISIBLE_DEVICES", "CUDA_VISIBLE_DEVICES"):
os.environ.pop(_v, None)
assert stack_mod._first_set_visible_mask() is None
os.environ["CUDA_VISIBLE_DEVICES"] = "1"
assert stack_mod._first_set_visible_mask() == "CUDA_VISIBLE_DEVICES"
os.environ["ROCR_VISIBLE_DEVICES"] = ""
assert stack_mod._first_set_visible_mask() == "ROCR_VISIBLE_DEVICES"
os.environ["HIP_VISIBLE_DEVICES"] = "0"
assert stack_mod._first_set_visible_mask() == "HIP_VISIBLE_DEVICES"
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (6, 4))
def test_explicit_gfx_index_honored_and_skips_strix_reroute(
self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""An explicit gfx wheel-index pin is authoritative: install from it verbatim
with torch 2.11, and the pin must not be second-guessed (host ROCm 6.4 would
otherwise pick the rocm6.4 wheel / trigger the Strix re-route). The gfx probe
may run for the bnb-skip flag, but returning a Strix arch must not reroute the
pinned torch index."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = "\n" # cpu torch -> reinstall
env = {"UNSLOTH_TORCH_INDEX_URL": "https://repo.amd.com/rocm/whl/gfx1151"}
with patch.dict(stack_mod.os.environ, env, clear = False):
stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
with patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx1151"]):
_ensure_rocm_torch()
assert mock_pip.call_count == 1
torch_call = str(mock_pip.call_args_list[0])
assert "gfx1151" in torch_call
assert "torch>=2.11.0,<2.12.0" in torch_call
def test_rocm_pin_family_mismatch_helper(self):
"""_rocm_pin_family_mismatch: exact rocm compare, else the 2.11 line."""
f = stack_mod._rocm_pin_family_mismatch
base = "https://download.pytorch.org/whl"
amd = "https://repo.amd.com/rocm/whl"
# Exact rocm version comparison.
assert f(f"{base}/rocm7.2", "2.11.0+rocm7.2") is False
assert f(f"{base}/rocm7.2", "2.10.0+rocm6.4") is True
assert f(f"{base}/rocm6.4", "2.10.0+rocm6.4") is False
# rocm7.2 is KNOWN-2.11. A +rocm7.2 wheel whose RELEASE drifted off 2.11 shares the
# tag but violates the spec -> mismatch (a plain version compare would accept it).
assert f(f"{base}/rocm7.2", "2.12.0+rocm7.2") is True
assert f(f"{base}/rocm7.2", "2.13.0+rocm7.2") is True
assert f(f"{base}/rocm7.2", "2.11.5+rocm7.2") is False # patch on 2.11 is in-spec
# An UNKNOWN newer rocm (not on the 2.11 allowlist) is not floored to 2.11, so a
# matching rocm version at any release line is NOT a mismatch on this branch.
assert f(f"{base}/rocm8.0", "2.12.0+rocm8.0") is False
# gfx pin (2.11 line) vs installed release line.
assert f(f"{amd}/gfx1151", "2.10.0+rocm6.4") is True
assert f(f"{amd}/gfx1151", "2.11.0+rocm7.13.0") is False
# rocm7.2 pin vs an untagged (no +rocm) wheel: a CPU/CUDA build never
# satisfies a ROCm pin, regardless of its release line -> always a mismatch.
assert f(f"{base}/rocm7.2", "2.10.0") is True
assert f(f"{base}/rocm7.2", "2.11.0") is True
assert f(f"{base}/rocm6.4", "2.10.0") is True
# A 2.11-allowlist gfx pin over a GENERIC (two-part +rocm7.2) 2.11 wheel mismatches:
# the user wants AMD's per-arch (three-part) wheel, not the generic one.
assert f(f"{amd}/gfx1151", "2.11.0+rocm7.2") is True
assert f(f"{amd}/gfx120X-all", "2.11.0+rocm7.2") is True
# ...but an already-installed per-arch (three-part) wheel is NOT re-flagged
# (no reinstall loop once the correct gfx wheel is present).
assert f(f"{amd}/gfx120X-all", "2.11.0+rocm7.13.0") is False
assert f(f"{amd}/gfx1150", "2.11.0+rocm7.13.0") is False
# A NON-2.11 gfx pin (gfx110X-all/gfx90a/gfx908) tracks the default <2.11 spec: a
# correct 2.10+rocm wheel is NOT a mismatch, a 2.11 build is.
assert f(f"{amd}/gfx110X-all", "2.10.0+rocm6.4") is False
assert f(f"{amd}/gfx90a", "2.10.0+rocm6.3") is False
assert f(f"{amd}/gfx908", "2.10.0+rocm7.0") is False
assert f(f"{amd}/gfx110X-all", "2.11.0+rocm7.2") is True
# A non-2.11 gfx pin over an untagged (no +rocm) wheel is a mismatch even
# when torch is already <2.11: a CPU/CUDA build never satisfies the ROCm pin.
assert f(f"{amd}/gfx110X-all", "2.10.0") is True
assert f(f"{amd}/gfx90a", "2.10.0") is True
# A major-only rocm pin (rocm7) compares on the major alone: rocm6.x mismatches,
# any rocm7.x satisfies it, an untagged wheel never does, a bare +rocm is lenient.
assert f(f"{base}/rocm7", "2.10.0+rocm6.4") is True
assert f(f"{base}/rocm7", "2.11.0+rocm7.2") is False
assert f(f"{base}/rocm7", "2.11.0+rocm7.13.0") is False
assert f(f"{base}/rocm7", "2.10.0") is True
assert f(f"{base}/rocm7", "2.10.0+rocm") is False
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
def test_rocm_pin_mismatch_over_installed_rocm_reinstalls(
self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""A rocm7.2 pin over an already-installed OLDER +rocm6.4 build must reinstall,
even though has_hip_torch is True (the ROCm analogue of the CUDA cuXXX mismatch)."""
mock_probe = MagicMock()
mock_probe.returncode = 0
# HIP marker present (has_hip_torch=True) + installed +rocm6.4 wheel.
mock_probe.stdout = _MARK + "2.10.0+rocm6.4|6.4.12345|\n"
env = {"UNSLOTH_TORCH_INDEX_FAMILY": "rocm7.2"}
with patch.dict(stack_mod.os.environ, env, clear = False):
stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_URL", None)
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
torch_call = str(mock_pip.call_args_list[0])
assert "rocm7.2" in torch_call
assert "torch>=2.11.0,<2.12.0" in torch_call
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (6, 4))
def test_gfx_pin_over_installed_pre211_rocm_reinstalls(
self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""A gfx* pin (2.11 line) over an installed pre-2.11 +rocm6.4 build reinstalls.
The gfx probe may run for the bnb-skip flag but must not alter the pinned index."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.10.0+rocm6.4|6.4.12345|\n"
env = {"UNSLOTH_TORCH_INDEX_URL": "https://repo.amd.com/rocm/whl/gfx1151"}
with patch.dict(stack_mod.os.environ, env, clear = False):
stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
with patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx1151"]):
_ensure_rocm_torch()
torch_call = str(mock_pip.call_args_list[0])
assert "gfx1151" in torch_call
assert "torch>=2.11.0,<2.12.0" in torch_call
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
def test_rocm_pin_matches_installed_no_torch_reinstall(
self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""A rocm7.2 pin over an already-matching +rocm7.2 build must NOT reinstall torch
(no false reinstall of a correct ROCm venv)."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.11.0+rocm7.2|7.2.12345|\n"
env = {"UNSLOTH_TORCH_INDEX_FAMILY": "rocm7.2"}
with patch.dict(stack_mod.os.environ, env, clear = False):
stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_URL", None)
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
# No torch reinstall: any pip_install call must not target a torch index.
for _call in mock_pip.call_args_list:
_args = [str(a) for a in _call.args]
if "--index-url" in _args:
_url = _args[_args.index("--index-url") + 1]
assert "rocm7.2" not in _url or "torch" not in " ".join(
_args
), "torch must not be reinstalled when the pin already matches"
# A torch reinstall would pass torch>=... as a positional; assert none did.
assert not any(
any(str(a).startswith("torch") for a in _c.args) for _c in mock_pip.call_args_list
)
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (6, 4))
def test_non211_gfx_pin_over_210_rocm_no_reinstall(
self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""A gfx110X-all pin (NOT in the 2.11 allowlist) over a correct 2.10+rocm
wheel must NOT be flagged stale -- the install path uses the default <2.11
specs for that arch, so re-flagging would reinstall-loop on every update."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.10.0+rocm6.4|6.4.12345|\n"
env = {"UNSLOTH_TORCH_INDEX_URL": "https://repo.amd.com/rocm/whl/gfx110X-all"}
with patch.dict(stack_mod.os.environ, env, clear = False):
stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
# has_hip_torch True + no mismatch -> torch must NOT be reinstalled.
assert not any(
any(str(a).startswith("torch") for a in _c.args) for _c in mock_pip.call_args_list
)
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
def test_gfx_pin_over_generic_rocm211_reinstalls(
self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""A gfx1151 pin over a GENERIC (two-part +rocm7.2) 2.11 wheel must reinstall
the AMD per-arch wheel -- even though both are torch 2.11, the generic wheel
is not the per-arch build the user pinned (Strix stays off the generic wheel)."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.11.0+rocm7.2|7.2.12345|\n"
env = {"UNSLOTH_TORCH_INDEX_URL": "https://repo.amd.com/rocm/whl/gfx1151"}
with patch.dict(stack_mod.os.environ, env, clear = False):
stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
# The gfx probe may run for the bnb-skip flag; returning a Strix
# arch must not reroute the pinned torch index.
with patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx1151"]):
_ensure_rocm_torch()
torch_call = str(mock_pip.call_args_list[0])
assert "gfx1151" in torch_call
assert "torch>=2.11.0,<2.12.0" in torch_call
def test_radeon_url_not_classified_as_pip_rocm_family(self):
"""A repo.radeon.com find-links dir (leaf rocm-rel-7.2.1) starts with "rocm" but is
NOT a pip --index-url ROCm family: it must route to the verbatim path, not a
--index-url reinstall that fails against a find-links listing."""
leaf_f = stack_mod._is_pip_rocm_family_leaf
# Real pip ROCm families (download.pytorch.org/whl/rocmX.Y, repo.amd.com gfx).
assert leaf_f("rocm7.2") is True
assert leaf_f("rocm6.4") is True
assert leaf_f("gfx120x-all") is True
assert leaf_f("gfx1151") is True
# A bare rocm<digits> (no minor) is still an exact family.
assert leaf_f("rocm7") is True
# A Radeon find-links dir leaf, a custom mirror, cpu and cuda are NOT pip rocm.
assert leaf_f("rocm-rel-7.2.1") is False
assert leaf_f("simple") is False
assert leaf_f("current") is False
assert leaf_f("cpu") is False
assert leaf_f("cu128") is False
# A rocm<digit>-SUFFIX private mirror shares the family prefix but is a custom pin
# the verbatim path owns: a ^rocm\d PREFIX match would wrongly treat it as a
# --index-url family. Match EXACTLY.
assert leaf_f("rocm7.2-private") is False
assert leaf_f("rocm7-current") is False
assert leaf_f("rocm7.2.1") is False # two-part local suffix -> custom, not rocm7.2
radeon = "https://repo.radeon.com/rocm/manylinux/rocm-rel-7.2.1"
pip_rocm = "https://download.pytorch.org/whl/rocm7.2"
amd_gfx = "https://repo.amd.com/rocm/whl/gfx120X-all"
def _classify(url, fn):
with patch.dict(stack_mod.os.environ, {"UNSLOTH_TORCH_INDEX_URL": url}, clear = False):
stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
return fn()
rocm_fn = stack_mod._explicit_rocm_torch_index_url
unk_fn = stack_mod._explicit_unknown_family_torch_index_url
# Real pip rocm/gfx pins ARE a ROCm family (reinstallable via --index-url) and
# are NOT "unknown".
assert _classify(pip_rocm, rocm_fn) == pip_rocm
assert _classify(amd_gfx, rocm_fn) == amd_gfx
assert _classify(pip_rocm, unk_fn) is None
assert _classify(amd_gfx, unk_fn) is None
# The Radeon find-links URL is NOT a pip ROCm family (so _ensure_rocm_torch skips
# it) and IS unknown, so the family repair helpers leave it alone.
assert _classify(radeon, rocm_fn) is None
assert _classify(radeon, unk_fn) == radeon
# A rocm<digit>-suffix private mirror routes the same way: NOT a pip rocm family,
# IS an unknown-family (verbatim) pin.
suffixed = "https://co.internal/whl/rocm7.2-private"
assert _classify(suffixed, rocm_fn) is None
assert _classify(suffixed, unk_fn) == suffixed
@patch.object(stack_mod, "pip_install")
def test_ensure_cpu_torch_broken_probe_reinstalls(self, mock_pip):
"""_ensure_cpu_torch: torch present but unimportable (probe exit != 0) under an
explicit CPU pin must reinstall from the pin, not return -- the base update does
not repair a broken installed torch, so returning would strand it (Codex P2)."""
mock_probe = MagicMock()
mock_probe.returncode = 1 # torch present but cannot import
mock_probe.stdout = ""
env = {"UNSLOTH_TORCH_INDEX_URL": "https://mirror.local/cpu"}
with patch.dict(stack_mod.os.environ, env, clear = False):
stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
with patch("subprocess.run", return_value = mock_probe):
with patch.object(stack_mod, "NO_TORCH", False):
stack_mod._ensure_cpu_torch()
assert mock_pip.call_count == 1
assert "https://mirror.local/cpu" in str(mock_pip.call_args)
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 1))
def test_probe_timeout_triggers_reinstall(
self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""Probe subprocess timeout should not crash; should proceed to reinstall."""
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", side_effect = subprocess.TimeoutExpired("python", 30)):
_ensure_rocm_torch()
# Probe timeout: treat torch as unusable and reinstall torch + bitsandbytes.
assert mock_pip.call_count == 1
assert "rocm7.1" in str(mock_pip.call_args_list[0])
assert mock_pip_try.call_count >= 1
assert mock_pip_try.call_args.kwargs["force_pip"] is True
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = False)
@patch.object(stack_mod, "_infer_linux_amd_gfx_arch", return_value = None)
def test_no_gpu_with_rocm_tools_skips(self, mock_infer, mock_gpu, mock_nvidia, mock_pip):
"""ROCm tools present but no actual AMD GPU should skip entirely."""
# Pin the Windows arch probe to None so a real AMD host's WMI fallback
# can't defeat the "no actual GPU" premise.
with patch.object(stack_mod, "_detect_windows_gfx_arch", return_value = None):
with patch("os.path.isdir", return_value = True):
_ensure_rocm_torch()
mock_pip.assert_not_called()
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = True)
def test_torch_backend_cuda_env_skips_entirely(self, mock_nvidia, mock_gpu, mock_pip):
"""UNSLOTH_TORCH_BACKEND=cuda must short-circuit before any GPU probe."""
with patch.dict(os.environ, {"UNSLOTH_TORCH_BACKEND": "cuda"}):
with patch.object(stack_mod, "_TORCH_BACKEND", "cuda"):
_ensure_rocm_torch()
mock_pip.assert_not_called()
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = True)
def test_torch_backend_cpu_env_skips_entirely(self, mock_nvidia, mock_gpu, mock_pip):
"""UNSLOTH_TORCH_BACKEND=cpu must short-circuit before any GPU probe."""
with patch.dict(os.environ, {"UNSLOTH_TORCH_BACKEND": "cpu"}):
with patch.object(stack_mod, "_TORCH_BACKEND", "cpu"):
_ensure_rocm_torch()
mock_pip.assert_not_called()
# TEST: gfx906 (MI50 / Radeon VII) legacy reroute -- generic wheels after rocm6.3
# lack gfx906 code objects, so torch must come from the rocm6.3 index.
class TestGfx906LegacyReroute:
"""gfx906 hosts on ROCm >= 6.4 must be rerouted to the rocm6.3 torch index;
hosts already on gfx906-capable wheels are left alone."""
@staticmethod
def _gfx906_reroute_block(source: str) -> str:
"""The MI50/gfx906 reroute block, bounded on the ';;' that closes its
rocm[0-9]* case arm -- robust to comment growth (no magic char offset)."""
start = source.find("MI50 / Radeon VII (gfx906")
assert start >= 0, "gfx906 reroute block not found in install.sh"
end = source.find("\n ;;", start)
assert end >= 0, "end of gfx906 case arm not found"
return source[start:end]
def test_gfx906_needs_legacy_index_floor(self):
f = stack_mod._gfx906_needs_legacy_index
# rocm6.0-6.3 tags still ship gfx906 kernels: no reroute.
assert f((6, 3)) is False
assert f((6, 0)) is False
assert f((5, 0)) is False # below any known tag
# Anything that picks a tag newer than rocm6.3 must reroute.
assert f((6, 4)) is True
assert f((7, 2)) is True
assert f((7, 14)) is True
def test_runtime_target_is_gfx906_selection(self, monkeypatch):
"""Env override wins; else gfx906 only when it is the SOLE distinct arch."""
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
# Sole gfx906 (one or several identical MI50s de-dup to {'gfx906'}).
with patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx906"]):
assert stack_mod._runtime_target_is_gfx906() is True
# Mixed host: gfx906 is NOT the sole arch -> not auto-selected (Codex #3:
# de-dup loses ordinals, so never downgrade a non-gfx906 selection).
with patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx906", "gfx1100"]):
assert stack_mod._runtime_target_is_gfx906() is False
with patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = []):
assert stack_mod._runtime_target_is_gfx906() is False
# Explicit override wins even when probes see nothing (Codex #2).
monkeypatch.setenv("UNSLOTH_ROCM_GFX_ARCH", "gfx906")
with patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = []):
assert stack_mod._runtime_target_is_gfx906() is True
# ...and a non-gfx906 override is honored on a gfx906-present host.
monkeypatch.setenv("UNSLOTH_ROCM_GFX_ARCH", "gfx1100")
with patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx906"]):
assert stack_mod._runtime_target_is_gfx906() is False
# A copied HIP gcnArchName (gfx906:sramecc-:xnack-) normalizes to gfx906
# (Codex #4: the feature-flag suffix must not defeat the exact comparison).
monkeypatch.setenv("UNSLOTH_ROCM_GFX_ARCH", "gfx906:sramecc-:xnack-")
with patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = []):
assert stack_mod._runtime_target_is_gfx906() is True
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
@patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx906"])
def test_gfx906_on_rocm72_routes_to_rocm63(
self, mock_gfx, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try, monkeypatch
):
"""CPU torch on a ROCm 7.2 MI50 host installs from rocm6.3, not rocm7.2."""
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = "\n" # cpu torch -> reinstall
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
torch_call = str(mock_pip.call_args_list[0])
assert "rocm6.3" in torch_call
assert "rocm7.2" not in torch_call
# The _default (<2.11) window: the rocm7.2 2.11 floor cannot be satisfied
# on the rocm6.3 index (torch <= 2.9.x there).
assert "torch>=2.4,<2.11.0" in torch_call
# gfx906 has no prebuilt bnb -- the generic wheel must not be installed.
assert not any("bitsandbytes" in str(c).lower() for c in mock_pip_try.call_args_list)
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
@patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx906"])
def test_gfx906_repairs_existing_rocm72_torch(
self, mock_gfx, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try, monkeypatch
):
"""An installed +rocm7.2 torch IS the broken combo: reinstall from rocm6.3
even though has_hip_torch is True."""
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.11.0+rocm7.2|7.2.12345|\n"
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
torch_call = str(mock_pip.call_args_list[0])
assert "rocm6.3" in torch_call
assert "torch>=2.4,<2.11.0" in torch_call
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
@patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx906"])
def test_gfx906_already_on_rocm63_left_alone(
self, mock_gfx, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try, monkeypatch
):
"""torch already on rocm6.3 wheels must not be reinstalled (no update loop),
and the generic bnb wheel must not clobber a source build."""
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = _MARK + "2.7.0+rocm6.3|6.3.42131|\n"
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
mock_pip.assert_not_called()
# gfx906: prebuilt bnb is skipped entirely (no torch reinstall, no bnb).
mock_pip_try.assert_not_called()
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
@patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx1100", "gfx906"])
def test_mixed_host_gfx906_not_sole_arch_skips_reroute(
self, mock_gfx, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try, monkeypatch
):
"""Mixed host (gfx906 + dGPU) with no explicit override: gfx906 is not the
sole arch, so the generic index is kept (Codex #3: never downgrade a
de-dup-ambiguous mixed host)."""
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = "\n" # cpu torch -> reinstall
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
torch_call = str(mock_pip.call_args_list[0])
assert "rocm7.2" in torch_call
assert "rocm6.3" not in torch_call
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
@patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = [])
def test_gfx906_env_override_forces_reroute(
self, mock_gfx, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try, monkeypatch
):
"""UNSLOTH_ROCM_GFX_ARCH=gfx906 reroutes even when probes emit no gfx token
(Codex #2: runtime-only ROCm hosts where rocminfo/amd-smi are absent)."""
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.setenv("UNSLOTH_ROCM_GFX_ARCH", "gfx906")
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = "\n" # cpu torch -> reinstall
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
torch_call = str(mock_pip.call_args_list[0])
assert "rocm6.3" in torch_call
assert "rocm7.2" not in torch_call
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
@patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx906"])
def test_gfx906_bnb_skipped_even_when_index_pinned(
self, mock_gfx, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try, monkeypatch
):
"""A gfx906 user who pins the ROCm index AND sets the arch override still
skips the generic bnb wheel: the pin suppresses the torch reroute, not the
gfx906 runtime flag used for the bnb skip."""
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.setenv("UNSLOTH_ROCM_GFX_ARCH", "gfx906")
monkeypatch.setenv("UNSLOTH_TORCH_INDEX_URL", "https://download.pytorch.org/whl/rocm6.3")
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = "\n" # cpu torch -> reinstall from the pinned index
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
# torch is (re)installed from the pinned rocm6.3 index...
assert any("rocm6.3" in str(c) for c in mock_pip.call_args_list)
# ...but the prebuilt bnb wheel is never installed.
assert not any("bitsandbytes" in str(c).lower() for c in mock_pip_try.call_args_list)
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
@patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx1151", "gfx906"])
def test_gfx906_override_wins_over_strix_probe(
self, mock_gfx, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try, monkeypatch
):
"""Mixed Strix + MI50 host: UNSLOTH_ROCM_GFX_ARCH=gfx906 suppresses the Strix
override (which probe order would otherwise pick) and routes to rocm6.3."""
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.setenv("UNSLOTH_ROCM_GFX_ARCH", "gfx906")
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = "\n" # cpu torch -> reinstall
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
torch_call = str(mock_pip.call_args_list[0])
assert "rocm6.3" in torch_call
assert "gfx1151" not in torch_call
# gfx906 target -> generic bnb wheel skipped.
assert not any("bitsandbytes" in str(c).lower() for c in mock_pip_try.call_args_list)
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
@patch.object(stack_mod, "_detect_amd_gfx_codes", return_value = ["gfx906"])
def test_gfx906_bnb_skipped_on_pinned_index_without_env_override(
self, mock_gfx, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try, monkeypatch
):
"""Codex #2: a real gfx906 host that pins the ROCm index but does NOT set
UNSLOTH_ROCM_GFX_ARCH must still skip the prebuilt bnb wheel -- the pin
suppresses only the torch reroute, not the probe-driven gfx906 detection
used for the bnb skip (otherwise `studio update` clobbers source-built bnb)."""
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
monkeypatch.setenv("UNSLOTH_TORCH_INDEX_URL", "https://download.pytorch.org/whl/rocm6.3")
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = "\n" # cpu torch -> reinstall from the pinned index
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
# torch is (re)installed from the pinned rocm6.3 index...
assert any("rocm6.3" in str(c) for c in mock_pip.call_args_list)
# ...but the prebuilt bnb wheel is never installed (probe saw sole gfx906).
assert not any("bitsandbytes" in str(c).lower() for c in mock_pip_try.call_args_list)
def test_install_sh_gfx906_env_suppresses_strix(self):
"""install.sh must skip the Strix reroute when UNSLOTH_ROCM_GFX_ARCH=gfx906."""
source = (PACKAGE_ROOT / "install.sh").read_text(encoding = "utf-8")
assert 'if [ "$_gfx906_env" != "gfx906" ]; then' in source
def test_install_sh_gfx906_normalizes_override_and_clears_radeon(self):
"""install.sh must (Codex #4) strip a gfx906:… feature suffix before the exact
comparison, and (Codex #3) clear the Radeon marketing flag for every gfx906
target -- not only when the >=6.4 reroute fires -- so a Radeon VII already on
rocm6.3 does not divert to the repo.radeon.com branch."""
source = (PACKAGE_ROOT / "install.sh").read_text(encoding = "utf-8")
# Override normalization (both the reroute block and the bnb-skip helper):
# strip the gfx906:… feature suffix and trim whitespace (mirror py .strip()).
assert "_gfx906_env=${_gfx906_env%%:*}" in source
assert "_bnb_gfx_env=${_bnb_gfx_env%%:*}" in source
assert source.count("tr -d '[:space:]'") >= 2
# Radeon flag cleared as soon as gfx906 is the target, before the leaf gate.
block = self._gfx906_reroute_block(source)
clear_pos = block.find("_amd_gpu_radeon=false")
leaf_gate_pos = block.find("_rocm_leaf_below")
assert clear_pos >= 0 and leaf_gate_pos >= 0
# the unconditional clear must precede the >=6.4 leaf-gated reroute.
assert clear_pos < leaf_gate_pos
def test_install_sh_bnb_skip_probes_under_pin(self):
"""install.sh _is_gfx906_bnb_skip must probe gfx906 when the index is pinned
(Codex #1): a pin skips the reroute block that sets _gfx906_target, so the
helper falls back to _probe_amd_gfx_arch to catch a real gfx906 host."""
source = (PACKAGE_ROOT / "install.sh").read_text(encoding = "utf-8")
start = source.find("_is_gfx906_bnb_skip() {")
assert start >= 0
body = source[start : start + 900]
assert "_torch_index_pinned" in body
assert "_probe_amd_gfx_arch" in body
def test_install_sh_has_gfx906_reroute(self):
"""install.sh must mirror the Python reroute: honor UNSLOTH_ROCM_GFX_ARCH,
gate on a gfx906 target, route to rocm6.3, with the same _default (<2.11)
trio, and skip the prebuilt bnb wheel."""
source = (PACKAGE_ROOT / "install.sh").read_text(encoding = "utf-8")
block = self._gfx906_reroute_block(source)
assert "_gfx906_target=" in block
assert "UNSLOTH_ROCM_GFX_ARCH" in block
assert "/rocm6.3" in block
for spec in stack_mod._ROCM_TORCH_PKG_SPECS["_default"]:
assert spec in block
# The bnb skip helper must exist and be wired at the install sites.
assert "_is_gfx906_bnb_skip" in source
def test_device_type_defaults_compile_off_on_gfx906(self):
"""unsloth/device_type.py must default Dynamo/compile off on gfx906
(user-overridable via setdefault)."""
source = (PACKAGE_ROOT / "unsloth" / "device_type.py").read_text(encoding = "utf-8")
gate_start = source.find("gfx906")
assert gate_start >= 0
gate_body = source[gate_start : gate_start + 800]
assert 'setdefault("TORCHDYNAMO_DISABLE", "1")' in gate_body
assert 'setdefault("TORCH_COMPILE_DISABLE", "1")' in gate_body
assert 'setdefault("UNSLOTH_COMPILE_DISABLE", "1")' in gate_body
# TEST: install_python_stack.py -- torch-index MARKER mechanism (PR #6692)
class TestHasRocmGpuKfdVendorGuard:
"""KFD sysfs fallback rejects non-AMD (NVIDIA) KFD nodes (source-level checks)."""
def _src(self) -> str:
"""Return the source of _has_rocm_gpu from install_python_stack.py."""
import inspect
return inspect.getsource(stack_mod._has_rocm_gpu)
def test_vendor_id_check_present(self):
"""_has_rocm_gpu sysfs fallback must check vendor_id 4098 (AMD 0x1002)."""
src = self._src()
assert "vendor_id" in src, (
"_has_rocm_gpu KFD sysfs fallback must read the properties file "
"to check vendor_id and exclude NVIDIA KFD nodes"
)
assert "4098" in src, (
"_has_rocm_gpu must require AMD vendor_id 4098 (0x1002) in the "
"KFD node properties to avoid false positives on NVIDIA systems"
)
def test_vendor_regex_pattern_anchored(self):
"""The vendor_id regex must use a word boundary to avoid partial matches."""
import re as _re
src = self._src()
# Word boundary so "vendor_id 41098" doesn't match "vendor_id 4098".
assert (
_re.search(r"\\b.*vendor_id.*\\b", src) or "\\bvendor_id" in src
), "_has_rocm_gpu vendor_id check should use word boundary anchors"
def test_sysfs_fallback_guarded_by_non_win32(self):
"""KFD sysfs fallback must be Linux-only (guarded by sys.platform != 'win32')."""
src = self._src()
assert "win32" in src, "_has_rocm_gpu sysfs fallback must be guarded by sys.platform check"
def test_cpu_node_excluded(self):
"""gpu_id == '0' must be excluded (CPU topology nodes)."""
src = self._src()
assert (
'!= "0"' in src or "== '0'" in src or "!= '0'" in src or '"0"' in src
), "_has_rocm_gpu must skip gpu_id 0 nodes (CPU nodes)"
def test_install_sh_has_vendor_check(self):
"""_has_amd_rocm_gpu in install.sh sysfs fallback must also check vendor_id 4098."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
func_start = source.find("_has_amd_rocm_gpu()")
func_end = source.find("\n}", func_start)
func_body = source[func_start:func_end]
assert "vendor_id" in func_body, "_has_amd_rocm_gpu sysfs fallback must check vendor_id"
assert "4098" in func_body, "_has_amd_rocm_gpu must require AMD vendor_id 4098 (0x1002)"
def test_has_rocm_gpu_returns_false_when_nvidia_present(self):
"""_has_rocm_gpu returns False when _has_usable_nvidia_gpu is True (NVIDIA always wins)."""
with patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = True):
with patch("shutil.which", return_value = "/usr/bin/rocminfo"):
# rocminfo claims an AMD GPU is present.
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = "Name: gfx1100\n"
with patch("subprocess.run", return_value = mock_result):
assert not stack_mod._has_rocm_gpu(), (
"_has_rocm_gpu must return False when NVIDIA GPU is detected, "
"regardless of what rocminfo reports"
)
def test_install_sh_has_rocm_gpu_nvidia_guard(self):
"""_has_amd_rocm_gpu in install.sh must call _has_usable_nvidia_gpu and return 1 if true."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
func_start = source.find("_has_amd_rocm_gpu()")
func_end = source.find("\n}", func_start)
func_body = source[func_start:func_end]
assert (
"_has_usable_nvidia_gpu" in func_body
), "_has_amd_rocm_gpu must call _has_usable_nvidia_gpu to block NVIDIA hosts"
assert (
"return 1" in func_body
), "_has_amd_rocm_gpu must return 1 (false) when NVIDIA GPU is detected"
def test_has_usable_nvidia_gpu_proc_fallback_present(self):
"""`_has_usable_nvidia_gpu` must have a /proc/driver/nvidia fallback."""
import inspect
src = inspect.getsource(stack_mod._has_usable_nvidia_gpu)
assert "/proc/driver/nvidia" in src, (
"_has_usable_nvidia_gpu must fall back to /proc/driver/nvidia/gpus when "
"nvidia-smi subprocess fails, to handle PATH gaps and driver init races"
)
def test_install_sh_has_usable_nvidia_gpu_proc_fallback(self):
"""_has_usable_nvidia_gpu in install.sh must also have a /proc/driver/nvidia fallback."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
func_start = source.find("_has_usable_nvidia_gpu()")
func_end = source.find("\n}", func_start)
func_body = source[func_start:func_end]
assert "/proc/driver/nvidia" in func_body, (
"_has_usable_nvidia_gpu in install.sh must fall back to "
"/proc/driver/nvidia/gpus when nvidia-smi fails"
)
# TEST: install_python_stack.py -- _ROCM_TORCH_INDEX mapping
class TestRocmTorchIndex:
"""Verify the ROCm version -> torch index tag mapping."""
def test_mapping_is_sorted_descending(self):
"""Keys should be in descending order for the next() iteration to work."""
keys = list(_ROCM_TORCH_INDEX.keys())
assert keys == sorted(keys, reverse = True)
def test_rocm_72_in_mapping(self):
"""ROCm 7.2 should be in the active mapping (torch 2.11.0 now supported)."""
assert (7, 2) in _ROCM_TORCH_INDEX
assert _ROCM_TORCH_INDEX[(7, 2)] == "rocm7.2"
def test_rocm_71_maps_correctly(self):
assert _ROCM_TORCH_INDEX[(7, 1)] == "rocm7.1"
def test_rocm_63_maps_correctly(self):
assert _ROCM_TORCH_INDEX[(6, 3)] == "rocm6.3"
def test_rocm_60_maps_correctly(self):
assert _ROCM_TORCH_INDEX[(6, 0)] == "rocm6.0"
def test_all_tags_use_download_pytorch(self):
"""All tags should be for download.pytorch.org, not repo.radeon.com."""
for tag in _ROCM_TORCH_INDEX.values():
assert tag.startswith("rocm")
assert "radeon" not in tag
def test_newer_rocm_selects_best_match(self):
"""ROCm 7.2 (now in map) should select rocm7.2 directly."""
ver = (7, 2)
tag = next(
(
t
for (maj, mn), t in sorted(_ROCM_TORCH_INDEX.items(), reverse = True)
if ver >= (maj, mn)
),
None,
)
assert tag == "rocm7.2"
def test_rocm_64_selects_64(self):
ver = (6, 4)
tag = next(
(
t
for (maj, mn), t in sorted(_ROCM_TORCH_INDEX.items(), reverse = True)
if ver >= (maj, mn)
),
None,
)
assert tag == "rocm6.4"
# TEST: hardware.py -- IS_ROCM flag and detect_hardware
class TestHardwareRocmFlag:
"""Verify IS_ROCM flag behavior without importing the full hardware module."""
def test_hardware_py_has_is_rocm(self):
"""hardware.py should define IS_ROCM."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
assert "IS_ROCM: bool" in source and "False" in source
def test_hardware_py_sets_is_rocm_on_hip(self):
"""detect_hardware() should set IS_ROCM when torch.version.hip is set."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
assert 'torch.version, "hip"' in source or "torch.version.hip" in source
def test_hardware_py_still_returns_cuda_for_rocm(self):
"""DeviceType should remain CUDA even on ROCm -- no DeviceType.ROCM."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
enum_section = source.split("class DeviceType")[1].split("\n\n")[0]
assert "ROCM" not in enum_section
def test_hardware_py_has_rocm_in_package_versions(self):
"""get_package_versions() should include 'rocm' key."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
assert '"rocm"' in source
def test_hardware_py_device_type_cuda_references_intact(self):
"""All existing DeviceType.CUDA references should still be present."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
assert "DeviceType.CUDA" in source
assert "DEVICE = DeviceType.CUDA" in source
def test_is_rocm_exported_from_init(self):
"""IS_ROCM should be exported from hardware __init__.py."""
init_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "__init__.py"
source = init_path.read_text(encoding = "utf-8")
assert "IS_ROCM" in source
def test_is_rocm_in_all_list(self):
"""IS_ROCM should be in __all__ list in __init__.py."""
init_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "__init__.py"
source = init_path.read_text(encoding = "utf-8")
assert '"IS_ROCM"' in source
def test_get_package_versions_returns_rocm_key(self):
"""get_package_versions() source should return both 'cuda' and 'rocm' keys."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
func_start = source.find("def get_package_versions")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert '"cuda"' in func_body
assert '"rocm"' in func_body
def test_distributed_stubs_cover_is_torchelastic_launched(self):
"""Must stub is_torchelastic_launched (Windows ROCm torch.distributed lacks it)."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
assert "is_torchelastic_launched" in source
def test_distributed_stubs_cover_core_helpers(self):
"""_determine_attention_impl_for_gpu_estimate must stub the four core distributed helpers."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
for attr in ("is_initialized", "is_available", "get_rank", "get_world_size"):
assert attr in source, f"distributed stub for '{attr}' missing from hardware.py"
# TEST: tokenizer_utils.py -- error message
class TestTokenizerErrorMessage:
"""Verify the AMD error message is updated."""
def test_no_old_amd_message(self):
"""Old 'We do not support AMD' message should be gone."""
tu_path = PACKAGE_ROOT / "unsloth" / "tokenizer_utils.py"
source = tu_path.read_text(encoding = "utf-8")
assert "We do not support AMD" not in source
def test_new_message_has_docs_link(self):
"""New message should point to Unsloth AMD docs."""
tu_path = PACKAGE_ROOT / "unsloth" / "tokenizer_utils.py"
source = tu_path.read_text(encoding = "utf-8")
assert "docs.unsloth.ai" in source or "No GPU detected" in source
# TEST: install.sh -- structural checks
class TestInstallShStructure:
"""Verify install.sh structural properties without running it."""
def test_no_here_strings(self):
"""install.sh must not use the bash-only `<<<` here-string operator (breaks dash)."""
import re
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
for i, line in enumerate(source.splitlines(), 1):
stripped = line.lstrip()
if stripped.startswith("#"):
continue
# Strip quoted literals so `<<<` inside them is ignored.
unquoted = re.sub(r"'[^']*'", "", line)
unquoted = re.sub(r'"[^"]*"', "", unquoted)
assert "<<<" not in unquoted, f"install.sh:{i} uses non-POSIX <<< here-string"
def test_rocm_detection_present(self):
"""install.sh should have ROCm detection in get_torch_index_url."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
assert "amd-smi" in source
def test_cpu_index_note_respects_explicit_pin(self):
"""An explicit UNSLOTH_TORCH_INDEX_URL/_FAMILY CPU pin is a request, not
a detection failure: the */cpu wheel note must report the pin instead of
claiming ROCm/HIP is unusable, the WSL setup guidance must be skipped,
and the gpu summary must not label a pinned AMD host "no usable ROCm".
The property is ORDER: the pin arm opens the chain the message sits in, so
a pinned host never reaches the message. This used to be approximated by a
character window before the message, which is a different claim -- it is a
budget on how much source may sit between the two, and every arm added to
the chain eats into it. #8529 added an unsupported-arch arm and pushed the
real distance to ~1121, so the window had to grow 400 -> 1400 for reasons
that had nothing to do with what the test is for. Walk the enclosing
if/elif chains instead: no distance, nothing to re-tune.
"""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
self._assert_guarded_by_pin_arm(
source,
'substep "AMD GPU detected, but no usable ROCm/HIP install',
'substep "CPU-only PyTorch (index pinned via UNSLOTH_TORCH_INDEX_URL / _FAMILY)."',
"elif _has_amd_rocm_gpu; then",
"the */cpu note must check the explicit pin before diagnosing ROCm",
)
assert (
'[ "$OS" = "wsl" ] && [ "$_torch_index_pinned" = false ]' in source
), "ROCm-on-WSL guidance is detection advice; skip it for pinned installs"
self._assert_guarded_by_pin_arm(
source,
'step "gpu" "AMD GPU (no usable ROCm -- CPU fallback)"',
'step "gpu" "AMD GPU (torch index pinned: $_torch_index_leaf)"',
"else",
"the gpu summary must not claim no usable ROCm for a pinned index",
)
_PIN_ARM = 'if [ "$_torch_index_pinned" = true ]'
# `fi` as a command of its own: at the start of the line, or after a `;`. install.sh
# closes chains inline too (`else :; fi`), and such a closure is invisible to a
# startswith("fi") test, leaving a closed chain looking open across the message.
_CLOSER = re.compile(r"(?:^|;)\s*fi\b")
# install.sh indents in 4-space steps, so a closure less than one step deeper than the
# arm closes the arm's own chain (or an enclosing one) rather than something nested in
# it. Comparing against the step, not the exact indent, keeps a `fi` re-indented by a
# space or two counting as the closure it is.
_INDENT_STEP = 4
@staticmethod
def _indent(line):
return len(line) - len(line.lstrip())
@classmethod
def _assert_guarded_by_pin_arm(cls, source, needle, pinned_note, arm, message):
"""Assert ``needle`` sits in the ``arm`` arm of a still-open chain whose
first arm is the pin check and whose pin arm says ``pinned_note``.
Deliberately local rather than a whole-file walk of if/elif/fi. install.sh
embeds awk programs and PowerShell in quoted blocks, all of which contain
their own `if (...)`, and it closes some chains with a trailing `; fi`; a
global keyword walk mis-nests on both and can end up satisfying this guard
from an unrelated chain. Indentation is the structure install.sh is written
in and needs no parse: find the pin arm above the message, then require the
chain it opens to still be open (nothing closes it before the message) and
to have moved on to another arm (an `elif`/`else` at its indent).
The pin arm has to be the exact whole line and the arm the message sits in
has to be the exact expected one, or "a chain opened by something that
mentions the pin" is enough to pass: a never-firing `&& [ ... ]` extension
of the condition, or a decoy chain left open across the message with the
real guard deleted, both read as guards otherwise. And a guard that routes
pinned hosts to an empty arm still suppresses the message, so the pin arm
must be checked to carry its own note.
"""
lines = source.splitlines()
# Skip comments: a commented-out `substep` is not a message users ever see.
hits = [
i
for i, line in enumerate(lines)
if needle in line and not line.lstrip().startswith("#")
]
assert len(hits) == 1, f"expected exactly one {needle!r} in install.sh, found {len(hits)}"
message_line = hits[0]
opens = [i for i in range(message_line) if lines[i].strip() == cls._PIN_ARM + "; then"]
assert opens, f"{message}\nno pin check appears anywhere above the message"
pin_line = opens[-1]
indent = cls._indent(lines[pin_line])
between = range(pin_line + 1, message_line)
closed = [
i + 1
for i in between
if cls._CLOSER.search(lines[i].strip())
and cls._indent(lines[i]) < indent + cls._INDENT_STEP
]
assert not closed, (
f"{message}\ninstall.sh:{pin_line + 1} opens on the pin check but the chain "
f"closes at install.sh:{closed[0]}, before the message at "
f"install.sh:{message_line + 1}, so the two are unrelated"
)
later_arms = [
i
for i in between
if re.match(r"^(elif|else)\b", lines[i].strip()) and cls._indent(lines[i]) == indent
]
assert later_arms, (
f"{message}\nthe message at install.sh:{message_line + 1} is inside the pin "
f"arm opened at install.sh:{pin_line + 1}, so a pinned host still reaches it"
)
# The last arm before the message is the one the message is actually in.
enclosing_arm = lines[later_arms[-1]].strip()
assert enclosing_arm == arm, (
f"{message}\nthe message at install.sh:{message_line + 1} sits in "
f"`{enclosing_arm}` at install.sh:{later_arms[-1] + 1}, not in the expected "
f"`{arm}`, so the chain guarding it is not the one this test describes"
)
pin_body = lines[pin_line + 1 : later_arms[0]]
assert any(
pinned_note in line and not line.lstrip().startswith("#") for line in pin_body
), (
f"{message}\nthe pin arm at install.sh:{pin_line + 1} no longer says "
f"{pinned_note!r}, so a pinned host is told nothing at all"
)
_ROCM_VERSION_SOURCES = (
"_rocm_tag_from_amd_smi",
"_rocm_tag_from_version_file",
"_rocm_tag_from_hipconfig",
"_rocm_tag_from_dpkg",
"_rocm_tag_from_rpm",
)
def _rocm_version_detection_script(self, source: str, rocm_prefix: str) -> str:
"""The version-source helpers plus the resolver and the guarded
assignment that get_torch_index_url makes, as a runnable script.
/opt/rocm is redirected to `rocm_prefix` so the version file is a source
the test controls: a real ROCm host would otherwise answer one of the
probes and make the assertions machine-dependent."""
parts = []
# _run_bounded first: _rocm_tag_from_rpm routes its query through it, so
# omitting it makes that source answer nothing and the per-position
# assertions below pass for the wrong reason.
for name in (
"_run_bounded",
*self._ROCM_VERSION_SOURCES,
"_highest_rocm_tag",
"_detect_rocm_version_tag",
):
body = _extract_sh_function_body(source, name)
assert body, f"install.sh no longer defines {name}()"
parts.append(body)
assignment = re.search(
r'^ _rocm_tag=\$\(_detect_rocm_version_tag.*?\|\| _rocm_tag=""\n',
source,
re.S | re.M,
)
assert assignment, "could not extract the guarded _rocm_tag assignment"
parts.append(assignment.group(0))
return "\n".join(parts).replace("/opt/rocm", rocm_prefix)
def test_rocm_version_chain_survives_no_source_under_set_e(self):
"""When every ROCm version source is missing (e.g. rocminfo present but
rocm-core not installed, so dpkg-query/rpm exit 1), detection must still
return empty AND succeed; otherwise set -e kills the installer BEFORE the
actionable no-version WARN it feeds. Executed, not text."""
shell = shutil.which("bash")
if not shell:
pytest.skip("bash needed to execute the version chain")
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
with tempfile.TemporaryDirectory() as d:
# Nothing is created under this prefix, so every source comes up empty.
script_body = self._rocm_version_detection_script(source, os.path.join(d, "rocm"))
# Tools exist on PATH but yield nothing usable, like a box with the
# probe tools installed and no rocm-core package.
for name in ("amd-smi", "hipconfig", "dpkg-query", "rpm"):
p = os.path.join(d, name)
with open(p, "w", encoding = "utf-8") as f:
f.write("#!/bin/sh\nexit 1\n")
os.chmod(p, 0o755)
script = (
"set -euo pipefail\n" + script_body + '\nprintf "SURVIVED:%s\\n" "$_rocm_tag"\n'
)
env = dict(os.environ, PATH = d + os.pathsep + os.environ.get("PATH", ""))
r = subprocess.run([shell, "-c", script], env = env, capture_output = True, text = True)
assert r.returncode == 0, f"version detection aborted under set -e: {r.stderr}"
assert r.stdout.strip() == "SURVIVED:", r.stdout
assert "rocm" in source.lower()
def test_rocm_version_detection_takes_the_highest_source(self):
"""Issue #8402: Debian 13 ships hipconfig 5.7 next to a 6.1 runtime, so
stopping at the first source that answered gated a working gfx1100 out to
CPU wheels. Every source is read and the highest wins. Executed, not
text, and asserted per source position so no single ordering passes by
accident."""
shell = shutil.which("bash")
if not shell:
pytest.skip("bash needed to execute the version chain")
source = (PACKAGE_ROOT / "install.sh").read_text(encoding = "utf-8")
# Each source answers; whichever position holds the 6.4 reading must win.
# "version-file" is /opt/rocm/.info/version rather than a PATH tool.
outputs = {
"amd-smi": "AMDSMI Tool: 25.0.1 | ROCm version: {v}",
"version-file": "{v}.1-98",
"hipconfig": "{v}.31921-0",
"dpkg-query": "1:{v}.4-1",
"rpm": "{v}.4-1",
}
for winner in outputs:
with tempfile.TemporaryDirectory() as d:
rocm_prefix = os.path.join(d, "rocm")
script_body = self._rocm_version_detection_script(source, rocm_prefix)
for name, template in outputs.items():
line = template.format(v = "6.4" if name == winner else "5.7")
if name == "version-file":
os.makedirs(os.path.join(rocm_prefix, ".info"), exist_ok = True)
with open(
os.path.join(rocm_prefix, ".info", "version"), "w", encoding = "utf-8"
) as f:
f.write(line + "\n")
continue
if name == "dpkg-query":
_write_dpkg_query_stub(os.path.join(d, name), line, "installed")
continue
p = os.path.join(d, name)
with open(p, "w", encoding = "utf-8") as f:
f.write(f"#!/bin/sh\necho '{line}'\n")
os.chmod(p, 0o755)
script = "set -euo pipefail\n" + script_body + '\nprintf "TAG:%s\\n" "$_rocm_tag"\n'
env = dict(os.environ, PATH = d + os.pathsep + os.environ.get("PATH", ""))
r = subprocess.run([shell, "-c", script], env = env, capture_output = True, text = True)
assert r.returncode == 0, r.stderr
assert r.stdout.strip() == "TAG:rocm6.4", (
f"{winner} reported 6.4 while every other source reported 5.7, "
f"but detection resolved {r.stdout.strip()}"
)
def test_removed_dpkg_rocm_core_cannot_pick_the_wheels(self):
"""Highest-wins fixed the undershoot in #8402 and opened the symmetric
hole: a source reading HIGHER than the runtime now wins outright.
Undershoot lands on CPU wheels, which work; overshoot installs wheels the
runtime cannot load. `dpkg-query -W` reports packages left in "deinstall
ok config-files" by `apt remove` without `apt purge`, still carrying the
version they had, so a host that ran ROCm 7.0 and went back to 6.1 hands
dpkg a 7.0 that beats every live source. Only "installed" counts.
Executed, not text, and paired with the installed case so the assertion
cannot pass on the version alone."""
shell = shutil.which("bash")
if not shell:
pytest.skip("bash needed to execute the version chain")
source = (PACKAGE_ROOT / "install.sh").read_text(encoding = "utf-8")
for status, expected in (("config-files", "TAG:rocm6.1"), ("installed", "TAG:rocm7.0")):
with tempfile.TemporaryDirectory() as d:
rocm_prefix = os.path.join(d, "rocm")
script_body = self._rocm_version_detection_script(source, rocm_prefix)
# The live runtime, via the version file: ROCm 6.1.
os.makedirs(os.path.join(rocm_prefix, ".info"), exist_ok = True)
with open(
os.path.join(rocm_prefix, ".info", "version"), "w", encoding = "utf-8"
) as f:
f.write("6.1.2-98\n")
_write_dpkg_query_stub(os.path.join(d, "dpkg-query"), "1:7.0.0-1", status)
# Silence the rest so a real ROCm on the test host cannot answer.
for name in ("amd-smi", "hipconfig", "rpm"):
p = os.path.join(d, name)
with open(p, "w", encoding = "utf-8") as f:
f.write("#!/bin/sh\nexit 1\n")
os.chmod(p, 0o755)
script = "set -euo pipefail\n" + script_body + '\nprintf "TAG:%s\\n" "$_rocm_tag"\n'
env = dict(os.environ, PATH = d + os.pathsep + os.environ.get("PATH", ""))
r = subprocess.run([shell, "-c", script], env = env, capture_output = True, text = True)
assert r.returncode == 0, r.stderr
assert r.stdout.strip() == expected, (
f"dpkg rocm-core 7.0 in state {status!r} next to a 6.1 version file "
f"resolved {r.stdout.strip()}, expected {expected}"
)
def test_cuda_precedence(self):
"""ROCm detection runs only when NVIDIA is absent (check runtime ordering in get_torch_index_url)."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
body = _extract_sh_function_body(source, "get_torch_index_url")
nvidia_call = body.find("_has_usable_nvidia_gpu")
# Gate uses _nvidia_detected (not -z "$_smi") to handle proc-only NVIDIA
# hosts where nvidia-smi is absent but the GPU is found via /proc.
no_nvidia_branch = body.find('if [ "$_nvidia_detected" -eq 0 ]')
if no_nvidia_branch < 0:
no_nvidia_branch = body.find('if [ -z "$_smi" ]')
rocm_call = body.find("_has_amd_rocm_gpu")
assert nvidia_call >= 0, "get_torch_index_url should call _has_usable_nvidia_gpu"
assert no_nvidia_branch >= 0, "get_torch_index_url should gate ROCm on no-nvidia branch"
assert (
rocm_call > no_nvidia_branch
), "ROCm detection should sit inside the 'no NVIDIA' branch"
assert (
nvidia_call < no_nvidia_branch
), "NVIDIA detection should run before the no-NVIDIA branch"
def test_bitsandbytes_amd_install(self):
"""install.sh should install bitsandbytes for AMD when ROCm detected."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
assert "bitsandbytes" in source
assert "rocm*)" in source # case pattern for ROCm URLs
def test_cpu_hint_mentions_amd(self):
"""CPU-only hint should mention AMD ROCm."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
assert "ROCm" in source
def test_rocm72_supported_future_capped(self):
"""ROCm 7.2 should pass through directly; 7.3+ falls back to rocm7.2."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
assert 'echo "$_base/rocm7.2"' in source # fallback for unknown future versions
assert "rocm6.*" in source
assert "rocm7.0" in source
assert "rocm7.1" in source
assert "rocm7.2" in source
def test_rocm_tag_validation_guard_exists(self):
"""install.sh should validate _rocm_tag with a case guard."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
assert "rocm[1-9]*.[0-9]*)" in source
assert '_rocm_tag=""' in source # rejection path
def test_dpkg_epoch_handling(self):
"""install.sh should strip Debian epoch prefix from dpkg-query output."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
assert "sed 's/^[0-9]*://' " in source or "sed 's/^[0-9]*://'" in source
def test_no_double_bracket_in_rocm_block(self):
"""ROCm block must not use bash-only [[ ]] (POSIX char classes [[:space:]] are fine)."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
func_start = source.find("get_torch_index_url()")
func_end = source.find("\n}", func_start)
func_body = source[func_start:func_end]
import re
for i, line in enumerate(func_body.splitlines(), 1):
stripped = line.lstrip()
if stripped.startswith("#"):
continue
# Strip POSIX char classes [[:foo:]] before checking for [[ ]].
cleaned = re.sub(r"\[\[:[a-z]+:\]\]", "", line)
assert "[[" not in cleaned, f"get_torch_index_url line {i} uses non-POSIX [["
def test_no_arithmetic_expansion_in_rocm_block(self):
"""ROCm detection block should not use (( )) (bash-only)."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
func_start = source.find("get_torch_index_url()")
func_end = source.find("\n}", func_start)
func_body = source[func_start:func_end]
for i, line in enumerate(func_body.splitlines(), 1):
stripped = line.lstrip()
if stripped.startswith("#"):
continue
assert (
"((" not in line or "))" not in line or "$(()" in line
), f"get_torch_index_url line {i} may use non-POSIX (( ))"
def test_macos_returns_cpu_before_rocm_check(self):
"""macOS should return CPU immediately (before any ROCm check)."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
func_start = source.find("get_torch_index_url()")
func_body = source[func_start:]
darwin_pos = func_body.find("Darwin")
rocm_pos = func_body.find("amd-smi")
assert darwin_pos < rocm_pos, "macOS check should come before ROCm detection"
def test_unsloth_torch_backend_exported_after_get_torch_index_url(self):
"""install.sh exports UNSLOTH_TORCH_BACKEND after TORCH_INDEX_URL (lets the stack skip GPU re-detection)."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
torch_url_pos = source.find("TORCH_INDEX_URL=$(get_torch_index_url)")
backend_pos = source.find("UNSLOTH_TORCH_BACKEND")
assert backend_pos > 0, "UNSLOTH_TORCH_BACKEND must be set in install.sh"
assert (
backend_pos > torch_url_pos
), "UNSLOTH_TORCH_BACKEND must be set AFTER TORCH_INDEX_URL is resolved"
assert '"cuda"' in source[backend_pos : backend_pos + 500]
assert '"rocm"' in source[backend_pos : backend_pos + 500]
assert '"cpu"' in source[backend_pos : backend_pos + 500]
# Must be exported so subprocesses see it.
assert "export UNSLOTH_TORCH_BACKEND" in source
def test_kfd_sysfs_amd_vendor_check_in_has_amd_rocm_gpu(self):
"""_has_amd_rocm_gpu sysfs fallback must require AMD vendor_id 4098 (nvidia-open registers KFD nodes too)."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
func_start = source.find("_has_amd_rocm_gpu()")
func_end = source.find("\n}", func_start)
func_body = source[func_start:func_end]
assert (
"vendor_id" in func_body
), "_has_amd_rocm_gpu sysfs fallback must check vendor_id to exclude NVIDIA KFD nodes"
assert (
"4098" in func_body
), "_has_amd_rocm_gpu sysfs fallback must require AMD vendor_id 4098 (0x1002)"
def test_kfd_awk_vendor_check_is_per_line(self):
"""KFD sysfs awk must decide on a single vendor_id line, with no cross-node state.
The old awk paired two per-node flags (gpu_id + vendor_id) and needed an FNR==1
reset so flags from different KFD nodes could not combine into a Ryzen+NVIDIA
false positive. gpu_id is a sibling sysfs file and never appears inside
properties, so that pairing also never matched at all (every ROCm-less AMD host
was reported as no-GPU). The replacement keys on one atomic line: only an AMD
GPU node reports `vendor_id 4098` (KFD CPU nodes report 0, NVIDIA's open kernel
module registers 4318), so there is no cross-file state left to reset.
"""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
func_start = source.find("_has_amd_rocm_gpu()")
func_end = source.find("\n}", func_start)
func_body = source[func_start:func_end]
assert "$2 == 4098" in func_body, (
"_has_amd_rocm_gpu KFD awk must match `vendor_id 4098` as a single-line "
"condition so no per-node state can leak across KFD nodes"
)
assert "/gpu_id/" not in func_body, (
"_has_amd_rocm_gpu KFD awk must not key on a gpu_id line: gpu_id is a "
"sibling sysfs file, not a line in properties, so it never matches there"
)
def test_setup_sh_kfd_awk_matches_install_sh(self):
"""setup.sh's KFD fallback must use the same per-line vendor_id check as install.sh.
setup.sh re-probes AMD detection independently of install.sh; if its copy keeps
the dead gpu_id-inside-properties pairing, a host that install.sh routes to ROCm
still gets a CPU llama.cpp from the setup step (_setup_amd_detected stays false).
"""
source = (PACKAGE_ROOT / "studio" / "setup.sh").read_text(encoding = "utf-8")
assert (
"$2 == 4098" in source
), "setup.sh KFD awk must match `vendor_id 4098` as a single-line condition"
assert (
"/gpu_id/" not in source
), "setup.sh KFD awk must not key on a gpu_id line inside properties"
def test_kfd_only_torch_falls_back_to_cpu(self):
"""An AMD host whose gfx arch can't be read (rocminfo/amd-smi missing, or
present but not enumerating the GPU) must route torch to CPU, not a generic
rocm index: a Strix box (gfx1150/1151) would otherwise get the broken
_grouped_mm wheels because the reroute has no gfx to correct it."""
source = (PACKAGE_ROOT / "install.sh").read_text(encoding = "utf-8")
body = _extract_sh_function_body(source, "get_torch_index_url")
probe = body.find("_amd_gfx_probe=$(_probe_amd_gfx_arch)")
assert probe >= 0, "get_torch_index_url must probe the gfx arch before picking a rocm index"
# The shared probe reads gfx (not just tests binary presence), from rocminfo
# AND amd-smi, so an installed-but-not-enumerating probe still falls to CPU.
helper = _extract_sh_function_body(source, "_probe_amd_gfx_arch")
assert helper, "install.sh must define the shared _probe_amd_gfx_arch helper"
assert (
"rocminfo 2>/dev/null) | grep -oE 'gfx" in helper
), "probe must read gfx from rocminfo"
assert (
"amd-smi list 2>/dev/null) | grep -oE 'gfx" in helper
), "probe must read gfx from amd-smi"
# The probe clears ROCR/HIP_VISIBLE_DEVICES so a container mask
# (ROCR_VISIBLE_DEVICES=-1) can't blind the env-independent KFD detection.
assert (
"unset ROCR_VISIBLE_DEVICES HIP_VISIBLE_DEVICES" in helper
), "the gfx probe must clear the visibility masks so a mask can't force CPU"
cpu_guard = body.find('if [ -z "$_amd_gfx_probe" ]')
assert cpu_guard >= 0, "unreadable gfx must fall back to CPU"
assert cpu_guard < body.find(
"_rocm_tag="
), "the gfx gate must run before the ROCm version/index selection"
def test_kfd_only_llama_requires_hipcc(self):
"""setup.sh must forward --has-rocm for a gfx-unknown (KFD-only) host only when
hipcc is present. With no gfx the prebuilt resolver finds no ROCm bundle and the
source build would fail, so without a HIP toolchain the host keeps the CPU
prebuilt rather than breaking the llama.cpp install."""
source = (PACKAGE_ROOT / "studio" / "setup.sh").read_text(encoding = "utf-8")
idx = source.find("_PREBUILT_CMD+=(--has-rocm)")
assert idx >= 0, "setup.sh must still be able to forward --has-rocm"
window = source[max(0, idx - 900) : idx]
assert (
"hipcc" in window
), "the gfx-unknown --has-rocm branch must gate on hipcc (a usable HIP toolchain)"
assert (
"command -v hipcc" in window or "/opt/rocm/bin/hipcc" in window
), "hipcc presence must be checked via command -v or the rocm bin path"
assert (
"/opt/rocm-*/bin/hipcc" in window
), "the hipcc gate must also accept a versioned /opt/rocm-*/bin/hipcc toolchain"
def test_gfx_unknown_guard_honors_override(self):
"""A user-set UNSLOTH_ROCM_GFX_ARCH must seed the gfx probe before the CPU
fallback: an air-gapped/rocminfo-less Strix host that names its arch should
still reach a rocm index instead of being forced to CPU."""
source = (PACKAGE_ROOT / "install.sh").read_text(encoding = "utf-8")
helper = _extract_sh_function_body(source, "_probe_amd_gfx_arch")
assert helper, "install.sh must define the shared _probe_amd_gfx_arch helper"
seed = helper.find("$(printf")
assert seed >= 0, "the gfx probe must seed from UNSLOTH_ROCM_GFX_ARCH"
assert "UNSLOTH_ROCM_GFX_ARCH" in helper[seed : seed + 80]
assert seed < helper.find(
"rocminfo 2>/dev/null) | grep -oE 'gfx"
), "the override must be read before probing rocminfo"
body = _extract_sh_function_body(source, "get_torch_index_url")
call = body.find("_amd_gfx_probe=$(_probe_amd_gfx_arch)")
assert call >= 0, "get_torch_index_url must call the shared probe"
assert call < body.find(
'if [ -z "$_amd_gfx_probe" ]; then'
), "the probe must run before the CPU fallback guard"
def test_gfx_override_seeds_reroute_without_tools(self):
"""The Strix reroute must honour UNSLOTH_ROCM_GFX_ARCH even when rocminfo and
amd-smi are absent, so a manual override reaches the arch index; with no
override and no tools it must stay empty (no false Strix routing)."""
shell = shutil.which("bash")
if not shell:
pytest.skip("bash needed to execute the probe block")
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
block = re.search(
r'^ _gfx_all=\$\(printf[^\n]*\n.*?(?=^ _strix_gfx="")',
source,
re.S | re.M,
)
assert block, "could not extract the gfx-detection block"
with tempfile.TemporaryDirectory() as d:
# Shim rocminfo/amd-smi to enumerate nothing, so only the override can
# supply a gfx (keeps coreutils on PATH for tr/grep/printf).
for name in ("rocminfo", "amd-smi"):
p = os.path.join(d, name)
with open(p, "w", encoding = "utf-8") as f:
f.write("#!/bin/sh\nexit 0\n")
os.chmod(p, 0o755)
script = (
'set -euo pipefail\nHIP_VISIBLE_DEVICES=""\nROCR_VISIBLE_DEVICES=""\n'
+ block.group(0)
+ '\nprintf "OK:%s\\n" "$_gfx_all"\n'
)
def run(**extra):
env = dict(os.environ, PATH = d + os.pathsep + os.environ.get("PATH", ""), **extra)
return subprocess.run(
[shell, "-c", script], env = env, capture_output = True, text = True
)
r = run(UNSLOTH_ROCM_GFX_ARCH = "GFX1151")
assert r.returncode == 0, f"override probe aborted: {r.stderr}"
assert "OK:gfx1151" in r.stdout, f"override not honoured/lowercased: {r.stdout!r}"
r2 = run()
assert r2.returncode == 0, f"empty probe aborted: {r2.stderr}"
assert (
"OK:\n" in r2.stdout or r2.stdout.strip() == "OK:"
), f"no override + no tools must leave gfx empty: {r2.stdout!r}"
def test_gfx_probe_ignores_visibility_mask(self):
"""A container visibility mask (ROCR_VISIBLE_DEVICES=-1) must not blind the
gfx probe: rocminfo honours the mask and would enumerate nothing, but KFD
detection is env-independent, so the probe clears the mask and still reads
the arch (else a masked host is wrongly forced to CPU)."""
shell = shutil.which("bash")
if not shell:
pytest.skip("bash needed to execute the probe block")
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
probe_fn = _extract_sh_function_body(source, "_probe_amd_gfx_arch")
assert probe_fn, "could not extract _probe_amd_gfx_arch"
with tempfile.TemporaryDirectory() as d:
# rocminfo that mimics ROCR_VISIBLE_DEVICES=-1 hiding all agents.
with open(os.path.join(d, "rocminfo"), "w", encoding = "utf-8") as f:
f.write(
"#!/bin/sh\n"
'if [ "${ROCR_VISIBLE_DEVICES:-}" = "-1" ]; then echo "no agents"; exit 0; fi\n'
'echo " Name: gfx1151"\n'
)
os.chmod(os.path.join(d, "rocminfo"), 0o755)
script = (
"set -euo pipefail\n"
"_ensure_rocm_probe_env() { :; }\n"
+ probe_fn
+ '\n_amd_gfx_probe=$(_probe_amd_gfx_arch)\nprintf "OK:%s\\n" "$_amd_gfx_probe"\n'
)
def run(**extra):
env = dict(os.environ, PATH = d + os.pathsep + os.environ.get("PATH", ""), **extra)
return subprocess.run(
[shell, "-c", script], env = env, capture_output = True, text = True
)
r = run(ROCR_VISIBLE_DEVICES = "-1")
assert r.returncode == 0, f"masked probe aborted: {r.stderr}"
assert (
"OK:gfx1151" in r.stdout
), f"a visibility mask must not blind the gfx probe: {r.stdout!r}"
def test_kfd_only_inferable_gfx_defers_to_reroute(self):
"""A KFD-only host (GPU detected, gfx unreadable) whose arch IS inferable
from hardware IDs must not print the 'installing CPU-only PyTorch' warning:
get_torch_index_url returns the cpu index quietly and the runtime-less
reroute upgrades it to AMD per-arch wheels. Only when inference also fails
(or maps to no supported family) is CPU final, with the actionable hint."""
shell = shutil.which("bash")
if not shell:
pytest.skip("bash needed to execute get_torch_index_url")
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
fn = _extract_sh_function_body(source, "get_torch_index_url")
probe_fn = _extract_sh_function_body(source, "_probe_amd_gfx_arch")
family_fn = _extract_sh_function_body(source, "_amd_arch_index_family_for_gfx")
assert fn and probe_fn and family_fn
with tempfile.TemporaryDirectory() as d:
# uname -> Linux/x86_64 so the AMD branch runs on any dev host; the
# rocminfo/amd-smi shims enumerate nothing (KFD-only host).
with open(os.path.join(d, "uname"), "w", encoding = "utf-8", newline = "\n") as f:
f.write('#!/bin/sh\ncase "${1:-}" in -m) echo x86_64 ;; *) echo Linux ;; esac\n')
for name in ("rocminfo", "amd-smi"):
with open(os.path.join(d, name), "w", encoding = "utf-8", newline = "\n") as f:
f.write("#!/bin/sh\nexit 0\n")
for name in ("uname", "rocminfo", "amd-smi"):
os.chmod(os.path.join(d, name), 0o755)
def run(infer_stub):
script = (
"set -euo pipefail\n"
"_ensure_rocm_probe_env() { :; }\n"
"_trim_index_path_slashes() { printf '%s\\n' \"$1\"; }\n"
"_has_usable_nvidia_gpu() { return 1; }\n"
"_has_amd_rocm_gpu() { return 0; }\n"
+ infer_stub
+ "\n"
+ probe_fn
+ "\n"
+ family_fn
+ "\n"
+ fn
+ "\n"
"get_torch_index_url\n"
)
# Run from a file, not -c: Windows bash mangles multi-KB -c strings.
sp = os.path.join(d, "gtiu.sh")
with open(sp, "w", encoding = "utf-8", newline = "\n") as f:
f.write(script)
env = dict(os.environ, PATH = d + os.pathsep + os.environ.get("PATH", ""))
for var in (
"UNSLOTH_ROCM_GFX_ARCH",
"UNSLOTH_TORCH_INDEX_URL",
"UNSLOTH_TORCH_INDEX_FAMILY",
"UNSLOTH_PYTORCH_MIRROR",
"ROCR_VISIBLE_DEVICES",
"HIP_VISIBLE_DEVICES",
):
env.pop(var, None)
return subprocess.run(
[shell, sp.replace("\\", "/")], env = env, capture_output = True, text = True
)
r = run("_infer_linux_amd_gfx_arch() { echo gfx1100; }")
assert r.returncode == 0, f"inferable case aborted: {r.stderr}"
assert r.stdout.strip().endswith(
"/cpu"
), f"must hand */cpu to the reroute: {r.stdout!r}"
assert (
"inferring gfx1100" in r.stderr
), f"must announce the inference handoff: {r.stderr!r}"
assert (
"installing CPU-only PyTorch" not in r.stderr
), f"must not promise a CPU-only install the reroute will override: {r.stderr!r}"
r2 = run("_infer_linux_amd_gfx_arch() { return 1; }")
assert r2.returncode == 0, f"uninferable case aborted: {r2.stderr}"
assert r2.stdout.strip().endswith("/cpu")
assert (
"installing CPU-only PyTorch" in r2.stderr
), f"uninferable gfx must keep the actionable CPU warning: {r2.stderr!r}"
r3 = run("_infer_linux_amd_gfx_arch() { echo gfx906; }")
assert r3.returncode == 0, f"unsupported-family case aborted: {r3.stderr}"
assert r3.stdout.strip().endswith("/cpu")
assert (
"installing CPU-only PyTorch" in r3.stderr
), f"an inferred arch with no wheel family must keep the CPU warning: {r3.stderr!r}"
def test_no_version_cpu_warning_respects_gfx_override(self):
"""With UNSLOTH_ROCM_GFX_ARCH set on a KFD-only host that has no ROCm
version sources, the gfx probe is seeded by the override, so the
no-version endpoint used to print 'falling back to CPU-only PyTorch'
even though the reroute then installs the per-arch wheels (Codex P3).
A supported override must defer; an unsupported override, or a
readable-gfx host without an override, keeps the CPU warning."""
shell = shutil.which("bash")
if not shell:
pytest.skip("bash needed to execute get_torch_index_url")
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
fn = _extract_sh_function_body(source, "get_torch_index_url")
probe_fn = _extract_sh_function_body(source, "_probe_amd_gfx_arch")
family_fn = _extract_sh_function_body(source, "_amd_arch_index_family_for_gfx")
# The version helpers must be extracted too: without them get_torch_index_url
# calls a missing command, the guarded assignment swallows the 127, and the
# no-version endpoint is reached for the wrong reason. Verified by mutation
# (making _detect_rocm_version_tag return a real tag left this test green).
version_fns = [
_extract_sh_function_body(source, name)
for name in (
"_rocm_tag_from_amd_smi",
"_rocm_tag_from_version_file",
"_rocm_tag_from_hipconfig",
"_rocm_tag_from_dpkg",
"_rocm_tag_from_rpm",
"_highest_rocm_tag",
"_detect_rocm_version_tag",
)
]
assert fn and probe_fn and family_fn
assert all(version_fns), "ROCm version helpers not found in install.sh"
with tempfile.TemporaryDirectory() as d:
# Neutralise the host's real ROCm: the version chain reads
# /opt/rocm/.info/version directly (no tool to shim), which resolves a tag on
# a real ROCm box and skips the no-version endpoint under test. The read lives
# in _rocm_tag_from_version_file, so the rewrite has to land on the helper
# text; applying it to fn alone is a no-op.
version_fns = [
body.replace("/opt/rocm/.info/version", Path(d, "no-rocm-version").as_posix())
for body in version_fns
]
with open(os.path.join(d, "uname"), "w", encoding = "utf-8", newline = "\n") as f:
f.write('#!/bin/sh\ncase "${1:-}" in -m) echo x86_64 ;; *) echo Linux ;; esac\n')
# Silence every ROCm version source, not just amd-smi: a dev box with
# a real hipconfig/dpkg would otherwise resolve a version and skip
# the no-version endpoint this test exercises.
with open(os.path.join(d, "amd-smi"), "w", encoding = "utf-8", newline = "\n") as f:
f.write("#!/bin/sh\nexit 0\n")
for name in ("hipconfig", "dpkg-query", "rpm"):
with open(os.path.join(d, name), "w", encoding = "utf-8", newline = "\n") as f:
f.write("#!/bin/sh\nexit 1\n")
for name in ("uname", "amd-smi", "hipconfig", "dpkg-query", "rpm"):
os.chmod(os.path.join(d, name), 0o755)
script = (
"set -euo pipefail\n"
"_ensure_rocm_probe_env() { :; }\n"
"_trim_index_path_slashes() { printf '%s\\n' \"$1\"; }\n"
"_has_usable_nvidia_gpu() { return 1; }\n"
"_has_amd_rocm_gpu() { return 0; }\n"
"_infer_linux_amd_gfx_arch() { return 1; }\n"
+ probe_fn
+ "\n"
+ family_fn
+ "\n"
+ "\n".join(version_fns)
+ "\n"
+ fn
+ "\n"
"get_torch_index_url\n"
)
sp = os.path.join(d, "gtiu.sh")
with open(sp, "w", encoding = "utf-8", newline = "\n") as f:
f.write(script)
def run(rocminfo_body, **extra):
with open(os.path.join(d, "rocminfo"), "w", encoding = "utf-8", newline = "\n") as f:
f.write("#!/bin/sh\n" + rocminfo_body)
os.chmod(os.path.join(d, "rocminfo"), 0o755)
env = dict(os.environ, PATH = d + os.pathsep + os.environ.get("PATH", ""), **extra)
for var in (
"UNSLOTH_TORCH_INDEX_URL",
"UNSLOTH_TORCH_INDEX_FAMILY",
"UNSLOTH_PYTORCH_MIRROR",
"ROCR_VISIBLE_DEVICES",
"HIP_VISIBLE_DEVICES",
):
env.pop(var, None)
if "UNSLOTH_ROCM_GFX_ARCH" not in extra:
env.pop("UNSLOTH_ROCM_GFX_ARCH", None)
return subprocess.run(
[shell, sp.replace("\\", "/")], env = env, capture_output = True, text = True
)
# Supported override on a tool-blind host: defer to the reroute.
r = run("exit 0\n", UNSLOTH_ROCM_GFX_ARCH = "gfx1151")
assert r.returncode == 0, f"override case aborted: {r.stderr}"
assert r.stdout.strip().endswith("/cpu")
assert (
"falling back to CPU-only PyTorch" not in r.stderr
), f"a supported override must not get the false CPU warning: {r.stderr!r}"
assert (
"UNSLOTH_ROCM_GFX_ARCH=gfx1151 is set" in r.stderr
), f"the override deferral must be announced: {r.stderr!r}"
# Unsupported override: the reroute can't map it -> CPU warning stays.
r2 = run("exit 0\n", UNSLOTH_ROCM_GFX_ARCH = "gfx906")
assert r2.returncode == 0, f"unsupported-override case aborted: {r2.stderr}"
assert (
"falling back to CPU-only PyTorch" in r2.stderr
), f"an unmappable override must keep the CPU warning: {r2.stderr!r}"
# Readable gfx, no override, no version: deliberate CPU fallback.
r3 = run('echo " Name: gfx1151"\n')
assert r3.returncode == 0, f"readable-gfx case aborted: {r3.stderr}"
assert (
"falling back to CPU-only PyTorch" in r3.stderr
), f"a readable-gfx host without a version keeps the CPU warning: {r3.stderr!r}"
def test_reroute_gate_covers_kfd_only(self):
"""The runtime-less reroute must fire for a KFD-only host: _has_amd_rocm_gpu
is now true via the KFD topology, so the gate also accepts a detected GPU
whose gfx probe is empty (unslothai#7314 P2). A */cpu index chosen with a
READABLE gfx (deliberate ROCm-version fallback) must stay un-rerouted."""
shell = shutil.which("bash")
if not shell:
pytest.skip("bash needed to execute the reroute block")
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
block = re.search(
r'^if \[ "\$_torch_index_pinned" = false \] && \[ "\$SKIP_TORCH" = false \] && \\\n'
r".*?^fi\n",
source,
re.S | re.M,
)
assert block, "could not extract the runtime-less reroute block"
family_fn = _extract_sh_function_body(source, "_amd_arch_index_family_for_gfx")
assert family_fn
with tempfile.TemporaryDirectory() as d:
with open(os.path.join(d, "uname"), "w", encoding = "utf-8", newline = "\n") as f:
f.write('#!/bin/sh\ncase "${1:-}" in -m) echo x86_64 ;; *) echo Linux ;; esac\n')
os.chmod(os.path.join(d, "uname"), 0o755)
def run(gpu_stub, probe_stub):
script = (
"set -euo pipefail\n"
"_has_usable_nvidia_gpu() { return 1; }\n"
f"_has_amd_rocm_gpu() {{ {gpu_stub}; }}\n"
f"_probe_amd_gfx_arch() {{ {probe_stub}; }}\n"
"_infer_linux_amd_gfx_arch() { echo gfx1100; }\n"
"_strip_index_url_credentials() { printf '%s\\n' \"$1\"; }\n" + family_fn + "\n"
"_torch_index_pinned=false\nSKIP_TORCH=false\n_ARCH=x86_64\n"
"TORCH_INDEX_URL=https://download.pytorch.org/whl/cpu\n"
+ block.group(0)
+ 'printf "URL:%s GFX:%s\\n" "$TORCH_INDEX_URL" "${UNSLOTH_ROCM_GFX_ARCH:-}"\n'
)
# Run from a file, not -c: Windows bash mangles multi-KB -c strings.
sp = os.path.join(d, "reroute.sh")
with open(sp, "w", encoding = "utf-8", newline = "\n") as f:
f.write(script)
env = dict(os.environ, PATH = d + os.pathsep + os.environ.get("PATH", ""))
for var in ("UNSLOTH_ROCM_GFX_ARCH", "UNSLOTH_AMD_ROCM_MIRROR"):
env.pop(var, None)
return subprocess.run(
[shell, sp.replace("\\", "/")], env = env, capture_output = True, text = True
)
# KFD-only: GPU detected, probe empty -> reroute to per-arch wheels.
r = run("return 0", "printf '\\n'")
assert r.returncode == 0, f"kfd-only reroute aborted: {r.stderr}"
assert (
"URL:https://repo.amd.com/rocm/whl/gfx110X-all/ GFX:gfx1100" in r.stdout
), f"KFD-only host must reach the AMD arch index: {r.stdout!r}"
# The diagnostic must not claim /dev/kfd is missing: KFD visibility is
# exactly what detected this host (Codex P3).
assert (
"ROCm runtime not visible" not in r.stderr
), f"KFD-only reroute must not claim /dev/kfd is missing: {r.stderr!r}"
assert (
"visible via the kernel driver (KFD)" in r.stderr
), f"KFD-only reroute must name the tooling gap: {r.stderr!r}"
# Readable gfx: the */cpu index is a deliberate fallback -> untouched.
r2 = run("return 0", "echo gfx1151")
assert r2.returncode == 0, f"readable-gfx case aborted: {r2.stderr}"
assert (
"URL:https://download.pytorch.org/whl/cpu GFX:" in r2.stdout
), f"a deliberate CPU fallback must not be rerouted: {r2.stdout!r}"
# No AMD GPU detected at all: the pre-KFD-fix path still reroutes.
r3 = run("return 1", "printf '\\n'")
assert r3.returncode == 0, f"undetected-GPU case aborted: {r3.stderr}"
assert (
"URL:https://repo.amd.com/rocm/whl/gfx110X-all/ GFX:gfx1100" in r3.stdout
), f"the original undetected-GPU reroute must keep working: {r3.stdout!r}"
assert (
"ROCm runtime not visible" in r3.stderr
), f"a truly runtime-invisible host keeps the original diagnostic: {r3.stderr!r}"
def test_get_torch_index_url_uses_nvidia_detected_flag(self):
"""get_torch_index_url must track NVIDIA via _nvidia_detected (proc-only NVIDIA still picks CUDA)."""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text(encoding = "utf-8")
func_start = source.find("get_torch_index_url()")
func_end = source.find("\n}", func_start)
func_body = source[func_start:func_end]
assert "_nvidia_detected" in func_body, (
"get_torch_index_url must use a _nvidia_detected flag (separate from "
"_smi) so that proc-only NVIDIA detection still selects CUDA wheels"
)
assert (
'_nvidia_detected" -eq 0' in func_body or "_nvidia_detected" in func_body
), "get_torch_index_url AMD branch must be skipped when _nvidia_detected=1"
# TEST: Live regression on current host (NVIDIA B200 expected)
class TestLiveRegression:
"""Live checks that run on the actual host -- skip if no NVIDIA GPU."""
def test_get_torch_index_url_returns_cuda_on_nvidia(self):
"""On an NVIDIA machine, get_torch_index_url should return a CUDA URL."""
import shutil
if not shutil.which("nvidia-smi"):
pytest.skip("No nvidia-smi available")
# Skip if nvidia-smi exists but lists no GPU (binary without driver).
check = subprocess.run(
[
"bash",
"-c",
"nvidia-smi -L 2>/dev/null | awk '/^GPU[[:space:]]+[0-9]+:/{f=1} END{exit !f}'",
],
capture_output = True,
)
if check.returncode != 0:
pytest.skip("nvidia-smi is on PATH but no GPU is listed")
sh_path = PACKAGE_ROOT / "install.sh"
# All three helper definitions must be in scope when we eval the extract.
extract_cmd = (
f"sed -n '/^_has_amd_rocm_gpu()/,/^}}$/p; "
f"/^_has_usable_nvidia_gpu()/,/^}}$/p; "
f"/^get_torch_index_url()/,/^}}$/p' '{sh_path}'"
)
result = subprocess.run(
["bash", "-c", f'eval "$({extract_cmd})"; get_torch_index_url'],
capture_output = True,
text = True,
timeout = 30,
)
if result.returncode != 0:
pytest.skip("Could not extract get_torch_index_url for live test")
url = result.stdout.strip()
assert "cu1" in url or "cuda" in url.lower(), f"Expected CUDA URL, got: {url}"
# TEST: worker.py -- ROCm Mamba/SSM source build path
_WORKER_PATH = PACKAGE_ROOT / "studio" / "backend" / "core" / "training" / "worker.py"
_EXPORT_WORKER_PATH = PACKAGE_ROOT / "studio" / "backend" / "core" / "export" / "worker.py"
# Shared torchao Windows-ROCm stub used by both workers.
_TORCHAO_STUB_PATH = PACKAGE_ROOT / "studio" / "backend" / "core" / "_torchao_stub.py"
# RAG embedder -- runs in the main backend process and also needs the stub.
_EMBEDDINGS_PATH = PACKAGE_ROOT / "studio" / "backend" / "core" / "rag" / "embeddings.py"
# Wheel-probe script literal lives in wheel_utils after the resolver refactor.
_WHEEL_UTILS_PATH = PACKAGE_ROOT / "studio" / "backend" / "utils" / "wheel_utils.py"
class TestWorkerRocmMambaSsm:
"""Verify worker.py Mamba/SSM install logic on ROCm."""
def test_probe_returns_hip_version_field(self):
"""The wheel probe should include hip_version, and worker.py consumes it."""
assert "hip_version" in _WHEEL_UTILS_PATH.read_text(encoding = "utf-8")
assert "hip_version" in _WORKER_PATH.read_text(encoding = "utf-8")
def test_probe_script_has_getattr_hip(self):
"""Probe script should use getattr for torch.version.hip (safe on CUDA)."""
source = _WHEEL_UTILS_PATH.read_text(encoding = "utf-8")
assert "getattr(torch.version, 'hip', None)" in source
def test_direct_wheel_url_returns_none_without_cuda_major(self, monkeypatch):
"""direct_wheel_url should return None when cuda_major is empty (ROCm)."""
_worker_spec = importlib.util.spec_from_file_location("test_worker", _WORKER_PATH)
assert _worker_spec is not None and _worker_spec.loader is not None
worker_mod = importlib.util.module_from_spec(_worker_spec)
# Stub worker.py imports via monkeypatch so the fake "utils" is undone
# and doesn't break later tests importing the real utils.* package.
loggers_mock = MagicMock()
loggers_mock.get_logger = MagicMock(return_value = MagicMock())
monkeypatch.setitem(sys.modules, "structlog", MagicMock())
monkeypatch.setitem(sys.modules, "loggers", loggers_mock)
monkeypatch.setitem(sys.modules, "utils", MagicMock())
monkeypatch.setitem(sys.modules, "utils.hardware", MagicMock())
try:
_worker_spec.loader.exec_module(worker_mod)
except Exception:
pytest.skip("Could not load worker module in test environment")
env_rocm = {
"python_tag": "cp312",
"torch_mm": "2.6",
"cuda_major": "",
"hip_version": "7.1.12345",
"cxx11abi": "TRUE",
}
result = worker_mod.direct_wheel_url(
filename_prefix = "causal_conv1d",
package_version = "1.6.1",
release_tag = "v1.6.1.post4",
release_base_url = "https://github.com/Dao-AILab/causal-conv1d/releases/download",
env = env_rocm,
)
assert result is None
def test_hipcc_check_exists_in_source(self):
"""worker.py should check for hipcc before ROCm source builds."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "hipcc" in source
def test_rocm_source_build_status_message(self):
"""worker.py should send a specific status for ROCm source compilation."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "Compiling" in source and "from source for ROCm" in source
def test_rocm_build_failure_message(self):
"""worker.py should send a clear error on ROCm build failure."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "Failed to compile" in source and "for ROCm" in source
def test_timeout_on_install(self):
"""worker.py should have a timeout on pip install subprocess."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "TimeoutExpired" in source
assert "timeout" in source
# TEST: amd.py -- AMD GPU monitoring
class TestAmdGpuMonitoring:
"""Verify amd.py module structure and mock behavior."""
def test_amd_py_exists(self):
"""amd.py should exist in the hardware directory."""
amd_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "amd.py"
assert amd_path.exists()
def test_amd_py_has_required_functions(self):
"""amd.py should export the same function signatures as nvidia.py."""
amd_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "amd.py"
source = amd_path.read_text(encoding = "utf-8")
assert "def get_physical_gpu_count" in source
assert "def get_primary_gpu_utilization" in source
assert "def get_visible_gpu_utilization" in source
def test_amd_smi_json_parsing(self, monkeypatch):
"""Verify _extract_gpu_metrics parses amd-smi JSON correctly."""
amd_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "amd.py"
_amd_spec = importlib.util.spec_from_file_location("test_amd", amd_path)
assert _amd_spec is not None and _amd_spec.loader is not None
amd_mod = importlib.util.module_from_spec(_amd_spec)
loggers_mock = MagicMock()
loggers_mock.get_logger = MagicMock(return_value = MagicMock())
monkeypatch.setitem(sys.modules, "loggers", loggers_mock)
try:
_amd_spec.loader.exec_module(amd_mod)
except Exception:
pytest.skip("Could not load amd module in test environment")
gpu_data = {
"usage": {"gfx_activity": "85"},
"temperature": {"edge": "72"},
"power": {
"current_socket_power": "200.5",
"power_cap": "300",
},
"vram": {
"vram_used": 8192, # MB
"vram_total": 16384, # MB
},
}
metrics = amd_mod._extract_gpu_metrics(gpu_data)
assert metrics["gpu_utilization_pct"] == 85.0
assert metrics["temperature_c"] == 72.0
assert metrics["power_draw_w"] == 200.5
assert metrics["power_limit_w"] == 300.0
assert metrics["vram_used_gb"] == round(8192 / 1024, 2)
assert metrics["vram_total_gb"] == round(16384 / 1024, 2)
assert metrics["vram_utilization_pct"] is not None
assert metrics["power_utilization_pct"] is not None
def test_amd_primary_gpu_with_mock(self, monkeypatch):
"""get_primary_gpu_utilization returns correct dict with mocked amd-smi."""
amd_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "amd.py"
_amd_spec = importlib.util.spec_from_file_location("test_amd2", amd_path)
assert _amd_spec is not None and _amd_spec.loader is not None
amd_mod = importlib.util.module_from_spec(_amd_spec)
loggers_mock = MagicMock()
loggers_mock.get_logger = MagicMock(return_value = MagicMock())
monkeypatch.setitem(sys.modules, "loggers", loggers_mock)
try:
_amd_spec.loader.exec_module(amd_mod)
except Exception:
pytest.skip("Could not load amd module")
# _first_visible_amd_gpu_id() returns None if HIP/ROCR/CUDA_VISIBLE_DEVICES
# is "" or "-1"; CI often sets CUDA_VISIBLE_DEVICES="", so clear them.
for var in (
"HIP_VISIBLE_DEVICES",
"ROCR_VISIBLE_DEVICES",
"CUDA_VISIBLE_DEVICES",
):
monkeypatch.delenv(var, raising = False)
# amd-smi is gated off on Windows w/o a HIP SDK; opt in so the mock is
# allowed on every platform.
monkeypatch.setenv("UNSLOTH_ENABLE_AMD_SMI", "1")
mock_json = json.dumps(
[
{
"usage": {"gfx_activity": "50"},
"temperature": {"edge": "65"},
"power": {"current_socket_power": "150", "power_cap": "250"},
"vram": {"vram_used": 4096, "vram_total": 16384},
}
]
)
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = mock_json
# Premise is "amd-smi exists and answers": the guard which()-checks
# before spawning, so mock which too for hosts lacking a real amd-smi.
with patch.object(amd_mod.shutil, "which", return_value = "/usr/bin/amd-smi"):
with patch.object(subprocess, "run", return_value = mock_result):
result = amd_mod.get_primary_gpu_utilization()
assert result["available"] is True
assert result["gpu_utilization_pct"] == 50.0
assert result["temperature_c"] == 65.0
def test_amd_smi_not_found_returns_unavailable(self, monkeypatch):
"""get_primary_gpu_utilization returns available=False when amd-smi is missing."""
amd_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "amd.py"
_amd_spec = importlib.util.spec_from_file_location("test_amd3", amd_path)
assert _amd_spec is not None and _amd_spec.loader is not None
amd_mod = importlib.util.module_from_spec(_amd_spec)
loggers_mock = MagicMock()
loggers_mock.get_logger = MagicMock(return_value = MagicMock())
monkeypatch.setitem(sys.modules, "loggers", loggers_mock)
try:
_amd_spec.loader.exec_module(amd_mod)
except Exception:
pytest.skip("Could not load amd module")
# Opt in so the call reaches subprocess.run (testing OSError handling).
with (
patch.dict(os.environ, {"UNSLOTH_ENABLE_AMD_SMI": "1"}),
patch.object(subprocess, "run", side_effect = OSError("amd-smi not found")),
):
result = amd_mod.get_primary_gpu_utilization()
assert result["available"] is False
def test_amd_timeout_returns_unavailable(self, monkeypatch):
"""get_primary_gpu_utilization handles timeout gracefully."""
amd_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "amd.py"
_amd_spec = importlib.util.spec_from_file_location("test_amd4", amd_path)
assert _amd_spec is not None and _amd_spec.loader is not None
amd_mod = importlib.util.module_from_spec(_amd_spec)
loggers_mock = MagicMock()
loggers_mock.get_logger = MagicMock(return_value = MagicMock())
monkeypatch.setitem(sys.modules, "loggers", loggers_mock)
try:
_amd_spec.loader.exec_module(amd_mod)
except Exception:
pytest.skip("Could not load amd module")
# Opt in so the call reaches subprocess.run (testing timeout handling).
with (
patch.dict(os.environ, {"UNSLOTH_ENABLE_AMD_SMI": "1"}),
patch.object(
subprocess,
"run",
side_effect = subprocess.TimeoutExpired("amd-smi", 5),
),
):
result = amd_mod.get_primary_gpu_utilization()
assert result["available"] is False
# TEST: hardware.py -- IS_ROCM branching to amd.py
class TestHardwareAmdBranching:
"""Verify hardware.py branches to amd.py when IS_ROCM is True."""
def test_hardware_imports_amd_module(self):
"""hardware.py should import from amd module when IS_ROCM."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
assert "from . import amd" in source
def test_hardware_branches_on_is_rocm_for_utilization(self):
"""get_gpu_utilization dispatches visible metrics through amd.py on ROCm."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
func_start = source.find("def get_gpu_utilization")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "_smi_query(" in func_body
assert '"get_visible_gpu_utilization"' in func_body
assert "_reconcile_rocm_unified_memory" in func_body
smi = source[
source.find("def _smi_query") : source.find("\ndef ", source.find("def _smi_query") + 1)
]
assert "IS_ROCM" in smi
assert "from . import amd" in smi
def test_hardware_branches_on_is_rocm_for_visible(self):
"""get_visible_gpu_utilization dispatches to amd.py via _smi_query when IS_ROCM."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
func_start = source.find("def get_visible_gpu_utilization")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
# The dispatcher call may wrap; allow whitespace before the func name arg.
import re as _re
assert _re.search(r'_smi_query\(\s*"get_visible_gpu_utilization"', func_body)
smi = source[
source.find("def _smi_query") : source.find("\ndef ", source.find("def _smi_query") + 1)
]
assert "IS_ROCM" in smi
assert "from . import amd" in smi
def test_hardware_branches_on_is_rocm_for_physical_count(self):
"""get_physical_gpu_count should try amd.py when IS_ROCM."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
func_start = source.find("def get_physical_gpu_count")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "IS_ROCM" in func_body
assert "from . import amd" in func_body
# TEST: hardware.py -- apply_gpu_ids ROCm fallback (issue #5180)
class TestApplyGpuIdsRocmFallback:
"""apply_gpu_ids sets HIP_VISIBLE_DEVICES on ROCm hosts even when IS_ROCM is still False (issue #5180)."""
def test_apply_gpu_ids_falls_back_to_torch_version_hip(self):
"""apply_gpu_ids probes torch.version.hip when IS_ROCM is False and no ROCm env vars set."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
func_start = source.find("def apply_gpu_ids")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert 'getattr(_torch.version, "hip", None)' in func_body
def test_apply_gpu_ids_sets_hip_but_not_rocr_visible_devices(self):
"""apply_gpu_ids sets HIP_VISIBLE_DEVICES but leaves ROCR_VISIBLE_DEVICES inherited (HSA indexing; issue #6118)."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
func_start = source.find("def apply_gpu_ids")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert 'os.environ["HIP_VISIBLE_DEVICES"] = value' in func_body
assert 'os.environ["ROCR_VISIBLE_DEVICES"] = value' not in func_body
def test_apply_gpu_ids_rocm_fallback_is_guarded_by_try_except(self):
"""torch import in apply_gpu_ids must be wrapped in try/except so a missing torch never crashes."""
hw_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
source = hw_path.read_text(encoding = "utf-8")
func_start = source.find("def apply_gpu_ids")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "import torch as _torch" in func_body
assert "except Exception" in func_body
# TEST: install_python_stack.py -- Windows AMD warning
class TestWindowsRocmWarning:
"""Verify Windows AMD GPU detection and warning message."""
def test_windows_amd_warning_in_source(self):
"""install_python_stack.py should warn Windows AMD users."""
source = _STACK_PATH.read_text(encoding = "utf-8")
assert "AMD GPU detected" in source
def test_windows_amd_warning_checks_hipinfo_or_amdsmi(self):
"""Warning should check for hipinfo or amd-smi."""
source = _STACK_PATH.read_text(encoding = "utf-8")
assert "hipinfo" in source
assert "amd-smi" in source
def test_windows_amd_warning_has_docs_link(self):
"""Warning should include AMD docs link."""
source = _STACK_PATH.read_text(encoding = "utf-8")
assert "docs.unsloth.ai/get-started/install-and-update/amd" in source
# TEST: unsloth/kernels/utils.py -- is_rdna() expansion
class TestIsRdnaExpansion:
"""Verify is_rdna() covers RDNA2, RDNA3, RDNA3.5, RDNA4 architectures."""
def test_is_rdna_source_has_rdna2(self):
"""is_rdna() should include RDNA2 architectures."""
utils_path = PACKAGE_ROOT / "unsloth" / "kernels" / "utils.py"
source = utils_path.read_text(encoding = "utf-8")
func_start = source.find("def is_rdna()")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "gfx1030" in func_body
assert "gfx1031" in func_body
assert "gfx1032" in func_body
assert "gfx1033" in func_body
assert "gfx1034" in func_body
assert "gfx1035" in func_body
assert "gfx1036" in func_body
def test_is_rdna_source_has_rdna3(self):
"""is_rdna() should include RDNA3 architectures."""
utils_path = PACKAGE_ROOT / "unsloth" / "kernels" / "utils.py"
source = utils_path.read_text(encoding = "utf-8")
func_start = source.find("def is_rdna()")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "gfx1100" in func_body
assert "gfx1101" in func_body
assert "gfx1102" in func_body
assert "gfx1103" in func_body
def test_is_rdna_source_has_rdna35(self):
"""is_rdna() should include RDNA3.5 architectures."""
utils_path = PACKAGE_ROOT / "unsloth" / "kernels" / "utils.py"
source = utils_path.read_text(encoding = "utf-8")
func_start = source.find("def is_rdna()")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "gfx1150" in func_body
assert "gfx1151" in func_body
assert "gfx1152" in func_body
def test_is_rdna_source_has_rdna4(self):
"""is_rdna() should include RDNA4 architectures."""
utils_path = PACKAGE_ROOT / "unsloth" / "kernels" / "utils.py"
source = utils_path.read_text(encoding = "utf-8")
func_start = source.find("def is_rdna()")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "gfx1200" in func_body
assert "gfx1201" in func_body
def test_is_cdna_not_changed(self):
"""is_cdna() should remain unchanged (no RDNA architectures added)."""
utils_path = PACKAGE_ROOT / "unsloth" / "kernels" / "utils.py"
source = utils_path.read_text(encoding = "utf-8")
func_start = source.find("def is_cdna()")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "gfx940" in func_body
assert "gfx941" in func_body
assert "gfx942" in func_body
assert "gfx950" in func_body
# RDNA architectures should NOT be in is_cdna
assert "gfx1030" not in func_body
assert "gfx1100" not in func_body
# TEST: install_python_stack.py -- _windows_rocm_index_url arch mapping
class TestWindowsRocmIndexUrl:
"""Verify GPU arch → AMD pip index URL mapping."""
def test_gfx1200_maps_to_gfx120x_all(self):
url = stack_mod._windows_rocm_index_url("gfx1200")
assert url is not None
assert "gfx120X-all" in url
def test_gfx1201_maps_to_gfx120x_all(self):
url = stack_mod._windows_rocm_index_url("gfx1201")
assert url is not None
assert "gfx120X-all" in url
def test_gfx1151_maps_to_gfx1151(self):
url = stack_mod._windows_rocm_index_url("gfx1151")
assert url is not None
assert "gfx1151" in url
def test_gfx1150_maps_to_gfx1150(self):
url = stack_mod._windows_rocm_index_url("gfx1150")
assert url is not None
assert "gfx1150" in url
def test_gfx1100_maps_to_gfx110x_all(self):
url = stack_mod._windows_rocm_index_url("gfx1100")
assert url is not None
assert "gfx110X-all" in url
def test_unknown_arch_returns_none(self):
assert stack_mod._windows_rocm_index_url("gfx9999") is None
def test_none_arch_returns_none(self):
assert stack_mod._windows_rocm_index_url(None) is None
def test_url_ends_with_slash(self):
"""AMD pip index URLs must end with / for --index-url compatibility."""
url = stack_mod._windows_rocm_index_url("gfx1200")
assert url is not None
assert url.endswith("/")
def test_base_url_uses_repo_amd_com_by_default(self):
url = stack_mod._windows_rocm_index_url("gfx1200")
assert url is not None
assert "repo.amd.com" in url
def test_mirror_env_var_overrides_base(self, monkeypatch):
monkeypatch.setenv("UNSLOTH_ROCM_WINDOWS_MIRROR", "https://my-mirror.example.com/rocm/whl")
# Reload module-level constant by calling helper directly
url = stack_mod._windows_rocm_index_url("gfx1200")
# The env var is read at module load time for _ROCM_WINDOWS_INDEX_BASE,
# so just verify the helper itself doesn't error.
assert url is not None
# TEST: install_python_stack.py -- _detect_windows_gfx_arch
class TestDetectWindowsGfxArch:
"""Verify hipinfo parsing for GPU arch detection on Windows."""
def test_returns_none_when_hipinfo_not_on_path(self):
# Neutralise the venv-hipInfo and WMI-name fallbacks too, since the
# suite may run on a real AMD host where WMI would answer.
with patch("shutil.which", return_value = None):
with patch("os.path.isfile", return_value = False):
with patch("subprocess.run", side_effect = FileNotFoundError):
result = stack_mod._detect_windows_gfx_arch()
assert result is None
def test_parses_gcnarchname_from_hipinfo_output(self):
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"gcnArchName : gfx1200\nsome other line\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1200"
def test_returns_arch_on_crash_with_gcnarchname_in_output(self):
# Regression #6043: hipinfo may crash (0xC0000005 on RDNA 4) after printing
# gcnArchName. Accept the arch whenever gcnArchName is in stdout, any exit code.
mock_result = MagicMock()
mock_result.returncode = -1073741819 # 0xC0000005 STATUS_ACCESS_VIOLATION
mock_result.stdout = b"gcnArchName : gfx1200\nsome other line\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1200"
def test_mixed_igpu_dgpu_prefers_discrete(self, monkeypatch):
# #7776: HIP enumerates the Raphael iGPU (gfx1036) before the discrete RX 9060 XT
# (gfx1200), so an index-0 pick installed gfx103X-all wheels and never used the dGPU.
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising = False)
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"gcnArchName : gfx1036\ngcnArchName : gfx1200\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1200"
def test_mixed_igpu_dgpu_reports_the_override_env_var(self, monkeypatch, capsys):
# The maintainer ask on #7776: say a second GPU is there and how to pick it.
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising = False)
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"gcnArchName : gfx1036\ngcnArchName : gfx1200\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
stack_mod._detect_windows_gfx_arch()
out = capsys.readouterr().out
assert "HIP_VISIBLE_DEVICES" in out
assert "gfx1036" in out and "gfx1200" in out
def test_explicit_visible_devices_still_wins_over_igpu_preference(self, monkeypatch):
# A user who pinned the iGPU on purpose must keep getting the iGPU.
monkeypatch.setenv("HIP_VISIBLE_DEVICES", "0")
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising = False)
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"gcnArchName : gfx1036\ngcnArchName : gfx1200\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1036"
def test_strix_igpu_selection_is_unchanged(self, monkeypatch):
# gfx1151 is a supported training target, not a shadowing APU, so a Strix host
# must keep resolving to its own arch-specific wheels.
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising = False)
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"gcnArchName : gfx1151\ngcnArchName : gfx1200\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1151"
def test_single_igpu_host_is_unchanged(self, monkeypatch):
# Nothing to prefer: an APU-only box still installs for the APU.
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising = False)
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"gcnArchName : gfx1036\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1036"
def test_unsupported_discrete_does_not_depose_a_supported_igpu(self, monkeypatch):
# A supported APU next to a discrete card with no Windows wheels (gfx1010 is absent
# from _GFX_TO_AMD_INDEX_ARCH): preferring the dGPU purely for being discrete
# resolves to no index and falls back to CPU, worse than the shadowing itself.
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising = False)
assert stack_mod._windows_rocm_index_url("gfx1010") is None
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"gcnArchName : gfx1036\ngcnArchName : gfx1010\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1036"
assert stack_mod._windows_rocm_index_url(result) is not None
def test_unsupported_igpu_still_yields_to_the_discrete_card(self, monkeypatch):
# Mirror case: with neither pick wheel-backed the swap costs nothing, so the
# discrete card still wins and the guard is not over-tightened.
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising = False)
assert stack_mod._windows_rocm_index_url("gfx1013") is None
assert stack_mod._windows_rocm_index_url("gfx1010") is None
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"gcnArchName : gfx1013\ngcnArchName : gfx1010\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1010"
def test_cuda_visible_devices_also_pins_the_igpu(self, monkeypatch):
# HIP honours CUDA_VISIBLE_DEVICES with the same semantics as its own masks, so a
# host that exposed only GPU 0 that way keeps getting the iGPU.
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "0")
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"gcnArchName : gfx1036\ngcnArchName : gfx1200\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1036"
def test_cuda_visible_devices_indexes_the_enumeration(self, monkeypatch):
# hipinfo is a HIP application, so a mask filters and renumbers its output before
# we see it: CUDA_VISIBLE_DEVICES=1 leaves the dGPU alone at logical 0. Re-indexing
# with the physical value applies the mask twice, which selected the #7776 iGPU.
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "1")
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"gcnArchName : gfx1200\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1200"
def test_a_reordering_mask_is_not_applied_twice(self, monkeypatch):
# CUDA_VISIBLE_DEVICES=1,0 makes HIP expose [physical 1, physical 0], so hipinfo
# prints the dGPU first; indexing that again read token 1, the shadowing iGPU.
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "1,0")
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"gcnArchName : gfx1200\ngcnArchName : gfx1036\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1200"
def test_cuda_visible_devices_index_matches_the_other_masks(self, monkeypatch):
# All three spellings resolve identically, comma-list and out-of-range included.
for _var in ("HIP_VISIBLE_DEVICES", "ROCR_VISIBLE_DEVICES", "CUDA_VISIBLE_DEVICES"):
monkeypatch.delenv(_var, raising = False)
for _value, _expected in (
("1", 1),
("1,0", 1),
("7", 0),
("-1", 0),
("", 0),
("GPU-abc", 0),
):
monkeypatch.setenv("CUDA_VISIBLE_DEVICES", _value)
assert stack_mod._pick_visible_index(2) == _expected, _value
@pytest.mark.parametrize("spelling", ["", "-1"])
def test_all_hiding_masks_count_as_a_selection(self, monkeypatch, spelling):
# "" and "-1" select NO GPU rather than meaning "unset": the runtime stores an
# empty var as " " (clr flags.cpp) and parseRequestedDeviceList then surfaces zero
# devices. So they are a deliberate choice and must suppress the shadowing skip.
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.setenv("CUDA_VISIBLE_DEVICES", spelling)
assert stack_mod._visible_devices_pinned() is True
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"gcnArchName : gfx1036\ngcnArchName : gfx1200\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1036"
def _hipinfo_pick(self, arches):
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = "".join(f"gcnArchName : {a}\n" for a in arches).encode()
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
return stack_mod._detect_windows_gfx_arch()
def test_an_all_hiding_hip_mask_shadows_a_later_cuda_mask(self, monkeypatch):
# clr picks the HIP mask whenever its first byte is not NUL and stores an empty var
# as " ", so an empty HIP_VISIBLE_DEVICES shadows CUDA_VISIBLE_DEVICES rather than
# deferring to it; falling through would resolve a device the runtime never exposes.
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
for _spelling in ("", "-1"):
monkeypatch.setenv("HIP_VISIBLE_DEVICES", _spelling)
monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "1")
assert stack_mod._pick_visible_index(2) == 0
assert self._hipinfo_pick(["gfx1036", "gfx1200"]) == "gfx1036"
def test_gcn_arch_feature_suffix_is_stripped(self, monkeypatch):
# hipinfo can print "gfx90a:sramecc+:xnack-"; setup.ps1 splits on ':' and the
# unsplit token matched neither the wheel table nor the skip set.
for _m in ("HIP_VISIBLE_DEVICES", "ROCR_VISIBLE_DEVICES", "CUDA_VISIBLE_DEVICES"):
monkeypatch.delenv(_m, raising = False)
assert self._hipinfo_pick(["gfx1036:xnack-", "gfx1200:xnack-"]) == "gfx1200"
def test_prefers_a_wheel_backed_discrete_over_an_unsupported_one(self, monkeypatch):
# gfx1010 has no Windows wheel index, so stopping at the first non-integrated
# token dropped a host with a perfectly good gfx1200 to CPU torch.
for _m in ("HIP_VISIBLE_DEVICES", "ROCR_VISIBLE_DEVICES", "CUDA_VISIBLE_DEVICES"):
monkeypatch.delenv(_m, raising = False)
assert stack_mod._windows_rocm_index_url("gfx1010") is None
assert self._hipinfo_pick(["gfx1036", "gfx1010", "gfx1200"]) == "gfx1200"
# Same when the iGPU itself has no wheels: still prefer the supported card.
assert self._hipinfo_pick(["gfx1013", "gfx1010", "gfx1200"]) == "gfx1200"
def test_pinning_a_wheelless_gpu_says_why_torch_will_be_cpu(self, monkeypatch, capsys):
# The pin is honoured, but silently installing CPU torch while another enumerated
# GPU does have wheels hides that the mask caused it. A reordering mask puts the
# wheel-less card at logical 0 with both devices still visible to hipinfo.
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "1,0")
assert self._hipinfo_pick(["gfx1010", "gfx1036"]) == "gfx1010"
out = capsys.readouterr().out
assert "no AMD Windows wheels" in out
assert "gfx1036" in out
def test_deduplicated_callers_do_not_get_a_bogus_range_warning(self, monkeypatch, capsys):
# _detect_amd_gfx_codes() dedupes, so on a dual same-arch Linux box the list is
# shorter than the device count and a valid index reads as out of range.
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising = False)
monkeypatch.setenv("HIP_VISIBLE_DEVICES", "1")
assert stack_mod._pick_visible_index(1, warn = False) == 0
assert capsys.readouterr().out == ""
# The Windows arch-selection caller still warns.
assert stack_mod._pick_visible_index(1) == 0
assert "out of range" in capsys.readouterr().out
def test_advisory_names_the_selected_gpus_real_index(self, monkeypatch, capsys):
# Not always device 1: here it is device 2, and naming 1 would expose the
# gfx1010 the installed wheels do not target.
for _m in ("HIP_VISIBLE_DEVICES", "ROCR_VISIBLE_DEVICES", "CUDA_VISIBLE_DEVICES"):
monkeypatch.delenv(_m, raising = False)
assert self._hipinfo_pick(["gfx1036", "gfx1010", "gfx1200"]) == "gfx1200"
out = capsys.readouterr().out
assert "HIP_VISIBLE_DEVICES 2" in out
assert "HIP_VISIBLE_DEVICES 1" not in out
def test_shadowing_set_holds_only_apu_arches(self):
# gfx1037 is not an AMDGPU target in LLVM, so no Windows tool emits it.
assert "gfx1037" not in stack_mod._SHADOWING_INTEGRATED_GFX
# gfx1033 (Van Gogh) has a wheel family, so omitting it let it pose as the
# "discrete" card and depose another APU.
assert "gfx1033" in stack_mod._SHADOWING_INTEGRATED_GFX
assert self._hipinfo_pick(["gfx1036", "gfx1033"]) == "gfx1036"
def test_returns_none_on_nonzero_returncode_without_gcnarchname(self):
# Non-zero exit without gcnArchName must return None (fall through to amd-smi/WMI).
mock_result = MagicMock()
mock_result.returncode = 1
mock_result.stdout = b"HIP runtime error: no device detected\n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result is None
def test_returns_none_when_no_gcnarchname_in_output(self):
# hipinfo answers without a gcnArchName line. The WMI fallback must get
# nothing (FileNotFoundError) so the mocked name can't resolve via the table.
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b"deviceName : SomeUnknownDevice\n"
def _run(cmd, **kwargs):
if cmd and "powershell" in str(cmd[0]).lower():
raise FileNotFoundError(cmd[0])
return mock_result
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", side_effect = _run):
result = stack_mod._detect_windows_gfx_arch()
assert result is None
def test_returns_none_on_timeout(self):
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch(
"subprocess.run",
side_effect = subprocess.TimeoutExpired("hipinfo", 10),
):
result = stack_mod._detect_windows_gfx_arch()
assert result is None
def test_strips_whitespace_from_arch(self):
mock_result = MagicMock()
mock_result.returncode = 0
mock_result.stdout = b" gcnArchName : gfx1201 \n"
with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
with patch("subprocess.run", return_value = mock_result):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1201"
# TEST: install_python_stack.py -- GPU-name / WMI fallback (no amd-smi, no hipinfo)
class TestGfxArchNameFallback:
"""With no amd-smi/hipinfo on Windows, arch must resolve from the GPU name via WMI (mirrors setup.ps1)."""
@pytest.mark.parametrize(
"name, expected",
[
("AMD Radeon(TM) 8060S Graphics", "gfx1151"),
("AMD Radeon(TM) 8065S Graphics", "gfx1151"),
("AMD Ryzen AI MAX+ 395 w/ Radeon 8060S", "gfx1151"),
("AMD Radeon(TM) 890M", "gfx1150"),
("AMD Ryzen AI 9 HX 370 w/ Radeon 890M", "gfx1150"),
("AMD Radeon RX 9070 XT", "gfx1201"),
("AMD Radeon RX 9070", "gfx1201"), # Navi 48 like the XT, not Navi 44
# Navi 48 workstation card, gfx1201 per #7624 / #7307. Its name holds neither
# "9070" nor "9080", so it matched nothing and the install fell back to CPU.
("AMD Radeon AI PRO R9700", "gfx1201"),
("ATI Radeon 9700 PRO", None), # a bare "9700" arm would claim this 2002 card
("AMD Radeon RX 9060 XT", "gfx1200"), # Navi 44
("AMD Radeon RX 7700S", "gfx1102"), # (?!S) lookahead must not hit gfx1101
("AMD Radeon RX 7700 XT", "gfx1101"), # Navi 32
("AMD Radeon RX 7900 XTX", "gfx1100"), # Navi 31
("AMD Radeon(TM) 780M", "gfx1103"),
("NVIDIA GeForce RTX 4090", None),
("Microsoft Basic Display Adapter", None),
("", None),
],
)
def test_name_to_arch_mapping(self, name, expected):
assert stack_mod._gfx_arch_from_gpu_name(name) == expected
def test_wmi_fallback_resolves_arch_without_any_tools(self):
"""hipinfo absent everywhere + amd-smi absent -> WMI name fallback."""
ps_result = MagicMock()
ps_result.returncode = 0
ps_result.stdout = b"AMD Radeon(TM) 8060S Graphics\r\nMicrosoft Basic Display Adapter\r\n"
def _run(cmd, **kwargs):
if cmd and "powershell.exe" in str(cmd[0]).lower():
return ps_result
raise FileNotFoundError(cmd[0])
with patch.dict(os.environ, {}, clear = False):
for _v in (
"HIP_PATH",
"ROCM_PATH",
"UNSLOTH_ROCM_GFX_ARCH",
"UNSLOTH_ENABLE_AMD_SMI",
):
os.environ.pop(_v, None)
with patch("shutil.which", return_value = None):
with patch("os.path.isfile", return_value = False):
with patch("subprocess.run", side_effect = _run):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1151"
def test_wmi_fallback_returns_none_for_non_amd_hosts(self):
ps_result = MagicMock()
ps_result.returncode = 0
ps_result.stdout = b"NVIDIA GeForce RTX 4090\r\n"
def _run(cmd, **kwargs):
if cmd and "powershell.exe" in str(cmd[0]).lower():
return ps_result
raise FileNotFoundError(cmd[0])
with patch.dict(os.environ, {}, clear = False):
for _v in ("HIP_PATH", "ROCM_PATH", "UNSLOTH_ROCM_GFX_ARCH"):
os.environ.pop(_v, None)
with patch("shutil.which", return_value = None):
with patch("os.path.isfile", return_value = False):
with patch("subprocess.run", side_effect = _run):
result = stack_mod._detect_windows_gfx_arch()
assert result is None
@staticmethod
def _wmi_arch(names, env):
"""Drive the WMI fallback over `names`, emulating the probe's AMD filter."""
ps_result = MagicMock()
ps_result.returncode = 0
amd = [n for n in names if re.search(r"AMD|Radeon", n, re.IGNORECASE)]
ps_result.stdout = ("\r\n".join(amd) + "\r\n").encode()
def _run(cmd, **kwargs):
if cmd and "powershell.exe" in str(cmd[0]).lower():
return ps_result
raise FileNotFoundError(cmd[0])
buf = io.StringIO()
with patch.dict(os.environ, {}, clear = False):
for _v in (
"HIP_PATH",
"ROCM_PATH",
"UNSLOTH_ROCM_GFX_ARCH",
"UNSLOTH_ENABLE_AMD_SMI",
"HIP_VISIBLE_DEVICES",
"ROCR_VISIBLE_DEVICES",
"CUDA_VISIBLE_DEVICES",
):
os.environ.pop(_v, None)
os.environ.update(env)
with contextlib.redirect_stdout(buf):
with patch("shutil.which", return_value = None):
with patch("os.path.isfile", return_value = False):
with patch("subprocess.run", side_effect = _run):
result = stack_mod._detect_windows_gfx_arch()
return result, buf.getvalue()
@pytest.mark.parametrize(
"name",
[
"NVIDIA GeForce RTX 4090",
"Intel(R) Arc(TM) A770 Graphics",
"Microsoft Basic Display Adapter",
"Apple M3 Max",
],
)
def test_non_amd_adapters_produce_no_amd_output(self, name):
# AMD-only path, so a CUDA or Intel user must never see it speak. The vendor filter
# lives in a PowerShell string, so Python re-applies it: a non-AMD adapter would
# shift the mask index and trigger ROCm advice on a host with no AMD GPU.
arch, out = self._wmi_arch([name], {"CUDA_VISIBLE_DEVICES": "0"})
assert arch is None
assert out == ""
@pytest.mark.parametrize("code", ["22", "45", "43", "1"])
def test_sole_unhealthy_adapter_still_yields_its_arch(self, code):
# The filter exists so a card Windows is not exposing cannot depose a live one.
# With one adapter there is nothing to depose, and filtering it out is the only
# reason the host looks GPU-less, handing a working Radeon CPU torch. Code 45
# ("not connected") is routine on muxless laptops.
arch, _ = self._wmi_arch([f"AMD Radeon RX 9060 XT|{code}"], {})
assert arch == "gfx1200"
def test_unhealthy_adapter_still_cannot_depose_a_healthy_one(self):
# Negative control: with a healthy iGPU present the disabled dGPU stays out, so
# the shadowing skip cannot install wheels for a GPU Windows never exposes.
arch, _ = self._wmi_arch(["AMD Radeon(TM) 780M Graphics|0", "AMD Radeon RX 9060 XT|22"], {})
assert arch == "gfx1103"
def test_mask_out_of_range_is_reported_once(self):
# _dedup_pick resolves the same masks again over its own list, so without
# warn=False the WMI path prints every out-of-range warning twice.
_, out = self._wmi_arch(["AMD Radeon RX 9060 XT|0"], {"HIP_VISIBLE_DEVICES": "1"})
assert out.count("is out of range") == 1
def test_wmi_probe_lists_only_amd_adapters(self):
# The masks index AMD devices, so an Intel or NVIDIA adapter ahead of the Radeon
# would shift every index. Same -match filter setup.ps1 applies to $wmiGpus.
source = _STACK_PATH.read_text(encoding = "utf-8")
assert "$_.Name -match 'AMD|Radeon'" in source
arch, _ = self._wmi_arch(
["Intel(R) UHD Graphics", "AMD Radeon RX 9060 XT"], {"HIP_VISIBLE_DEVICES": "0"}
)
assert arch == "gfx1200"
def test_wmi_unmappable_adapter_suppresses_the_repick(self):
# A driver-only host can report a name the table does not know ("AMD Radeon
# Graphics"), which drops out of the arch list: skipping the iGPU could then land
# on the wrong card and the advisory would count arches, not GPUs. Like setup.ps1,
# repick only when every adapter mapped.
arch, out = self._wmi_arch(
["AMD Radeon Graphics", "AMD Radeon 780M Graphics", "AMD Radeon RX 9060 XT"], {}
)
assert arch == "gfx1103"
assert "setx HIP_VISIBLE_DEVICES" not in out
def test_wmi_repicks_the_dgpu_when_every_adapter_maps(self):
# Negative control: fully mapped, the #7776 preference runs and the advisory
# names the dGPU's real index.
arch, out = self._wmi_arch(["AMD Radeon 780M Graphics", "AMD Radeon RX 9060 XT"], {})
assert arch == "gfx1200"
assert "setx HIP_VISIBLE_DEVICES 1" in out
def test_wmi_masked_unmappable_adapter_warns_instead_of_substituting(self):
# Under a mask the selected adapter is the user's choice, so another card's arch
# must not stand in (same rule as setup.ps1). That leaves no arch and torch goes
# CPU-only, which has to be said out loud.
arch, out = self._wmi_arch(
["AMD Radeon Graphics", "AMD Radeon RX 9060 XT"], {"HIP_VISIBLE_DEVICES": "0"}
)
assert arch is None
assert "UNSLOTH_ROCM_GFX_ARCH" in out and "AMD Radeon Graphics" in out
def test_wmi_unpinned_unmappable_adapter_still_infers(self):
# Negative control: unpinned there is no choice to respect, so the adapter that
# did map still decides and the host keeps its wheels.
arch, _ = self._wmi_arch(["AMD Radeon Graphics", "AMD Radeon RX 9060 XT"], {})
assert arch == "gfx1200"
def test_wmi_mask_selects_the_adapter_it_names(self):
arch, _ = self._wmi_arch(
["AMD Radeon 780M Graphics", "AMD Radeon RX 9060 XT"], {"HIP_VISIBLE_DEVICES": "0"}
)
assert arch == "gfx1103"
def test_wmi_fallback_skips_disabled_adapters(self):
# WMI keeps listing a Radeon that is disabled or on a driver error, and the
# shadowing skip would let one depose the working iGPU: a laptop with a live 780M
# and a disabled RX 9060 installed gfx120X-all wheels for a GPU Windows never
# exposes. setup.ps1 filters on ConfigManagerErrorCode, so this path does too.
captured = {}
ps_result = MagicMock()
ps_result.returncode = 0
ps_result.stdout = b"AMD Radeon(TM) 780M Graphics\r\n"
def _run(cmd, **kwargs):
if cmd and "powershell.exe" in str(cmd[0]).lower():
captured["cmd"] = cmd[-1]
return ps_result
raise FileNotFoundError(cmd[0])
with patch.dict(os.environ, {}, clear = False):
for _v in (
"HIP_PATH",
"ROCM_PATH",
"UNSLOTH_ROCM_GFX_ARCH",
"UNSLOTH_ENABLE_AMD_SMI",
"HIP_VISIBLE_DEVICES",
"ROCR_VISIBLE_DEVICES",
"CUDA_VISIBLE_DEVICES",
):
os.environ.pop(_v, None)
with patch("shutil.which", return_value = None):
with patch("os.path.isfile", return_value = False):
with patch("subprocess.run", side_effect = _run):
result = stack_mod._detect_windows_gfx_arch()
assert result == "gfx1103"
assert "ConfigManagerErrorCode" in captured["cmd"]
@pytest.mark.skipif(shutil.which("pwsh") is None, reason = "pwsh not installed")
def test_wmi_adapter_filter_is_valid_powershell(self):
# The filter lives in a Python string, so nothing else checks it. Run the real
# expression over stand-ins: a null error code (absent on some drivers) and 0
# stay, anything else goes.
source = _STACK_PATH.read_text(encoding = "utf-8")
assert "ConfigManagerErrorCode" in source
script = """
class FakeAdapter { [string]$Name; [object]$ConfigManagerErrorCode }
$adapters = @(
[FakeAdapter]@{ Name = 'AMD Radeon(TM) 780M Graphics'; ConfigManagerErrorCode = 0 }
[FakeAdapter]@{ Name = 'AMD Radeon RX 9060 XT'; ConfigManagerErrorCode = 22 }
[FakeAdapter]@{ Name = 'AMD Radeon RX 7900 XTX'; ConfigManagerErrorCode = $null }
)
($adapters | Where-Object { ($null -eq $_.ConfigManagerErrorCode) -or (
$_.ConfigManagerErrorCode -eq 0) }).Name
"""
# run_pwsh, not subprocess.run: this asserts the WMI adapter filter is valid
# PowerShell, so an interpreter that died before parsing it would be recorded as
# the filter expression itself being malformed.
# See tests/_shared/unsloth_pwsh_runner.py.
out = run_pwsh(
["pwsh", "-NoProfile", "-NonInteractive", "-Command", script],
check = True,
capture_output = True,
text = True,
).stdout.split()
assert "9060" not in " ".join(out)
assert "780M" in " ".join(out) and "7900" in " ".join(out)
def test_stack_probes_venv_hipinfo(self):
"""venv Scripts hipInfo.exe (from AMD torch wheels) must be a probe candidate for driver-only hosts."""
source = _STACK_PATH.read_text(encoding = "utf-8")
assert 'os.path.join(os.path.dirname(sys.executable), "hipInfo.exe")' in source
def test_prebuilt_resolve_exe_probes_venv_dir(self):
"""_resolve_exe must include the venv Scripts candidate for driver-only standalone reruns."""
source = _PREBUILT_PATH.read_text(encoding = "utf-8")
assert "_venv_candidate" in source
def test_runtime_monitor_guards_amd_smi_absence(self):
"""amd.py must which()-check amd-smi before spawning (absence disables the poller)."""
amd_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "amd.py"
source = amd_path.read_text(encoding = "utf-8")
assert 'shutil.which("amd-smi") is None' in source
class TestSetupPs1ShadowingParity:
"""setup.ps1 resolves the arch and builds $ROCmIndexUrl itself, before the Python
stack installer ever runs, so both halves of the #7776 preference have to exist in
the PowerShell mirror too or a fresh Windows install still picks the iGPU."""
@staticmethod
def _setup_ps1() -> str:
return (PACKAGE_ROOT / "studio" / "setup.ps1").read_text(encoding = "utf-8")
def test_repick_is_wheel_aware(self):
# Mirrors _dedup_pick()'s guard: deposing a supported APU for a discrete card with
# no Windows wheels resolves to no index and drops the host to CPU.
source = self._setup_ps1()
body = source[source.index("function Resolve-ShadowingGfxPick") :]
body = body[: body.index("\n}\n")]
assert "archFamilyMap" in body, "repick must consult the wheel index map"
assert "pickedHasWheels" in body
def test_arch_family_map_precedes_the_repick(self):
# Consumed during detection now, not just at install time, so it must be
# defined before the function that reads it.
source = self._setup_ps1()
assert source.index("$archFamilyMap = @{") < source.index(
"function Resolve-ShadowingGfxPick"
)
def test_wmi_fallback_keeps_every_amd_adapter(self):
# WMI orders controllers as the driver stack enumerated them, so taking the first
# AMD match reintroduced #7776 on Adrenalin-only hosts (a 780M ahead of an RX
# 9060 XT inferred gfx1103 and installed gfx110X-all wheels).
source = self._setup_ps1()
block = source[source.index("$amdGpus = @(Get-CimInstance Win32_VideoController") :]
block = block[: block.index("$ROCmGpuLabel = $script:ROCmGpuLabels[0]")]
assert "Select-Object -First 1" not in block
# A disabled or driver-errored Radeon must not depose a working iGPU.
assert "ConfigManagerErrorCode" in block
# ... but when the filter would leave nothing, keep the parked card rather than
# reporting no AMD GPU (matches the Python WMI path).
# @() around the WHOLE if, not each branch: an if-expression unrolls a one-element array,
# and a CimInstance scalar's .Count is $null on 5.1, so a host with exactly one Radeon
# reported no AMD GPU at all (#8335).
assert "@(if ($healthyGpus.Count -gt 0) { $healthyGpus } else { $amdGpus })" in block
def test_index_picks_read_the_same_masks_as_the_pin_check(self):
# Resolve-ShadowingGfxPick honours CUDA_VISIBLE_DEVICES as a pin, so every site
# that turns a mask into an index must read it too, or the skip is suppressed
# while the index still resolves to GPU 0. One shared helper keeps them agreeing.
source = self._setup_ps1()
helper = source[source.index("function Resolve-VisibleGpuIndex") :]
helper = helper[: helper.index("\n}\n")]
for _mask in ("HIP_VISIBLE_DEVICES", "ROCR_VISIBLE_DEVICES", "CUDA_VISIBLE_DEVICES"):
assert _mask in helper
# hipinfo, amd-smi list, amd-smi static --asic, WMI name inference.
assert (
source.count("Resolve-VisibleGpuIndex $") + source.count("Resolve-VisibleGpuIndex ")
>= 4
)
# No pick site may still resolve a mask on its own.
assert "$_hipVisIdx = if (" not in source
assert "$visGpu = if (" not in source
def test_hipinfo_is_not_reindexed(self):
# hipinfo already applied the mask; only amd-smi and WMI may call the index
# resolver, or the mask lands twice.
source = self._setup_ps1()
hipinfo = source[source.index("$_hipAllArches.Count -gt 0") :]
hipinfo = hipinfo[: hipinfo.index("Resolve-ShadowingGfxPick")]
assert "Resolve-VisibleGpuIndex" not in hipinfo
assert "$_hipAllArches[0]" in hipinfo
def test_unknown_selected_wmi_adapter_is_not_substituted(self):
# Borrowing another adapter's arch is right when nothing was selected, but under
# a mask it installs for a GPU the user masked away.
source = self._setup_ps1()
block = source[source.index("$pickedName = Get-GfxArchFromGpuName") :]
block = block[: block.index("if ($pickedName) {")]
assert "Test-VisibleDevicesPinned" in block
def test_name_inference_runs_the_shadowing_pick(self):
# Every AMD adapter name gets an arch, then the same preference chooses.
source = self._setup_ps1()
assert "Get-GfxArchFromGpuName" in source
infer = source[source.index("$nameArches = @()") :]
infer = infer[: infer.index("Tip: set UNSLOTH_ROCM_GFX_ARCH")]
assert "Resolve-ShadowingGfxPick" in infer
@pytest.mark.skipif(shutil.which("pwsh") is None, reason = "pwsh not installed")
class TestSetupPs1ShadowingBehaviour:
"""Runs the real Resolve-ShadowingGfxPick / Resolve-VisibleGpuIndex. The parity
class above only greps text, so a rename fails it while a semantic bug passes.
Both helpers are sliced out by AST -- setup.ps1 is never dot-sourced."""
_HARNESS = r"""
$ErrorActionPreference = "Stop"
$errors = $null; $tokens = $null
$ast = [System.Management.Automation.Language.Parser]::ParseFile($env:SETUP_PS1, [ref]$tokens, [ref]$errors)
if ($errors -and $errors.Count -gt 0) { throw "setup.ps1 has $($errors.Count) parse errors" }
function substep { param([string]$Message, [string]$Color = "DarkGray") }
foreach ($n in @("Test-VisibleDevicesPinned", "Resolve-VisibleGpuIndex", "Resolve-ShadowingGfxPick")) {
$f = $ast.Find({ param($x) $x -is [System.Management.Automation.Language.FunctionDefinitionAst] -and $x.Name -eq $n }, $true)
if (-not $f) { throw "$n not found" }
Invoke-Expression $f.Extent.Text
}
foreach ($v in @('$script:ShadowingIntegratedGfx', '$archFamilyMap')) {
$a = $ast.FindAll({ param($x) $x -is [System.Management.Automation.Language.AssignmentStatementAst] }, $true) |
Where-Object { $_.Left.Extent.Text -eq $v } | Select-Object -First 1
if (-not $a) { throw "$v not found" }
Invoke-Expression $a.Extent.Text
}
"""
def _run(
self,
body: str,
env: "dict[str, str] | None" = None,
) -> str:
merged = {k: v for k, v in os.environ.items() if not k.endswith("_VISIBLE_DEVICES")}
merged["SETUP_PS1"] = str(PACKAGE_ROOT / "studio" / "setup.ps1")
merged.update(env or {})
# run_pwsh, not subprocess.run: every Resolve-ShadowingGfxPick case goes through
# here, and a crashed interpreter returns no stdout at all, which the callers would
# read as the shadowing pick returning the wrong gfx arch.
# See tests/_shared/unsloth_pwsh_runner.py.
out = run_pwsh(
["pwsh", "-NoProfile", "-NonInteractive", "-Command", self._HARNESS + body],
check = True,
capture_output = True,
text = True,
env = merged,
)
return out.stdout.strip()
_MIXED = 'Resolve-ShadowingGfxPick -Picked "gfx1036" -AllArches @("gfx1036","gfx1200")'
def test_unpinned_mixed_host_picks_the_discrete_gpu(self):
assert self._run(self._MIXED) == "gfx1200"
def test_explicit_pin_is_honoured(self):
assert self._run(self._MIXED, {"HIP_VISIBLE_DEVICES": "0"}) == "gfx1036"
def test_strix_is_not_deposed(self):
body = 'Resolve-ShadowingGfxPick -Picked "gfx1151" -AllArches @("gfx1151","gfx1200")'
assert self._run(body) == "gfx1151"
def test_all_hiding_spellings_count_as_a_selection(self):
# "" / "-1" surface zero devices, so they are a deliberate choice and suppress
# the repick rather than invite it.
for _val in ("", "-1"):
assert self._run(self._MIXED, {"HIP_VISIBLE_DEVICES": _val}) == "gfx1036"
@pytest.mark.parametrize("mask", ["1", "1,0", " 1 "])
def test_index_resolves_every_mask_spelling(self, mask):
# A comma list and a padded value used to resolve differently per branch, so a
# pinned host silently landed on the iGPU at index 0.
for _var in ("HIP_VISIBLE_DEVICES", "ROCR_VISIBLE_DEVICES", "CUDA_VISIBLE_DEVICES"):
assert self._run("Resolve-VisibleGpuIndex 2", {_var: mask}) == "1"
def test_out_of_range_and_unparseable_masks_fall_back_to_zero(self):
for _val in ("7", "GPU-abc123", "-1", ""):
assert self._run("Resolve-VisibleGpuIndex 2", {"HIP_VISIBLE_DEVICES": _val}) == "0"
def test_an_all_hiding_hip_mask_shadows_a_later_cuda_mask(self):
# First-set-wins, as in Python: an empty HIP mask shadows CUDA_VISIBLE_DEVICES
# in the runtime rather than deferring to it.
env = {"HIP_VISIBLE_DEVICES": "", "CUDA_VISIBLE_DEVICES": "1"}
assert self._run("Resolve-VisibleGpuIndex 2", env) == "0"
def test_wheelless_discrete_does_not_depose_a_supported_igpu(self):
body = 'Resolve-ShadowingGfxPick -Picked "gfx1036" -AllArches @("gfx1036","gfx1010")'
assert self._run(body) == "gfx1036"
def test_wheelless_igpu_still_prefers_a_wheel_backed_discrete(self):
# PowerShell mirror of test_prefers_a_wheel_backed_discrete_over_an_unsupported_one.
# gfx90c (Cezanne) has no wheel index, so the guard above is inactive and every
# discrete arch qualifies: taking the first landed on the wheel-less gfx1010,
# $ROCmIndexUrl stayed null, and the install went CPU-only despite the gfx1200.
for _igpu in ("gfx90c", "gfx1013", "gfx1153"):
body = (
f'Resolve-ShadowingGfxPick -Picked "{_igpu}" '
f'-AllArches @("{_igpu}","gfx1010","gfx1200")'
)
assert self._run(body) == "gfx1200", _igpu
# Not over-tightened: with nothing wheel-backed to move to, a wheel-less iGPU
# still yields to the discrete card, as _dedup_pick() does.
body = 'Resolve-ShadowingGfxPick -Picked "gfx90c" -AllArches @("gfx90c","gfx1010")'
assert self._run(body) == "gfx1010"
def test_advisory_names_the_selected_gpus_real_index(self):
body = (
'$null = Resolve-ShadowingGfxPick -Picked "gfx1036" '
'-AllArches @("gfx1036","gfx1010","gfx1200"); '
'($script:Msgs -join " ")'
)
harness = self._HARNESS.replace(
'function substep { param([string]$Message, [string]$Color = "DarkGray") }',
'function substep { param([string]$Message, [string]$Color = "DarkGray") '
"$script:Msgs += $Message }\n$script:Msgs = @()",
)
merged = {k: v for k, v in os.environ.items() if not k.endswith("_VISIBLE_DEVICES")}
merged["SETUP_PS1"] = str(PACKAGE_ROOT / "studio" / "setup.ps1")
# run_pwsh, not subprocess.run: the advisory text is collected from this run's
# stdout, so a pwsh that aborted at startup would look like substep never naming
# the real HIP_VISIBLE_DEVICES index. See tests/_shared/unsloth_pwsh_runner.py.
out = run_pwsh(
["pwsh", "-NoProfile", "-NonInteractive", "-Command", harness + body],
check = True,
capture_output = True,
text = True,
env = merged,
).stdout
assert "HIP_VISIBLE_DEVICES 2" in out
assert "HIP_VISIBLE_DEVICES 1" not in out
def test_degenerate_inputs_do_not_throw(self):
for _all in ("$null", "@()", '@("gfx1036")'):
body = f'Resolve-ShadowingGfxPick -Picked "gfx1036" -AllArches {_all}'
assert self._run(body) == "gfx1036"
# TEST: install_python_stack.py -- _install_bnb_windows_rocm
class TestInstallBnbWindowsRocm:
"""Verify AMD Windows BNB wheel install helper."""
@pytest.fixture(autouse = True)
def _isolate_sitecustomize_persistence(self, monkeypatch, request):
"""Keep helper tests from writing to the active interpreter site-packages."""
if request.node.name.startswith("test_persist"):
return
monkeypatch.setattr(
stack_mod,
"_persist_bnb_rocm_version",
lambda version: True,
)
def test_calls_pip_install_try_with_win_amd64_url(self):
"""Should call pip_install_try with the win_amd64 wheel URL via plain pip."""
with patch.object(stack_mod, "pip_install_try", return_value = True) as mock_pip:
stack_mod._install_bnb_windows_rocm()
assert mock_pip.call_count == 1
call_args = str(mock_pip.call_args_list[0])
assert "bitsandbytes" in call_args
assert "win_amd64" in call_args
# Force plain pip (uv mangles the bitsandbytes wheel) -- see
# https://unsloth.ai/docs/get-started/install/amd/amd-hackathon
assert mock_pip.call_args.kwargs.get("force_pip") is True
def test_forces_plain_pip_not_uv(self):
"""The bnb wheel must be installed with plain pip, never uv."""
with patch.object(stack_mod, "pip_install_try", return_value = True) as mock_pip:
stack_mod._install_bnb_windows_rocm()
assert mock_pip.call_args.kwargs.get("force_pip") is True
def test_does_not_touch_uv_skip_env_var(self):
"""The UV_SKIP_WHEEL_FILENAME_CHECK hack is gone; the env must be untouched."""
observed = {}
def _capture(*args, **kwargs):
observed["during"] = os.environ.get("UV_SKIP_WHEEL_FILENAME_CHECK")
return True
with patch.dict(os.environ, {}, clear = False):
os.environ.pop("UV_SKIP_WHEEL_FILENAME_CHECK", None)
with patch.object(stack_mod, "pip_install_try", side_effect = _capture):
stack_mod._install_bnb_windows_rocm()
assert observed.get("during") is None
assert "UV_SKIP_WHEEL_FILENAME_CHECK" not in os.environ
def test_returns_false_on_pip_failure(self):
"""A failed pip_install_try must surface as a False return, not BNB_ROCM_VERSION."""
with patch.dict(os.environ, {}, clear = False):
os.environ.pop("BNB_ROCM_VERSION", None)
with patch.object(stack_mod, "pip_install_try", return_value = False):
result = stack_mod._install_bnb_windows_rocm()
assert result is False
assert "BNB_ROCM_VERSION" not in os.environ
def test_falls_back_to_pypi_when_win_amd64_url_missing(self):
"""No win_amd64 pre-release wheel must not mean no bitsandbytes: PyPI
>=0.50.0 ships libbitsandbytes_rocm{714,72}.dll, so it is a real ROCm build."""
with patch.object(stack_mod, "_BNB_ROCM_PRERELEASE_URLS", {}):
with patch.object(stack_mod, "pip_install_try", return_value = True) as mock_pip:
stack_mod._install_bnb_windows_rocm()
assert mock_pip.call_count == 1
assert stack_mod._BNB_ROCM_PYPI_FALLBACK in mock_pip.call_args.args
def test_falls_back_to_pypi_when_prerelease_install_fails(self):
"""A blocked GitHub pre-release URL must fall through to the PyPI floor rather
than leaving Windows ROCm with no working bitsandbytes."""
with patch.object(stack_mod, "pip_install_try", side_effect = [False, True]) as mock_pip:
with patch.object(stack_mod, "_detect_bnb_rocm_dll_ver", return_value = "72"):
result = stack_mod._install_bnb_windows_rocm()
assert result is True
assert mock_pip.call_count == 2
assert "win_amd64" in str(mock_pip.call_args_list[0])
assert stack_mod._BNB_ROCM_PYPI_FALLBACK in mock_pip.call_args_list[1].args
def test_returns_false_only_when_both_paths_fail(self):
"""Both the pre-release wheel and the PyPI fallback must fail before the
helper reports failure."""
with patch.dict(os.environ, {}, clear = False):
os.environ.pop("BNB_ROCM_VERSION", None)
with patch.object(stack_mod, "pip_install_try", return_value = False) as mock_pip:
result = stack_mod._install_bnb_windows_rocm()
assert result is False
assert mock_pip.call_count == 2
def test_sets_bnb_rocm_version_from_detected_dll(self):
"""BNB_ROCM_VERSION is set from the DLL detected after install."""
with patch.dict(os.environ, {}, clear = False):
os.environ.pop("BNB_ROCM_VERSION", None)
os.environ.pop(stack_mod._BNB_ROCM_VERSION_SOURCE_ENV, None)
with patch.object(stack_mod, "pip_install_try", return_value = True):
with patch.object(stack_mod, "_detect_bnb_rocm_dll_ver", return_value = "72"):
stack_mod._install_bnb_windows_rocm()
assert os.environ.get("BNB_ROCM_VERSION") == "72"
def test_sets_bnb_rocm_version_from_newer_dll(self):
"""If AMD ships a newer DLL (e.g. rocm713.dll), that version is used."""
with patch.dict(os.environ, {}, clear = False):
os.environ.pop("BNB_ROCM_VERSION", None)
os.environ.pop(stack_mod._BNB_ROCM_VERSION_SOURCE_ENV, None)
with patch.object(stack_mod, "pip_install_try", return_value = True):
with patch.object(stack_mod, "_detect_bnb_rocm_dll_ver", return_value = "713"):
stack_mod._install_bnb_windows_rocm()
assert os.environ.get("BNB_ROCM_VERSION") == "713"
def test_falls_back_to_72_when_detection_fails(self):
"""Falls back to '72' when DLL detection returns None."""
with patch.dict(os.environ, {}, clear = False):
os.environ.pop("BNB_ROCM_VERSION", None)
os.environ.pop(stack_mod._BNB_ROCM_VERSION_SOURCE_ENV, None)
with patch.object(stack_mod, "pip_install_try", return_value = True):
with patch.object(stack_mod, "_detect_bnb_rocm_dll_ver", return_value = None):
stack_mod._install_bnb_windows_rocm()
assert os.environ.get("BNB_ROCM_VERSION") == "72"
def test_does_not_override_existing_bnb_rocm_version(self):
"""An explicit BNB_ROCM_VERSION in the caller's env must not be clobbered."""
with patch.dict(os.environ, {"BNB_ROCM_VERSION": "60"}):
os.environ.pop(stack_mod._BNB_ROCM_VERSION_SOURCE_ENV, None)
with patch.object(stack_mod, "pip_install_try", return_value = True):
stack_mod._install_bnb_windows_rocm()
assert os.environ.get("BNB_ROCM_VERSION") == "60"
def test_does_not_persist_existing_bnb_rocm_version(self):
"""A caller override must not become the venv's managed default."""
with patch.dict(os.environ, {"BNB_ROCM_VERSION": "60"}):
os.environ.pop(stack_mod._BNB_ROCM_VERSION_SOURCE_ENV, None)
with patch.object(stack_mod, "pip_install_try", return_value = True):
with patch.object(stack_mod, "_detect_bnb_rocm_dll_ver") as mock_detect:
with patch.object(
stack_mod, "_persist_bnb_rocm_version", return_value = True
) as mock_persist:
stack_mod._install_bnb_windows_rocm()
assert os.environ.get("BNB_ROCM_VERSION") == "60"
mock_detect.assert_not_called()
mock_persist.assert_not_called()
def test_redetects_when_bnb_rocm_version_came_from_sitecustomize(self):
"""Persisted defaults should not mask a newer DLL suffix after reinstall."""
with patch.dict(
os.environ,
{
"BNB_ROCM_VERSION": "72",
stack_mod._BNB_ROCM_VERSION_SOURCE_ENV: (
stack_mod._BNB_ROCM_VERSION_SOURCE_SITECUSTOMIZE
),
},
):
with patch.object(stack_mod, "pip_install_try", return_value = True):
with patch.object(stack_mod, "_detect_bnb_rocm_dll_ver", return_value = "713"):
with patch.object(
stack_mod, "_persist_bnb_rocm_version", return_value = True
) as mock_persist:
stack_mod._install_bnb_windows_rocm()
assert os.environ.get("BNB_ROCM_VERSION") == "713"
assert (
os.environ.get(stack_mod._BNB_ROCM_VERSION_SOURCE_ENV)
== stack_mod._BNB_ROCM_VERSION_SOURCE_DETECTED
)
mock_persist.assert_called_once_with("713")
def test_persists_bnb_rocm_version_for_direct_venv_python(self, tmp_path):
"""BNB_ROCM_VERSION must apply to a fresh Python process in the venv."""
site_packages = tmp_path / "site-packages"
with patch.dict(os.environ, {}, clear = False):
os.environ.pop("BNB_ROCM_VERSION", None)
os.environ.pop(stack_mod._BNB_ROCM_VERSION_SOURCE_ENV, None)
with patch.object(stack_mod, "pip_install_try", return_value = True):
with patch.object(stack_mod, "_detect_bnb_rocm_dll_ver", return_value = "72"):
with patch.object(
stack_mod.sysconfig, "get_path", return_value = str(site_packages)
):
stack_mod._install_bnb_windows_rocm()
sitecustomize = site_packages / "sitecustomize.py"
source = sitecustomize.read_text(encoding = "utf-8")
assert "BNB_ROCM_VERSION" in source
assert stack_mod._BNB_ROCM_VERSION_SOURCE_ENV in source
assert "'72'" in source
probe_env = os.environ.copy()
probe_env.pop("BNB_ROCM_VERSION", None)
probe_env.pop(stack_mod._BNB_ROCM_VERSION_SOURCE_ENV, None)
probe_env["PYTHONPATH"] = str(site_packages)
result = subprocess.run(
[
sys.executable,
"-c",
(
"import os; "
"print(os.environ.get('BNB_ROCM_VERSION', ''), "
"os.environ.get('UNSLOTH_BNB_ROCM_VERSION_SOURCE', ''))"
),
],
env = probe_env,
stdout = subprocess.PIPE,
stderr = subprocess.PIPE,
text = True,
check = True,
)
assert result.stdout.strip() == "72 sitecustomize"
def test_persist_bnb_rocm_version_replaces_existing_managed_block(self, tmp_path):
"""Updating sitecustomize.py must not duplicate the managed BNB block."""
site_packages = tmp_path / "site-packages"
site_packages.mkdir()
sitecustomize = site_packages / "sitecustomize.py"
sitecustomize.write_text(
"EXISTING = True\n"
"# BEGIN Unsloth BNB_ROCM_VERSION\n"
"import os as _unsloth_os\n"
"_unsloth_os.environ.setdefault('BNB_ROCM_VERSION', '72')\n"
"# END Unsloth BNB_ROCM_VERSION\n",
encoding = "utf-8",
)
with patch.object(stack_mod.sysconfig, "get_path", return_value = str(site_packages)):
assert stack_mod._persist_bnb_rocm_version("713") is True
source = sitecustomize.read_text(encoding = "utf-8")
assert source.count("# BEGIN Unsloth BNB_ROCM_VERSION") == 1
assert "EXISTING = True" in source
assert "'713'" in source
assert "'72'" not in source
def test_persist_bnb_rocm_version_handles_non_utf8_sitecustomize(self, tmp_path):
"""A legacy non-UTF-8 sitecustomize.py should not abort installation."""
site_packages = tmp_path / "site-packages"
site_packages.mkdir()
sitecustomize = site_packages / "sitecustomize.py"
sitecustomize.write_bytes(b"\xff\xfe\x00")
with patch.object(stack_mod.sysconfig, "get_path", return_value = str(site_packages)):
assert stack_mod._persist_bnb_rocm_version("72") is False
def test_persist_bnb_rocm_version_repairs_truncated_block(self, tmp_path):
"""A managed block missing its END marker is replaced, not duplicated."""
site_packages = tmp_path / "site-packages"
site_packages.mkdir()
sitecustomize = site_packages / "sitecustomize.py"
sitecustomize.write_text(
"EXISTING = True\n"
"# BEGIN Unsloth BNB_ROCM_VERSION\n"
"import os as _unsloth_os\n"
"_unsloth_os.environ.setdefault('BNB_ROCM_VERSION', '72')\n",
encoding = "utf-8",
)
with patch.object(stack_mod.sysconfig, "get_path", return_value = str(site_packages)):
assert stack_mod._persist_bnb_rocm_version("713") is True
source = sitecustomize.read_text(encoding = "utf-8")
assert source.count("# BEGIN Unsloth BNB_ROCM_VERSION") == 1
assert source.count("# END Unsloth BNB_ROCM_VERSION") == 1
assert "EXISTING = True" in source
assert "'713'" in source
assert "'72'" not in source
def test_persist_bnb_rocm_version_dedupes_duplicate_blocks(self, tmp_path):
"""Multiple managed blocks collapse to one while preserving user content."""
site_packages = tmp_path / "site-packages"
site_packages.mkdir()
sitecustomize = site_packages / "sitecustomize.py"
block = (
"# BEGIN Unsloth BNB_ROCM_VERSION\n"
"import os as _unsloth_os\n"
"_unsloth_os.environ.setdefault('BNB_ROCM_VERSION', '72')\n"
"# END Unsloth BNB_ROCM_VERSION\n"
)
sitecustomize.write_text(block + "USER_MID = 1\n" + block, encoding = "utf-8")
with patch.object(stack_mod.sysconfig, "get_path", return_value = str(site_packages)):
assert stack_mod._persist_bnb_rocm_version("713") is True
source = sitecustomize.read_text(encoding = "utf-8")
assert source.count("# BEGIN Unsloth BNB_ROCM_VERSION") == 1
assert source.count("# END Unsloth BNB_ROCM_VERSION") == 1
assert "USER_MID = 1" in source
assert "'713'" in source
assert "'72'" not in source
def test_persist_bnb_rocm_version_atomic_no_leftover_tmp(self, tmp_path):
"""The write-then-rename path must not leave its temp file behind."""
site_packages = tmp_path / "site-packages"
site_packages.mkdir()
with patch.object(stack_mod.sysconfig, "get_path", return_value = str(site_packages)):
assert stack_mod._persist_bnb_rocm_version("72") is True
leftovers = [p.name for p in site_packages.iterdir() if "unsloth-tmp" in p.name]
assert leftovers == []
assert (site_packages / "sitecustomize.py").exists()
class TestRuntimeBnbRocmSourceGuards:
"""Runtime entrypoints redetect managed defaults but keep caller overrides."""
_MAIN_PATH = PACKAGE_ROOT / "studio" / "backend" / "main.py"
_TRAINING_WORKER_PATH = PACKAGE_ROOT / "studio" / "backend" / "core" / "training" / "worker.py"
def test_main_gate_redetects_persisted_default(self):
source = self._MAIN_PATH.read_text(encoding = "utf-8")
assert 'os.environ.get("UNSLOTH_BNB_ROCM_VERSION_SOURCE") == "sitecustomize"' in source
assert 'os.environ["UNSLOTH_BNB_ROCM_VERSION_SOURCE"] = "detected"' in source
def test_worker_gate_redetects_persisted_default(self):
source = self._TRAINING_WORKER_PATH.read_text(encoding = "utf-8")
assert 'os.environ.get("UNSLOTH_BNB_ROCM_VERSION_SOURCE") == "sitecustomize"' in source
assert 'os.environ["UNSLOTH_BNB_ROCM_VERSION_SOURCE"] = "detected"' in source
def test_fallback_prefers_seeded_value_over_hardcoded_72(self):
"""A failed redetect must not downgrade a persisted suffix to '72'."""
for path in (self._MAIN_PATH, self._TRAINING_WORKER_PATH):
source = path.read_text(encoding = "utf-8")
assert (
'_bnb_rocm_ver or os.environ.get("BNB_ROCM_VERSION") or "72"' in source
), path.name
def test_main_requires_found_rocm_dll(self):
"""HIP_PATH/ROCM_PATH alone (HIP SDK on a CUDA/CPU box) must not force
a ROCm backend onto a non-ROCm bitsandbytes."""
source = self._MAIN_PATH.read_text(encoding = "utf-8")
assert "if _found_rocm_bnb:" in source
assert "_hip_env" not in source
def test_worker_requires_found_rocm_dll(self):
"""No DLL found: the worker must not write any override or touch the
seeded marker (later import fixes must still see sitecustomize)."""
source = self._TRAINING_WORKER_PATH.read_text(encoding = "utf-8")
assert "if _found_rocm_bnb:" in source
class TestDetectBnbRocmDllVer:
"""Unit tests for _detect_bnb_rocm_dll_ver()."""
def test_returns_none_when_bnb_not_installed(self):
"""Returns None if bitsandbytes is not importable."""
import importlib.util
with patch.object(importlib.util, "find_spec", return_value = None):
assert stack_mod._detect_bnb_rocm_dll_ver() is None
def test_detects_rocm72_dll(self, tmp_path):
"""Returns '72' when libbitsandbytes_rocm72.dll is present."""
(tmp_path / "libbitsandbytes_rocm72.dll").write_text("")
mock_spec = MagicMock()
mock_spec.submodule_search_locations = [str(tmp_path)]
import importlib.util
with patch.object(importlib.util, "find_spec", return_value = mock_spec):
assert stack_mod._detect_bnb_rocm_dll_ver() == "72"
def test_detects_rocm713_dll(self, tmp_path):
"""Returns '713' when libbitsandbytes_rocm713.dll is present."""
(tmp_path / "libbitsandbytes_rocm713.dll").write_text("")
mock_spec = MagicMock()
mock_spec.submodule_search_locations = [str(tmp_path)]
import importlib.util
with patch.object(importlib.util, "find_spec", return_value = mock_spec):
assert stack_mod._detect_bnb_rocm_dll_ver() == "713"
def test_returns_none_when_only_cuda_dlls(self, tmp_path):
"""Returns None when only CUDA DLLs are present (no ROCm DLL)."""
(tmp_path / "libbitsandbytes_cuda121.dll").write_text("")
mock_spec = MagicMock()
mock_spec.submodule_search_locations = [str(tmp_path)]
import importlib.util
with patch.object(importlib.util, "find_spec", return_value = mock_spec):
assert stack_mod._detect_bnb_rocm_dll_ver() is None
def test_picks_highest_suffix_when_multiple_dlls(self, tmp_path):
"""Returns the highest numeric suffix across ROCm DLL variants (glob order is not guaranteed)."""
(tmp_path / "libbitsandbytes_rocm72.dll").write_text("")
(tmp_path / "libbitsandbytes_rocm713.dll").write_text("")
mock_spec = MagicMock()
mock_spec.submodule_search_locations = [str(tmp_path)]
import importlib.util
with patch.object(importlib.util, "find_spec", return_value = mock_spec):
assert stack_mod._detect_bnb_rocm_dll_ver() == "713"
# TEST: install_python_stack.py -- UNSLOTH_ROCM_TORCH_INSTALLED early-return path
class TestRocmTorchInstalledEnvVar:
"""Verify UNSLOTH_ROCM_TORCH_INSTALLED=1 skips main install but still installs BNB."""
@staticmethod
def _ok_torch_probe(*a, **kw):
# The shared probe exits 0 when torch imports and reports the build on stdout;
# a non-empty HIP version is what marks it as ROCm.
rv = MagicMock()
rv.returncode = 0
rv.stdout = _MARK + "2.10.0+rocm7.1|7.1|\n"
return rv
@patch.object(stack_mod, "_install_bnb_windows_rocm")
@patch.object(stack_mod, "pip_install")
def test_env_var_skips_main_pip_install(self, mock_pip, mock_bnb):
"""UNSLOTH_ROCM_TORCH_INSTALLED=1 should not trigger torch pip_install."""
with (
patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}),
patch.object(stack_mod.subprocess, "run", side_effect = self._ok_torch_probe),
):
stack_mod._ensure_rocm_torch()
mock_pip.assert_not_called()
@patch.object(stack_mod, "_install_bnb_windows_rocm")
@patch.object(stack_mod, "pip_install")
def test_env_var_calls_bnb_install(self, mock_pip, mock_bnb):
"""UNSLOTH_ROCM_TORCH_INSTALLED=1 should still call _install_bnb_windows_rocm."""
with (
patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}),
patch.object(stack_mod.subprocess, "run", side_effect = self._ok_torch_probe),
):
stack_mod._ensure_rocm_torch()
mock_bnb.assert_called_once()
@patch.object(stack_mod, "_install_bnb_windows_rocm")
@patch.object(stack_mod, "pip_install")
def test_env_var_sets_rocm_windows_flag(self, mock_pip, mock_bnb):
"""UNSLOTH_ROCM_TORCH_INSTALLED=1 should set _rocm_windows_torch_installed."""
stack_mod._rocm_windows_torch_installed = False
with (
patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}),
patch.object(stack_mod.subprocess, "run", side_effect = self._ok_torch_probe),
):
stack_mod._ensure_rocm_torch()
assert stack_mod._rocm_windows_torch_installed is True
@patch.object(stack_mod, "_install_bnb_windows_rocm")
@patch.object(stack_mod, "pip_install")
def test_env_var_falls_through_when_torch_missing(self, mock_pip, mock_bnb):
"""If the venv was wiped between runs, the stale env-var must not suppress reinstall."""
stack_mod._rocm_windows_torch_installed = False
def _bad_probe(*a, **kw):
rv = MagicMock()
rv.returncode = 1
return rv
with (
patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}),
patch.object(stack_mod.subprocess, "run", side_effect = _bad_probe),
patch.object(stack_mod, "IS_WINDOWS", False),
patch.object(stack_mod, "IS_MACOS", True),
):
stack_mod._ensure_rocm_torch()
# macOS branch is the next exit; the point is the early-return did NOT fire.
mock_bnb.assert_not_called()
class TestWindowsRocmTorchaoGuard:
"""Verify the torchao skip can detect an installed Windows ROCm torch build."""
def test_installed_torch_is_windows_rocm_accepts_rocm_probe(self):
rv = MagicMock()
rv.returncode = 0
rv.stdout = _MARK + "2.10.0+rocm7.1|7.1|"
with (
patch.object(stack_mod, "IS_WINDOWS", True),
patch.object(stack_mod.subprocess, "run", return_value = rv),
):
assert stack_mod._installed_torch_is_windows_rocm() is True
def test_installed_torch_is_windows_rocm_rejects_non_rocm_probe(self):
rv = MagicMock()
rv.returncode = 0
rv.stdout = ""
with (
patch.object(stack_mod, "IS_WINDOWS", True),
patch.object(stack_mod.subprocess, "run", return_value = rv),
):
assert stack_mod._installed_torch_is_windows_rocm() is False
def test_installed_torch_is_windows_rocm_is_non_windows_noop(self):
with patch.object(stack_mod, "IS_WINDOWS", False):
assert stack_mod._installed_torch_is_windows_rocm() is False
@patch.object(stack_mod, "_repair_bad_anyio")
@patch.object(stack_mod, "_ensure_rocm_torch")
@patch.object(stack_mod, "_ensure_cuda_torch")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = True)
@patch.object(stack_mod, "run")
@patch.object(stack_mod, "pip_install")
def test_install_python_stack_skips_torchao_when_windows_rocm_torch_is_installed(
self, mock_pip, mock_run, mock_has_nvidia, mock_cuda, mock_rocm, mock_anyio, tmp_path
):
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 = ""
with (
patch.dict(os.environ, {"SKIP_STUDIO_BASE": "1"}),
patch.object(stack_mod, "IS_WINDOWS", True),
patch.object(stack_mod, "IS_MACOS", False),
patch.object(stack_mod, "IS_MAC_ARM", False),
patch.object(stack_mod, "NO_TORCH", False),
patch.object(stack_mod, "_rocm_windows_torch_installed", False),
patch.object(stack_mod, "_bootstrap_uv", return_value = False),
patch.object(stack_mod, "_installed_torch_is_windows_rocm", return_value = True),
patch.object(stack_mod, "LOCAL_DD_UNSTRUCTURED_PLUGIN", unstructured_plugin),
patch.object(stack_mod, "LOCAL_DD_GITHUB_PLUGIN", github_plugin),
patch.object(stack_mod.subprocess, "run", return_value = subprocess_result),
):
assert stack_mod.install_python_stack() == 0
installed_specs = [str(arg) for call in mock_pip.call_args_list for arg in call.args]
assert not any("torchao" in arg for arg in installed_specs)
class TestProgressStepCountMatchesTotal:
"""The progress bar must reach exactly _TOTAL: every _progress() step is counted in
base_total. Regression for a repair step added without incrementing base_total,
which pushed _STEP past _TOTAL (Codex P2)."""
def _run_stack(
self,
tmp_path,
*,
is_windows,
is_macos,
is_mac_arm,
skip_base = True,
):
unstructured_plugin = tmp_path / "unstructured"
github_plugin = tmp_path / "github"
unstructured_plugin.mkdir()
github_plugin.mkdir()
sub = MagicMock()
sub.returncode = 0
sub.stdout = ""
with (
patch.dict(os.environ, {"SKIP_STUDIO_BASE": "1" if skip_base else "0"}),
patch.object(stack_mod, "_report_mlx_stack_health"),
patch.object(stack_mod, "_bitsandbytes_installed", return_value = False),
patch.object(stack_mod, "IS_WINDOWS", is_windows),
patch.object(stack_mod, "IS_MACOS", is_macos),
patch.object(stack_mod, "IS_MAC_ARM", is_mac_arm),
patch.object(stack_mod, "NO_TORCH", False),
patch.object(stack_mod, "_rocm_windows_torch_installed", False),
patch.object(stack_mod, "_bootstrap_uv", return_value = False),
patch.object(stack_mod, "_installed_torch_is_windows_rocm", return_value = False),
patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = True),
patch.object(stack_mod, "_repair_bad_anyio"),
patch.object(stack_mod, "_ensure_cuda_torch"),
patch.object(stack_mod, "_ensure_rocm_torch"),
patch.object(stack_mod, "_ensure_cpu_torch"),
patch.object(stack_mod, "LOCAL_DD_UNSTRUCTURED_PLUGIN", unstructured_plugin),
patch.object(stack_mod, "LOCAL_DD_GITHUB_PLUGIN", github_plugin),
patch.object(stack_mod.subprocess, "run", return_value = sub),
):
assert stack_mod.install_python_stack() == 0
return stack_mod._STEP, stack_mod._TOTAL
def test_windows_progress_reaches_total(self, tmp_path):
step, total = self._run_stack(tmp_path, is_windows = True, is_macos = False, is_mac_arm = False)
assert step == total, f"Windows progress {step} != total {total} (final step uncounted)"
def test_linux_progress_reaches_total(self, tmp_path):
step, total = self._run_stack(tmp_path, is_windows = False, is_macos = False, is_mac_arm = False)
assert step == total, f"Linux progress {step} != total {total}"
@pytest.mark.parametrize("skip_base", [True, False])
@pytest.mark.parametrize(
"is_windows, is_macos, is_mac_arm",
[(False, False, False), (True, False, False), (False, True, True), (False, True, False)],
ids = ["linux", "windows", "macos_arm", "macos_intel"],
)
def test_progress_reaches_total_on_both_core_paths(
self, tmp_path, is_windows, is_macos, is_mac_arm, skip_base
):
"""Both the installer handoff and `studio update`, on every platform.
The update path on Apple Silicon runs the MLX step, which the earlier
SKIP_STUDIO_BASE=1-only cases never reached, so its slot went uncounted.
"""
step, total = self._run_stack(
tmp_path,
is_windows = is_windows,
is_macos = is_macos,
is_mac_arm = is_mac_arm,
skip_base = skip_base,
)
assert step == total, f"progress {step} != total {total}"
# TEST: worker.py -- Windows ROCm patches (source-level checks)
class TestWorkerWindowsRocmPatches:
"""Verify worker.py contains the required Windows ROCm runtime patches."""
def test_grouped_mm_dispatch_patch_present(self):
"""worker.py must register a _grouped_mm CUDA dispatch override."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert '_gm_lib.impl("_grouped_mm"' in source
def test_grouped_mm_patch_targets_cuda_dispatch_key(self):
"""The dispatch override must target the CUDA key (not CompositeImplicitAutograd)."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert '"_grouped_mm", _grouped_mm_safe_impl, "CUDA"' in source
def test_grouped_mm_lib_kept_alive(self):
"""The Library object must be stored to prevent GC clearing the registration."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "_WINDOWS_ROCM_GROUPED_MM_LIB" in source
def test_grouped_mm_handles_offs_grouped_case(self):
"""_grouped_mm fallback must handle the grouped (offs!=None) variant."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "offs_list" in source
assert "offs.tolist()" in source
def test_worker_calls_shared_torchao_stub(self):
"""worker.py must invoke the shared torchao stub entrypoint."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "install_torchao_windows_rocm_stub()" in source
def test_export_worker_calls_shared_torchao_stub(self):
"""export/worker.py must invoke the same shared torchao stub entrypoint."""
source = _EXPORT_WORKER_PATH.read_text(encoding = "utf-8")
assert "install_torchao_windows_rocm_stub()" in source
def test_embedder_calls_shared_torchao_stub(self):
"""embeddings.py must install the stub before importing sentence-transformers:
it runs in the main process (not a stubbed worker), so otherwise transformers
-> torchao crashes on Windows ROCm and the embedder drops to llama-server."""
source = _EMBEDDINGS_PATH.read_text(encoding = "utf-8")
assert "install_torchao_windows_rocm_stub()" in source
def test_torchao_stub_uses_stub_type_meta(self):
"""Torchao stub must use _StubTypeMeta so isinstance() returns False not TypeError."""
source = _TORCHAO_STUB_PATH.read_text(encoding = "utf-8")
assert "_StubTypeMeta" in source
def test_stub_type_meta_has_instancecheck(self):
"""_StubTypeMeta must define __instancecheck__ returning False."""
source = _TORCHAO_STUB_PATH.read_text(encoding = "utf-8")
assert "__instancecheck__" in source
def test_stub_subpackage_finder_registered(self):
"""_StubSubpackageFinder must be appended to sys.meta_path."""
source = _TORCHAO_STUB_PATH.read_text(encoding = "utf-8")
assert "sys.meta_path.append(_StubSubpackageFinder())" in source
def test_torchao_key_submodules_pre_stubbed(self):
"""Key torchao submodules (dtypes, quantization) must be pre-stubbed."""
source = _TORCHAO_STUB_PATH.read_text(encoding = "utf-8")
assert "torchao.dtypes" in source
assert "torchao.quantization" in source
def test_torchdynamo_disabled_on_windows_rocm(self):
"""worker.py should disable dynamo on Windows ROCm as belt-and-suspenders."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "TORCHDYNAMO_DISABLE" in source
def test_bnb_rocm_version_set_on_windows_rocm(self):
"""worker.py must set BNB_ROCM_VERSION from the detected DLL suffix (BNB's auto-detect can mismatch)."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "BNB_ROCM_VERSION" in source
assert "_detect_bnb_rocm_dll_ver" in source or "libbitsandbytes_rocm" in source
# Falls back to the seeded value, never a blind "72".
assert '_bnb_rocm_ver or os.environ.get("BNB_ROCM_VERSION")' in source
def test_bnb_rocm_version_set_before_ml_imports(self):
"""BNB_ROCM_VERSION must appear in section 1f, before section 2 ML imports."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
idx_bnb = source.find("BNB_ROCM_VERSION")
# Use the entry-point section-2 marker (not the trainer helper's own "# ── 2.").
idx_sec2 = source.find("# ── 2. Now import ML libraries")
assert idx_bnb != -1, "BNB_ROCM_VERSION not found in worker.py"
assert idx_sec2 != -1, "'# ── 2. Now import ML libraries' marker not found in worker.py"
assert idx_bnb < idx_sec2, (
"BNB_ROCM_VERSION must be set before section 2 ML imports "
f"(found at {idx_bnb}, section 2 at {idx_sec2})"
)
def test_grouped_mm_patch_guarded_by_windows_and_hip_check(self):
"""_grouped_mm patch must only apply on Windows + HIP torch."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert 'sys.platform == "win32"' in source
# Gates on HIP version via a getattr chain ("version", "hip").
assert '"version"' in source and '"hip"' in source
def test_hip_ver_at_least_helper_defined(self):
"""_hip_ver_at_least helper must be defined inside the Windows ROCm block."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "def _hip_ver_at_least(major: int, minor: int)" in source
def test_grouped_mm_patch_gated_on_hip_lt_713(self):
"""_grouped_mm patch must be skipped on HIP >= 7.13 (AMD fixed the bug in ROCm 7.13)."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "_hip_ver_at_least(7, 13)" in source
# Patch must be inside the negated `if not` guard.
assert "if not _hip_ver_at_least(7, 13):" in source
def test_grouped_mm_hip_713_skip_message_present(self):
"""worker.py must log a message when skipping the patch on HIP >= 7.13."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "HIP >= 7.13" in source
assert "7.13" in source
def test_grouped_mm_patch_else_branch_present(self):
"""An else branch must follow the _hip_ver_at_least gate (skip path for 7.13+)."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
gate_idx = source.find("if not _hip_ver_at_least(7, 13):")
assert gate_idx != -1, "Version gate not found in worker.py"
else_idx = source.find("else:", gate_idx)
assert else_idx != -1, "else: branch after _hip_ver_at_least gate not found"
def test_hip_ver_at_least_handles_amd_version_format(self):
"""_hip_ver_at_least must split on '.' and compare only major.minor (handles '7.13.99004')."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert 'split(".")[:2]' in source or ".split('.')[:2]" in source
# TEST: install_python_stack.py -- _ROCM_TORCH_PKG_SPECS mapping
class TestRocmTorchPkgSpecs:
"""Verify per-tag torch version specs are correct."""
def test_rocm72_has_torch_211(self):
"""rocm7.2 should specify torch 2.11.x."""
specs = stack_mod._ROCM_TORCH_PKG_SPECS.get("rocm7.2")
assert specs is not None
torch_spec = specs[0]
assert "2.11" in torch_spec
def test_default_caps_below_211(self):
"""Default spec (rocm7.1 and earlier) should cap below 2.11."""
specs = stack_mod._ROCM_TORCH_PKG_SPECS.get("_default")
assert specs is not None
torch_spec = specs[0]
assert "<2.11" in torch_spec
def test_specs_have_torch_vision_audio(self):
"""Each entry should be a 3-tuple: torch, torchvision, torchaudio."""
for tag, specs in stack_mod._ROCM_TORCH_PKG_SPECS.items():
assert len(specs) == 3, f"{tag}: expected (torch, torchvision, torchaudio)"
assert "torch" in specs[0]
assert "torchvision" in specs[1]
assert "torchaudio" in specs[2]
def test_gfx_to_amd_index_covers_rdna4(self):
"""_GFX_TO_AMD_INDEX_ARCH must cover gfx1200 and gfx1201 (RDNA 4)."""
mapping = stack_mod._GFX_TO_AMD_INDEX_ARCH
assert mapping.get("gfx1200") == "gfx120X-all"
assert mapping.get("gfx1201") == "gfx120X-all"
def test_gfx_to_amd_index_covers_strix_halo(self):
"""_GFX_TO_AMD_INDEX_ARCH must cover gfx1151 and gfx1150 (RDNA 3.5)."""
mapping = stack_mod._GFX_TO_AMD_INDEX_ARCH
assert mapping.get("gfx1151") == "gfx1151"
assert mapping.get("gfx1150") == "gfx1150"
def test_gfx_to_amd_index_covers_rdna3(self):
"""_GFX_TO_AMD_INDEX_ARCH must cover gfx1100-gfx1103 (RDNA 3)."""
mapping = stack_mod._GFX_TO_AMD_INDEX_ARCH
for arch in ("gfx1100", "gfx1101", "gfx1102", "gfx1103"):
assert mapping.get(arch) == "gfx110X-all", f"{arch} missing from mapping"
# TEST: setup.ps1 / install.ps1 -- Strix Halo gfx arch detection
_SETUP_PS1_PATH = PACKAGE_ROOT / "studio" / "setup.ps1"
_INSTALL_PS1_PATH = PACKAGE_ROOT / "install.ps1"
class TestStrixHaloGfxArchDetection:
"""setup.ps1 / install.ps1 gfx arch detection for Strix Halo / iGPU (HIP runtime only, no hipinfo)."""
def test_amd_smi_static_asic_attempted_in_setup(self):
"""setup.ps1 must try 'amd-smi static --asic' when list output lacks gfx arch."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "static --asic" in source
def test_amd_smi_static_asic_attempted_in_install(self):
"""install.ps1 must try 'amd-smi static --asic' when list output lacks gfx arch."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "static --asic" in source
def test_env_var_override_in_setup(self):
"""setup.ps1 must honour UNSLOTH_ROCM_GFX_ARCH as a manual arch override."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "UNSLOTH_ROCM_GFX_ARCH" in source
def test_env_var_override_in_install(self):
"""install.ps1 must honour UNSLOTH_ROCM_GFX_ARCH as a manual arch override."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "UNSLOTH_ROCM_GFX_ARCH" in source
def test_name_arch_table_covers_strix_halo_in_setup(self):
"""setup.ps1 name→arch table must map 890M / Strix Halo to gfx1151."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "gfx1151" in source
assert "890M" in source or "Strix Halo" in source
def test_name_arch_table_covers_strix_halo_in_install(self):
"""install.ps1 name→arch table must map 890M / Strix Halo to gfx1151."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "gfx1151" in source
assert "890M" in source or "Strix Halo" in source
def test_name_arch_table_covers_strix_point_in_setup(self):
"""setup.ps1 name→arch table must map 880M / Strix Point to gfx1150."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "gfx1150" in source
assert "880M" in source or "Strix Point" in source
def test_name_arch_table_covers_strix_point_in_install(self):
"""install.ps1 name→arch table must map 880M / Strix Point to gfx1150."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "gfx1150" in source
assert "880M" in source or "Strix Point" in source
def test_name_arch_table_covers_rdna3_phoenix_in_setup(self):
"""setup.ps1 name→arch table must map 780M / Phoenix to gfx1103."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "gfx1103" in source
assert "780M" in source or "Phoenix" in source
def test_wmi_does_not_set_hasrocm_in_setup(self):
"""WMI block in setup.ps1 must NOT set $HasROCm = $true (no runtime confirmation)."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
wmi_idx = source.find("Win32_VideoController")
assert wmi_idx != -1, "WMI block not found in setup.ps1"
# $HasROCm = $true must not appear within 300 chars of the WMI call.
wmi_context = source[wmi_idx : wmi_idx + 300]
assert "$HasROCm = $true" not in wmi_context
def test_gfx_arch_regex_parses_from_amd_smi_output(self):
"""Both files must use the gfx\\d+[a-z]? regex to parse arch from amd-smi output."""
for path in (_SETUP_PS1_PATH, _INSTALL_PS1_PATH):
source = path.read_text(encoding = "utf-8")
assert (
"gfx\\d+" in source or r"gfx\d+" in source
), f"gfx arch regex not found in {path.name}"
# TEST: HIP SDK tool path resolution via HIP_PATH / ROCM_PATH env vars
class TestHipSdkEnvPathResolution:
"""Both install scripts resolve hipinfo/hipconfig via HIP_PATH/ROCM_PATH off $PATH, and warn."""
@staticmethod
def _assert_accepts_partial_hipinfo_output(source: str):
hipout_idx = source.find("$hipOut = & $hipinfoExe.Source")
assert hipout_idx != -1
hipinfo_block = source[hipout_idx : hipout_idx + 1600]
assert 'if ($hipOut -match "(?i)gcnArchName")' in hipinfo_block
assert "$LASTEXITCODE -eq 0 -and $hipOut -match" not in hipinfo_block
assert "but reported gcnArchName" in hipinfo_block
# ── hipinfo resolution ────────────────────────────────────────────────────
def test_setup_checks_hip_path_for_hipinfo(self):
"""setup.ps1 must reference HIP_PATH when resolving hipinfo."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "HIP_PATH" in source
assert "hipinfo" in source
def test_install_checks_hip_path_for_hipinfo(self):
"""install.ps1 must reference HIP_PATH when resolving hipinfo."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "HIP_PATH" in source
assert "hipinfo" in source
def test_setup_checks_rocm_path_as_hipinfo_fallback(self):
"""setup.ps1 must also check ROCM_PATH as a secondary hipinfo fallback."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "ROCM_PATH" in source
assert "ROCM_PATH" in source and "HIP_PATH" in source
def test_install_checks_rocm_path_as_hipinfo_fallback(self):
"""install.ps1 must also check ROCM_PATH as a secondary hipinfo fallback."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "ROCM_PATH" in source
assert "ROCM_PATH" in source and "HIP_PATH" in source
def test_setup_resolves_hipinfo_via_bin_subdir(self):
"""setup.ps1 must join the env var root with 'bin\\hipinfo.exe'."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert r"bin\hipinfo.exe" in source
def test_install_resolves_hipinfo_via_bin_subdir(self):
"""install.ps1 must join the env var root with 'bin\\hipinfo.exe'."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert r"bin\hipinfo.exe" in source
# ── hipinfo not-on-PATH warning ───────────────────────────────────────────
def test_setup_warns_when_hipinfo_not_on_path(self):
"""setup.ps1 must warn when hipinfo is found via env var but not on PATH."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "hipinfo not on PATH" in source
def test_install_warns_when_hipinfo_not_on_path(self):
"""install.ps1 must warn when hipinfo is found via env var but not on PATH."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "hipinfo not on PATH" in source
# ── warn when HIP_PATH set but exe missing ────────────────────────────────
def test_setup_warns_when_hip_path_set_but_exe_missing(self):
"""setup.ps1 must warn when HIP_PATH is set but hipinfo.exe is not present."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "incomplete" in source or "not found at" in source
def test_install_warns_when_hip_path_set_but_exe_missing(self):
"""install.ps1 must warn when HIP_PATH is set but hipinfo.exe is not present."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "incomplete" in source or "not found at" in source
# ── hipinfo runtime error warning ─────────────────────────────────────────
def test_setup_warns_on_hipinfo_nonzero_exit(self):
"""setup.ps1 must warn when hipinfo runs but returns a non-zero exit code."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "HIP runtime error" in source or "runtime error" in source.lower()
def test_install_warns_on_hipinfo_nonzero_exit(self):
"""install.ps1 must warn when hipinfo runs but returns a non-zero exit code."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "HIP runtime error" in source or "runtime error" in source.lower()
def test_setup_accepts_hipinfo_gcnarchname_on_nonzero_exit(self):
"""setup.ps1 must accept partial hipinfo output from the #6043 crash path."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
self._assert_accepts_partial_hipinfo_output(source)
def test_install_accepts_hipinfo_gcnarchname_on_nonzero_exit(self):
"""install.ps1 must accept partial hipinfo output from the #6043 crash path."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
self._assert_accepts_partial_hipinfo_output(source)
# ── hipconfig resolution ──────────────────────────────────────────────────
def test_setup_resolves_hipconfig_via_bin_subdir(self):
"""setup.ps1 must also fall back to HIP_PATH/bin/hipconfig.exe for version detection."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert r"bin\hipconfig.exe" in source
def test_install_resolves_hipconfig_via_bin_subdir(self):
"""install.ps1 must also fall back to HIP_PATH/bin/hipconfig.exe for version detection."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert r"bin\hipconfig.exe" in source
def test_setup_warns_when_hipconfig_not_on_path(self):
"""setup.ps1 must warn when hipconfig is found via env var but not on PATH."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "hipconfig not on PATH" in source
def test_install_warns_when_hipconfig_not_on_path(self):
"""install.ps1 must warn when hipconfig is found via env var but not on PATH."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "hipconfig not on PATH" in source
# ── PATH fix hint ─────────────────────────────────────────────────────────
def test_setup_provides_path_fix_hint(self):
"""setup.ps1 must tell the user how to add the HIP bin dir to PATH."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "PATH" in source and ("SetEnvironmentVariable" in source or "Add" in source)
def test_install_provides_path_fix_hint(self):
"""install.ps1 must tell the user how to add the HIP bin dir to PATH."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "PATH" in source and ("SetEnvironmentVariable" in source or "Add" in source)
# TEST: HIP SDK detected substep -- path + hipconfig version shown in terminal
class TestHipSdkDetectedSubstep:
"""Both scripts print HIP SDK path and full hipconfig version as substeps when ROCm is detected."""
def test_setup_prints_hip_sdk_path_substep(self):
"""setup.ps1 must print an 'HIP SDK:' substep showing the resolved path."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "HIP SDK:" in source
def test_install_prints_hip_sdk_path_substep(self):
"""install.ps1 must print an 'HIP SDK:' substep showing the resolved path."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "HIP SDK:" in source
def test_setup_shows_hipconfig_full_version(self):
"""setup.ps1 must capture and display the full hipconfig version string."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "ROCmVersionFull" in source or "hipconfig:" in source
def test_install_shows_hipconfig_full_version(self):
"""install.ps1 must capture and display the full hipconfig version string."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "ROCmVersionFull" in source or "hipconfig:" in source
def test_setup_captures_full_version_not_just_major_minor(self):
"""setup.ps1 must store the raw hipconfig output line, not just major.minor."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "ROCmVersionFull" in source
def test_install_captures_full_version_not_just_major_minor(self):
"""install.ps1 must store the raw hipconfig output line, not just major.minor."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "ROCmVersionFull" in source
def test_setup_uses_hip_path_or_rocm_path_for_sdk_display(self):
"""setup.ps1 HIP SDK path substep must check HIP_PATH then ROCM_PATH."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "HIP_PATH" in source and "ROCM_PATH" in source
def test_install_uses_hip_path_or_rocm_path_for_sdk_display(self):
"""install.ps1 HIP SDK path substep must check HIP_PATH then ROCM_PATH."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "HIP_PATH" in source and "ROCM_PATH" in source
def test_setup_rocm_step_uses_full_version(self):
"""setup.ps1 'rocm' step label must prefer the full version string."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "ROCmVersionFull" in source and "rocm" in source
# TEST: install.sh -- Strix Halo rocm7.1 → rocm7.2 override
_INSTALL_SH_PATH = PACKAGE_ROOT / "install.sh"
_SETUP_SH_PATH = PACKAGE_ROOT / "studio" / "setup.sh"
class TestStrixRocm71Override:
"""install.sh routes gfx1151/gfx1150 to AMD's arch index instead of ROCm 7.1 (_grouped_mm segfault)."""
def test_linux_gfx_inference_helpers_present(self):
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
assert "_infer_linux_amd_gfx_arch" in source
assert "_amd_arch_index_family_for_gfx" in source
assert "_amd_gpu_present_via_pci" in source
assert "unslothai#7301" in source
def test_infer_linux_amd_gfx_from_cpuinfo(self):
assert stack_mod._linux_amd_gfx_from_cpuinfo is not None
with patch.object(
Path,
"read_text",
return_value = "model name : AMD Ryzen AI Max+ 395 w/ Radeon 8060S\n",
):
assert stack_mod._linux_amd_gfx_from_cpuinfo() == "gfx1151"
# 8065S (Gorgon Halo) must match on the Radeon name alone, even without the
# "Ryzen AI Max" branding (mirrors setup.sh / setup.ps1 which list 8065S).
with patch.object(Path, "read_text", return_value = "model name : AMD Radeon 8065S\n"):
assert stack_mod._linux_amd_gfx_from_cpuinfo() == "gfx1151"
def test_infer_gfx_gated_out_of_wsl_without_runtime(self):
"""On WSL the cpuinfo/lspci inference must be skipped unless the WSL ROCDXG
runtime (librocdxg) is present: a bare `unsloth studio update` must not
install per-arch ROCm wheels into an env that still can't expose the GPU.
An explicit UNSLOTH_ROCM_GFX_ARCH override stays authoritative regardless."""
m = stack_mod
with (
patch.object(m, "_linux_amd_gfx_from_cpuinfo", return_value = "gfx1151"),
patch.object(m, "_linux_amd_gfx_from_lspci", return_value = None),
# PCI evidence present (the WSL branch never consults it anyway).
patch.object(m, "_linux_amd_display_device_present", return_value = True),
patch.dict(os.environ, {"UNSLOTH_ROCM_GFX_ARCH": ""}),
):
# WSL + no runtime -> inference suppressed (CPU torch stays).
with (
patch.object(m, "_is_wsl", return_value = True),
patch.object(m, "_wsl_rocm_runtime_present", return_value = False),
):
assert m._infer_linux_amd_gfx_arch() is None
# WSL + runtime present (this dev box) -> inference still runs.
with (
patch.object(m, "_is_wsl", return_value = True),
patch.object(m, "_wsl_rocm_runtime_present", return_value = True),
):
assert m._infer_linux_amd_gfx_arch() == "gfx1151"
# Native Linux (not WSL) -> the gate never applies.
with (
patch.object(m, "_is_wsl", return_value = False),
patch.object(m, "_wsl_rocm_runtime_present", return_value = False),
):
assert m._infer_linux_amd_gfx_arch() == "gfx1151"
# Explicit override wins even on a bare WSL box (no runtime).
with (
patch.object(m, "_is_wsl", return_value = True),
patch.object(m, "_wsl_rocm_runtime_present", return_value = False),
patch.dict(os.environ, {"UNSLOTH_ROCM_GFX_ARCH": "gfx1151"}),
):
assert m._infer_linux_amd_gfx_arch() == "gfx1151"
def test_infer_gfx_requires_amd_display_device_on_native_linux(self):
"""A VM/container on a Strix host still shows the host CPU model in
/proc/cpuinfo while receiving no AMD GPU, so on native Linux the
CPU-model inference must require an AMD PCI display device (#7305
review). WSL is exempt (no PCI enumeration there; the librocdxg gate is
the evidence) and the explicit override stays authoritative."""
m = stack_mod
with (
patch.object(m, "_linux_amd_gfx_from_cpuinfo", return_value = "gfx1151"),
patch.object(m, "_linux_amd_gfx_from_lspci", return_value = None),
patch.object(m, "_is_wsl", return_value = False),
patch.dict(os.environ, {"UNSLOTH_ROCM_GFX_ARCH": ""}),
):
# No AMD display device -> the CPU-model text alone must not infer.
with patch.object(m, "_linux_amd_display_device_present", return_value = False):
assert m._infer_linux_amd_gfx_arch() is None
# Device present -> inference unchanged.
with patch.object(m, "_linux_amd_display_device_present", return_value = True):
assert m._infer_linux_amd_gfx_arch() == "gfx1151"
# Explicit override needs no device evidence (headless/cross-install).
with (
patch.object(m, "_is_wsl", return_value = False),
patch.object(m, "_linux_amd_display_device_present", return_value = False),
patch.dict(os.environ, {"UNSLOTH_ROCM_GFX_ARCH": "GFX1151"}),
):
assert m._infer_linux_amd_gfx_arch() == "gfx1151"
def test_install_sh_cpuinfo_inference_requires_pci_evidence(self):
"""install.sh mirror of the VM/container guard: every cpuinfo grep must be
gated on _gpu_evidence (AMD PCI display device via _amd_gpu_present_via_pci,
or the WSL librocdxg gate), and the gate must sit before the first grep."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
body = _extract_sh_function_body(source, "_infer_linux_amd_gfx_arch")
assert body, "could not extract _infer_linux_amd_gfx_arch"
pci = body.find("_amd_gpu_present_via_pci")
infer = body.find("grep -qiE 'Ryzen AI Max")
assert pci >= 0 and infer >= 0
assert pci < infer, "the PCI evidence check must run before the cpuinfo inference"
assert body.count("grep -qiE") == body.count(
'[ -n "$_gpu_evidence" ] && grep -qiE'
), "every cpuinfo grep (gfx1151/gfx1150/gfx1152) must be gated on _gpu_evidence"
def test_lspci_scan_covers_all_display_controllers(self):
"""The lspci fallback must scan every display-class line, not just the
first: a non-AMD controller (Intel iGPU, ASPEED BMC) often enumerates
before the AMD dGPU. Non-AMD vendors must never map (an NVIDIA GeForce
GTX 860M would otherwise hit the AMD 860M pattern), and a 0000: PCI
domain prefix must not break matching."""
m = stack_mod
def fake_lspci(stdout):
result = SimpleNamespace(returncode = 0, stdout = stdout)
return (
patch.object(m.shutil, "which", return_value = "/usr/bin/lspci"),
patch.object(m.subprocess, "run", return_value = result),
)
intel_then_amd = (
"00:02.0 VGA compatible controller [0300]: Intel Corporation Raptor Lake-S GT1 [8086:a780]\n"
"03:00.0 VGA compatible controller [0300]: Advanced Micro Devices, Inc. [AMD/ATI]"
" Navi 31 [Radeon RX 7900 XT] [1002:744c]\n"
)
nvidia_only = "01:00.0 3D controller [0302]: NVIDIA Corporation GM107M [GeForce GTX 860M] [10de:1392]\n"
domain_prefixed = (
"0000:c5:00.0 VGA compatible controller [0300]: Advanced Micro Devices, Inc. [AMD/ATI]"
" Strix Halo [Radeon Graphics / Radeon 8060S] [1002:150e]\n"
)
unmapped_then_mapped = (
"03:00.0 Display controller [0380]: Advanced Micro Devices, Inc. [AMD/ATI]"
" Cape Verde [FirePro W600] [1002:6821]\n"
"04:00.0 VGA compatible controller [0300]: Advanced Micro Devices, Inc. [AMD/ATI]"
" Navi 33 [Radeon RX 7600] [1002:7480]\n"
)
for stdout, expected in (
(intel_then_amd, "gfx1100"),
(nvidia_only, None),
(domain_prefixed, "gfx1151"),
(unmapped_then_mapped, "gfx1102"),
):
w, r = fake_lspci(stdout)
with w, r:
assert m._linux_amd_gfx_from_lspci() == expected, stdout
def test_install_sh_lspci_scan_covers_all_display_controllers(self):
"""install.sh mirror of the scan-all behaviour, executed with a shimmed
lspci: Intel-first still finds the AMD dGPU, NVIDIA-only maps nothing
(860M collision), a domain-prefixed AMD line still maps."""
shell = shutil.which("bash")
if not shell:
pytest.skip("bash needed to execute the probe block")
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
name_fn = re.search(
r"^_infer_amd_gfx_arch_from_gpu_name\(\) \{\n.*?\n\}\n", source, re.S | re.M
)
scan = re.search(
r"^ if command -v lspci[^\n]*\n.*?\nEOF\n fi\n return 1\n", source, re.S | re.M
)
assert name_fn and scan, "could not extract the lspci scan block"
cases = (
(
"00:02.0 VGA compatible controller [0300]: Intel Corporation UHD [8086:a780]\n"
"03:00.0 VGA compatible controller [0300]: Advanced Micro Devices, Inc. [AMD/ATI]"
" Navi 31 [Radeon RX 7900 XT] [1002:744c]",
"OK:gfx1100",
),
(
"01:00.0 3D controller [0302]: NVIDIA Corporation GM107M [GeForce GTX 860M] [10de:1392]",
"OK:",
),
(
"0000:c5:00.0 VGA compatible controller [0300]: Advanced Micro Devices, Inc."
" [AMD/ATI] Strix Halo [Radeon 8060S] [1002:150e]",
"OK:gfx1151",
),
)
for lspci_out, expected in cases:
with tempfile.TemporaryDirectory() as d:
p = os.path.join(d, "lspci")
with open(p, "w", encoding = "utf-8") as f:
f.write(f'#!/bin/sh\ncat <<"EOT"\n{lspci_out}\nEOT\n')
os.chmod(p, 0o755)
script = (
"set -euo pipefail\n"
+ name_fn.group(0)
+ "probe() {\n"
+ scan.group(0)
+ "}\nprintf 'OK:%s\\n' \"$(probe || true)\"\n"
)
env = dict(os.environ, PATH = d + os.pathsep + os.environ.get("PATH", ""))
r = subprocess.run([shell, "-c", script], env = env, capture_output = True, text = True)
assert r.returncode == 0, f"scan aborted: {r.stderr}"
assert (
r.stdout.splitlines()[-1] == expected
), f"lspci scan wrong for {lspci_out!r}: {r.stdout!r}"
def test_install_sh_infer_gfx_gated_on_wsl_runtime(self):
"""install.sh's _infer_linux_amd_gfx_arch must, like the Python side, skip
the cpuinfo/lspci inference on WSL unless librocdxg is present -- the
override still returns first, so it stays authoritative."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
body = _extract_sh_function_body(source, "_infer_linux_amd_gfx_arch")
assert body, "could not extract _infer_linux_amd_gfx_arch"
override = body.find("UNSLOTH_ROCM_GFX_ARCH")
dxg = body.find("/dev/dxg")
rocdxg = body.find("librocdxg")
# Anchor on the first cpuinfo *inference* (the grep), not a comment mention.
infer = body.find("grep -qiE 'Ryzen AI Max")
assert override >= 0 and dxg >= 0 and rocdxg >= 0 and infer >= 0
assert "microsoft" in body, "WSL gate must also detect WSL via /proc/version"
assert override < dxg, "the explicit override must return before the WSL gate"
assert (
dxg < infer and rocdxg < infer
), "the WSL/librocdxg gate must run before the cpuinfo/lspci inference"
def test_install_sh_reroute_is_x86_64_only(self):
"""The Linux inferred-gfx reroute must be x86_64-only: ROCm torch wheels are
not published for arm64, so an inferred/overridden gfx must not push an
arm64 host to the AMD arch index (get_torch_index_url returns CPU there)."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
idx = source.find("_linux_inferred_gfx=$(_infer_linux_amd_gfx_arch")
assert idx >= 0, "reroute consumer not found"
window = source[max(0, idx - 400) : idx]
assert (
'case "$_ARCH" in x86_64|amd64)' in window
), "the inferred-gfx reroute must guard on x86_64|amd64 arch"
def test_install_sh_reroute_skips_visible_rocm_gpu(self):
"""A */cpu index on a host whose AMD GPU IS visible to the ROCm probes is a
deliberate fallback (unsupported/unreadable ROCm version, warned about in
get_torch_index_url), not a missing runtime: the reroute must not override
it with inferred per-arch wheels. The explicit UNSLOTH_ROCM_GFX_ARCH
override must still win either way."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
idx = source.find("_linux_inferred_gfx=$(_infer_linux_amd_gfx_arch")
assert idx >= 0, "reroute consumer not found"
window = source[max(0, idx - 700) : idx]
assert (
"! _has_amd_rocm_gpu" in window
), "the reroute must be gated on _has_amd_rocm_gpu being false"
assert (
'[ -n "${UNSLOTH_ROCM_GFX_ARCH:-}" ] || ! _has_amd_rocm_gpu' in window
), "an explicit UNSLOTH_ROCM_GFX_ARCH override must bypass the visible-GPU gate"
def test_install_sh_reroute_exports_gfx_for_setup_sh(self):
"""The inferred arch must be exported as UNSLOTH_ROCM_GFX_ARCH so the
downstream setup.sh run (which re-probes ROCm independently and finds
nothing on these runtime-less hosts) routes llama.cpp to the matching
ROCm prebuilt instead of the CPU one -- setup.sh and
install_llama_prebuilt.py both read that env var."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
assign = source.find('TORCH_INDEX_URL="${_amd_mirror}/${_amd_family}/"')
assert assign >= 0, "inferred-gfx index assignment not found"
block_end = source.find("esac", assign)
assert (
'export UNSLOTH_ROCM_GFX_ARCH="$_linux_inferred_gfx"' in source[assign:block_end]
), "the reroute must export the inferred gfx for the setup.sh handoff"
# setup.sh's side of the handoff must still exist.
setup_source = (PACKAGE_ROOT / "studio" / "setup.sh").read_text(encoding = "utf-8")
assert "UNSLOTH_ROCM_GFX_ARCH" in setup_source
def test_amd_arch_index_url_linux_honors_amd_mirror(self):
"""On Linux the inferred-gfx repair must honour UNSLOTH_AMD_ROCM_MIRROR (the
var install.sh uses), not the Windows mirror var, so a mirrored/air-gapped
Linux install does not silently fall back to repo.amd.com. Windows still
delegates to the Windows mirror path."""
m = stack_mod
with (
patch.object(m, "IS_WINDOWS", False),
patch.dict(os.environ, {"UNSLOTH_AMD_ROCM_MIRROR": "https://mirror.local/rocm"}),
):
assert m._amd_arch_index_url("gfx1151") == "https://mirror.local/rocm/gfx1151/"
with (
patch.object(m, "IS_WINDOWS", False),
patch.dict(os.environ, {"UNSLOTH_AMD_ROCM_MIRROR": ""}),
):
assert m._amd_arch_index_url("gfx1151") == "https://repo.amd.com/rocm/whl/gfx1151/"
assert m._amd_arch_index_url("gfx9999") is None
# Windows path is unchanged: delegate to the Windows mirror helper.
with patch.object(m, "IS_WINDOWS", True):
assert m._amd_arch_index_url("gfx1151") == m._windows_rocm_index_url("gfx1151")
def test_strix_gfx_detection_in_install_sh(self):
"""install.sh must detect gfx1151 and gfx1150 for the override."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
assert "gfx1151" in source and "gfx1150" in source
def test_rocm71_override_to_amd_arch_index_in_install_sh(self):
"""install.sh must override TORCH_INDEX_URL to AMD arch-specific index for Strix."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
assert "repo.amd.com/rocm/whl" in source
assert "_strix_gfx" in source
# URL must incorporate the detected gfx arch (gfx1151 -> .../gfx1151/).
strix_idx = source.find("_amd_strix_base")
assert strix_idx != -1
ctx = source[strix_idx : strix_idx + 500]
assert "_strix_gfx" in ctx
def test_radeon_repo_bypassed_for_strix_in_install_sh(self):
"""install.sh must set _amd_gpu_radeon=false when Strix + ROCm 7.1 detected."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
assert "_amd_gpu_radeon=false" in source
def test_strix_override_warns_with_moe_utils_reference(self):
"""install.sh must emit a [WARN] mentioning the moe_utils segfault."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
assert "moe_utils" in source or "_grouped_mm" in source
def test_strix_override_scoped_below_arch_floor(self):
"""Strix reroute must fire for rocm leaves BELOW the arch floor (7.13) and
NOT at/above it. Executed via _rocm_leaf_below so it verifies the actual
version comparison, not a text match that a comment could satisfy."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
# Selector + gate must switch on the index LEAF, not the whole URL (a mirror
# base path with its own rocm token would false-positive otherwise).
assert 'case "$_torch_index_leaf" in' in source
assert '_rocm_leaf_below "$_torch_index_leaf" 7 13' in source
shell = shutil.which("sh") or shutil.which("bash")
if not shell:
pytest.skip("no POSIX shell to execute _rocm_leaf_below")
match = re.search(r"^_rocm_leaf_below\(\) \{.*?^\}", source, re.S | re.M)
assert match, "could not extract _rocm_leaf_below from install.sh"
fn = match.group(0)
def below(leaf):
return (
subprocess.run(
[shell, "-c", f'{fn}\n_rocm_leaf_below "$1" 7 13', "_", leaf]
).returncode
== 0
)
for leaf in ("rocm6.0", "rocm7.0", "rocm7.1", "rocm7.2", "rocm7.12"):
assert below(leaf), f"{leaf} must reroute (below arch floor 7.13)"
for leaf in ("rocm7.13", "rocm7.14", "rocm8.0", "gfx1151", "cu128", "cpu"):
assert not below(leaf), f"{leaf} must NOT reroute (>= floor or non-rocm)"
def test_gfx_probe_survives_no_match_under_set_e(self):
"""A gfx probe whose grep finds no match must not abort install.sh under
set -euo pipefail before the amd-smi fallback runs. The reroute case now
matches every rocm* index, so this would break ordinary 6.x/7.2 installs
with a flaky rocminfo. Executed with shimmed tools, not a text match."""
shell = shutil.which("bash")
if not shell:
pytest.skip("bash needed to execute the probe block")
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
block = re.search(
r'^ _gfx_all=\$\(printf[^\n]*\n.*?(?=^ _strix_gfx="")',
source,
re.S | re.M,
)
assert block, "could not extract the gfx-detection block"
with tempfile.TemporaryDirectory() as d:
# rocminfo emits no gfx token; amd-smi supplies gfx1151 (the fallback)
for name, out in (("rocminfo", "no gpu here"), ("amd-smi", "GPU: gfx1151")):
p = os.path.join(d, name)
with open(p, "w", encoding = "utf-8") as f:
f.write(f'#!/bin/sh\ncat <<"EOT"\n{out}\nEOT\n')
os.chmod(p, 0o755)
script = (
'set -euo pipefail\nHIP_VISIBLE_DEVICES=""\nROCR_VISIBLE_DEVICES=""\n'
+ block.group(0)
+ '\nprintf "OK:%s\\n" "$_gfx_all"\n'
)
env = dict(os.environ, PATH = d + os.pathsep + os.environ.get("PATH", ""))
r = subprocess.run([shell, "-c", script], env = env, capture_output = True, text = True)
assert r.returncode == 0, f"probe aborted under set -e: {r.stderr}"
assert "OK:gfx1151" in r.stdout, f"amd-smi fallback not reached: {r.stdout!r}"
def test_strix_reroute_reprobes_when_mask_hides_all(self):
"""A visibility mask hiding every agent (ROCR_VISIBLE_DEVICES=-1) must not
skip the Strix reroute: get_torch_index_url reads the arch unmasked, so
the reroute must re-probe unmasked too or a masked Strix box gets the
broken generic wheels. A partial mask must keep its per-GPU selection.
Executed with mask-honouring shims, not a text match."""
shell = shutil.which("bash")
if not shell:
pytest.skip("bash needed to execute the probe block")
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
block = re.search(
r'^ _gfx_all=\$\(printf[^\n]*\n.*?(?=^ _strix_gfx="")',
source,
re.S | re.M,
)
assert block, "could not extract the gfx-detection block"
with tempfile.TemporaryDirectory() as d:
# rocminfo honours ROCR_VISIBLE_DEVICES like the real tool: -1 and
# set-but-empty hide both agents, 1 renumbers to the dGPU only,
# unset shows both.
rocminfo = (
"#!/bin/sh\n"
'case "${ROCR_VISIBLE_DEVICES-__unset__}" in\n'
' __unset__) printf "Name: gfx1151\\nName: gfx1201\\n" ;;\n'
' ""|-1) echo "no visible agents" ;;\n'
' 1) printf "Name: gfx1201\\n" ;;\n'
' *) printf "Name: gfx1151\\nName: gfx1201\\n" ;;\n'
"esac\n"
)
for name, body in (("rocminfo", rocminfo), ("amd-smi", "#!/bin/sh\nexit 0\n")):
p = os.path.join(d, name)
with open(p, "w", encoding = "utf-8") as f:
f.write(body)
os.chmod(p, 0o755)
script = (
"set -euo pipefail\n" + block.group(0) + '\nprintf "OK:%s\\n" "$_runtime_gfx"\n'
)
def run(**extra):
env = dict(os.environ, PATH = d + os.pathsep + os.environ.get("PATH", ""), **extra)
env.pop("UNSLOTH_ROCM_GFX_ARCH", None)
env.pop("HIP_VISIBLE_DEVICES", None)
return subprocess.run(
[shell, "-c", script], env = env, capture_output = True, text = True
)
# Mask hides everything: re-probe must recover the first GPU (Strix).
r = run(ROCR_VISIBLE_DEVICES = "-1")
assert r.returncode == 0, f"masked probe aborted: {r.stderr}"
assert "OK:gfx1151" in r.stdout, f"reroute blinded by full mask: {r.stdout!r}"
# A SET-but-empty mask also hides every agent and must re-probe too
# (the ${VAR+x} guard, not ${VAR:-}).
r0 = run(ROCR_VISIBLE_DEVICES = "")
assert r0.returncode == 0, f"empty-mask probe aborted: {r0.stderr}"
assert "OK:gfx1151" in r0.stdout, f"reroute blinded by empty mask: {r0.stdout!r}"
# Partial mask: enumeration already reflects it; the dGPU selection
# must survive (no unmasked re-probe overriding the user's pick).
r2 = run(ROCR_VISIBLE_DEVICES = "1")
assert r2.returncode == 0, f"partial-mask probe aborted: {r2.stderr}"
assert "OK:gfx1201" in r2.stdout, f"partial mask selection lost: {r2.stdout!r}"
def test_strix_routing_helpers_cover_rocm714(self):
# Reroute for any generic pytorch.org index below the 7.13 arch floor (7.0,
# 7.2, a future 7.3+), never at/above it -- mirrors install.sh _rocm_leaf_below.
assert stack_mod._generic_pytorch_rocm_tag((7, 14)) == "rocm7.2"
assert stack_mod._strix_needs_amd_arch_index((7, 14)) is True
assert stack_mod._strix_needs_amd_arch_index((7, 0)) is True
assert stack_mod._strix_needs_amd_arch_index((6, 0)) is True
assert stack_mod._strix_needs_amd_arch_index((5, 0)) is False
def test_torch_constraint_updated_for_strix_amd_index(self):
"""install.sh must set TORCH_CONSTRAINT>=2.11 when routing Strix to AMD index."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
assert "TORCH_CONSTRAINT" in source and "2.11" in source
def test_torch_constraint_211_matches_leaf_not_whole_url(self):
"""The 2.11 constraint case must match the index LEAF, not the whole URL.
A custom UNSLOTH_PYTORCH_MIRROR whose base path contains a gfx/rocm7.2
segment (e.g. https://mirror.local/gfx-cache) with a cu*/cpu family must
not be pushed to the torch 2.11 line -- same leaf-only reasoning the
UNSLOTH_TORCH_BACKEND classification uses.
"""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
# The 2.11 constraint block must switch on $_torch_index_leaf, not the full
# $TORCH_INDEX_URL (a */gfx* match false-positives on a mirror base path). Only the
# _grouped_mm-bug gfx families (gfx120X-all / gfx1151 / gfx1150 / gfx1152) go to 2.11;
# a bare gfx* would also floor gfx110X-all/gfx90a/gfx908, left bare on purpose.
assert (
'case "$_torch_index_leaf" in\n rocm7.2|gfx120x-all|gfx1151|gfx1150|gfx1152)'
in source
), (
"the torch>=2.11 constraint must match the specific gfx leaves that need "
"it (rocm7.2|gfx120x-all|gfx1151|gfx1150|gfx1152), not a bare gfx* or the URL"
)
def test_amd_rocm_mirror_env_var_respected(self):
"""install.sh must honour UNSLOTH_AMD_ROCM_MIRROR for air-gapped installs."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
assert "UNSLOTH_AMD_ROCM_MIRROR" in source
def test_tauri_family_recognises_amd_arch_url(self):
"""_tauri_torch_index_family must return a rocm* family for AMD arch-specific URLs."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
assert "rocm/whl/gfx" in source
# TEST: setup.sh -- gcc-install-dir fix for Ubuntu 24.04 + ROCm 7.x clang-20
class TestSetupShGccInstallDir:
"""setup.sh applies --gcc-install-dir for HIP builds on Ubuntu 24.04+ (ROCm 7.x clang-20 header bug)."""
def test_gcc_install_dir_search_loop_present(self):
"""setup.sh must iterate gcc versions 14→11 to find one with C++ headers."""
source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
assert "_GCC_INSTALL_DIR" in source
assert "/usr/lib/gcc/x86_64-linux-gnu" in source
def test_gcc_install_dir_checks_include_dir(self):
"""setup.sh must check that the gcc dir has an 'include' subdirectory."""
source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
assert "include" in source and "_GCC_INSTALL_DIR" in source
def test_gcc_install_dir_appended_to_cmake_hip_flags(self):
"""setup.sh must pass --gcc-install-dir via CMAKE_HIP_FLAGS."""
source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
assert "CMAKE_HIP_FLAGS" in source
assert "gcc-install-dir" in source
def test_gcc_install_dir_only_applied_in_hip_build_block(self):
"""The --gcc-install-dir fix must only apply in the HIP/ROCm build branch."""
source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
hip_idx = source.find("GGML_HIP=ON")
gcc_idx = source.find("gcc-install-dir")
assert hip_idx != -1 and gcc_idx != -1
assert hip_idx < gcc_idx
def test_gcc_install_dir_logs_substep(self):
"""setup.sh must print a substep when the gcc install dir is resolved."""
source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
assert "gcc install dir" in source or "GCC_INSTALL_DIR" in source
# TEST: main.py -- BNB_ROCM_VERSION server startup + distributed stubs
_MAIN_PY_PATH = PACKAGE_ROOT / "studio" / "backend" / "main.py"
_HARDWARE_PY_PATH = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
class TestServerStartupRocmFixes:
"""main.py sets BNB_ROCM_VERSION pre-bnb-import; hardware.py stubs _distributed_c10d pre-torch.distributed."""
# ── BNB_ROCM_VERSION in server process ────────────────────────────────────
def test_main_py_sets_bnb_rocm_version(self):
"""main.py must set BNB_ROCM_VERSION in the server process before imports."""
source = _MAIN_PY_PATH.read_text(encoding = "utf-8")
assert "BNB_ROCM_VERSION" in source
def test_main_py_bnb_detection_scoped_to_win32(self):
"""main.py BNB_ROCM_VERSION logic must be inside the win32 platform guard."""
source = _MAIN_PY_PATH.read_text(encoding = "utf-8")
win32_idx = source.find('sys.platform == "win32"')
bnb_idx = source.find("BNB_ROCM_VERSION")
assert win32_idx != -1 and bnb_idx != -1
assert win32_idx < bnb_idx
def test_main_py_bnb_dll_detection_uses_glob(self):
"""main.py must scan for libbitsandbytes_rocm*.dll to find the right version."""
source = _MAIN_PY_PATH.read_text(encoding = "utf-8")
assert "libbitsandbytes_rocm" in source
def test_main_py_bnb_falls_back_to_72(self):
"""main.py must fall back to BNB_ROCM_VERSION='72' when no DLL is found."""
source = _MAIN_PY_PATH.read_text(encoding = "utf-8")
assert '"72"' in source or "'72'" in source
def test_main_py_bnb_only_set_when_not_already_in_env(self):
"""main.py must not override an existing BNB_ROCM_VERSION env var."""
source = _MAIN_PY_PATH.read_text(encoding = "utf-8")
assert '"BNB_ROCM_VERSION" not in os.environ' in source
# ── hipInfo.exe PATH prepend (bitsandbytes arch-probe fix) ────────────────
# bnb's get_rocm_gpu_arch() runs hipinfo.exe via PATH at import; the AMD wheel ships it
# in venv Scripts (on PATH only for activated venvs), so without the prepend bnb logs
# "[WinError 2]" when launched directly.
def test_main_py_prepends_hipinfo_dir_to_path(self):
"""main.py must make hipInfo.exe resolvable before bnb imports."""
source = _MAIN_PY_PATH.read_text(encoding = "utf-8")
assert "hipInfo.exe" in source
# Prepend must precede the BNB_ROCM_VERSION block so bnb sees the fixed PATH.
assert source.find("hipInfo.exe") < source.find("BNB_ROCM_VERSION")
def test_main_py_hipinfo_prepend_gated_on_file_presence(self):
"""Prepend must check hipInfo.exe exists first (only AMD wheels ship it; leave NVIDIA/CPU untouched)."""
source = _MAIN_PY_PATH.read_text(encoding = "utf-8")
assert 'os.path.isfile(os.path.join(_scripts_dir, "hipInfo.exe"))' in source
def test_worker_py_prepends_hipinfo_dir_to_path(self):
"""worker.py must mirror the prepend for standalone-spawned workers."""
source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "hipInfo.exe" in source
def test_install_stack_prepends_hipinfo_dir_to_path(self):
"""install_python_stack.py must prepend so child import checks inherit a PATH where bnb's probe works."""
source = _STACK_PATH.read_text(encoding = "utf-8")
assert "hipInfo.exe" in source
# ── torch._C._distributed_c10d stubs in hardware.py ──────────────────────
def test_hardware_py_injects_distributed_c10d_stub(self):
"""hardware.py must inject torch._C._distributed_c10d into sys.modules."""
source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
assert "_distributed_c10d" in source
def test_hardware_py_stub_injected_before_distributed_import(self):
"""The sys.modules stub must be injected BEFORE import torch.distributed."""
source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
c10d_idx = source.find("_distributed_c10d")
dist_idx = source.find("import torch.distributed")
assert c10d_idx != -1 and dist_idx != -1
assert c10d_idx < dist_idx
def test_hardware_py_stub_uses_types_moduletype(self):
"""hardware.py must create the stub with types.ModuleType."""
source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
assert "ModuleType" in source
def test_hardware_py_stub_scoped_to_win32(self):
"""hardware.py distributed stub injection must be gated on win32."""
source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
assert 'platform == "win32"' in source or "win32" in source
def test_hardware_py_stub_exposes_fake_process_group(self):
"""hardware.py stub must set FakeProcessGroup so torch.distributed doesn't raise AttributeError."""
source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
assert "FakeProcessGroup" in source
def test_hardware_py_stub_exposes_process_group(self):
"""hardware.py stub must set ProcessGroup on the c10d stub."""
source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
assert "ProcessGroup" in source
def test_hardware_py_stub_uses_setattr_for_symbols(self):
"""hardware.py must use setattr to populate stub symbols dynamically."""
source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
assert "setattr" in source
def test_hardware_py_stub_all_c10d_siblings_covered(self):
"""hardware.py must stub all three torch._C._distributed_* submodules."""
source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
assert "_distributed_c10d" in source
assert "_distributed_autograd" in source
assert "_distributed_rpc" in source
# TEST: install.ps1 / setup.ps1 -- HipSdkInstalled flag (SDK found, device inaccessible)
class TestHipSdkInstalledButDeviceInaccessible:
"""When hipinfo is found but exits non-zero, both scripts distinguish device-inaccessible from SDK-not-found."""
def test_install_ps1_has_hip_sdk_installed_flag(self):
"""install.ps1 must track HipSdkInstalled separately from HasROCm."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "HipSdkInstalled" in source
def test_setup_ps1_has_hip_sdk_installed_flag(self):
"""setup.ps1 must track HipSdkInstalled separately from HasROCm."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "HipSdkInstalled" in source
def test_install_ps1_sets_flag_when_hipinfo_binary_found(self):
"""install.ps1 must set HipSdkInstalled=true inside the 'if ($hipinfoExe)' block."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
hipinfo_block_idx = source.find("if ($hipinfoExe)")
sdk_flag_idx = source.find("$HipSdkInstalled = $true", hipinfo_block_idx)
assert hipinfo_block_idx != -1 and sdk_flag_idx != -1
assert sdk_flag_idx > hipinfo_block_idx
def test_setup_ps1_sets_flag_when_hipinfo_binary_found(self):
"""setup.ps1 must set HipSdkInstalled=true inside the 'if ($hipinfoExe)' block."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
hipinfo_block_idx = source.find("if ($hipinfoExe)")
sdk_flag_idx = source.find("$HipSdkInstalled = $true", hipinfo_block_idx)
assert hipinfo_block_idx != -1 and sdk_flag_idx != -1
assert sdk_flag_idx > hipinfo_block_idx
def test_install_ps1_version_capture_runs_when_sdk_installed(self):
"""install.ps1 must capture hipconfig version when HipSdkInstalled even if HasROCm is false."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "HasROCm -or $HipSdkInstalled" in source or "$HipSdkInstalled" in source
def test_setup_ps1_version_capture_runs_when_sdk_installed(self):
"""setup.ps1 must capture hipconfig version when HipSdkInstalled even if HasROCm is false."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "HasROCm -or $HipSdkInstalled" in source or "$HipSdkInstalled" in source
def test_install_ps1_distinct_message_for_sdk_found_but_device_inaccessible(self):
"""install.ps1 must show 'not ROCm-accessible' message (not 'HIP SDK not found') when SDK present."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "not ROCm-accessible" in source
def test_setup_ps1_distinct_message_for_sdk_found_but_device_inaccessible(self):
"""setup.ps1 must show 'not ROCm-accessible' message (not 'HIP SDK not found') when SDK present."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "not ROCm-accessible" in source
def test_install_ps1_driver_guidance_in_sdk_found_branch(self):
"""install.ps1 must tell user this is a driver issue, not an SDK issue."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "driver issue" in source
def test_setup_ps1_driver_guidance_in_sdk_found_branch(self):
"""setup.ps1 must tell user this is a driver issue, not an SDK issue."""
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "driver issue" in source
def test_install_ps1_cpu_hint_distinguishes_driver_vs_no_sdk(self):
"""install.ps1 CPU-only hint must say 'GPU not ROCm-accessible' not 'require the HIP SDK' when SDK found."""
source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
assert "GPU not ROCm-accessible" in source
# TEST: --rocm-gfx forwarding -- setup.sh/setup.ps1 forward their resolved gfx
# arch to install_llama_prebuilt.py so the per-gfx prebuilt is picked.
_SETUP_SH_PATH = PACKAGE_ROOT / "studio" / "setup.sh"
class TestNormalizeForwardedGfx:
"""A forwarded gfx string is reduced to a single clean gfx token."""
def test_plain_token(self):
assert _normalize_forwarded_gfx("gfx1151") == "gfx1151"
def test_uppercase_normalized(self):
assert _normalize_forwarded_gfx("GFX1151") == "gfx1151"
def test_extracts_from_noise(self):
assert _normalize_forwarded_gfx("gcnArchName: gfx942") == "gfx942"
def test_malformed_is_ignored(self):
assert _normalize_forwarded_gfx("not-a-gpu") is None
def test_empty_and_none(self):
assert _normalize_forwarded_gfx("") is None
assert _normalize_forwarded_gfx(None) is None
class TestApplyHostOverrides:
"""Forwarded ROCm detection is folded into the host profile correctly."""
def test_forwarded_gfx_fills_empty_probe(self):
# Installer probe found no gfx (amd-smi-only / name-inferred host).
host = rocm_host(rocm_gfx_target = None)
out = _apply_host_overrides(host, override_rocm_gfx = "gfx1151")
assert out.has_rocm is True
assert out.rocm_gfx_target == "gfx1151"
def test_forwarded_gfx_implies_rocm(self):
# A CPU-looking host with a forwarded gfx is an AMD host.
out = _apply_host_overrides(cpu_host(), override_rocm_gfx = "gfx1200")
assert out.has_rocm is True
assert out.rocm_gfx_target == "gfx1200"
def test_forwarded_gfx_does_not_clobber_probed_arch(self, monkeypatch):
# setup.ps1's pick is not fully visible-device aware (its amd-smi branch drops comma
# masks), so when it resolved the host's OTHER physical GPU it must not replace the
# arch detect_host() picked for the visible one.
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
host = rocm_host(rocm_gfx_target = "gfx1010", rocm_gfx_targets = ["gfx1100", "gfx1010"])
out = _apply_host_overrides(host, override_rocm_gfx = "gfx1100")
assert out.rocm_gfx_target == "gfx1010"
assert out.rocm_gfx_targets == ["gfx1100", "gfx1010"]
assert out.has_rocm is True
def test_shadowing_igpu_does_not_veto_the_forwarded_dgpu(self, monkeypatch):
# #7776: unmasked, setup skips a leading APU for the discrete card, but
# _pick_rocm_gfx_target() still reads that APU as device 0. Discarding the forward
# as "advisory" would give torch gfx120X and llama.cpp the gfx1103 iGPU bundle.
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
for _v in ("HIP_VISIBLE_DEVICES", "ROCR_VISIBLE_DEVICES", "CUDA_VISIBLE_DEVICES"):
monkeypatch.delenv(_v, raising = False)
host = rocm_host(rocm_gfx_target = "gfx1103", rocm_gfx_targets = ["gfx1103", "gfx1200"])
out = _apply_host_overrides(host, override_rocm_gfx = "gfx1200")
assert out.rocm_gfx_target == "gfx1200"
assert out.rocm_gfx_targets == ["gfx1103", "gfx1200"]
def test_masked_host_still_keeps_the_probed_arch(self, monkeypatch):
# Negative control: the repick never overrides a pin, so under a mask a differing
# forward is a mispick again and the probe's pick survives.
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
monkeypatch.setenv("HIP_VISIBLE_DEVICES", "0")
host = rocm_host(rocm_gfx_target = "gfx1103", rocm_gfx_targets = ["gfx1103", "gfx1200"])
out = _apply_host_overrides(host, override_rocm_gfx = "gfx1200")
assert out.rocm_gfx_target == "gfx1103"
def test_shadowing_exception_needs_a_probed_dgpu(self, monkeypatch):
# The exception only covers a card the probe saw; a family label is a bundle
# name, so gfx103X over a gfx1033 APU stays advisory as before.
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
for _v in ("HIP_VISIBLE_DEVICES", "ROCR_VISIBLE_DEVICES", "CUDA_VISIBLE_DEVICES"):
monkeypatch.delenv(_v, raising = False)
host = rocm_host(rocm_gfx_target = "gfx1033", rocm_gfx_targets = ["gfx1033"])
out = _apply_host_overrides(host, override_rocm_gfx = "gfx103X")
assert out.rocm_gfx_target == "gfx1033"
def test_forwarded_gfx_absent_from_host_stays_authoritative(self, monkeypatch):
# An arch no probe here ever reported is not a setup mispick: it is an explicit
# --rocm-gfx for a host whose probe is wrong or stale, so it must still win.
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
host = rocm_host(rocm_gfx_target = "gfx1100", rocm_gfx_targets = ["gfx1100"])
out = _apply_host_overrides(host, override_rocm_gfx = "gfx1151")
assert out.rocm_gfx_target == "gfx1151"
# ... but it says which arch HIP targets, not which cards exist, so the probed
# gfx1100 is still in the box and stays in the per-GPU list.
assert out.rocm_gfx_targets == ["gfx1100", "gfx1151"]
assert out.has_rocm is True
def test_forwarded_family_label_never_overrides_a_probed_arch(self, monkeypatch):
# The update path re-derives --rocm-gfx from the marker's family-named asset, so a
# family label is a bundle name, not a real arch, and must stay advisory: gfx1033 is
# in-generation but unbuilt, so gfx103X winning would serve a bundle it cannot run.
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
host = rocm_host(rocm_gfx_target = "gfx1033", rocm_gfx_targets = ["gfx1033"])
out = _apply_host_overrides(host, override_rocm_gfx = "gfx103X")
assert out.rocm_gfx_target == "gfx1033"
assert out.has_rocm is True
def test_forwarded_family_label_still_fills_an_unprobed_arch(self, monkeypatch):
# Negative control: with no probed arch the forward is the only source, so it
# applies.
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
out = _apply_host_overrides(cpu_host(), override_rocm_gfx = "gfx110X")
assert out.rocm_gfx_target == "gfx110x"
assert out.has_rocm is True
def test_forwarded_gfx_matching_active_keeps_physical_gfx_list(self, monkeypatch):
# When the forward agrees with the probe the per-GPU list must survive: collapsing
# it would hide the host's other AMD cards from the Windows auto-Vulkan floor
# check.
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
host = rocm_host(rocm_gfx_target = "gfx1010", rocm_gfx_targets = ["gfx1100", "gfx1010"])
out = _apply_host_overrides(host, override_rocm_gfx = "gfx1010")
assert out.rocm_gfx_target == "gfx1010"
assert out.rocm_gfx_targets == ["gfx1100", "gfx1010"]
def test_forwarded_gfx_never_drops_a_probed_physical_gpu(self, monkeypatch):
# The per-GPU list is the PHYSICAL inventory the Windows auto-Vulkan floor check
# reads, so a forwarded arch the probe never saw must be ADDED, not replace it:
# dropping the probe-confirmed gfx1100 would tell that check no AMD GPU on the box
# reaches the HIP floor when one plainly does.
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
host = rocm_host(rocm_gfx_target = "gfx803", rocm_gfx_targets = ["gfx1100", "gfx803"])
out = _apply_host_overrides(host, override_rocm_gfx = "gfx900")
assert out.rocm_gfx_target == "gfx900"
assert out.rocm_gfx_targets == ["gfx1100", "gfx803", "gfx900"]
def test_forwarded_gfx_not_duplicated_when_already_probed(self, monkeypatch):
# UNSLOTH_ROCM_GFX_ARCH makes the forward win over the probe's visible-device
# pick, so this reaches the same branch; the list must stay deduplicated.
monkeypatch.setenv("UNSLOTH_ROCM_GFX_ARCH", "gfx803")
host = rocm_host(rocm_gfx_target = "gfx1100", rocm_gfx_targets = ["gfx1100", "gfx803"])
out = _apply_host_overrides(host, override_rocm_gfx = "gfx803")
assert out.rocm_gfx_target == "gfx803"
assert out.rocm_gfx_targets == ["gfx1100", "gfx803"]
def test_forwarded_gfx_on_an_unprobed_host_lists_only_itself(self, monkeypatch):
# Negative control for the two above: nothing probed means no inventory to
# preserve, so the driver-only host keeps a single-entry list and auto-Vulkan.
monkeypatch.delenv("UNSLOTH_ROCM_GFX_ARCH", raising = False)
out = _apply_host_overrides(cpu_host(), override_rocm_gfx = "gfx803")
assert out.rocm_gfx_target == "gfx803"
assert out.rocm_gfx_targets == ["gfx803"]
def test_manual_env_override_still_wins_over_probe(self, monkeypatch):
# UNSLOTH_ROCM_GFX_ARCH is the manual escape hatch for hosts whose arch the probes
# get wrong, so it stays authoritative.
monkeypatch.setenv("UNSLOTH_ROCM_GFX_ARCH", "gfx1151")
host = rocm_host(rocm_gfx_target = "gfx1100", rocm_gfx_targets = ["gfx1100"])
out = _apply_host_overrides(host, override_rocm_gfx = "gfx1151")
assert out.rocm_gfx_target == "gfx1151"
assert out.rocm_gfx_targets == ["gfx1100", "gfx1151"]
def test_has_rocm_only_keeps_probe_gfx(self):
out = _apply_host_overrides(cpu_host(), override_has_rocm = True)
assert out.has_rocm is True
assert out.rocm_gfx_target is None
def test_malformed_forwarded_gfx_falls_back_to_has_rocm(self):
out = _apply_host_overrides(cpu_host(), override_has_rocm = True, override_rocm_gfx = "junk")
assert out.has_rocm is True
assert out.rocm_gfx_target is None
def test_no_overrides_leaves_host_unchanged(self):
host = nvidia_host()
assert _apply_host_overrides(host) is host
class TestRocmGfxForwarding:
"""setup.sh / setup.ps1 forward their resolved gfx; the installer accepts it."""
def test_installer_exposes_rocm_gfx_arg(self):
source = _PREBUILT_PATH.read_text(encoding = "utf-8")
assert '"--rocm-gfx"' in source
# Defaults to the env override for standalone runs.
assert 'os.environ.get("UNSLOTH_ROCM_GFX_ARCH")' in source
def test_setup_sh_forwards_rocm_gfx(self):
source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
assert "--rocm-gfx" in source
assert '"$_setup_gfx"' in source
def test_setup_sh_forwards_has_rocm(self):
# If AMD is detected but gfx resolution fails, --has-rocm is still forwarded.
source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
assert "--has-rocm" in source
assert "_setup_amd_detected" in source
def test_setup_ps1_forwards_rocm_gfx(self):
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert "--rocm-gfx" in source
assert "$script:ROCmGfxArch" in source
def test_setup_sh_routes_unconditionally_to_fork(self):
# CPU-only hosts no longer fall back to ggml-org -- the release-repo
# decision is an unconditional fork assignment now. Pin the line text.
source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
assert '_HELPER_RELEASE_REPO="unslothai/llama.cpp"' in source
assert '_HELPER_RELEASE_REPO="ggml-org/llama.cpp"' not in source
def test_setup_ps1_routes_unconditionally_to_fork(self):
# Same on Windows: the fork now ships the windows-cpu / windows-arm64
# bundles, so $HelperReleaseRepo is an unconditional fork assignment.
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
assert '$HelperReleaseRepo = "unslothai/llama.cpp"' in source
assert "$HelperReleaseRepo = if (" not in source
# The text pins above guard the literal. The tests below execute the real routing line
# from setup.sh / setup.ps1 and assert the resolved release repo, so a refactor that
# reintroduces a conditional (or a ggml-org branch) is still caught. Inputs vary
# (CPU-only, inferred/forwarded gfx, usable NVIDIA) to prove no host hits ggml-org.
@staticmethod
def _resolve_setup_sh_repo(
host_machine,
nvidia_usable,
setup_gfx,
rocm_gfx_arch_env = "",
):
"""Run setup.sh's release-repo routing block under bash and return the
resolved _HELPER_RELEASE_REPO. PATH is emptied so any stray tooling probe
misses; routing is unconditional, so the GPU inputs only prove no branch
reroutes a host to ggml-org."""
import shutil
bash = shutil.which("bash")
if bash is None:
pytest.skip("bash not available")
source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
start = source.index('\n_HELPER_RELEASE_REPO="unslothai/llama.cpp"\n') + 1
end = source.index("\n_LLAMA_PR=", start)
block = source[start:end]
assert "_HELPER_RELEASE_REPO" in block, "setup.sh routing anchors not found"
env = {
"PATH": "", # no ROCm tooling discoverable
"ROUTING_BLOCK": block,
"_HOST_SYSTEM": "Linux",
"_HOST_MACHINE": host_machine,
"_setup_nvidia_usable": "true" if nvidia_usable else "false",
"_setup_gfx": setup_gfx,
"UNSLOTH_ROCM_GFX_ARCH": rocm_gfx_arch_env,
}
result = subprocess.run(
[bash, "-c", 'eval "$ROUTING_BLOCK"; printf "%s" "$_HELPER_RELEASE_REPO"'],
capture_output = True,
text = True,
timeout = 30,
env = env,
)
assert result.returncode == 0, result.stderr
return result.stdout.strip()
@pytest.mark.parametrize(
"machine, nvidia_usable, setup_gfx, env_gfx",
[
("x86_64", False, "", ""), # plain CPU host (used to take ggml-org)
("aarch64", False, "", ""), # plain CPU arm64 host (used to take ggml-org)
("x86_64", False, "gfx1100", ""), # name-inferred gfx
("x86_64", False, "", "gfx1100"), # env-forwarded gfx
("x86_64", True, "", ""), # usable NVIDIA
],
)
def test_setup_sh_routing_block_always_resolves_to_fork(
self, machine, nvidia_usable, setup_gfx, env_gfx
):
assert (
self._resolve_setup_sh_repo(
machine, nvidia_usable, setup_gfx, rocm_gfx_arch_env = env_gfx
)
== "unslothai/llama.cpp"
)
@staticmethod
def _resolve_setup_ps1_repo():
"""Run setup.ps1's $HelperReleaseRepo assignment under pwsh and return the
resolved repo. The assignment is unconditional now, so there are no host
inputs to vary."""
import shutil
pwsh = shutil.which("pwsh")
if pwsh is None:
pytest.skip("pwsh not available")
source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
line = next(
(ln for ln in source.splitlines() if ln.strip().startswith("$HelperReleaseRepo =")),
None,
)
assert line is not None, "$HelperReleaseRepo selection not found in setup.ps1"
harness = f"{line}\nWrite-Output $HelperReleaseRepo"
# run_pwsh, not subprocess.run: the returncode assertion just below treats a
# non-zero exit as $HelperReleaseRepo failing to resolve, and a signal-killed pwsh
# would be blamed the same way. See tests/_shared/unsloth_pwsh_runner.py.
result = run_pwsh(
[pwsh, "-NoProfile", "-Command", harness],
capture_output = True,
text = True,
timeout = 60,
)
assert result.returncode == 0, result.stderr
return result.stdout.strip()
def test_setup_ps1_routing_resolves_to_fork(self):
# Windows routing is unconditional now: CPU-only Windows (x64 and arm64)
# uses the fork's windows-cpu / windows-arm64 bundles, not ggml-org.
assert self._resolve_setup_ps1_repo() == "unslothai/llama.cpp"
# TEST: _pick_rocm_gfx_target -- visible-device selection from rocminfo output.
# Honours CUDA/HIP_VISIBLE_DEVICES so a mixed-arch host installs the prebuilt for the
# selected GPU, not GPU 0.
_pick_rocm_gfx_target = prebuilt_mod._pick_rocm_gfx_target
def test_pick_rocm_gfx_target_honors_cuda_visible_devices(monkeypatch):
"""CUDA_VISIBLE_DEVICES=1 must select gfx1100 on a gfx1151 + gfx1100 host (HIP honours CUDA var)."""
# rocminfo reports each token twice (as in the real tool output).
probe_out = "gfx1151\ngfx1151\ngfx1100\ngfx1100"
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "1")
assert _pick_rocm_gfx_target(probe_out) == "gfx1100"
def test_pick_rocm_gfx_target_cuda_visible_devices_minus_one_returns_none(monkeypatch):
"""CUDA_VISIBLE_DEVICES=-1 means no GPU visible; resolver must return None."""
probe_out = "gfx1151\ngfx1100"
monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "-1")
assert _pick_rocm_gfx_target(probe_out) is None
def test_pick_rocm_gfx_target_same_arch_multi_gpu(monkeypatch):
"""Regression: [gfx1100, gfx1100, gfx1151] with HIP_VISIBLE_DEVICES=2 must return gfx1151 (no dict.fromkeys collapse)."""
# rocminfo output for 3 GPUs (2x gfx1100 + 1x gfx1151), one Agent section each.
probe_out = (
"***\nAgent 1\n***\n gfx1100 some info\n gfx1100\n"
"***\nAgent 2\n***\n gfx1100 some info\n gfx1100\n"
"***\nAgent 3\n***\n gfx1151 some info\n gfx1151\n"
)
monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising = False)
monkeypatch.setenv("HIP_VISIBLE_DEVICES", "2")
assert _pick_rocm_gfx_target(probe_out) == "gfx1151"
# TEST: WSL ROCDXG fixes -- drop-in persistence + system-HIP-before-bundle
_INSTALL_SH_PATH = PACKAGE_ROOT / "install.sh"
_LLAMA_CPP_PATH = PACKAGE_ROOT / "studio" / "backend" / "core" / "inference" / "llama_cpp.py"
class TestWslSystemRocmLibDirs:
"""_wsl_system_rocm_lib_dirs: no-op off a ROCDXG WSL host; else returns the system ROCm lib dir for binary_env."""
def test_empty_without_dev_dxg(self):
with patch("os.path.exists", return_value = False):
assert prebuilt_mod._wsl_system_rocm_lib_dirs() == []
def test_empty_on_bare_metal_linux(self):
# /dev/dxg present but /proc/version is not a WSL kernel.
with patch("os.path.exists", lambda p: p == "/dev/dxg"):
with patch(
"builtins.open",
mock_open(read_data = "Linux version 6.8.0-generic"),
):
assert prebuilt_mod._wsl_system_rocm_lib_dirs() == []
def test_returns_system_lib_on_wsl_with_librocdxg(self):
# Normalize separators: os.path.join uses "\" on the Windows test host.
def _exists(p):
p = str(p).replace("\\", "/")
return p in ("/dev/dxg", "/opt/rocm/lib/librocdxg.so")
with patch("os.path.exists", _exists):
with patch(
"builtins.open",
mock_open(read_data = "Linux version 5.15.0-microsoft-standard-WSL2"),
):
assert prebuilt_mod._wsl_system_rocm_lib_dirs() == ["/opt/rocm/lib"]
def test_empty_on_wsl_without_librocdxg(self):
# WSL kernel + /dev/dxg but no librocdxg -> not a ROCDXG ROCm install.
with patch("os.path.exists", lambda p: p == "/dev/dxg"):
with patch(
"builtins.open",
mock_open(read_data = "microsoft-standard-WSL2"),
):
assert prebuilt_mod._wsl_system_rocm_lib_dirs() == []
class TestBinaryEnvWslOrdering:
"""binary_env puts system ROCm lib ahead of the bundle dir + sets HSA_ENABLE_DXG_DETECTION on WSL; no-op bare-metal."""
@staticmethod
def _linux_host():
return HostInfo(
system = "Linux",
machine = "x86_64",
is_windows = False,
is_linux = True,
is_macos = False,
is_x86_64 = True,
is_arm64 = False,
nvidia_smi = None,
driver_cuda_version = None,
compute_caps = [],
visible_cuda_devices = None,
has_physical_nvidia = False,
has_usable_nvidia = False,
has_rocm = True,
)
def test_wsl_prepends_system_rocm_and_sets_hsa(self, tmp_path):
binary = tmp_path / "bundle" / "llama-server"
binary.parent.mkdir(parents = True)
binary.write_text("")
# dedupe_existing_dirs drops non-existent dirs, so use a real dir.
sys_rocm = tmp_path / "sysrocm"
sys_rocm.mkdir()
with patch.object(prebuilt_mod, "_wsl_system_rocm_lib_dirs", return_value = [str(sys_rocm)]):
with patch.dict(os.environ, {}, clear = True):
env = prebuilt_mod.binary_env(binary, tmp_path, self._linux_host())
ld = env["LD_LIBRARY_PATH"].split(os.pathsep)
# Compare resolved paths (dedupe_existing_dirs calls Path.resolve()).
ld_resolved = [str(Path(p).resolve()) for p in ld]
assert ld_resolved[0] == str(sys_rocm.resolve())
assert str(binary.parent.resolve()) in ld_resolved
assert ld_resolved.index(str(sys_rocm.resolve())) < ld_resolved.index(
str(binary.parent.resolve())
)
assert env.get("HSA_ENABLE_DXG_DETECTION") == "1"
def test_bare_metal_linux_unchanged(self, tmp_path):
binary = tmp_path / "bundle" / "llama-server"
binary.parent.mkdir(parents = True)
binary.write_text("")
with patch.object(prebuilt_mod, "_wsl_system_rocm_lib_dirs", return_value = []):
with patch.dict(os.environ, {}, clear = True):
env = prebuilt_mod.binary_env(binary, tmp_path, self._linux_host())
ld = env["LD_LIBRARY_PATH"].split(os.pathsep)
assert ld[0] == str(binary.parent) # bundle dir first, as before
assert "HSA_ENABLE_DXG_DETECTION" not in env
class TestInstallShDropinPersistence:
"""install.sh persists the ROCm-on-WSL drop-in even when rocminfo already enumerates the GPU (reinstall safety)."""
def test_has_persist_helper(self):
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
assert "_persist_rocm_wsl_dropin()" in source
def test_gate5_early_return_persists_dropin(self):
"""The rocminfo-already-works early return must call the persist helper before returning."""
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
# The persist call must precede `return 0` at the rocminfo GPU-agent gate
# (uniquely identified by the `!/generic/` clause the other probes lack).
gate = source.find("Name:[[:space:]]*gfx[1-9]/ && !/generic/")
assert gate != -1
window = source[gate : gate + 900]
assert "_persist_rocm_wsl_dropin" in window
assert window.find("_persist_rocm_wsl_dropin") < window.find("return 0")
def test_persist_helper_gated_on_librocdxg(self):
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
body_start = source.find("_persist_rocm_wsl_dropin()")
body = source[body_start : body_start + 1200]
assert "librocdxg.so" in body
assert "profile.d/unsloth-rocm-wsl.sh" in body
_STRIXHALO_WSL_PATH = PACKAGE_ROOT / "scripts" / "install_rocm_wsl_strixhalo.sh"
class TestWslRerouteNvidiaGuard:
"""_maybe_reroute_strixhalo_to_2404 must skip the AMD reroute on hybrid AMD+NVIDIA hosts by
reusing _has_usable_nvidia_gpu (CUDA_VISIBLE_DEVICES-aware + /proc/driver/nvidia fallback),
which must be defined before the reroute's call site so it is actually available."""
def test_reroute_calls_nvidia_helper_before_amd_signal(self):
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
start = source.find("_maybe_reroute_strixhalo_to_2404()")
assert start != -1
# Slice the WHOLE function body (to its closing brace at column 0), not a
# fixed-length window: preamble growth must not push the signals out of view.
end = source.find("\n}", start)
assert end != -1
body = source[start:end]
nv = body.find("_has_usable_nvidia_gpu")
wmi = body.find("_wsl_amd_gpu_name")
assert nv != -1, "reroute must consult _has_usable_nvidia_gpu before deciding to reroute"
assert wmi != -1
# The NVIDIA guard must precede the AMD/WMI signal and return early.
assert nv < wmi
assert body.find("return 0", nv) < wmi
def test_nvidia_helper_and_deps_defined_before_reroute_callsite(self):
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
call = source.find("\n_maybe_reroute_strixhalo_to_2404 || true")
assert call != -1
for fn in ("_run_bounded() {", "_cvd_hides_nvidia() {", "_has_usable_nvidia_gpu() {"):
idx = source.find(fn)
assert idx != -1 and idx < call, f"{fn} must be defined before the reroute call"
class TestStrixhaloGfxOverridePipefail:
"""The UNSLOTH_WSL_GFX override check must use a consuming grep, not grep -q: under
`set -o pipefail` an early -q exit SIGPIPEs printf and misreports the arch on large output."""
def test_gfx_override_uses_consuming_grep(self):
source = _STRIXHALO_WSL_PATH.read_text(encoding = "utf-8")
idx = source.find('grep -E "Name:[[:space:]]*${GFX}')
assert idx != -1, "GFX override must use a consuming grep -E (not grep -q)"
line = source[idx : source.find("\n", idx)]
assert ">/dev/null" in line
assert 'grep -qE "Name:[[:space:]]*${GFX}' not in source
class TestLlamaCppRuntimeWslOrdering:
"""The serve-time launcher mirrors binary_env: system HIP before the bundle dir on WSL."""
def test_has_wsl_helper(self):
source = _LLAMA_CPP_PATH.read_text(encoding = "utf-8")
assert "_wsl_system_rocm_lib_dirs" in source
def test_prepends_before_binary_dir(self):
source = _LLAMA_CPP_PATH.read_text(encoding = "utf-8")
idx_helper = source.find("lib_dirs.extend(_wsl_system_rocm_lib_dirs())")
idx_binary = source.find("lib_dirs.append(binary_dir)")
assert idx_helper != -1 and idx_binary != -1
assert idx_helper < idx_binary
class TestRocmWslSupplyChainPins:
"""The ROCm-on-WSL bootstrap runs unattended and installs with sudo, so every piece of
code it fetches must come from an immutable ref: a moving branch would turn any rewrite
of that branch into root code execution on affected WSL hosts."""
def test_helper_is_fetched_from_a_pinned_commit_not_a_branch(self):
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
start = source.find("_maybe_bootstrap_rocm_wsl()")
assert start != -1
body = source[start : source.find("\n}", start)]
assert (
"unsloth/main/scripts/install_rocm_wsl_strixhalo.sh" not in body
), "the helper must not be fetched from the mutable main branch"
ref = re.search(r'_ROCM_WSL_HELPER_REF="([0-9a-f]+)"', body)
assert ref is not None, "the helper fetch must pin a full commit SHA"
assert len(ref.group(1)) == 40
assert "${_ROCM_WSL_HELPER_REF}/scripts/install_rocm_wsl_strixhalo.sh" in body
# And whatever that pin (or a user's older checkout) supplies, a helper without the
# pinned-source check is refused rather than trusted.
assert 'grep -q "^UNSLOTH_ROCM_WSL_HELPER_CONTRACT=2$" "$_rw_helper"' in body
assert "predates the pinned-source check" in body
def test_librocdxg_ref_is_pinned_and_verified_before_build(self):
source = _STRIXHALO_WSL_PATH.read_text(encoding = "utf-8")
assert (
"UNSLOTH_LIBROCDXG_REF:-develop" not in source
), "librocdxg must not default to a moving branch"
sha = re.search(r'LIBROCDXG_SHA="\$\{UNSLOTH_LIBROCDXG_SHA:-([0-9a-f]{40})\}"', source)
assert sha is not None, "librocdxg must default to a pinned commit SHA"
# The ref itself must be that commit, so even a helper with no SHA check (an older
# fetched copy) resolves the same immutable revision rather than a movable tag.
ref = re.search(r'LIBROCDXG_REF="\$\{UNSLOTH_LIBROCDXG_REF:-([0-9a-f]{40})\}"', source)
assert ref is not None and ref.group(1) == sha.group(1)
# The SHA check must run before anything from the clone is built or installed.
# A failed checkout must not fall through to the default branch.
assert "_co_failed=1" in source
assert "Refusing to build the repository's default branch instead" in source
idx_check = source.find("git rev-parse HEAD")
idx_cmake = source.find("cmake .. -DWIN_SDK")
idx_install = source.find("$SUDO make install")
assert idx_check != -1 and idx_check < idx_cmake < idx_install
def test_install_sh_forwards_the_same_librocdxg_pin_as_the_helper(self):
# install.sh forwards the pin so it also applies to a helper fetched from an older
# commit. The two copies must never disagree, or a bumped helper would be handed a
# stale pin by install.sh and build the wrong revision.
install_sh = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
helper = _STRIXHALO_WSL_PATH.read_text(encoding = "utf-8")
fwd_ref = re.search(r'_rw_dxg_ref="([0-9a-f]{40})"', install_sh)
fwd_sha = re.search(r'_rw_dxg_sha="\$_rw_dxg_ref"', install_sh)
own_ref = re.search(r'LIBROCDXG_REF="\$\{UNSLOTH_LIBROCDXG_REF:-([0-9a-f]{40})\}"', helper)
own_sha = re.search(r'LIBROCDXG_SHA="\$\{UNSLOTH_LIBROCDXG_SHA:-([0-9a-f]{40})\}"', helper)
assert None not in (fwd_ref, fwd_sha, own_ref, own_sha)
assert fwd_ref.group(1) == own_ref.group(1) == own_sha.group(1)
# A SHA, whether ours or the operator's, is forwarded AS the ref, so even a helper
# that ignores UNSLOTH_LIBROCDXG_SHA resolves exactly the authorised commit.
assert 'if [ -n "$_rw_dxg_sha" ]; then\n _rw_dxg_ref="$_rw_dxg_sha"' in install_sh
# A ref-only override must reach the helper untouched, so that path still works.
assert '_rw_dxg_ref="${UNSLOTH_LIBROCDXG_REF:-}"' in install_sh
assert (
'UNSLOTH_LIBROCDXG_REF="$_rw_dxg_ref" UNSLOTH_LIBROCDXG_SHA="$_rw_dxg_sha"'
in install_sh
)
def test_helper_contract_matches_what_install_sh_requires(self):
# install.sh runs only a helper declaring this contract, so the level it requires and
# the level the helper declares must move together, and the helper must actually
# implement both guarantees the level stands for.
install_sh = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
helper = _STRIXHALO_WSL_PATH.read_text(encoding = "utf-8")
required = re.search(r'grep -q "\^UNSLOTH_ROCM_WSL_HELPER_CONTRACT=(\d+)\$"', install_sh)
declared = re.search(r"^UNSLOTH_ROCM_WSL_HELPER_CONTRACT=(\d+)$", helper, re.M)
assert required is not None and declared is not None
assert required.group(1) == declared.group(1)
assert "git rev-parse HEAD" in helper
assert "_co_failed=1" in helper
if __name__ == "__main__":
pytest.main([__file__, "-v"])