Merge remote-tracking branch 'origin/main' into feature/rag
# Conflicts: # studio/install_python_stack.py
This commit is contained in:
commit
2af60c5480
37 changed files with 6854 additions and 375 deletions
|
|
@ -12,8 +12,10 @@ PATH to point at the venv.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import glob
|
||||
import os
|
||||
import platform
|
||||
import re
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
|
|
@ -54,12 +56,9 @@ PLATFORM_LACKS_TORCHCODEC_WHEEL = (
|
|||
# ── ROCm / AMD GPU support ─────────────────────────────────────────────────────
|
||||
# Mapping from detected ROCm (major, minor) to the best PyTorch wheel tag on
|
||||
# download.pytorch.org. Entries are checked newest-first (>=).
|
||||
# ROCm 7.2 only has torch 2.11.0 on download.pytorch.org, which exceeds the
|
||||
# current torch upper bound (<2.11.0). Fall back to rocm7.1 (torch 2.10.0).
|
||||
# TODO: uncomment rocm7.2 when torch upper bound is bumped to >=2.11.0
|
||||
_ROCM_TORCH_INDEX: dict[tuple[int, int], str] = {
|
||||
# (7, 2): "rocm7.2", # torch 2.11.0 -- requires torch>=2.11
|
||||
(7, 1): "rocm7.1",
|
||||
(7, 2): "rocm7.2", # torch 2.11.0
|
||||
(7, 1): "rocm7.1", # torch 2.10.0
|
||||
(7, 0): "rocm7.0",
|
||||
(6, 4): "rocm6.4",
|
||||
(6, 3): "rocm6.3",
|
||||
|
|
@ -67,10 +66,47 @@ _ROCM_TORCH_INDEX: dict[tuple[int, int], str] = {
|
|||
(6, 1): "rocm6.1",
|
||||
(6, 0): "rocm6.0",
|
||||
}
|
||||
|
||||
# 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": (
|
||||
"torch>=2.11.0,<2.12.0",
|
||||
"torchvision>=0.26.0,<0.27.0",
|
||||
"torchaudio>=2.11.0,<2.12.0",
|
||||
),
|
||||
# Default for rocm7.1 and earlier: torch 2.x below 2.11
|
||||
"_default": (
|
||||
"torch>=2.4,<2.11.0",
|
||||
"torchvision>=0.19,<0.26.0",
|
||||
"torchaudio>=2.4,<2.11.0",
|
||||
),
|
||||
}
|
||||
_PYTORCH_WHL_BASE = (
|
||||
os.environ.get("UNSLOTH_PYTORCH_MIRROR") or "https://download.pytorch.org/whl"
|
||||
).rstrip("/")
|
||||
|
||||
# AMD Windows ROCm wheels — repo.amd.com (arch-specific pip index)
|
||||
# Format: https://repo.amd.com/rocm/whl/{arch_family}/
|
||||
# Override with UNSLOTH_ROCM_WINDOWS_MIRROR for air-gapped / mirror installs.
|
||||
_ROCM_WINDOWS_INDEX_BASE = (
|
||||
os.environ.get("UNSLOTH_ROCM_WINDOWS_MIRROR") or "https://repo.amd.com/rocm/whl"
|
||||
).rstrip("/")
|
||||
|
||||
# Maps gfx arch → AMD index arch-family suffix.
|
||||
# Each family is a separate pip index on repo.amd.com.
|
||||
_GFX_TO_AMD_INDEX_ARCH: dict[str, str] = {
|
||||
"gfx1201": "gfx120X-all",
|
||||
"gfx1200": "gfx120X-all", # RDNA 4
|
||||
"gfx1151": "gfx1151",
|
||||
"gfx1150": "gfx1150", # RDNA 3.5 (Strix Halo/Point)
|
||||
"gfx1103": "gfx110X-all",
|
||||
"gfx1102": "gfx110X-all", # RDNA 3
|
||||
"gfx1101": "gfx110X-all",
|
||||
"gfx1100": "gfx110X-all",
|
||||
"gfx90a": "gfx90a",
|
||||
"gfx908": "gfx908", # MI200/MI100
|
||||
}
|
||||
|
||||
# bitsandbytes continuous-release_main wheels with the ROCm 4-bit GEMV fix
|
||||
# (bnb PR #1887, post-0.49.2). bnb <= 0.49.2 NaNs at decode shape on every
|
||||
# AMD GPU. Drop the pin once bnb 0.50+ ships on PyPI.
|
||||
|
|
@ -85,6 +121,16 @@ _BNB_ROCM_PRERELEASE_URLS: dict[str, str] = {
|
|||
"download/continuous-release_main/"
|
||||
"bitsandbytes-1.33.7.preview-py3-none-manylinux_2_24_aarch64.whl"
|
||||
),
|
||||
# Windows ROCm wheel — ships libbitsandbytes_rocm{VER}.dll.
|
||||
# BNB auto-detects HIP version from torch.version.hip, which does not always
|
||||
# match the DLL suffix in this prerelease wheel (e.g. torch 7.13 with a rocm72
|
||||
# DLL). We scan the installed wheel for the actual DLL name and set
|
||||
# BNB_ROCM_VERSION accordingly in _install_bnb_windows_rocm() and worker.py.
|
||||
"win_amd64": (
|
||||
"https://github.com/bitsandbytes-foundation/bitsandbytes/releases/"
|
||||
"download/continuous-release_main/"
|
||||
"bitsandbytes-1.33.7.preview-py3-none-win_amd64.whl"
|
||||
),
|
||||
}
|
||||
_BNB_ROCM_PYPI_FALLBACK = "bitsandbytes>=0.49.1"
|
||||
|
||||
|
|
@ -165,8 +211,6 @@ def _detect_rocm_version() -> tuple[int, int] | None:
|
|||
# for the rocm-core package version. Matches the chain in
|
||||
# install.sh::get_torch_index_url so `unsloth studio update` behaves
|
||||
# the same as a fresh `curl | sh` install.
|
||||
import re as _re_pkg
|
||||
|
||||
for cmd in (
|
||||
["dpkg-query", "-W", "-f=${Version}\n", "rocm-core"],
|
||||
["rpm", "-q", "--qf", "%{VERSION}\n", "rocm-core"],
|
||||
|
|
@ -188,18 +232,157 @@ def _detect_rocm_version() -> tuple[int, int] | None:
|
|||
continue
|
||||
raw = result.stdout.strip()
|
||||
# dpkg can prepend an epoch ("1:6.3.0-1"); strip it before parsing.
|
||||
raw = _re_pkg.sub(r"^\d+:", "", raw)
|
||||
m = _re_pkg.match(r"(\d+)[.-](\d+)", raw)
|
||||
raw = re.sub(r"^\d+:", "", raw)
|
||||
m = re.match(r"(\d+)[.-](\d+)", raw)
|
||||
if m:
|
||||
return int(m.group(1)), int(m.group(2))
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _pick_visible_index(num_tokens: int) -> int:
|
||||
"""Resolve HIP_VISIBLE_DEVICES / ROCR_VISIBLE_DEVICES to an integer
|
||||
index into a list of length num_tokens. Returns 0 (first GPU) for
|
||||
unset, empty, '-1', UUID-style, or out-of-range values."""
|
||||
for _env in ("HIP_VISIBLE_DEVICES", "ROCR_VISIBLE_DEVICES"):
|
||||
_val = os.environ.get(_env)
|
||||
if _val is None:
|
||||
continue
|
||||
_val = _val.strip()
|
||||
if _val == "" or _val == "-1":
|
||||
return 0
|
||||
_first = _val.split(",")[0].strip()
|
||||
try:
|
||||
_idx = int(_first)
|
||||
if 0 <= _idx < num_tokens:
|
||||
return _idx
|
||||
except ValueError:
|
||||
pass
|
||||
return 0
|
||||
return 0
|
||||
|
||||
|
||||
def _detect_windows_gfx_arch() -> str | None:
|
||||
"""Return the gcnArchName on Windows (e.g. 'gfx1200'), or None.
|
||||
|
||||
Probe order matches the PowerShell installer: env-var override first,
|
||||
then hipinfo (PATH or HIP_PATH / ROCM_PATH bin), then amd-smi. Without
|
||||
the amd-smi fallback, runtime-only AMD installs without hipinfo on PATH
|
||||
return early and `studio update` cannot repair a CPU-only venv.
|
||||
|
||||
On multi-GPU hosts, all detected gfx tokens are deduplicated (preserving
|
||||
enumeration order) and HIP_VISIBLE_DEVICES / ROCR_VISIBLE_DEVICES selects
|
||||
which one to install for. The first GPU is used when no env var is set.
|
||||
"""
|
||||
# 1. Explicit override (matches PowerShell installer's env-var path).
|
||||
_override = os.environ.get("UNSLOTH_ROCM_GFX_ARCH")
|
||||
if _override and _override.strip():
|
||||
return _override.strip().lower()
|
||||
|
||||
def _dedup_pick(tokens: list[str]) -> "str | None":
|
||||
if not tokens:
|
||||
return None
|
||||
# Index into the full (ordered) list first so HIP_VISIBLE_DEVICES
|
||||
# correctly addresses GPU N on mixed-arch hosts, then return that arch.
|
||||
return tokens[_pick_visible_index(len(tokens))]
|
||||
|
||||
# 2. hipinfo via PATH, then HIP_PATH\bin / ROCM_PATH\bin.
|
||||
hipinfo = shutil.which("hipinfo")
|
||||
if not hipinfo:
|
||||
for _env_var in ("HIP_PATH", "ROCM_PATH"):
|
||||
_root = os.environ.get(_env_var)
|
||||
if _root:
|
||||
_candidate = os.path.join(_root, "bin", "hipinfo.exe")
|
||||
if os.path.isfile(_candidate):
|
||||
hipinfo = _candidate
|
||||
break
|
||||
if hipinfo:
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[hipinfo],
|
||||
stdout = subprocess.PIPE,
|
||||
stderr = subprocess.DEVNULL,
|
||||
timeout = 10,
|
||||
)
|
||||
if result.returncode == 0:
|
||||
text = result.stdout.decode(errors = "replace")
|
||||
# findall picks every gcnArchName line so multi-GPU hosts
|
||||
# are enumerable and HIP_VISIBLE_DEVICES selects correctly.
|
||||
_tokens = [
|
||||
t.strip().lower()
|
||||
for t in re.findall(r"(?im)^\s*gcnArchName\s*:\s*(\S+)", text)
|
||||
]
|
||||
_pick = _dedup_pick(_tokens)
|
||||
if _pick:
|
||||
return _pick
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 3. amd-smi fallback -- runtime-only Radeon installs ship amd-smi but no hipinfo.
|
||||
amd_smi = shutil.which("amd-smi")
|
||||
if amd_smi:
|
||||
for _args in (("static", "--asic"), ("list",)):
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[amd_smi, *_args],
|
||||
stdout = subprocess.PIPE,
|
||||
stderr = subprocess.DEVNULL,
|
||||
timeout = 10,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
continue
|
||||
text = result.stdout.decode(errors = "replace")
|
||||
# Prefer labelled gfx lines; fall back to bare tokens.
|
||||
_labelled = re.findall(
|
||||
r"(?im)^\s*(?:target_graphics_version|gfx|arch|asic)\b[^:\r\n]*:\s*(gfx[1-9][0-9a-z]{2,3})\b",
|
||||
text,
|
||||
)
|
||||
_tokens = [t.lower() for t in _labelled]
|
||||
if not _tokens:
|
||||
_tokens = re.findall(r"\bgfx[1-9][0-9a-z]{2,3}\b", text.lower())
|
||||
_pick = _dedup_pick(_tokens)
|
||||
if _pick:
|
||||
return _pick
|
||||
except Exception:
|
||||
continue
|
||||
return None
|
||||
|
||||
|
||||
def _windows_rocm_index_url(gfx_arch: str | None) -> str | None:
|
||||
"""Return the AMD pip index URL for the given GPU arch, or None if unsupported."""
|
||||
arch_family = _GFX_TO_AMD_INDEX_ARCH.get(gfx_arch or "")
|
||||
if arch_family is None:
|
||||
return None
|
||||
return f"{_ROCM_WINDOWS_INDEX_BASE}/{arch_family}/"
|
||||
|
||||
|
||||
def _detect_bnb_rocm_dll_ver() -> str | None:
|
||||
"""Scan the installed bitsandbytes package for libbitsandbytes_rocm{VER}.dll.
|
||||
|
||||
Returns the version suffix string (e.g. ``"72"``, ``"713"``) or ``None``
|
||||
if bitsandbytes is not installed or no ROCm DLL is found. Does NOT import
|
||||
bitsandbytes — uses importlib.util.find_spec so it is safe to call before
|
||||
BNB is imported.
|
||||
"""
|
||||
import importlib.util
|
||||
|
||||
spec = importlib.util.find_spec("bitsandbytes")
|
||||
if spec is None or not spec.submodule_search_locations:
|
||||
return None
|
||||
all_vers: list[str] = []
|
||||
for pkg_dir in spec.submodule_search_locations:
|
||||
for dll in glob.glob(os.path.join(pkg_dir, "libbitsandbytes_rocm*.dll")):
|
||||
m = re.search(r"libbitsandbytes_rocm(\d+)\.dll", os.path.basename(dll))
|
||||
if m:
|
||||
all_vers.append(m.group(1))
|
||||
# Pick the highest numeric suffix so that e.g. "713" wins over "72" when
|
||||
# both variants are present in the wheel. Filesystem glob order is not
|
||||
# guaranteed, so always sort rather than stopping at the first match.
|
||||
return max(all_vers, key = lambda v: int(v)) if all_vers else None
|
||||
|
||||
|
||||
def _has_rocm_gpu() -> bool:
|
||||
"""Return True only if an actual AMD GPU is visible (not just ROCm tools installed)."""
|
||||
import re
|
||||
|
||||
for cmd, check_fn in (
|
||||
# rocminfo: look for a real gfx GPU id (3-4 chars, nonzero first digit).
|
||||
# gfx000 is the CPU agent; ROCm 6.1+ also emits generic ISA lines like
|
||||
|
|
@ -231,6 +414,26 @@ def _has_rocm_gpu() -> bool:
|
|||
if result.returncode == 0 and result.stdout.strip():
|
||||
if check_fn(result.stdout):
|
||||
return True
|
||||
# sysfs KFD topology fallback (Linux only) -- matches install.sh's
|
||||
# runtime-only detection. On minimal package-managed installs (no
|
||||
# rocminfo / no amd-smi GUI tools), the kernel exposes AMD GPUs via
|
||||
# /sys/class/kfd so `studio update` can still detect the GPU and
|
||||
# repair the venv.
|
||||
if sys.platform != "win32":
|
||||
try:
|
||||
kfd_nodes = "/sys/class/kfd/kfd/topology/nodes"
|
||||
if os.path.isdir(kfd_nodes):
|
||||
for entry in os.listdir(kfd_nodes):
|
||||
gpu_id_path = os.path.join(kfd_nodes, entry, "gpu_id")
|
||||
try:
|
||||
with open(gpu_id_path) as fh:
|
||||
gpu_id = fh.read().strip()
|
||||
except OSError:
|
||||
continue
|
||||
if gpu_id and gpu_id != "0": # gpu_id 0 = CPU node
|
||||
return True
|
||||
except OSError:
|
||||
pass
|
||||
return False
|
||||
|
||||
|
||||
|
|
@ -252,23 +455,194 @@ def _has_usable_nvidia_gpu() -> bool:
|
|||
return result.returncode == 0 and "GPU " in result.stdout
|
||||
|
||||
|
||||
def _detect_amd_gfx_codes() -> list[str]:
|
||||
"""Return the list of AMD gfx ISA strings visible to ROCm (e.g. ['gfx1151']).
|
||||
|
||||
Probes rocminfo first, then falls back to ``amd-smi list`` and
|
||||
``amd-smi static --asic`` for runtime-only Radeon hosts that ship
|
||||
amd-smi but no rocminfo. Returns an empty list when no probe yields
|
||||
a gfx target.
|
||||
"""
|
||||
|
||||
def _extract(text: str) -> list[str]:
|
||||
codes = re.findall(r"gfx([1-9][0-9a-z]{2,3})", text.lower())
|
||||
return list(dict.fromkeys(f"gfx{c}" for c in codes))
|
||||
|
||||
probes: list[list[str]] = []
|
||||
if shutil.which("rocminfo"):
|
||||
probes.append(["rocminfo"])
|
||||
if shutil.which("amd-smi"):
|
||||
probes.append(["amd-smi", "list"])
|
||||
probes.append(["amd-smi", "static", "--asic"])
|
||||
for cmd in probes:
|
||||
try:
|
||||
result = subprocess.run(
|
||||
cmd,
|
||||
stdout = subprocess.PIPE,
|
||||
stderr = subprocess.DEVNULL,
|
||||
text = True,
|
||||
timeout = 15,
|
||||
)
|
||||
except Exception:
|
||||
continue
|
||||
if result.returncode != 0 or not result.stdout.strip():
|
||||
continue
|
||||
codes = _extract(result.stdout)
|
||||
if codes:
|
||||
return codes
|
||||
return []
|
||||
|
||||
|
||||
# Set by _ensure_rocm_torch() on success; suppresses the post-install AMD warning.
|
||||
_rocm_windows_torch_installed: bool = False
|
||||
|
||||
|
||||
def _install_bnb_windows_rocm() -> bool:
|
||||
"""Install the AMD Windows BNB prerelease wheel. Returns True on success.
|
||||
|
||||
The continuous-release wheel is intentionally mismatched: the filename
|
||||
encodes version 1.33.7.preview (parsed as 1.33.7rc0 by PEP 440) while the
|
||||
wheel metadata reports 0.50.0.dev0. uv rejects this filename/metadata
|
||||
mismatch -- and bypassing it with UV_SKIP_WHEEL_FILENAME_CHECK still leaves
|
||||
uv mangling the bitsandbytes install. Per the AMD install guide
|
||||
(https://unsloth.ai/docs/get-started/install/amd/amd-hackathon) the wheel
|
||||
must be installed with plain pip, not uv, so we force pip here
|
||||
(force_pip=True). plain pip performs no wheel filename/metadata check.
|
||||
"""
|
||||
_bnb_win_url = _BNB_ROCM_PRERELEASE_URLS.get("win_amd64")
|
||||
if _bnb_win_url is None:
|
||||
return False
|
||||
_ok = pip_install_try(
|
||||
"bitsandbytes (AMD Windows, pre-release main)",
|
||||
"--force-reinstall",
|
||||
"--no-cache-dir",
|
||||
"--no-deps",
|
||||
_bnb_win_url,
|
||||
constrain = False,
|
||||
force_pip = True,
|
||||
)
|
||||
if not _ok:
|
||||
return False
|
||||
# After install: detect the actual ROCm DLL suffix shipped in the wheel and
|
||||
# set BNB_ROCM_VERSION so bitsandbytes loads the correct DLL regardless of
|
||||
# what torch.version.hip reports. The wheel may ship an older suffix (e.g.
|
||||
# "72") while torch reports a newer HIP version (e.g. 7.13); the env var
|
||||
# override ensures bitsandbytes does not fail looking for a non-existent DLL.
|
||||
# The worker subprocess inherits this env var automatically.
|
||||
# Fall back to "72" if detection fails (e.g. install was a no-op / dry-run).
|
||||
if "BNB_ROCM_VERSION" not in os.environ:
|
||||
_ver = _detect_bnb_rocm_dll_ver() or "72"
|
||||
os.environ["BNB_ROCM_VERSION"] = _ver
|
||||
return True
|
||||
|
||||
|
||||
def _ensure_rocm_torch() -> None:
|
||||
"""Reinstall torch with ROCm wheels when the venv received CPU-only torch.
|
||||
|
||||
Runs only on Linux x86_64 hosts where an AMD GPU is present and the
|
||||
ROCm runtime is detectable (rocminfo / amd-smi / hipconfig /
|
||||
rocm-core package). No-op when torch already links against HIP
|
||||
(ROCm), on Windows / macOS, on non-x86_64 Linux (PyTorch does not
|
||||
publish ROCm wheels for aarch64 / arm64), or on mixed AMD+NVIDIA
|
||||
hosts (NVIDIA takes precedence).
|
||||
On Linux x86_64: uses pytorch.org ROCm wheel index tags.
|
||||
On Windows: uses AMD's repo.amd.com arch-specific pip index.
|
||||
No-op on macOS, non-x86_64 Linux, NVIDIA-primary hosts, or when torch
|
||||
already links against HIP.
|
||||
Uses pip_install() to respect uv, constraints, and --python targeting.
|
||||
"""
|
||||
# Explicit OS / architecture guards so the helper is safe to call
|
||||
# from any context -- PyTorch only publishes ROCm wheels for
|
||||
# linux_x86_64, so aarch64 / arm64 hosts must skip this repair path
|
||||
# instead of failing the update with a missing-wheel error.
|
||||
if IS_WINDOWS or IS_MACOS:
|
||||
global _rocm_windows_torch_installed
|
||||
# setup.ps1 sets this when it already installed AMD wheels; skip the probe
|
||||
# only when torch is actually importable as ROCm. If the venv was wiped
|
||||
# between runs, the stale env-var would suppress a needed reinstall.
|
||||
if os.environ.get("UNSLOTH_ROCM_TORCH_INSTALLED") == "1":
|
||||
_torch_ok = False
|
||||
try:
|
||||
_probe = subprocess.run(
|
||||
[
|
||||
sys.executable,
|
||||
"-c",
|
||||
(
|
||||
"import torch; "
|
||||
"hip=getattr(torch.version,'hip','') or ''; "
|
||||
"import sys; "
|
||||
"sys.exit(0 if (hip or 'rocm' in torch.__version__.lower()) else 1)"
|
||||
),
|
||||
],
|
||||
stdout = subprocess.DEVNULL,
|
||||
stderr = subprocess.DEVNULL,
|
||||
timeout = 90,
|
||||
)
|
||||
_torch_ok = _probe.returncode == 0
|
||||
except (OSError, subprocess.TimeoutExpired):
|
||||
pass
|
||||
if _torch_ok:
|
||||
_rocm_windows_torch_installed = True
|
||||
# setup.ps1 already installed ROCm torch, but we still need to install
|
||||
# the AMD Windows BNB wheel here -- the PyPI bitsandbytes wheel ships
|
||||
# only CUDA DLLs and will fail to load on ROCm.
|
||||
_install_bnb_windows_rocm()
|
||||
return
|
||||
# torch was wiped between runs; fall through to the full install path
|
||||
if IS_MACOS:
|
||||
return
|
||||
|
||||
if IS_WINDOWS:
|
||||
if _has_usable_nvidia_gpu():
|
||||
return
|
||||
gfx_arch = _detect_windows_gfx_arch()
|
||||
if not gfx_arch:
|
||||
return # no AMD GPU visible via hipinfo
|
||||
# Probe whether torch already links against HIP.
|
||||
_torch_already_rocm = False
|
||||
try:
|
||||
probe = subprocess.run(
|
||||
[
|
||||
sys.executable,
|
||||
"-c",
|
||||
(
|
||||
"import torch; "
|
||||
"hip=getattr(torch.version,'hip','') or ''; "
|
||||
"ver=torch.__version__; "
|
||||
"print('yes' if hip or 'rocm' in ver.lower() else '')"
|
||||
),
|
||||
],
|
||||
stdout = subprocess.PIPE,
|
||||
stderr = subprocess.DEVNULL,
|
||||
timeout = 90,
|
||||
)
|
||||
if probe.returncode == 0 and probe.stdout.decode().strip() == "yes":
|
||||
_torch_already_rocm = True
|
||||
except (OSError, subprocess.TimeoutExpired):
|
||||
pass
|
||||
if not _torch_already_rocm:
|
||||
index_url = _windows_rocm_index_url(gfx_arch)
|
||||
if index_url is None:
|
||||
print(
|
||||
f" No AMD Windows torch index for GPU arch {gfx_arch} -- skipping"
|
||||
)
|
||||
return
|
||||
print(f" {gfx_arch} (Windows) -- installing torch from {index_url}")
|
||||
pip_install(
|
||||
f"ROCm torch (Windows, {gfx_arch})",
|
||||
"--force-reinstall",
|
||||
"--index-url",
|
||||
index_url,
|
||||
"torch",
|
||||
"torchvision",
|
||||
"torchaudio",
|
||||
constrain = False,
|
||||
)
|
||||
# ROCm torch is installed (or already was); flag it so later install
|
||||
# phases do not overwrite it with the generic CPU torch wheel. BNB is
|
||||
# a separate dependency -- a BNB install failure must NOT roll the
|
||||
# torch ROCm install back.
|
||||
_rocm_windows_torch_installed = True
|
||||
# Always install AMD Windows bitsandbytes -- the PyPI wheel ships only
|
||||
# CUDA DLLs and will fail to load on ROCm. Install even when torch was
|
||||
# already a ROCm build so that `studio update` repairs a broken bnb.
|
||||
if not _install_bnb_windows_rocm():
|
||||
print(
|
||||
" Warning: AMD Windows bitsandbytes install failed; "
|
||||
"ROCm torch is installed but bitsandbytes may need manual install"
|
||||
)
|
||||
return
|
||||
|
||||
# ── Linux x86_64 only: PyTorch ROCm wheels are not published for aarch64 ──
|
||||
if platform.machine().lower() not in {"x86_64", "amd64"}:
|
||||
return
|
||||
# NVIDIA takes precedence on mixed hosts -- but only if an actual GPU is usable
|
||||
|
|
@ -297,11 +671,19 @@ def _ensure_rocm_torch() -> None:
|
|||
[
|
||||
sys.executable,
|
||||
"-c",
|
||||
"import torch; print(getattr(torch.version,'hip','') or '')",
|
||||
(
|
||||
"import torch; "
|
||||
"hip=getattr(torch.version,'hip','') or ''; "
|
||||
"ver=getattr(torch,'__version__','').lower(); "
|
||||
# Print the HIP version when present (back-compat), else
|
||||
# "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 ''))"
|
||||
),
|
||||
],
|
||||
stdout = subprocess.PIPE,
|
||||
stderr = subprocess.DEVNULL,
|
||||
timeout = 30,
|
||||
timeout = 90,
|
||||
)
|
||||
except (OSError, subprocess.TimeoutExpired):
|
||||
probe = None
|
||||
|
|
@ -313,7 +695,83 @@ def _ensure_rocm_torch() -> None:
|
|||
|
||||
rocm_torch_ready = has_hip_torch
|
||||
|
||||
if not has_hip_torch:
|
||||
# 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
|
||||
# with the real fix, so route those hosts there instead of the generic
|
||||
# pytorch.org rocm7.1 wheel. Mirrors install.sh's Strix override.
|
||||
# On mixed hosts (Strix iGPU + non-Strix dGPU), only route to the AMD
|
||||
# per-gfx index when the GPU HIP will actually run on is the Strix one --
|
||||
# otherwise the dGPU would get an incompatible wheel. Use HIP_VISIBLE_DEVICES
|
||||
# to determine the runtime target.
|
||||
_strix_override_url: "str | None" = None
|
||||
_strix_override_pkgs: "tuple[str, str, str] | None" = None
|
||||
if ver < (7, 2):
|
||||
gfx_codes = _detect_amd_gfx_codes()
|
||||
_strix_gfx = {"gfx1151", "gfx1150"}
|
||||
_detected_strix = _strix_gfx.intersection(gfx_codes)
|
||||
if _detected_strix:
|
||||
# Pick the runtime-visible GPU. If HIP_VISIBLE_DEVICES selects a
|
||||
# specific index into gfx_codes, use that gfx; else default to the
|
||||
# first listed GPU. Skip the override unless the resolved GPU is
|
||||
# Strix.
|
||||
_runtime_gfx = (
|
||||
gfx_codes[_pick_visible_index(len(gfx_codes))] if gfx_codes else None
|
||||
)
|
||||
if _runtime_gfx in _strix_gfx:
|
||||
_selected_gfx = _runtime_gfx
|
||||
_amd_mirror = (
|
||||
os.environ.get("UNSLOTH_AMD_ROCM_MIRROR")
|
||||
or "https://repo.amd.com/rocm/whl"
|
||||
).rstrip("/")
|
||||
_strix_override_url = f"{_amd_mirror}/{_selected_gfx}/"
|
||||
_strix_override_pkgs = (
|
||||
"torch>=2.11.0,<2.12.0",
|
||||
# Pin torchvision/torchaudio to the 2.11.x-compatible range.
|
||||
# The install uses --index-url (exclusive, no PyPI fallback),
|
||||
# so bare unversioned names risk resolving a build from AMD's
|
||||
# index that targets a different torch major (e.g. 0.27 built
|
||||
# against torch 2.12), which would fail at runtime with an
|
||||
# ABI/version mismatch. Matches _ROCM_TORCH_CONSTRAINT["rocm7.2"].
|
||||
"torchvision>=0.26.0,<0.27.0",
|
||||
"torchaudio>=2.11.0,<2.12.0",
|
||||
)
|
||||
print(
|
||||
f"\n {_selected_gfx} (AMD Strix) is the runtime target with ROCm "
|
||||
f"{ver[0]}.{ver[1]}.\n"
|
||||
f" ROCm 7.1 has a known _grouped_mm segfault on this GPU;\n"
|
||||
f" routing torch install to AMD's arch-specific index\n"
|
||||
f" ({_strix_override_url}) which serves torch 2.11.0+rocm7.13.0\n"
|
||||
f" with the upstream fix.\n"
|
||||
)
|
||||
else:
|
||||
_gfx_str = ", ".join(sorted(_detected_strix))
|
||||
print(
|
||||
f"\n Strix GPU ({_gfx_str}) present but HIP_VISIBLE_DEVICES "
|
||||
f"selects a non-Strix runtime target ({_runtime_gfx});\n"
|
||||
f" skipping AMD per-gfx index override.\n"
|
||||
)
|
||||
|
||||
# Strix override on ROCm 7.1 must fire even when has_hip_torch is True --
|
||||
# an existing torch with `torch.version.hip == "7.1"` is exactly the broken
|
||||
# combo the override is meant to repair, so skipping it leaves users on
|
||||
# the known _grouped_mm segfault.
|
||||
if _strix_override_url is not None and _strix_override_pkgs is not None:
|
||||
index_url = _strix_override_url
|
||||
_torch_pkg, _vision_pkg, _audio_pkg = _strix_override_pkgs
|
||||
print(f" Strix ROCm 7.1 override -- installing torch from {index_url}")
|
||||
pip_install(
|
||||
"ROCm torch (Strix arch-specific)",
|
||||
"--force-reinstall",
|
||||
"--no-cache-dir",
|
||||
_torch_pkg,
|
||||
_vision_pkg,
|
||||
_audio_pkg,
|
||||
"--index-url",
|
||||
index_url,
|
||||
constrain = False,
|
||||
)
|
||||
rocm_torch_ready = True
|
||||
elif not has_hip_torch:
|
||||
# Select best matching wheel tag (newest ROCm version <= installed)
|
||||
tag = next(
|
||||
(
|
||||
|
|
@ -331,13 +789,16 @@ def _ensure_rocm_torch() -> None:
|
|||
else:
|
||||
index_url = f"{_PYTORCH_WHL_BASE}/{tag}"
|
||||
print(f" ROCm {ver[0]}.{ver[1]} -- installing torch from {index_url}")
|
||||
_torch_pkg, _vision_pkg, _audio_pkg = _ROCM_TORCH_PKG_SPECS.get(
|
||||
tag, _ROCM_TORCH_PKG_SPECS["_default"]
|
||||
)
|
||||
pip_install(
|
||||
f"ROCm torch ({tag})",
|
||||
"--force-reinstall",
|
||||
"--no-cache-dir",
|
||||
"torch>=2.4,<2.11.0",
|
||||
"torchvision<0.26.0",
|
||||
"torchaudio<2.11.0",
|
||||
_torch_pkg,
|
||||
_vision_pkg,
|
||||
_audio_pkg,
|
||||
"--index-url",
|
||||
index_url,
|
||||
constrain = False,
|
||||
|
|
@ -346,7 +807,9 @@ def _ensure_rocm_torch() -> None:
|
|||
|
||||
# Install bitsandbytes only when torch links against ROCm. Prefers the
|
||||
# continuous-release_main wheel (bnb PR #1887 4-bit GEMV fix) and falls
|
||||
# back to PyPI when the pre-release URL is unreachable.
|
||||
# back to PyPI when the pre-release wheel cannot be installed. Use pip for
|
||||
# the pre-release wheel because uv rejects the wheel's filename/metadata
|
||||
# version mismatch.
|
||||
if rocm_torch_ready:
|
||||
_bnb_url = _bnb_rocm_prerelease_url()
|
||||
_bnb_installed = False
|
||||
|
|
@ -358,11 +821,12 @@ def _ensure_rocm_torch() -> None:
|
|||
"--no-deps",
|
||||
_bnb_url,
|
||||
constrain = False,
|
||||
force_pip = True,
|
||||
)
|
||||
if not _bnb_installed:
|
||||
print(
|
||||
_red(
|
||||
" bnb pre-release unreachable; falling back to PyPI "
|
||||
" bnb pre-release install failed; falling back to PyPI "
|
||||
"(4-bit decode will be broken on ROCm)"
|
||||
)
|
||||
)
|
||||
|
|
@ -809,6 +1273,7 @@ def pip_install_try(
|
|||
label: str,
|
||||
*args: str,
|
||||
constrain: bool = True,
|
||||
force_pip: bool = False,
|
||||
) -> bool:
|
||||
"""Like pip_install but returns False on failure instead of exiting.
|
||||
For optional installs with a follow-up fallback.
|
||||
|
|
@ -819,7 +1284,7 @@ def pip_install_try(
|
|||
constraint_args_pip = ["-c", str(CONSTRAINTS)]
|
||||
constraint_args_uv = ["-c", _uv_safe_path(CONSTRAINTS)]
|
||||
|
||||
if USE_UV:
|
||||
if USE_UV and not force_pip:
|
||||
cmd = _build_uv_cmd(args) + constraint_args_uv
|
||||
else:
|
||||
cmd = _build_pip_cmd(args) + constraint_args_pip
|
||||
|
|
@ -948,8 +1413,12 @@ def install_python_stack() -> int:
|
|||
base_total = 10 if IS_WINDOWS else 11
|
||||
if IS_MACOS:
|
||||
base_total -= 1 # triton step is skipped on macOS
|
||||
if not IS_WINDOWS and not IS_MACOS and not NO_TORCH:
|
||||
base_total += 3
|
||||
if not IS_MACOS and not NO_TORCH:
|
||||
base_total += 1 # ROCm torch check (line 1526) -- all non-macOS platforms
|
||||
if not IS_WINDOWS:
|
||||
base_total += (
|
||||
2 # flash-attn (line 1620) + ROCm torch final (line 1705) -- Linux only
|
||||
)
|
||||
if not NO_TORCH:
|
||||
base_total += 1 # studio RAG deps (rag.txt)
|
||||
_TOTAL = (base_total - 1) if skip_base else base_total
|
||||
|
|
@ -1123,12 +1592,12 @@ def install_python_stack() -> int:
|
|||
# 2b. AMD ROCm: reinstall torch with HIP wheels if the host has ROCm but the
|
||||
# venv received CPU-only torch (common when pip resolves torch from PyPI).
|
||||
# Must come immediately after base packages so torch is present for inspection.
|
||||
if not IS_WINDOWS and not IS_MACOS and not NO_TORCH:
|
||||
if not IS_MACOS and not NO_TORCH:
|
||||
_progress("ROCm torch check")
|
||||
_ensure_rocm_torch()
|
||||
|
||||
# Windows + AMD GPU: PyTorch does not publish ROCm wheels for Windows.
|
||||
# Detect and warn so users know manual steps are needed for GPU training.
|
||||
# Windows + AMD GPU: if ROCm torch was not installed (wrong Python version
|
||||
# or unknown ROCm version), warn the user.
|
||||
if IS_WINDOWS and not NO_TORCH and not _has_usable_nvidia_gpu():
|
||||
# Validate actual AMD GPU presence (not just tool existence)
|
||||
import re as _re_win
|
||||
|
|
@ -1157,14 +1626,14 @@ def install_python_stack() -> int:
|
|||
if _wr.returncode == 0 and _check_fn(_wr.stdout):
|
||||
_win_amd_gpu = True
|
||||
break
|
||||
if _win_amd_gpu:
|
||||
if _win_amd_gpu and not _rocm_windows_torch_installed:
|
||||
_safe_print(
|
||||
_dim(" Note:"),
|
||||
"AMD GPU detected on Windows. ROCm-enabled PyTorch must be",
|
||||
"AMD GPU detected but ROCm PyTorch could not be auto-installed.",
|
||||
)
|
||||
_safe_print(
|
||||
" " * 8,
|
||||
"installed manually. See: https://docs.unsloth.ai/get-started/install-and-update/amd",
|
||||
"Manual install may be required. See: https://docs.unsloth.ai/get-started/install-and-update/amd",
|
||||
)
|
||||
|
||||
# 3. Extra dependencies
|
||||
|
|
@ -1191,10 +1660,17 @@ def install_python_stack() -> int:
|
|||
_progress("dependency overrides (skipped, no torch)")
|
||||
else:
|
||||
_progress("dependency overrides")
|
||||
_override_extra_args: tuple[str, ...] = ()
|
||||
if _rocm_windows_torch_installed:
|
||||
# torchao in overrides.txt declares torch as a dependency; without
|
||||
# --no-deps uv would resolve and install CPU torch from PyPI,
|
||||
# overwriting the AMD ROCm wheels we just installed.
|
||||
_override_extra_args = ("--no-deps",)
|
||||
pip_install(
|
||||
"Installing dependency overrides",
|
||||
"--force-reinstall",
|
||||
"--no-cache-dir",
|
||||
*_override_extra_args,
|
||||
req = REQ_ROOT / "overrides.txt",
|
||||
)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue