refactor(rocm/win): switch to repo.amd.com arch-aware index, remove stubs
AMD recommends repo.amd.com/rocm/whl/{arch}/ as the Windows ROCm wheel
source. These wheels bundle their own ROCm runtime, support all Python
versions (not just cp312), and include the full torch._C extension set
(including _distributed_c10d) that the old repo.radeon.com wheel omitted.
Changes:
- install.ps1: remove Select-ROCmWheelRelease + hardcoded cp312 wheel
URLs; remove Python 3.12 forced-preference logic; install via
--index-url repo.amd.com/rocm/whl/{arch-family}/
- studio/setup.ps1: same -- remove Select-ROCmWheelRelease, switch to
repo.amd.com arch-aware index URL
- studio/install_python_stack.py: replace _ROCM_WINDOWS_RELEASES /
_select_windows_rocm_release with _windows_rocm_index_url() using the
_GFX_TO_AMD_INDEX_ARCH map; drop Python 3.12 restriction
- studio/backend/core/training/worker.py: remove all stub machinery
(_make_mod_stub, _StubSubpackageFinder, _StubSubpackageLoader,
_StubClassMeta, torchao/fsdp/dtensor stubs, _c10d_functional ops
stubs, BNB DLL detection) -- no longer needed with new wheel source
This commit is contained in:
parent
32643198b2
commit
d731c5fc63
4 changed files with 85 additions and 613 deletions
|
|
@ -69,38 +69,22 @@ _PYTORCH_WHL_BASE = (
|
|||
os.environ.get("UNSLOTH_PYTORCH_MIRROR") or "https://download.pytorch.org/whl"
|
||||
).rstrip("/")
|
||||
|
||||
# AMD Windows ROCm wheels — repo.radeon.com (cp312 only)
|
||||
_ROCM_WINDOWS_WHEEL_BASE = (
|
||||
# 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.radeon.com/rocm/windows"
|
||||
or "https://repo.amd.com/rocm/whl"
|
||||
).rstrip("/")
|
||||
# Maps (major, minor) → (release_folder, [wheel_filename, ...])
|
||||
_ROCM_WINDOWS_RELEASES: dict[tuple[int, int], tuple[str, list[str]]] = {
|
||||
(7, 2): (
|
||||
"rocm-rel-7.2.1",
|
||||
[
|
||||
# rocm tarball provides the 'rocm_sdk' Python namespace package
|
||||
"rocm-7.2.1.tar.gz",
|
||||
"rocm_sdk_core-7.2.1-py3-none-win_amd64.whl",
|
||||
"rocm_sdk_devel-7.2.1-py3-none-win_amd64.whl",
|
||||
"rocm_sdk_libraries_custom-7.2.1-py3-none-win_amd64.whl",
|
||||
"torch-2.9.1+rocm7.2.1-cp312-cp312-win_amd64.whl",
|
||||
"torchvision-0.24.1+rocm7.2.1-cp312-cp312-win_amd64.whl",
|
||||
"torchaudio-2.9.1+rocm7.2.1-cp312-cp312-win_amd64.whl",
|
||||
],
|
||||
),
|
||||
(7, 1): (
|
||||
"rocm-rel-7.1.1",
|
||||
[
|
||||
# rocm tarball provides the 'rocm_sdk' Python namespace package
|
||||
"rocm-0.1.dev0.tar.gz",
|
||||
"rocm_sdk_core-0.1.dev0-py3-none-win_amd64.whl",
|
||||
"rocm_sdk_libraries_custom-0.1.dev0-py3-none-win_amd64.whl",
|
||||
"torch-2.9.0+rocmsdk20251116-cp312-cp312-win_amd64.whl",
|
||||
"torchvision-0.24.0+rocmsdk20251116-cp312-cp312-win_amd64.whl",
|
||||
"torchaudio-2.9.0+rocmsdk20251116-cp312-cp312-win_amd64.whl",
|
||||
],
|
||||
),
|
||||
|
||||
# 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
|
||||
|
|
@ -235,20 +219,6 @@ def _detect_rocm_version() -> tuple[int, int] | None:
|
|||
return None
|
||||
|
||||
|
||||
# GPU arch → minimum (major, minor) ROCm release that supports it on Windows.
|
||||
# Wheels bundle their own ROCm runtime, so the installed HIP SDK version does
|
||||
# not constrain selection — only the GPU's architecture minimum matters.
|
||||
_GFX_MIN_ROCM_WINDOWS: dict[str, tuple[int, int]] = {
|
||||
"gfx1201": (7, 1), "gfx1200": (7, 1), # RDNA 4
|
||||
"gfx1151": (7, 1), "gfx1150": (7, 1), # RDNA 3.5 (Strix Halo/Point)
|
||||
"gfx1103": (6, 4), "gfx1102": (6, 4), "gfx1101": (6, 4), "gfx1100": (6, 4), # RDNA 3
|
||||
"gfx1036": (6, 4), "gfx1035": (6, 4), "gfx1034": (6, 4), "gfx1033": (6, 4), # RDNA 2
|
||||
"gfx1032": (6, 4), "gfx1031": (6, 4), "gfx1030": (6, 4),
|
||||
"gfx1011": (6, 4), "gfx1010": (6, 4), # RDNA 1
|
||||
"gfx906": (6, 4), "gfx908": (6, 4), "gfx90a": (6, 4), # Vega/MI
|
||||
}
|
||||
|
||||
|
||||
def _detect_windows_gfx_arch() -> str | None:
|
||||
"""Return the gcnArchName from hipinfo on Windows (e.g. 'gfx1200'), or None."""
|
||||
import re
|
||||
|
|
@ -272,17 +242,12 @@ def _detect_windows_gfx_arch() -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
def _select_windows_rocm_release(gfx_arch: str | None) -> tuple[str, list[str]] | None:
|
||||
"""Pick the best available Windows ROCm release for the given GPU arch.
|
||||
|
||||
Always selects the newest available release whose ROCm version meets the
|
||||
GPU's minimum requirement. Returns None when no release qualifies.
|
||||
"""
|
||||
min_ver = _GFX_MIN_ROCM_WINDOWS.get(gfx_arch or "", (6, 4))
|
||||
for (maj, mn), entry in sorted(_ROCM_WINDOWS_RELEASES.items(), reverse = True):
|
||||
if (maj, mn) >= min_ver:
|
||||
return entry
|
||||
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 _has_rocm_gpu() -> bool:
|
||||
|
|
@ -378,7 +343,7 @@ def _ensure_rocm_torch() -> None:
|
|||
"""Reinstall torch with ROCm wheels when the venv received CPU-only torch.
|
||||
|
||||
On Linux x86_64: uses pytorch.org ROCm wheel index tags.
|
||||
On Windows (cp312 only): uses AMD's repo.radeon.com direct wheel releases.
|
||||
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.
|
||||
|
|
@ -392,13 +357,6 @@ def _ensure_rocm_torch() -> None:
|
|||
return
|
||||
|
||||
if IS_WINDOWS:
|
||||
# AMD only publishes Windows ROCm wheels for Python 3.12 (cp312)
|
||||
if sys.version_info[:2] != (3, 12):
|
||||
print(
|
||||
f" ROCm torch on Windows requires Python 3.12 "
|
||||
f"(current: {sys.version_info[0]}.{sys.version_info[1]}) -- skipping"
|
||||
)
|
||||
return
|
||||
if _has_usable_nvidia_gpu():
|
||||
return
|
||||
gfx_arch = _detect_windows_gfx_arch()
|
||||
|
|
@ -425,34 +383,19 @@ def _ensure_rocm_torch() -> None:
|
|||
return # already ROCm torch
|
||||
except (OSError, subprocess.TimeoutExpired):
|
||||
pass
|
||||
entry = _select_windows_rocm_release(gfx_arch)
|
||||
if entry is None:
|
||||
print(f" No AMD Windows torch wheel for GPU arch {gfx_arch} -- skipping")
|
||||
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
|
||||
rel_tag, wheel_files = entry
|
||||
base = f"{_ROCM_WINDOWS_WHEEL_BASE}/{rel_tag}"
|
||||
wheel_urls = [f"{base}/{fn}" for fn in wheel_files]
|
||||
print(f" {gfx_arch} (Windows) -- installing torch from {base}/")
|
||||
# Install rocm namespace tarball first (torch/_rocm_init.py imports it)
|
||||
tarball_url = next((u for u in wheel_urls if u.endswith(".tar.gz")), None)
|
||||
whl_urls = [u for u in wheel_urls if not u.endswith(".tar.gz")]
|
||||
if tarball_url:
|
||||
pip_install(
|
||||
f"ROCm namespace ({rel_tag})",
|
||||
"--force-reinstall",
|
||||
"--no-deps",
|
||||
tarball_url,
|
||||
constrain = False,
|
||||
)
|
||||
print(f" {gfx_arch} (Windows) -- installing torch from {index_url}")
|
||||
pip_install(
|
||||
f"ROCm torch (Windows, {rel_tag})",
|
||||
f"ROCm torch (Windows, {gfx_arch})",
|
||||
"--force-reinstall",
|
||||
"--no-deps",
|
||||
*whl_urls,
|
||||
"--index-url", index_url,
|
||||
"torch", "torchvision", "torchaudio",
|
||||
constrain = False,
|
||||
)
|
||||
# bitsandbytes Windows ROCm wheel (ships libbitsandbytes_rocm72.dll).
|
||||
# BNB_ROCM_VERSION=72 is set in worker.py before the bnb import.
|
||||
# bitsandbytes Windows ROCm wheel.
|
||||
_bnb_win_url = _BNB_ROCM_PRERELEASE_URLS.get("win_amd64")
|
||||
if _bnb_win_url is not None:
|
||||
pip_install_try(
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue