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

@ -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"]