install: tighten pinned torch-index override edge cases

- install.sh: trim whitespace-only UNSLOTH_TORCH_INDEX_URL/_FAMILY before the
  _torch_index_pinned guard, matching get_torch_index_url, so a blank override no
  longer skips the WSL bootstrap and Radeon/Strix reroutes while detection still
  picks the normal index.
- install.sh / install.ps1 / setup.ps1 / install_python_stack.py: force the torch
  2.11 floor only for the gfx families with the <2.11 _grouped_mm bug (gfx120X-all,
  gfx1151, gfx1150). A pinned override to gfx110X-all/gfx90a/gfx908 stays on the
  default range, matching the automatic AMD path.
- install_python_stack.py _ensure_cuda_torch: treat an untagged CUDA build under a
  CUDA pin as a family mismatch (reinstall), and match cuXXX pins narrowly (cu +
  digits) so a custom/current mirror leaf no longer forces CUDA over a CPU/ROCm venv.
- install_python_stack.py _ensure_rocm_torch: reinstall when an explicit ROCm pin
  names a different ROCm family than the already-installed ROCm torch (the ROCm
  analogue of the CUDA cuXXX mismatch repair).

Adds tests for each case.
This commit is contained in:
Daniel Han 2026-07-05 23:36:24 +00:00
commit 7814e2c261
6 changed files with 332 additions and 35 deletions

View file

@ -2115,15 +2115,23 @@ exit 0
if ($_pinLeaf -match '^rocm(\d+)\.(\d+)') {
$_pinRocm211 = ([int]$Matches[1] -gt 7) -or ([int]$Matches[1] -eq 7 -and [int]$Matches[2] -ge 2)
}
if ($_pinLeaf -like 'gfx*' -or $_pinRocm211) {
# Only the gfx families the AMD arch map above pins to torch 2.11 need the
# floor here (gfx120X-all, gfx1151, gfx1150 -- the _grouped_mm bug arches).
# Other per-arch indexes (gfx110X-all, gfx90a, gfx908) publish <2.11 wheels
# and the automatic path leaves them bare, so an override to one of those
# must NOT force a 2.11 floor the normal path intentionally avoids (which
# would fail resolution or drag in mismatched companions).
$_pinGfx211 = @('gfx120x-all', 'gfx1151', 'gfx1150') -contains $_pinLeaf
if ($_pinGfx211 -or $_pinRocm211) {
$ROCmIndexUrl = $TorchIndexUrl
$ROCmTorchFloor = "torch>=2.11.0,<2.12.0"
$PinnedRocmVisionSpec = "torchvision>=0.26.0,<0.27.0"
$PinnedRocmAudioSpec = "torchaudio>=2.11.0,<2.12.0"
substep "pinned ROCm index ($_pinLeaf) -- enforcing $ROCmTorchFloor" "Cyan"
} elseif ($_pinLeaf -like 'rocm*') {
# Older rocm (<=7.1) ships torch <2.11; route via the ROCm path with the
# default floor so the pinned family resolves its own wheels.
} elseif ($_pinLeaf -like 'gfx*' -or $_pinLeaf -like 'rocm*') {
# Other gfx per-arch indexes and older rocm (<=7.1) ship torch <2.11;
# route via the ROCm path with bare specs so the pinned family resolves
# its own wheels (matches the automatic path's bare floor for these).
$ROCmIndexUrl = $TorchIndexUrl
}
}

View file

@ -2451,8 +2451,18 @@ _maybe_bootstrap_rocm_wsl() {
# sudo + large downloads after probing /dev/dxg) and the Radeon/Strix rerouting
# below (which would re-probe the GPU and overwrite the pinned URL). A headless /
# container / CI build must get exactly the index it asked for.
# Trim leading/trailing whitespace before the -n check (parity with
# get_torch_index_url): a whitespace-only override (e.g. from a CI template) is
# treated as unset there and falls through to normal CPU/ROCm detection, so it
# must NOT flip _torch_index_pinned true here -- otherwise the WSL bootstrap and
# Radeon/Strix reroutes below are skipped while detection still picks the normal
# index, silently disabling them.
_torch_index_pinned=false
if [ -n "${UNSLOTH_TORCH_INDEX_URL:-}" ] || [ -n "${UNSLOTH_TORCH_INDEX_FAMILY:-}" ]; then
_ti_url_trim="${UNSLOTH_TORCH_INDEX_URL:-}"
_ti_url_trim="${_ti_url_trim#"${_ti_url_trim%%[![:space:]]*}"}"; _ti_url_trim="${_ti_url_trim%"${_ti_url_trim##*[![:space:]]}"}"
_ti_family_trim="${UNSLOTH_TORCH_INDEX_FAMILY:-}"
_ti_family_trim="${_ti_family_trim#"${_ti_family_trim%%[![:space:]]*}"}"; _ti_family_trim="${_ti_family_trim%"${_ti_family_trim##*[![:space:]]}"}"
if [ -n "$_ti_url_trim" ] || [ -n "$_ti_family_trim" ]; then
_torch_index_pinned=true
fi
[ "$_torch_index_pinned" = true ] || _maybe_bootstrap_rocm_wsl || true
@ -2474,21 +2484,25 @@ case "$_torch_index_leaf" in
*) export UNSLOTH_TORCH_BACKEND="cuda" ;;
esac
# rocm7.2 and the AMD per-gfx indexes (repo.amd.com/.../gfxNNNN) ship torch
# 2.11.0 -- adjust the constraint to allow it. This also covers a pinned full-URL
# or family override (e.g. UNSLOTH_TORCH_INDEX_URL=.../gfx1151) that returns early
# above and so never hits the Strix reroute that otherwise raises this constraint.
# Pin the companions to the matching 2.11 range too: the per-gfx index publishes
# torchvision/torchaudio independently and a bare name can resolve a 2.12-built
# wheel (ABI mismatch). Matches setup.ps1's *FloorMap and _ROCM_TORCH_PKG_SPECS.
# All other ROCm tags and CUDA stay within <2.11.0.
# rocm7.2 and the AMD per-gfx indexes with the torch._C._grouped_mm bug on <2.11
# (repo.amd.com/.../gfx120X-all, gfx1151, gfx1150) ship torch 2.11.0 -- raise the
# constraint to allow it. This also covers a pinned full-URL or family override
# (e.g. UNSLOTH_TORCH_INDEX_URL=.../gfx1151) that returns early above and so never
# hits the Strix reroute that otherwise raises this constraint. Pin the companions
# to the matching 2.11 range too: the per-gfx index publishes torchvision/torchaudio
# independently and a bare name can resolve a 2.12-built wheel (ABI mismatch).
# Matches setup.ps1's *FloorMap and _ROCM_GFX_TORCH211_LEAVES in install_python_stack.py.
# Only these gfx families need the floor: the other per-arch indexes (gfx110X-all,
# gfx90a, gfx908) publish <2.11 wheels, so a pinned override to one of those must NOT
# be forced onto the 2.11 line (which would fail resolution / pull mismatched
# companions). All other ROCm tags and CUDA stay within <2.11.0.
# Match on the FINAL path segment ($_torch_index_leaf, computed above), NOT the
# whole URL: a custom UNSLOTH_PYTORCH_MIRROR whose base path contains a "gfx" or
# "rocm7.2" segment (e.g. https://mirror.local/gfx-cache) with a cu*/cpu family
# must not be treated as an AMD per-arch index and pushed to the 2.11 line. This
# mirrors the leaf-only backend classification just above.
case "$_torch_index_leaf" in
rocm7.2|gfx*)
rocm7.2|gfx120x-all|gfx1151|gfx1150)
TORCH_CONSTRAINT="torch>=2.11.0,<2.12.0"
TORCHVISION_CONSTRAINT="torchvision>=0.26.0,<0.27.0"
TORCHAUDIO_CONSTRAINT="torchaudio>=2.11.0,<2.12.0"

View file

@ -74,6 +74,13 @@ _ROCM_TORCH_INDEX: dict[tuple[int, int], str] = {
(6, 0): "rocm6.0",
}
# AMD per-arch index leaves that need the torch 2.11 floor (the torch._C._grouped_mm
# null-ptr bug lives in the <2.11 wheels for these arches). Mirrors the gfx keys in
# _WINDOWS_ROCM_TORCH_PKG_SPECS and the *FloorMap sets in install.ps1 / setup.ps1.
# Other per-arch indexes (gfx110X-all, gfx90a, gfx908) publish <2.11 wheels and must
# stay bare, so an override to one of them must NOT be forced onto the 2.11 line.
_ROCM_GFX_TORCH211_LEAVES: frozenset[str] = frozenset({"gfx120x-all", "gfx1151", "gfx1150"})
# Per-tag pip specs; rocm7.2 ships torch 2.11.0 (older tags cap at 2.10.x).
_ROCM_TORCH_PKG_SPECS: dict[str, tuple[str, str, str]] = {
"rocm7.2": (
@ -1080,6 +1087,41 @@ def _explicit_rocm_torch_index_url() -> "str | None":
return url if leaf.startswith(("rocm", "gfx")) else None
def _rocm_pin_family_mismatch(pin_url: str, installed_ver: str) -> bool:
"""True when an explicit ROCm pin names a different ROCm family than the
already-installed ROCm torch, so the pin needs a reinstall to be applied.
Mirrors setup.ps1's stale-venv ROCm comparison:
- both +rocmX.Y versions readable -> compare them exactly
- gfx pin, or an unreadable installed version -> compare on the torch 2.11
line (gfx*/rocm>=7.2 serve 2.11+, older ROCm does not)
A pin that resolves to the same family as what is installed is NOT a mismatch,
so a correct ROCm venv is never needlessly reinstalled. Pure function.
"""
leaf = pin_url.rstrip("/").rsplit("/", 1)[-1].lower()
# Pinned ROCm version (from a rocmX.Y leaf) and whether the pin serves 2.11+.
_pin_rocm = re.match(r"^rocm(\d+)\.(\d+)", leaf)
_pin_ver = (int(_pin_rocm.group(1)), int(_pin_rocm.group(2))) if _pin_rocm else None
if leaf.startswith("gfx"):
_pin_is_211 = True
elif _pin_ver is not None:
_pin_is_211 = _pin_ver >= (7, 2)
else:
_pin_is_211 = False
# Installed ROCm version (+rocmX.Y) and whether the installed torch is 2.11+.
_inst_rocm = re.search(r"\+rocm(\d+)\.(\d+)", installed_ver)
_inst_ver = (int(_inst_rocm.group(1)), int(_inst_rocm.group(2))) if _inst_rocm else None
_inst_rel = re.match(r"^(\d+)\.(\d+)", installed_ver)
_inst_is_211 = (
(int(_inst_rel.group(1)), int(_inst_rel.group(2))) >= (2, 11) if _inst_rel else False
)
if _pin_ver is not None and _inst_ver is not None:
# Both ROCm versions readable: exact comparison.
return _pin_ver != _inst_ver
# gfx pin or unreadable version: compare on the torch 2.11 line.
return _pin_is_211 != _inst_is_211
def _explicit_cpu_torch_index_url() -> "str | None":
"""The pinned wheel index URL when it names the CPU family (leaf == cpu), else None.
@ -1093,19 +1135,32 @@ def _explicit_cpu_torch_index_url() -> "str | None":
return url if leaf == "cpu" else None
def _is_cuda_family_leaf(leaf: str) -> bool:
"""True only for a real CUDA wheel-family leaf: "cu" followed by digits
(cu118, cu126, cu128, cu130, ...).
A bare startswith("cu") wrongly matches arbitrary mirror leaves like "custom"
or "current", which would let _ensure_cuda_torch treat a generic mirror pin as
CUDA authority and force a CUDA reinstall over a CPU/ROCm venv on a non-NVIDIA
host -- exactly what _explicit_cuda_torch_index_url's contract forbids.
"""
return re.match(r"^cu[0-9]", leaf) is not None
def _explicit_cuda_torch_index_url() -> "str | None":
"""The pinned wheel index URL when it names a CUDA family (leaf cu*), else None.
"""The pinned wheel index URL when it names a CUDA family (leaf cuXXX), else None.
Mirrors _explicit_rocm/cpu_torch_index_url so _ensure_cuda_torch only treats a
*CUDA* pin as authority to override the NVIDIA-presence gate. An arbitrary
mirror URL (or a ROCm/CPU pin) must not force a CUDA reinstall over a working
ROCm/CPU venv on a non-NVIDIA host.
ROCm/CPU venv on a non-NVIDIA host, so match cuXXX (cu + digits) narrowly
rather than any leaf starting with "cu" (which would catch custom/current).
"""
url = _explicit_torch_index_url()
if url is None:
return None
leaf = url.rstrip("/").rsplit("/", 1)[-1].lower()
return url if leaf.startswith("cu") else None
return url if _is_cuda_family_leaf(leaf) else None
def _ensure_cuda_torch() -> None:
@ -1196,13 +1251,20 @@ def _ensure_cuda_torch() -> None:
# with no CUDA pin, is deliberate and left alone.
_pin = _explicit_torch_index_url()
_pin_leaf = _pin.rstrip("/").rsplit("/", 1)[-1].lower() if _pin else ""
_pinned_cuda = _pin_leaf.startswith("cu")
_pinned_cuda = _is_cuda_family_leaf(_pin_leaf)
if _marker == "hip":
_why = "torch is a ROCm build on an NVIDIA host"
elif _marker == "cpu" and _pinned_cuda:
_why = "torch is a CPU build but an explicit CUDA index is pinned"
elif _marker == "cuda" and _pinned_cuda and _installed_cu and _installed_cu != _pin_leaf:
_why = f"torch is {_installed_cu} but the pinned CUDA index is {_pin_leaf}"
elif _marker == "cuda" and _pinned_cuda and _installed_cu != _pin_leaf:
# Mismatch when the installed cuXXX differs from the pin. An UNTAGGED cuda
# build (empty _installed_cu -- e.g. torch re-resolved from default PyPI
# into a CUDA wheel with no +cuXXX local tag) also counts: the family
# cannot be confirmed to match the pin, so reinstall to enforce it. The
# reinstall targets the pinned family, so an already-matching untagged
# build simply re-lands on the same family (idempotent).
_installed_desc = _installed_cu if _installed_cu else "an untagged CUDA build"
_why = f"torch is {_installed_desc} but the pinned CUDA index is {_pin_leaf}"
else:
return # healthy CUDA torch matching the pin, or a deliberate CPU wheel
@ -1458,9 +1520,13 @@ def _ensure_rocm_torch() -> None:
# pins anyway) and any ver comparisons stay well-defined.
ver = (0, 0)
# Probe whether torch already links against HIP (ROCm already working).
# Probe whether torch already links against HIP (ROCm already working), and
# capture the installed ROCm build tag so a pin mismatch can be detected.
# Do NOT skip for CUDA-only builds: they are unusable on AMD-only hosts
# (the NVIDIA check above already handled mixed AMD+NVIDIA setups).
# Line 1: the HIP presence marker (HIP version, "rocm" sentinel, or "").
# Line 2: the installed wheel version string (e.g. "2.10.0+rocm6.4"), used to
# compare the installed ROCm family against an explicit pin below.
try:
probe = subprocess.run(
[
@ -1473,7 +1539,8 @@ def _ensure_rocm_torch() -> None:
# Print the HIP version when present (back-compat), else a
# "rocm" sentinel when only torch.__version__ flags ROCm
# (AMD SDK / Radeon wheels). Empty string = CPU/CUDA.
"print(hip if hip else ('rocm' if 'rocm' in ver else ''))"
"print(hip if hip else ('rocm' if 'rocm' in ver else '')); "
"print(ver)"
),
],
stdout = subprocess.PIPE,
@ -1482,11 +1549,28 @@ def _ensure_rocm_torch() -> None:
)
except (OSError, subprocess.TimeoutExpired):
probe = None
has_hip_torch = (
probe is not None and probe.returncode == 0 and probe.stdout.decode().strip() != ""
_probe_lines = (
[ln.strip() for ln in probe.stdout.decode(errors = "replace").splitlines() if ln.strip()]
if (probe is not None and probe.returncode == 0)
else []
)
has_hip_torch = bool(_probe_lines) and _probe_lines[0] != ""
_installed_torch_ver = _probe_lines[1] if len(_probe_lines) > 1 else ""
rocm_torch_ready = has_hip_torch
# An explicit ROCm pin whose family differs from the already-installed ROCm
# torch must reinstall, mirroring _ensure_cuda_torch (installed cuXXX != pin).
# Without this, `studio update` with UNSLOTH_TORCH_INDEX_FAMILY=rocm7.2 (or a
# gfx* URL) on a venv that already carries an OLDER ROCm build (+rocm6.4 /
# +rocm7.1) short-circuits on has_hip_torch and never applies the override.
# Compare exact +rocmX.Y versions when both are readable; otherwise (gfx pin,
# or an unreadable installed version) fall back to the torch 2.11 line, which
# is what distinguishes the gfx/rocm>=7.2 wheels from older ROCm. Matches the
# stale-venv comparison in setup.ps1.
_rocm_pin_mismatch = False
if has_hip_torch and _rocm_pin is not None:
_rocm_pin_mismatch = _rocm_pin_family_mismatch(_rocm_pin, _installed_torch_ver)
rocm_torch_ready = has_hip_torch and not _rocm_pin_mismatch
# Strix Halo / Strix Point (gfx1151 / gfx1150) segfault under ROCm 7.1
# in torch._grouped_mm. AMD's per-gfx repo ships torch 2.11.0+rocm7.13.0
@ -1561,7 +1645,10 @@ def _ensure_rocm_torch() -> None:
constrain = False,
)
rocm_torch_ready = True
elif not has_hip_torch:
elif not has_hip_torch or _rocm_pin_mismatch:
# Reinstall when torch is not ROCm yet, OR when a ROCm build is present but
# its family differs from an explicit pin (_rocm_pin_mismatch -- the ROCm
# analogue of _ensure_cuda_torch's installed-cuXXX != pin reinstall).
# Honour an explicit ROCm wheel-index pin verbatim instead of re-detecting
# the host ROCm version; otherwise select the best wheel tag (newest ROCm
# version <= installed). gfx*/rocm7.2 indexes serve torch 2.11+, so match
@ -1585,8 +1672,15 @@ def _ensure_rocm_torch() -> None:
if _override_idx is None:
index_url = f"{_PYTORCH_WHL_BASE}/{tag}"
print(f" ROCm torch -- installing from {index_url}")
if tag.startswith("gfx"):
# Only the gfx arches with the _grouped_mm bug (gfx120X-all, gfx1151,
# gfx1150) need the torch 2.11 spec; other gfx per-arch indexes
# (gfx110X-all, gfx90a, gfx908) publish <2.11 wheels, so a pinned
# override to one of those stays on the default range. Matches the
# gfx floor gating in install.ps1 / setup.ps1.
if tag in _ROCM_GFX_TORCH211_LEAVES:
_torch_pkg, _vision_pkg, _audio_pkg = _ROCM_TORCH_PKG_SPECS["rocm7.2"]
elif tag.startswith("gfx"):
_torch_pkg, _vision_pkg, _audio_pkg = _ROCM_TORCH_PKG_SPECS["_default"]
else:
_torch_pkg, _vision_pkg, _audio_pkg = _ROCM_TORCH_PKG_SPECS.get(
tag, _ROCM_TORCH_PKG_SPECS["_default"]

View file

@ -2924,13 +2924,22 @@ if ($TorchIndexPinned -and -not $ROCmIndexUrl -and $PinnedTorchIndexUrl) {
if ($_pinLeaf -match '^rocm(\d+)\.(\d+)') {
$_pinRocm211 = ([int]$Matches[1] -gt 7) -or ([int]$Matches[1] -eq 7 -and [int]$Matches[2] -ge 2)
}
if ($_pinLeaf -like 'gfx*' -or $_pinRocm211) {
# Only the gfx families the AMD arch map above pins to torch 2.11 need the
# floor here (gfx120X-all, gfx1151, gfx1150 -- the _grouped_mm bug arches).
# Other per-arch indexes (gfx110X-all, gfx90a, gfx908) publish <2.11 wheels
# and the automatic path leaves them bare, so an override to one of those
# must NOT force a 2.11 floor the normal path intentionally avoids.
$_pinGfx211 = @('gfx120x-all', 'gfx1151', 'gfx1150') -contains $_pinLeaf
if ($_pinGfx211 -or $_pinRocm211) {
$ROCmIndexUrl = $PinnedTorchIndexUrl
$ROCmTorchSpec = "torch>=2.11.0,<2.12.0"
$ROCmVisionSpec = "torchvision>=0.26.0,<0.27.0"
$ROCmAudioSpec = "torchaudio>=2.11.0,<2.12.0"
substep "pinned ROCm index ($_pinLeaf) -- enforcing $ROCmTorchSpec" "Cyan"
} elseif ($_pinLeaf -like 'rocm*') {
} elseif ($_pinLeaf -like 'gfx*' -or $_pinLeaf -like 'rocm*') {
# Other gfx per-arch indexes and older rocm (<=7.1) ship torch <2.11;
# route via the ROCm path with bare specs (matches the automatic path's
# bare floor for these arches).
$ROCmIndexUrl = $PinnedTorchIndexUrl
$ROCmTorchSpec = "torch"
$ROCmVisionSpec = "torchvision"

View file

@ -65,11 +65,13 @@ def _run_cuda_repair(
smi_path = "/usr/bin/nvidia-smi",
cvd = None,
index_family = None,
index_url = None,
):
"""Invoke _ensure_cuda_torch under a fully mocked host; return the pip mock.
cvd controls CUDA_VISIBLE_DEVICES: None removes it from the env, any string sets it.
index_family sets UNSLOTH_TORCH_INDEX_FAMILY (the explicit wheel-index pin)."""
index_family sets UNSLOTH_TORCH_INDEX_FAMILY (the explicit wheel-index pin).
index_url sets UNSLOTH_TORCH_INDEX_URL (the full-URL pin form)."""
env = {}
if rocm_marker:
env["UNSLOTH_ROCM_TORCH_INSTALLED"] = "1"
@ -77,6 +79,8 @@ def _run_cuda_repair(
env["CUDA_VISIBLE_DEVICES"] = cvd
if index_family is not None:
env["UNSLOTH_TORCH_INDEX_FAMILY"] = index_family
if index_url is not None:
env["UNSLOTH_TORCH_INDEX_URL"] = index_url
def _which(name, *a, **k):
if name == "nvidia-smi":
@ -105,6 +109,7 @@ def _run_cuda_repair(
stack_mod.os.environ.pop("CUDA_VISIBLE_DEVICES", None)
if index_family is None:
stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
if index_url is None:
stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_URL", None)
_ensure_cuda_torch()
return mock_pip
@ -164,6 +169,28 @@ class TestCudaRepairFires:
assert mock_pip.call_count == 1
assert "cu128" in _index_url(mock_pip)
def test_tagged_cuda_mismatch_repairs(self):
# A healthy CUDA torch whose +cuXXX differs from the pin is repaired.
mock_pip = _run_cuda_repair(
index_family = "cu128",
torch_state = "cuda|cu126",
cuda_version = "12.8",
)
assert mock_pip.call_count == 1
assert "cu128" in _index_url(mock_pip)
def test_untagged_cuda_build_under_pin_repairs(self):
# An untagged CUDA build (torch re-resolved from default PyPI, no +cuXXX
# local tag -> empty installed cu) cannot be confirmed to match the pin,
# so the pin is enforced with a reinstall to the pinned family.
mock_pip = _run_cuda_repair(
index_family = "cu128",
torch_state = "cuda", # marker cuda, empty installed cu
cuda_version = "12.8",
)
assert mock_pip.call_count == 1
assert "cu128" in _index_url(mock_pip)
# No-op cases.
@ -227,6 +254,47 @@ class TestCudaRepairSkips:
mock_pip = _run_cuda_repair(cvd = "0", torch_state = "hip")
assert mock_pip.call_count == 1
def test_matching_tagged_cuda_pin_no_repair(self):
# Healthy CUDA torch whose +cuXXX already matches the pin: no reinstall.
mock_pip = _run_cuda_repair(
index_family = "cu128",
torch_state = "cuda|cu128",
cuda_version = "12.8",
)
mock_pip.assert_not_called()
def test_custom_mirror_leaf_not_treated_as_cuda_pin(self):
# A generic mirror URL whose leaf starts with "cu" but is not cuXXX
# (e.g. .../custom, .../current) must NOT be treated as a CUDA pin, so it
# cannot bypass the NVIDIA gate to force CUDA over a CPU/ROCm venv.
for _leaf in ("custom", "current"):
mock_pip = _run_cuda_repair(
nvidia = False,
backend = "cuda",
index_url = f"https://mymirror.example/{_leaf}",
torch_state = "hip",
)
mock_pip.assert_not_called()
def test_explicit_cuda_family_leaf_helper(self):
# _explicit_cuda_torch_index_url matches cuXXX narrowly, not any cu* leaf.
import contextlib
def _with(url):
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 stack_mod._explicit_cuda_torch_index_url()
assert _with("https://download.pytorch.org/whl/cu128") is not None
assert _with("https://download.pytorch.org/whl/cu126") is not None
assert _with("https://mymirror.example/custom") is None
assert _with("https://mymirror.example/current") is None
assert _with("https://download.pytorch.org/whl/cpu") is None
with contextlib.suppress(Exception):
stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_URL", None)
# CUDA index ladder.

View file

@ -710,6 +710,105 @@ class TestEnsureRocmTorch:
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_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 = b"6.4.12345\n2.10.0+rocm6.4\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."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = b"6.4.12345\n2.10.0+rocm6.4\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", side_effect = AssertionError
):
_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 = b"7.2.12345\n2.11.0+rocm7.2\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
)
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"
# 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
# gfx pin (2.11 line) vs installed release line.
assert f("https://repo.amd.com/rocm/whl/gfx1151", "2.10.0+rocm6.4") is True
assert f("https://repo.amd.com/rocm/whl/gfx1151", "2.11.0+rocm7.13.0") is False
# rocm7.2 pin vs unreadable installed rocm version -> compare on 2.11 line.
assert f(f"{base}/rocm7.2", "2.10.0") is True
assert f(f"{base}/rocm7.2", "2.11.0") 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")
@ -2887,10 +2986,15 @@ class TestStrixRocm71Override:
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
# The 2.11 constraint block must switch on $_torch_index_leaf, not on the
# full $TORCH_INDEX_URL (which the earlier, buggy version matched with
# */gfx* and would false-positive on a mirror base path).
assert 'case "$_torch_index_leaf" in\n rocm7.2|gfx*)' in source, (
"the torch>=2.11 constraint must match the index leaf (rocm7.2|gfx*), "
"not the whole URL"
# */gfx* and would false-positive on a mirror base path). Only the gfx
# families with the <2.11 _grouped_mm bug (gfx120X-all / gfx1151 / gfx1150)
# are pushed to 2.11 -- a bare gfx* would also floor gfx110X-all/gfx90a/
# gfx908, which the automatic AMD path intentionally leaves bare.
assert (
'case "$_torch_index_leaf" in\n rocm7.2|gfx120x-all|gfx1151|gfx1150)' in source
), (
"the torch>=2.11 constraint must match the specific gfx leaves that need "
"it (rocm7.2|gfx120x-all|gfx1151|gfx1150), not a bare gfx* or the whole URL"
)
def test_amd_rocm_mirror_env_var_respected(self):