fix(studio/rocm): gate ROCm-only side-effects on active torch runtime
Address five edge cases flagged during PR review:
1. studio/backend/main.py: BNB_ROCM_VERSION was set whenever HIP_PATH or
ROCM_PATH was present in the environment. A Windows CUDA user who once
installed the HIP SDK and reverted to a CUDA torch wheel still has those
env vars set, so bitsandbytes would try to load libbitsandbytes_rocm72.dll
against a CUDA torch and crash. Now probe torch.version.hip inside the
env-var guard (worker.py already does this).
2. studio/backend/main.py: os.add_dll_directory returned handles were
discarded. Per CPython docs, the directory leaves the DLL search list when
the handle is garbage collected. Retain handles in module-level
_ROCM_DLL_HANDLES list so they survive process lifetime.
3. studio/install_python_stack.py: _install_bnb_windows_rocm() returned None
regardless of pip_install_try outcome, and the caller flipped
_rocm_windows_torch_installed to True unconditionally. On a failed BNB
install the post-install "manual install may be required" warning was
suppressed and the user was misled. Helper now returns bool; caller gates
on it.
4. studio/install_python_stack.py: _detect_windows_gfx_arch returned the raw
capture group, so mixed-case hipinfo output ("Gfx1151") missed the
lowercase keys in _GFX_TO_AMD_INDEX_ARCH and silently fell back to CPU
torch. Lowercase the token.
5. studio/install_python_stack.py: UNSLOTH_ROCM_TORCH_INSTALLED=1 early-
return trusted the env var even when the venv was wiped between runs.
Subprocess-probe torch importability first; fall through to the full
install path if the probe fails.
Tests: 231 passed, 1 skipped in tests/studio/install/test_rocm_support.py
(adds one new test for case 5 fall-through).
This commit is contained in:
parent
692e876127
commit
0c2020d5a1
3 changed files with 94 additions and 23 deletions
|
|
@ -259,7 +259,10 @@ def _detect_windows_gfx_arch() -> str | None:
|
|||
return None
|
||||
text = result.stdout.decode(errors = "replace")
|
||||
m = re.search(r"(?im)^\s*gcnArchName\s*:\s*(\S+)", text)
|
||||
return m.group(1).strip() if m else None
|
||||
# Lowercase the captured token -- some hipinfo builds emit "Gfx1151"
|
||||
# which would miss the lowercase keys in _GFX_TO_AMD_INDEX_ARCH and
|
||||
# silently fall back to CPU torch.
|
||||
return m.group(1).strip().lower() if m else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
|
@ -384,8 +387,8 @@ def _detect_amd_gfx_codes() -> list[str]:
|
|||
_rocm_windows_torch_installed: bool = False
|
||||
|
||||
|
||||
def _install_bnb_windows_rocm() -> None:
|
||||
"""Install the AMD Windows BNB prerelease wheel.
|
||||
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
|
||||
|
|
@ -395,11 +398,11 @@ def _install_bnb_windows_rocm() -> None:
|
|||
"""
|
||||
_bnb_win_url = _BNB_ROCM_PRERELEASE_URLS.get("win_amd64")
|
||||
if _bnb_win_url is None:
|
||||
return
|
||||
return False
|
||||
_prev = os.environ.get("UV_SKIP_WHEEL_FILENAME_CHECK")
|
||||
os.environ["UV_SKIP_WHEEL_FILENAME_CHECK"] = "1"
|
||||
try:
|
||||
pip_install_try(
|
||||
_ok = pip_install_try(
|
||||
"bitsandbytes (AMD Windows, pre-release main)",
|
||||
"--force-reinstall",
|
||||
"--no-cache-dir",
|
||||
|
|
@ -412,6 +415,8 @@ def _install_bnb_windows_rocm() -> None:
|
|||
os.environ.pop("UV_SKIP_WHEEL_FILENAME_CHECK", None)
|
||||
else:
|
||||
os.environ["UV_SKIP_WHEEL_FILENAME_CHECK"] = _prev
|
||||
if not _ok:
|
||||
return False
|
||||
# After install: detect the actual ROCm DLL suffix from the wheel so any
|
||||
# post-install BNB import in this process loads the correct DLL.
|
||||
# The worker subprocess does the same detection independently (worker.py §1f).
|
||||
|
|
@ -419,6 +424,7 @@ def _install_bnb_windows_rocm() -> None:
|
|||
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:
|
||||
|
|
@ -431,14 +437,38 @@ def _ensure_rocm_torch() -> None:
|
|||
Uses pip_install() to respect uv, constraints, and --python targeting.
|
||||
"""
|
||||
global _rocm_windows_torch_installed
|
||||
# setup.ps1 sets this when it already installed AMD wheels; skip the probe.
|
||||
# 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":
|
||||
_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 (no libbitsandbytes_rocm72.dll).
|
||||
_install_bnb_windows_rocm()
|
||||
return
|
||||
_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 = 30,
|
||||
)
|
||||
_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
|
||||
|
||||
|
|
@ -488,11 +518,13 @@ def _ensure_rocm_torch() -> None:
|
|||
"torchaudio",
|
||||
constrain = False,
|
||||
)
|
||||
# Always install AMD Windows bitsandbytes — the PyPI wheel ships only
|
||||
# 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.
|
||||
_install_bnb_windows_rocm()
|
||||
_rocm_windows_torch_installed = True
|
||||
# Only flip the success flag when the install actually succeeds; otherwise
|
||||
# the post-install "manual install may be required" warning is suppressed.
|
||||
if _install_bnb_windows_rocm():
|
||||
_rocm_windows_torch_installed = True
|
||||
return
|
||||
|
||||
# ── Linux x86_64 only: PyTorch ROCm wheels are not published for aarch64 ──
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue