From 96b9e4659b40a95e668d4936a61055d66fe81691 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Tue, 19 May 2026 10:54:19 +0000 Subject: [PATCH] fix(studio/rocm): multi-GPU selection, Strix sibling handling, defensive cleanups Round 4 robustness pass based on 5 parallel Opus reviewers of head 21773215. Seven items from across regression / edge-case / error-paths / architecture reviews: 1. studio/backend/main.py BNB gate: aligned with the broad ROCm check used everywhere else in this PR (torch.version.hip OR 'rocm' in __version__). AMD SDK / Radeon Linux wheels do not always populate torch.version.hip; without this, main.py would silently skip BNB_ROCM_VERSION while worker.py set it. 2. studio/install_python_stack.py _install_bnb_windows_rocm: init _ok = False before the try block. Without this, if pip_install_try itself raises (e.g. OSError on uv binary missing), the finally block restored env vars correctly but the subsequent `if not _ok:` raised UnboundLocalError, masking the original exception. 3. studio/install_python_stack.py _detect_windows_gfx_arch: - Rewrote to use re.findall (not re.search) on both hipinfo and amd-smi output, dedup tokens preserving order, and select via new _pick_visible_index() helper. - HIP_VISIBLE_DEVICES / ROCR_VISIBLE_DEVICES (first comma entry, integer) now picks the right GPU on multi-AMD-GPU hosts. Out-of-range or non-int values fall back to the first GPU (matches detect_host behaviour in install_llama_prebuilt.py). 4. studio/install_python_stack.py Strix override now consults the runtime target before flipping: - Previous behaviour intersected gfx_codes with {gfx1151, gfx1150} and picked the first Strix arch, ignoring whether HIP_VISIBLE_DEVICES selected a non-Strix sibling (e.g. discrete RX 7900 in a mixed APU+dGPU box). Could install Strix-specific wheels onto a gfx1100 dGPU. - Now resolves the runtime gfx via _pick_visible_index() and only overrides when that runtime target is in the Strix set. 5. studio/backend/main.py + studio/backend/core/training/worker.py: ROCm version dir scan no longer sorts lexically. Previous sort placed "10.0" before "7.0" alphabetically, which would mis-prioritise ROCm 10.x bin dirs once AMD ships them. New _ver_key() splits on "." and sorts numerically with a string fallback. 6. install.sh Strix override URL: replaced ${var%/} (strips one trailing slash) with a while-loop that strips all trailing slashes, matching Python's .rstrip("/"). A user setting UNSLOTH_AMD_ROCM_MIRROR with "http://corp/whl///" no longer ends up with "http://corp/whl///gfx1151/" which strict pip proxies (artifactory, sonatype) 404 on. 7. studio/install_python_stack.py: bumped torch import probe timeout from 30s to 90s. PyTorch's lazy .so loading can take 60-90s on cold NFS or USB-backed venvs. The shorter timeout was producing a false "torch missing" classification and reinstalling a working ROCm torch. Tests: 231 passed, 1 skipped. sim_5301 30 cases pass (added 7 new sims for multi-GPU detection, Strix sibling handling, and _ok-init regression). --- install.sh | 8 +- studio/backend/core/training/worker.py | 13 ++- studio/backend/main.py | 20 +++- studio/install_python_stack.py | 130 ++++++++++++++++------ tests/studio/install/test_rocm_support.py | 2 +- 5 files changed, 132 insertions(+), 41 deletions(-) diff --git a/install.sh b/install.sh index a13ed0a3d0..92531a1f88 100755 --- a/install.sh +++ b/install.sh @@ -1852,7 +1852,13 @@ case "$TORCH_INDEX_URL" in # over the pytorch.org rocm7.2 fallback because it exercises the real GPU # kernel path. Set UNSLOTH_AMD_ROCM_MIRROR to override for air-gapped installs. _amd_strix_base="${UNSLOTH_AMD_ROCM_MIRROR:-https://repo.amd.com/rocm/whl}" - TORCH_INDEX_URL="${_amd_strix_base%/}/${_strix_gfx}/" + # Strip ALL trailing slashes to match Python's .rstrip("/") -- a + # double-/triple-slash mirror URL would otherwise produce 404s on + # strict pip proxies (artifactory, sonatype). + while [ "${_amd_strix_base%/}" != "$_amd_strix_base" ]; do + _amd_strix_base="${_amd_strix_base%/}" + done + TORCH_INDEX_URL="${_amd_strix_base}/${_strix_gfx}/" TORCH_CONSTRAINT="torch>=2.11.0,<2.12.0" _amd_gpu_radeon=false fi diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py index aa98a20366..3bcf605323 100644 --- a/studio/backend/core/training/worker.py +++ b/studio/backend/core/training/worker.py @@ -91,9 +91,20 @@ if sys.platform == "win32": _default_root = os.path.join( os.environ.get("ProgramFiles", r"C:\Program Files"), "AMD", "ROCm" ) + + def _ver_key(name: str) -> tuple: + # Numeric tuple key so "10.0" sorts after "7.0"; non-numeric chunks fall back to string. + parts = [] + for chunk in name.split("."): + try: + parts.append((0, int(chunk))) + except ValueError: + parts.append((1, chunk)) + return tuple(parts) + try: if os.path.isdir(_default_root): - for _ver in sorted(os.listdir(_default_root), reverse = True): + for _ver in sorted(os.listdir(_default_root), key = _ver_key, reverse = True): _bin = os.path.join(_default_root, _ver, "bin") if os.path.isdir(_bin): _candidates.append(_bin) diff --git a/studio/backend/main.py b/studio/backend/main.py index 5d1670c182..576e15e21e 100644 --- a/studio/backend/main.py +++ b/studio/backend/main.py @@ -32,9 +32,19 @@ if sys.platform == "win32": _default_root = os.path.join( os.environ.get("ProgramFiles", r"C:\Program Files"), "AMD", "ROCm" ) + def _ver_key(name: str) -> tuple: + # Numeric tuple key so "10.0" sorts after "7.0"; non-numeric chunks fall back to string. + parts = [] + for chunk in name.split("."): + try: + parts.append((0, int(chunk))) + except ValueError: + parts.append((1, chunk)) + return tuple(parts) + try: if os.path.isdir(_default_root): - for _ver in sorted(os.listdir(_default_root), reverse = True): + for _ver in sorted(os.listdir(_default_root), key = _ver_key, reverse = True): _bin = os.path.join(_default_root, _ver, "bin") if os.path.isdir(_bin): candidates.append(_bin) @@ -66,7 +76,13 @@ if sys.platform == "win32": try: import torch as _torch_probe - _is_rocm_host = bool(getattr(_torch_probe.version, "hip", None)) + # Broad check: torch.version.hip OR "rocm" in torch.__version__ -- + # AMD SDK / Radeon wheels may not populate torch.version.hip but + # still encode "rocm" in __version__. Matches worker.py + hardware.py. + _is_rocm_host = bool( + getattr(getattr(_torch_probe, "version", None), "hip", None) + or "rocm" in getattr(_torch_probe, "__version__", "").lower() + ) del _torch_probe except Exception: pass diff --git a/studio/install_python_stack.py b/studio/install_python_stack.py index c2bdd30370..43c972b5ba 100644 --- a/studio/install_python_stack.py +++ b/studio/install_python_stack.py @@ -227,6 +227,28 @@ def _detect_rocm_version() -> tuple[int, int] | None: 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. @@ -234,6 +256,10 @@ def _detect_windows_gfx_arch() -> str | None: 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. """ import re @@ -242,6 +268,15 @@ def _detect_windows_gfx_arch() -> str | None: if _override and _override.strip(): return _override.strip().lower() + def _dedup_pick(tokens: list[str]) -> "str | None": + if not tokens: + return None + _seen: list[str] = [] + for _t in tokens: + if _t not in _seen: + _seen.append(_t) + return _seen[_pick_visible_index(len(_seen))] + # 2. hipinfo via PATH, then HIP_PATH\bin / ROCM_PATH\bin. hipinfo = shutil.which("hipinfo") if not hipinfo: @@ -262,10 +297,14 @@ def _detect_windows_gfx_arch() -> str | None: ) if result.returncode == 0: text = result.stdout.decode(errors = "replace") - m = re.search(r"(?im)^\s*gcnArchName\s*:\s*(\S+)", text) - if m: - # Lowercase -- hipinfo sometimes emits "Gfx1151". - return m.group(1).strip().lower() + # 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 @@ -283,20 +322,17 @@ def _detect_windows_gfx_arch() -> str | None: if result.returncode != 0: continue text = result.stdout.decode(errors = "replace") - # Anchor on a labelled gfx line (e.g. "TARGET_GRAPHICS_VERSION: gfx1151" - # or "Arch: gfx1151") to avoid catching stray gfx mentions in - # warnings or device-name strings. Fall back to a bare token match - # only if no labelled line is found. - m = re.search( + # 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, ) - if not m: - m = re.search(r"\bgfx[1-9][0-9a-z]{2,3}\b", text.lower()) - if m: - return m.group(0) - continue - return m.group(1).lower() + _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 @@ -436,6 +472,7 @@ def _install_bnb_windows_rocm() -> bool: return False _prev = os.environ.get("UV_SKIP_WHEEL_FILENAME_CHECK") os.environ["UV_SKIP_WHEEL_FILENAME_CHECK"] = "1" + _ok = False # init so a raise inside pip_install_try does not produce UnboundLocalError try: _ok = pip_install_try( "bitsandbytes (AMD Windows, pre-release main)", @@ -491,7 +528,7 @@ def _ensure_rocm_torch() -> None: ], stdout = subprocess.DEVNULL, stderr = subprocess.DEVNULL, - timeout = 30, + timeout = 90, ) _torch_ok = _probe.returncode == 0 except (OSError, subprocess.TimeoutExpired): @@ -529,7 +566,7 @@ def _ensure_rocm_torch() -> None: ], stdout = subprocess.PIPE, stderr = subprocess.DEVNULL, - timeout = 30, + timeout = 90, ) if probe.returncode == 0 and probe.stdout.decode().strip() == "yes": _torch_already_rocm = True @@ -609,7 +646,7 @@ def _ensure_rocm_torch() -> None: ], stdout = subprocess.PIPE, stderr = subprocess.DEVNULL, - timeout = 30, + timeout = 90, ) except (OSError, subprocess.TimeoutExpired): probe = None @@ -625,6 +662,10 @@ def _ensure_rocm_torch() -> None: # 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): @@ -632,25 +673,42 @@ def _ensure_rocm_torch() -> None: _strix_gfx = {"gfx1151", "gfx1150"} _detected_strix = _strix_gfx.intersection(gfx_codes) if _detected_strix: - _gfx_str = ", ".join(sorted(_detected_strix)) - _selected_gfx = sorted(_detected_strix)[0] - _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", - "torchvision", - "torchaudio", - ) - print( - f"\n ⚠️ {_gfx_str} (AMD Strix) detected with ROCm {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" + # 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", + "torchvision", + "torchaudio", + ) + 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" + ) if not has_hip_torch: if _strix_override_url is not None and _strix_override_pkgs is not None: diff --git a/tests/studio/install/test_rocm_support.py b/tests/studio/install/test_rocm_support.py index bf6ab2a9bc..3b3ebd206d 100644 --- a/tests/studio/install/test_rocm_support.py +++ b/tests/studio/install/test_rocm_support.py @@ -2319,7 +2319,7 @@ class TestStrixRocm71Override: # The URL must incorporate the detected gfx arch so gfx1151 → .../gfx1151/ strix_idx = source.find("_amd_strix_base") assert strix_idx != -1 - ctx = source[strix_idx : strix_idx + 200] + ctx = source[strix_idx : strix_idx + 500] assert "_strix_gfx" in ctx def test_radeon_repo_bypassed_for_strix_in_install_sh(self):