diff --git a/.github/scripts/assert-llama-loads.sh b/.github/scripts/assert-llama-loads.sh
new file mode 100755
index 0000000000..c2ffe27469
--- /dev/null
+++ b/.github/scripts/assert-llama-loads.sh
@@ -0,0 +1,57 @@
+#!/usr/bin/env bash
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
+#
+# Assert Studio installed a llama.cpp that loads and runs on THIS macOS. Tests
+# the contract that matters (binaries load and their minimum-OS is <= this host)
+# instead of the old "did install.sh fall back to a source build?" grep, since a
+# source build with a correct deployment target is a valid outcome.
+set -uo pipefail
+
+UNSLOTH_HOME="${STUDIO_HOME:-$HOME/.unsloth}"
+LLAMA_DIR="${LLAMA_CPP_DIR:-$UNSLOTH_HOME/llama.cpp}"
+BIN_DIR="$LLAMA_DIR/build/bin"
+
+fail() {
+ echo "::error::$*"
+ if [ -f logs/install.log ]; then
+ echo "---- install.log (llama.cpp lines) ----"
+ grep -E "llama-prebuilt|llama\.cpp|macos prebuilt|falling back" logs/install.log | tail -80 || true
+ fi
+ exit 1
+}
+
+SERVER="$(find "$LLAMA_DIR" -type f -name 'llama-server' 2>/dev/null | head -1)"
+QUANT="$(find "$LLAMA_DIR" -type f -name 'llama-quantize' 2>/dev/null | head -1)"
+[ -n "$SERVER" ] || fail "llama-server not found under $LLAMA_DIR after install"
+[ -n "$QUANT" ] || fail "llama-quantize not found under $LLAMA_DIR after install"
+
+HOST_VER="$(sw_vers -productVersion 2>/dev/null || echo '0')"
+HOST_MAJOR="${HOST_VER%%.*}"
+
+# Static minimum-OS check on every Mach-O we ship. vtool ships with the Xcode
+# command line tools, which GitHub macOS runners always have; if it is somehow
+# missing we skip the static check and rely on the runtime launch below.
+if command -v vtool >/dev/null 2>&1; then
+ while IFS= read -r macho; do
+ [ -n "$macho" ] || continue
+ minos="$(vtool -show-build "$macho" 2>/dev/null | awk '/minos/{print $2; exit}')"
+ [ -n "$minos" ] || continue
+ min_major="${minos%%.*}"
+ if [ "$min_major" -gt "$HOST_MAJOR" ] 2>/dev/null; then
+ fail "$(basename "$macho") is built for macOS $minos but this runner is macOS $HOST_VER (prebuilt is newer than the host)"
+ fi
+ done < <(find "$BIN_DIR" -type f \( -name '*.dylib' -o -name 'llama-server' -o -name 'llama-quantize' \) 2>/dev/null)
+fi
+
+# Runtime launch: --version forces dyld to load every linked dylib (including
+# libggml-metal.dylib). A missing Metal symbol or too-new binary fails here.
+if ! "$SERVER" --version >/tmp/llama-server-version.txt 2>&1; then
+ echo "---- llama-server --version output ----"
+ cat /tmp/llama-server-version.txt || true
+ fail "llama-server failed to launch on macOS $HOST_VER (dyld load / symbol error)"
+fi
+
+echo "llama.cpp load validation passed on macOS $HOST_VER"
+echo " server: $SERVER"
+sed -n '1,4p' /tmp/llama-server-version.txt 2>/dev/null || true
diff --git a/.github/workflows/lint-ci.yml b/.github/workflows/lint-ci.yml
index 00e6e357e2..8cb3535075 100644
--- a/.github/workflows/lint-ci.yml
+++ b/.github/workflows/lint-ci.yml
@@ -79,6 +79,56 @@ jobs:
run: |
ruff check unsloth unsloth_cli studio tests cli.py unsloth-cli.py
+ - name: Import-hoist verifier self-test
+ # scripts/verify_import_hoist.py is a scope-aware (LEGB) AST
+ # resolver that gates import-hoisting / alias-rename refactors
+ # against two bugs ruff and pyflakes both miss:
+ # 1. dangling alias -- `from a import b as _b` hoisted to
+ # `from a import b` but a leftover `_b` reference now
+ # resolves to nothing (or to some other module-level `_b`).
+ # 2. rename clash -- `_b -> b` silently re-points at a
+ # different object already named `b` in that scope.
+ # This step runs the tool's 8 negative-control cases so a
+ # regression in the verifier itself fails before we trust it on
+ # a diff. Hermetic, stdlib-only, sub-second. Hard gate.
+ run: |
+ python scripts/verify_import_hoist.py --self-test
+
+ - name: Import-hoist / alias-rename safety (changed Python files)
+ # Runs the verifier in compare mode on every in-place-modified
+ # .py in the PR: parses each file BEFORE (base branch) and AFTER
+ # (this diff), resolves every name load, and fails on a BLOCKER
+ # (dangling alias / rename clash / re-pointed import). INFO
+ # findings (a helper relocated to another file) do not fail.
+ #
+ # --diff-filter=M (in-place edits only) is deliberate: that is
+ # exactly where a hoist refactor lives, and it skips brand-new
+ # files whose re-export imports would otherwise look "unused".
+ #
+ # actions/checkout uses fetch-depth: 1, so the base branch is not
+ # present locally. Fetch the single base commit with an explicit
+ # refspec so origin/ is reliably created (a bare
+ # `git fetch origin [` only updates FETCH_HEAD in some
+ # configs). Two-dot diff avoids needing a merge-base on a shallow
+ # clone.
+ if: github.event_name == 'pull_request'
+ run: |
+ git fetch --no-tags --depth=1 origin \
+ "${{ github.base_ref }}:refs/remotes/origin/${{ github.base_ref }}"
+ mapfile -t CHANGED < <(
+ git diff --name-only --diff-filter=M \
+ "origin/${{ github.base_ref }}" HEAD -- '*.py' \
+ | grep -vE '(^|/)(unsloth_compiled_cache|node_modules|build|dist)/' || true
+ )
+ if [ "${#CHANGED[@]}" -eq 0 ]; then
+ echo "no in-place-modified Python files to check"
+ exit 0
+ fi
+ printf 'checking %d file(s):\n' "${#CHANGED[@]}"
+ printf ' %s\n' "${CHANGED[@]}"
+ python scripts/verify_import_hoist.py \
+ --before "origin/${{ github.base_ref }}" --after HEAD "${CHANGED[@]}"
+
- name: No leftover debugger / pdb / breakpoint calls
# Catches the "I'll just stick a breakpoint() here" mistake
# before it ships. AST-based so commented-out debugger
diff --git a/.github/workflows/studio-mac-api-smoke.yml b/.github/workflows/studio-mac-api-smoke.yml
index b4e274155e..8839b559fd 100644
--- a/.github/workflows/studio-mac-api-smoke.yml
+++ b/.github/workflows/studio-mac-api-smoke.yml
@@ -89,13 +89,8 @@ jobs:
set -o pipefail
bash install.sh --local --no-torch 2>&1 | tee logs/install.log
- - name: Assert install.sh used the Mac llama.cpp prebuilt
- run: |
- if grep -q "falling back to source build" logs/install.log; then
- echo "::error::install.sh fell back to source-build llama.cpp on Mac. Studio must install the prebuilt llama-bNNNN-bin-macos-arm64 on Apple Silicon."
- grep -E "llama-prebuilt|llama.cpp" logs/install.log | tail -60
- exit 1
- fi
+ - name: Assert llama.cpp loads on this macOS
+ run: bash .github/scripts/assert-llama-loads.sh
- name: Install pyjwt for the JWT-expiry forge test
run: pip install 'pyjwt>=2.6'
diff --git a/.github/workflows/studio-mac-inference-smoke.yml b/.github/workflows/studio-mac-inference-smoke.yml
index fab0a36bd1..1096b1abb4 100644
--- a/.github/workflows/studio-mac-inference-smoke.yml
+++ b/.github/workflows/studio-mac-inference-smoke.yml
@@ -114,13 +114,8 @@ jobs:
set -o pipefail
bash install.sh --local --no-torch 2>&1 | tee logs/install.log
- - name: Assert install.sh used the Mac llama.cpp prebuilt
- run: |
- if grep -q "falling back to source build" logs/install.log; then
- echo "::error::install.sh fell back to source-build llama.cpp on Mac. Studio must install the prebuilt llama-bNNNN-bin-macos-arm64 on Apple Silicon."
- grep -E "llama-prebuilt|llama.cpp" logs/install.log | tail -60
- exit 1
- fi
+ - name: Assert llama.cpp loads on this macOS
+ run: bash .github/scripts/assert-llama-loads.sh
- name: Install OpenAI + Anthropic Python SDKs
run: pip install 'openai>=1.50' 'anthropic>=0.40'
@@ -369,13 +364,8 @@ jobs:
set -o pipefail
bash install.sh --local --no-torch 2>&1 | tee logs/install.log
- - name: Assert install.sh used the Mac llama.cpp prebuilt
- run: |
- if grep -q "falling back to source build" logs/install.log; then
- echo "::error::install.sh fell back to source-build llama.cpp on Mac. Studio must install the prebuilt llama-bNNNN-bin-macos-arm64 on Apple Silicon."
- grep -E "llama-prebuilt|llama.cpp" logs/install.log | tail -60
- exit 1
- fi
+ - name: Assert llama.cpp loads on this macOS
+ run: bash .github/scripts/assert-llama-loads.sh
- name: Reset auth + boot Studio (API-only, default tool policy)
# We deliberately use the API-only mode rather than
@@ -760,13 +750,8 @@ jobs:
set -o pipefail
bash install.sh --local --no-torch 2>&1 | tee logs/install.log
- - name: Assert install.sh used the Mac llama.cpp prebuilt
- run: |
- if grep -q "falling back to source build" logs/install.log; then
- echo "::error::install.sh fell back to source-build llama.cpp on Mac. Studio must install the prebuilt llama-bNNNN-bin-macos-arm64 on Apple Silicon."
- grep -E "llama-prebuilt|llama.cpp" logs/install.log | tail -60
- exit 1
- fi
+ - name: Assert llama.cpp loads on this macOS
+ run: bash .github/scripts/assert-llama-loads.sh
- name: Install OpenAI + Anthropic Python SDKs
run: pip install 'openai>=1.50' 'anthropic>=0.40'
diff --git a/.github/workflows/studio-mac-install-matrix.yml b/.github/workflows/studio-mac-install-matrix.yml
new file mode 100644
index 0000000000..4e2722d1cd
--- /dev/null
+++ b/.github/workflows/studio-mac-install-matrix.yml
@@ -0,0 +1,80 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
+
+# Proves Studio's llama.cpp install loads on every supported macOS. The heavy
+# app smokes stay single-OS; this matrix covers the OS-version dimension cheaply
+# (install.sh + binary-load assert). Regression guard for the macOS-version
+# selection in studio/install_llama_prebuilt.py.
+
+name: Mac Studio Install Matrix CI
+
+on:
+ pull_request:
+ paths:
+ - 'studio/install_llama_prebuilt.py'
+ - 'studio/setup.sh'
+ - 'install.sh'
+ - '.github/scripts/assert-llama-loads.sh'
+ - '.github/workflows/studio-mac-install-matrix.yml'
+ push:
+ branches: [main, pip]
+ workflow_dispatch:
+
+concurrency:
+ group: ${{ github.workflow }}-${{ github.ref }}
+ cancel-in-progress: true
+
+permissions:
+ contents: read
+
+jobs:
+ install-load:
+ name: Install + load (${{ matrix.os }})
+ runs-on: ${{ matrix.os }}
+ timeout-minutes: 25
+ continue-on-error: ${{ matrix.experimental }}
+ strategy:
+ fail-fast: false
+ matrix:
+ include:
+ - os: macos-14 # Apple Silicon, macOS 14 Sonoma
+ experimental: false
+ - os: macos-15 # Apple Silicon, macOS 15 Sequoia
+ experimental: false
+ - os: macos-26 # Apple Silicon, macOS 26 Tahoe
+ experimental: false
+ - os: macos-15-intel # Intel x86_64, macOS 15 (informational)
+ experimental: true
+ - os: macos-26-intel # Intel x86_64, macOS 26 (last Intel macOS)
+ experimental: true
+ steps:
+ - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
+ with:
+ persist-credentials: false
+
+ - uses: actions/setup-node@48b55a011bda9f5d6aeb4c2d9c7362e8dae4041e # v6.4.0
+ with:
+ node-version: '22'
+
+ - uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6.2.0
+ with:
+ python-version: '3.12'
+
+ - name: Install Studio (--local, --no-torch)
+ env:
+ GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
+ run: |
+ mkdir -p logs
+ set -o pipefail
+ bash install.sh --local --no-torch 2>&1 | tee logs/install.log
+
+ - name: Assert llama.cpp loads on this macOS
+ run: bash .github/scripts/assert-llama-loads.sh
+
+ - name: Upload install log
+ if: always()
+ uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1
+ with:
+ name: mac-install-matrix-${{ matrix.os }}-log
+ path: logs/install.log
+ retention-days: 7
diff --git a/.github/workflows/studio-mac-ui-smoke.yml b/.github/workflows/studio-mac-ui-smoke.yml
index b353f0ec83..0176b0a168 100644
--- a/.github/workflows/studio-mac-ui-smoke.yml
+++ b/.github/workflows/studio-mac-ui-smoke.yml
@@ -89,13 +89,8 @@ jobs:
set -o pipefail
bash install.sh --local --no-torch 2>&1 | tee logs/install.log
- - name: Assert install.sh used the Mac llama.cpp prebuilt
- run: |
- if grep -q "falling back to source build" logs/install.log; then
- echo "::error::install.sh fell back to source-build llama.cpp on Mac. Studio must install the prebuilt llama-bNNNN-bin-macos-arm64 on Apple Silicon."
- grep -E "llama-prebuilt|llama.cpp" logs/install.log | tail -60
- exit 1
- fi
+ - name: Assert llama.cpp loads on this macOS
+ run: bash .github/scripts/assert-llama-loads.sh
- name: Install Playwright + Chromium
# No --with-deps on Mac: that flag installs Linux apt packages.
diff --git a/.github/workflows/studio-mac-update-smoke.yml b/.github/workflows/studio-mac-update-smoke.yml
index b65439f174..1eca227096 100644
--- a/.github/workflows/studio-mac-update-smoke.yml
+++ b/.github/workflows/studio-mac-update-smoke.yml
@@ -67,21 +67,8 @@ jobs:
set -o pipefail
bash install.sh --local --no-torch 2>&1 | tee logs/install.log
- - name: Assert install.sh used the Mac llama.cpp prebuilt
- run: |
- # Mac install must take the prebuilt path. Source-build
- # fallback here is an Unsloth bug.
- if grep -q "falling back to source build" logs/install.log; then
- echo "::error::install.sh fell back to source-build llama.cpp on Mac. Studio must install the prebuilt llama-bNNNN-bin-macos-arm64 on Apple Silicon."
- grep -E "llama-prebuilt|llama.cpp" logs/install.log | tail -60
- exit 1
- fi
- if ! grep -qE "prebuilt installed and validated|prebuilt up to date and validated|bin-macos-arm64" logs/install.log; then
- echo "::error::no Mac prebuilt llama.cpp marker in install.log."
- grep -E "llama-prebuilt|llama.cpp" logs/install.log | tail -60
- exit 1
- fi
- echo "install.sh installed the Mac prebuilt llama.cpp"
+ - name: Assert llama.cpp loads on this macOS
+ run: bash .github/scripts/assert-llama-loads.sh
- name: First update should be a no-op (prebuilt already validated)
env:
diff --git a/README.md b/README.md
index ecf0f8a7a3..948d84a789 100644
--- a/README.md
+++ b/README.md
@@ -202,7 +202,7 @@ unsloth studio -p 8888
#### Nightly: Windows:
Run in Windows Powershell:
-```bash
+```powershell
git clone https://github.com/unslothai/unsloth.git
cd unsloth
git checkout nightly
diff --git a/install.ps1 b/install.ps1
index c4e8d8d522..cab66f5ae1 100644
--- a/install.ps1
+++ b/install.ps1
@@ -976,6 +976,7 @@ shell.Run cmd, 0, False
# Find-CompatiblePython returns @{ Version = "3.13"; Path = "C:\...\python.exe" } or $null.
Write-TauriLog "STEP" "Installing Python"
$DetectedPython = Find-CompatiblePython
+
if ($DetectedPython) {
step "python" "Python $($DetectedPython.Version) already installed"
}
@@ -1239,11 +1240,196 @@ shell.Run cmd, 0, False
}
}
}
+ # ── AMD ROCm detection (Windows) — mirrors setup.ps1 ──
+ $HasROCm = $false
+ $HipSdkInstalled = $false # HIP SDK binary found (independent of device accessibility)
+ $ROCmGpuLabel = $null
+ $ROCmVersion = $null
+ $ROCmGfxArch = $null
+ if (-not $HasNvidiaSmi) {
+ # hipinfo: PATH first, then HIP_PATH/ROCM_PATH bin fallback (mirrors NVIDIA smi path resolution).
+ # AMD HIP SDK sets HIP_PATH but may not add the bin dir to PATH depending on install type.
+ $hipinfoExe = Get-Command hipinfo -ErrorAction SilentlyContinue
+ if (-not $hipinfoExe) {
+ $hipRoot = if ($env:HIP_PATH) { $env:HIP_PATH } elseif ($env:ROCM_PATH) { $env:ROCM_PATH } else { $null }
+ $hipEnvLabel = if ($env:HIP_PATH) { "HIP_PATH" } else { "ROCM_PATH" }
+ if ($hipRoot) {
+ $hipinfoCandidate = Join-Path $hipRoot "bin\hipinfo.exe"
+ if (Test-Path $hipinfoCandidate) {
+ Write-Host " [WARN] hipinfo not on PATH -- located via ${hipEnvLabel}: $hipinfoCandidate" -ForegroundColor Yellow
+ Write-Host " Add '$(Join-Path $hipRoot 'bin')' to your PATH to suppress this warning" -ForegroundColor Yellow
+ Write-Host " Quick fix: [Environment]::SetEnvironmentVariable('PATH',`$env:PATH+';$(Join-Path $hipRoot 'bin')','User')" -ForegroundColor Yellow
+ $hipinfoExe = [PSCustomObject]@{ Source = $hipinfoCandidate }
+ } else {
+ Write-Host " [WARN] ${hipEnvLabel}=$hipRoot is set but hipinfo.exe not found at $hipinfoCandidate" -ForegroundColor Yellow
+ Write-Host " HIP SDK install may be incomplete -- re-install from:" -ForegroundColor Yellow
+ Write-Host " https://rocm.docs.amd.com/en/latest/deploy/windows/index.html" -ForegroundColor Yellow
+ }
+ }
+ }
+ if ($hipinfoExe) {
+ $HipSdkInstalled = $true # binary found → SDK is installed regardless of device state
+ try {
+ $hipOut = & $hipinfoExe.Source 2>&1 | Out-String
+ if ($LASTEXITCODE -eq 0 -and $hipOut -match "(?i)gcnArchName") {
+ $HasROCm = $true
+ $_hipAllArches = @([regex]::Matches($hipOut, "(?im)^\s*gcnArchName\s*:\s*(\S+)") | ForEach-Object { ($_.Groups[1].Value -split ':')[0].Trim().ToLower() })
+ $_hipVisIdx = if ($env:HIP_VISIBLE_DEVICES -match '^\d') { [int]($env:HIP_VISIBLE_DEVICES -split ',')[0] } elseif ($env:ROCR_VISIBLE_DEVICES -match '^\d') { [int]($env:ROCR_VISIBLE_DEVICES -split ',')[0] } else { 0 }
+ if ($_hipAllArches.Count -gt 0) {
+ $ROCmGfxArch = if ($_hipVisIdx -lt $_hipAllArches.Count) { $_hipAllArches[$_hipVisIdx] } else { $_hipAllArches[0] }
+ $ROCmGpuLabel = "AMD ROCm ($ROCmGfxArch)"
+ } else {
+ $ROCmGpuLabel = "AMD ROCm"
+ }
+ } elseif ($LASTEXITCODE -ne 0) {
+ # hipinfo ran but returned a HIP runtime error (e.g. "no ROCm-capable device detected")
+ $firstLine = ($hipOut -split '\r?\n' | Where-Object { $_.Trim() } | Select-Object -First 1)
+ Write-Host " [WARN] hipinfo returned a HIP runtime error (exit $LASTEXITCODE)" -ForegroundColor Yellow
+ Write-Host " $firstLine" -ForegroundColor Yellow
+ Write-Host " Ensure ROCm drivers are installed: https://rocm.docs.amd.com/en/latest/deploy/windows/index.html" -ForegroundColor Yellow
+ }
+ } catch {}
+ }
+ if (-not $HasROCm) {
+ $amdSmiExe = Get-Command "amd-smi" -ErrorAction SilentlyContinue
+ if ($amdSmiExe) {
+ try {
+ $smiOut = & $amdSmiExe.Source list 2>&1 | Out-String
+ if ($LASTEXITCODE -eq 0 -and $smiOut -match "(?im)^GPU\s*[:\[]\s*\d") {
+ $HasROCm = $true
+ # Mirror the hipinfo path: collect all gfx tokens in enumeration
+ # order and pick the runtime-visible one via HIP_VISIBLE_DEVICES.
+ $_smiVisIdx = if ($env:HIP_VISIBLE_DEVICES -match '^\d') { [int]($env:HIP_VISIBLE_DEVICES -split ',')[0] } elseif ($env:ROCR_VISIBLE_DEVICES -match '^\d') { [int]($env:ROCR_VISIBLE_DEVICES -split ',')[0] } else { 0 }
+ # Attempt 1: newer amd-smi versions embed the gfx arch in list output.
+ $_smiGfxTokens = @([regex]::Matches($smiOut, "(?i)\b(gfx\d+[a-z]?)\b") | ForEach-Object { $_.Groups[1].Value.ToLower() })
+ if ($_smiGfxTokens.Count -gt 0) {
+ $ROCmGfxArch = if ($_smiVisIdx -lt $_smiGfxTokens.Count) { $_smiGfxTokens[$_smiVisIdx] } else { $_smiGfxTokens[0] }
+ $ROCmGpuLabel = "AMD ROCm ($ROCmGfxArch)"
+ } else {
+ # Attempt 2: 'static --asic' exposes ASIC details on ROCm 6+,
+ # including the GFX target needed for wheel index selection.
+ $smiAsicOut = ""
+ try { $smiAsicOut = & $amdSmiExe.Source static --asic 2>&1 | Out-String } catch {}
+ $_asicGfxTokens = @([regex]::Matches($smiAsicOut, "(?i)\b(gfx\d+[a-z]?)\b") | ForEach-Object { $_.Groups[1].Value.ToLower() })
+ if ($_asicGfxTokens.Count -gt 0) {
+ $ROCmGfxArch = if ($_smiVisIdx -lt $_asicGfxTokens.Count) { $_asicGfxTokens[$_smiVisIdx] } else { $_asicGfxTokens[0] }
+ $ROCmGpuLabel = "AMD ROCm ($ROCmGfxArch)"
+ } elseif ($smiAsicOut -match "(?im)Market.?Name\s*[:\|]\s*([^\r\n]+)") {
+ $ROCmGpuLabel = "AMD ROCm ($($Matches[1].Trim()))"
+ } else {
+ $ROCmGpuLabel = "AMD ROCm"
+ }
+ }
+ }
+ } catch {}
+ }
+ }
+ if (-not $HasROCm) {
+ try {
+ $wmiGpu = Get-WmiObject Win32_VideoController -ErrorAction SilentlyContinue |
+ Where-Object { $_.Name -match "AMD|Radeon" } |
+ Select-Object -First 1
+ if ($wmiGpu) { $ROCmGpuLabel = $wmiGpu.Name }
+ } catch {}
+ }
+ # ── Arch resolution: env-var override → name inference ──────────────
+ # Covers users whose amd-smi is too old to report the GFX target and
+ # who don't have hipinfo (HIP-runtime-only, common on Strix Halo / iGPU).
+ if ($HasROCm -and -not $ROCmGfxArch) {
+ # 1. Manual override: set UNSLOTH_ROCM_GFX_ARCH=gfx1151 before running.
+ if ($env:UNSLOTH_ROCM_GFX_ARCH) {
+ $ROCmGfxArch = $env:UNSLOTH_ROCM_GFX_ARCH.Trim().ToLower()
+ $ROCmGpuLabel = "AMD ROCm ($ROCmGfxArch)"
+ substep "gfx arch from UNSLOTH_ROCM_GFX_ARCH env override: $ROCmGfxArch" "Cyan"
+ }
+ # 2. Best-effort name → arch lookup from marketing name (amd-smi / WMI).
+ elseif ($ROCmGpuLabel) {
+ $nameArchTable = @(
+ @{ P = "9070 XT|9080"; A = "gfx1201" } # RDNA 4
+ @{ P = "9070|9060"; A = "gfx1200" } # RDNA 4
+ @{ P = "8060S|890M|Strix Halo|HX 37[05]|HX 38[05]|AI 9 HX"; A = "gfx1151" } # RDNA 3.5 iGPU (Strix Halo / Radeon 8060S retail)
+ @{ P = "880M|Strix Point|AI 9 36[05]|AI 7 35[05]|AI 5 34[05]"; A = "gfx1150" } # RDNA 3.5 iGPU (Strix Point)
+ @{ P = "RX 7900|RX 7800|RX 7700(?! S)"; A = "gfx1100" } # RDNA 3 desktop
+ @{ P = "RX 7600"; A = "gfx1102" } # RDNA 3
+ @{ P = "780M|760M|740M|Phoenix"; A = "gfx1103" } # RDNA 3 iGPU (Phoenix)
+ )
+ foreach ($row in $nameArchTable) {
+ if ($ROCmGpuLabel -match $row.P) {
+ $ROCmGfxArch = $row.A
+ $ROCmGpuLabel = "AMD ROCm ($ROCmGfxArch)"
+ substep "gfx arch inferred from GPU name: $ROCmGfxArch" "Cyan"
+ substep "Tip: set UNSLOTH_ROCM_GFX_ARCH=$ROCmGfxArch to skip inference next time" "Cyan"
+ break
+ }
+ }
+ }
+ }
+ # Capture ROCm version for wheel selection (hipconfig, then amd-smi).
+ # Run whenever the HIP SDK binary is present, not just when the device is accessible --
+ # hipconfig --version works even when hipinfo reports no ROCm device (driver issue).
+ if ($HasROCm -or $HipSdkInstalled) {
+ $hipConfigExe = Get-Command hipconfig -ErrorAction SilentlyContinue
+ if (-not $hipConfigExe) {
+ $hipRoot = if ($env:HIP_PATH) { $env:HIP_PATH } elseif ($env:ROCM_PATH) { $env:ROCM_PATH } else { $null }
+ if ($hipRoot) {
+ $hipConfigCandidate = Join-Path $hipRoot "bin\hipconfig.exe"
+ if (Test-Path $hipConfigCandidate) {
+ $hipConfigEnvLabel = if ($env:HIP_PATH) { "HIP_PATH" } else { "ROCM_PATH" }
+ Write-Host " [WARN] hipconfig not on PATH -- located via ${hipConfigEnvLabel}: $hipConfigCandidate" -ForegroundColor Yellow
+ $hipConfigExe = [PSCustomObject]@{ Source = $hipConfigCandidate }
+ }
+ }
+ }
+ if ($hipConfigExe) {
+ try {
+ $hipVerOut = & $hipConfigExe.Source --version 2>&1 | Out-String
+ if ($LASTEXITCODE -eq 0) {
+ $hipVerLine = ($hipVerOut -split '\r?\n' | Where-Object { $_.Trim() } | Select-Object -First 1).Trim()
+ if ($hipVerLine -match '(\d+\.\d+)') {
+ $ROCmVersion = $Matches[1]
+ $ROCmVersionFull = $hipVerLine
+ }
+ }
+ } catch {}
+ }
+ if (-not $ROCmVersion) {
+ $amdSmiVer = Get-Command "amd-smi" -ErrorAction SilentlyContinue
+ if ($amdSmiVer) {
+ try {
+ $smiVerOut = & $amdSmiVer.Source version 2>&1 | Out-String
+ if ($LASTEXITCODE -eq 0 -and $smiVerOut -match 'ROCm version:\s*(\d+\.\d+)') {
+ $ROCmVersion = $Matches[1]
+ }
+ } catch {}
+ }
+ }
+ }
+ }
+
if ($HasNvidiaSmi) {
step "gpu" "NVIDIA GPU detected"
+ } elseif ($HasROCm) {
+ step "gpu" $ROCmGpuLabel
+ $hipSdkPath = if ($env:HIP_PATH) { $env:HIP_PATH } elseif ($env:ROCM_PATH) { $env:ROCM_PATH } else { "on system PATH" }
+ substep "HIP SDK: $hipSdkPath"
+ if ($ROCmVersionFull) { substep "hipconfig: $ROCmVersionFull" }
+ } elseif ($HipSdkInstalled -and $ROCmGpuLabel) {
+ # HIP SDK is installed but ROCm can't see the device (driver issue, not SDK issue)
+ $sdkVer = if ($ROCmVersionFull) { " (HIP $ROCmVersionFull)" } else { "" }
+ step "gpu" "AMD GPU detected -- not ROCm-accessible$sdkVer" "Yellow"
+ substep "Detected: $ROCmGpuLabel" "Yellow"
+ substep "[WARN] HIP SDK is installed but hipinfo reports no ROCm-capable device." "Yellow"
+ substep " This is a driver issue, not an SDK issue." "Yellow"
+ substep " Ensure the ROCm compute driver is installed alongside the display driver:" "Yellow"
+ substep " https://rocm.docs.amd.com/en/latest/deploy/windows/index.html" "Yellow"
+ } elseif ($ROCmGpuLabel) {
+ step "gpu" "AMD GPU detected -- HIP SDK not found" "Yellow"
+ substep "Detected: $ROCmGpuLabel" "Yellow"
+ substep "Install the HIP SDK for ROCm GPU inference:" "Yellow"
+ substep "https://rocm.docs.amd.com/en/latest/deploy/windows/index.html" "Yellow"
} else {
step "gpu" "none (chat-only / GGUF)" "Yellow"
- substep "Training and GPU inference require an NVIDIA GPU with drivers installed." "Yellow"
+ substep "Training and GPU inference require an NVIDIA or AMD ROCm GPU." "Yellow"
}
# ── Choose the correct PyTorch index URL based on driver CUDA version ──
@@ -1270,14 +1456,73 @@ shell.Run cmd, 0, False
return "$baseUrl/cu126"
}
$TorchIndexUrl = Get-TorchIndexUrl
- $TorchIndexFamily = Get-TauriTorchIndexFamily $TorchIndexUrl
+
+ # ── GPU arch → newest compatible Windows ROCm wheel release ──
+ # Wheels bundle their own ROCm runtime; the installed HIP SDK version does
+ # not constrain which release to use. Always picks the newest release that
+ # supports the GPU architecture.
+ # ── AMD Windows ROCm: arch-aware pip index (repo.amd.com) ──
+ # Wheels bundle their own ROCm runtime and support all Python versions.
+ # Override with UNSLOTH_ROCM_WINDOWS_MIRROR for air-gapped / mirror installs.
+ $ROCmIndexUrl = $null
+ $ROCmTorchFloor = $null
+ if ($HasROCm -and $TorchIndexUrl -like "*/cpu" -and -not $SkipTorch) {
+ $amdIndexBase = if ($env:UNSLOTH_ROCM_WINDOWS_MIRROR) { $env:UNSLOTH_ROCM_WINDOWS_MIRROR.TrimEnd('/') } else { "https://repo.amd.com/rocm/whl" }
+ $archFamilyMap = @{
+ "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
+ }
+ # gfx120X (RDNA 4) and gfx1151/gfx1150 (Strix) have a null-pointer bug in
+ # torch._C._grouped_mm on torch <2.11.0 (rocm7.12 and rocm7.1 respectively).
+ # TheRock issues #5284 and #3284. Force torch>=2.11.0 so pip never resolves
+ # to the broken 2.10.0 wheels even though they exist on the AMD index.
+ # The <2.12.0 ceiling matches the Linux install_python_stack.py constraint
+ # for the same arches: AMD actively publishes new versions on their index,
+ # so without a ceiling a future 2.12.0+rocmX.Y wheel would be pulled in
+ # automatically before it has been validated on these architectures.
+ # Bump the ceiling here (and in install_python_stack.py) when 2.12.x is
+ # confirmed working on gfx120X / Strix.
+ $torchFloorMap = @{
+ "gfx1201" = "torch>=2.11.0,<2.12.0"; "gfx1200" = "torch>=2.11.0,<2.12.0"
+ "gfx1151" = "torch>=2.11.0,<2.12.0"; "gfx1150" = "torch>=2.11.0,<2.12.0"
+ }
+ $archFamily = if ($ROCmGfxArch -and $archFamilyMap.ContainsKey($ROCmGfxArch)) { $archFamilyMap[$ROCmGfxArch] } else { $null }
+ if ($archFamily) {
+ $ROCmIndexUrl = "$amdIndexBase/$archFamily/"
+ $ROCmTorchFloor = if ($ROCmGfxArch -and $torchFloorMap.ContainsKey($ROCmGfxArch)) { $torchFloorMap[$ROCmGfxArch] } else { $null }
+ $archLabel = if ($ROCmGfxArch) { $ROCmGfxArch } else { "AMD GPU" }
+ substep "$archLabel -- AMD repo.amd.com index selected" "Cyan"
+ if ($ROCmTorchFloor) {
+ substep " enforcing $ROCmTorchFloor (known _grouped_mm bug in older wheels)" "Cyan"
+ }
+ } elseif ($ROCmGfxArch) {
+ substep "AMD GPU ($ROCmGfxArch) not in supported arch list -- falling back to CPU-only PyTorch" "Yellow"
+ } else {
+ substep "AMD GPU detected but arch unknown -- falling back to CPU-only PyTorch" "Yellow"
+ }
+ }
+
+ if ($ROCmIndexUrl) {
+ $TorchIndexFamily = "rocm"
+ } else {
+ $TorchIndexFamily = Get-TauriTorchIndexFamily $TorchIndexUrl
+ }
$GpuBranch = Get-TauriGpuBranch $TorchIndexFamily
Write-TauriDiag -GpuBranch $GpuBranch -TorchIndexFamily $TorchIndexFamily -PythonVersionForDiag $DetectedPython.Version
# ── Print CPU-only hint when no GPU detected ──
- if (-not $SkipTorch -and $TorchIndexUrl -like "*/cpu") {
+ if (-not $SkipTorch -and -not $ROCmIndexUrl -and $TorchIndexUrl -like "*/cpu") {
Write-Host ""
- substep "No NVIDIA GPU detected." "Yellow"
+ if ($HipSdkInstalled -and -not $HasROCm) {
+ substep "Installing CPU-only PyTorch (HIP SDK found but GPU not ROCm-accessible)." "Yellow"
+ } elseif ($ROCmGpuLabel) {
+ substep "Installing CPU-only PyTorch (ROCm wheels require the HIP SDK)." "Yellow"
+ } else {
+ substep "No NVIDIA GPU detected." "Yellow"
+ }
substep "Installing CPU-only PyTorch. If you only need GGUF chat/inference," "Yellow"
substep "re-run with --no-torch for a faster, lighter install:" "Yellow"
substep ".\install.ps1 --no-torch" "Yellow"
@@ -1355,9 +1600,18 @@ shell.Run cmd, 0, False
return (Exit-InstallFailure "Failed to overlay unsloth-zoo (exit code $zooOverlayExit)" $zooOverlayExit)
}
}
- } elseif ($TorchIndexUrl) {
+ } elseif ($TorchIndexUrl -or $ROCmIndexUrl) {
if ($SkipTorch) {
substep "skipping PyTorch (--no-torch flag set)." "Yellow"
+ } elseif ($ROCmIndexUrl) {
+ Write-TauriLog "STEP" "Installing PyTorch (AMD ROCm Windows)"
+ substep "installing PyTorch from $ROCmIndexUrl..."
+ $torchSpec = if ($ROCmTorchFloor) { $ROCmTorchFloor } else { "torch" }
+ $torchInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --force-reinstall --index-url $ROCmIndexUrl $torchSpec torchvision torchaudio }
+ if ($torchInstallExit -ne 0) {
+ Write-Host "[ERROR] Failed to install AMD ROCm PyTorch (exit code $torchInstallExit)" -ForegroundColor Red
+ return (Exit-InstallFailure "Failed to install AMD ROCm PyTorch (exit code $torchInstallExit)" $torchInstallExit)
+ }
} else {
Write-TauriLog "STEP" "Installing PyTorch"
substep "installing PyTorch ($TorchIndexUrl)..."
diff --git a/install.sh b/install.sh
index a3b76a3011..f0af60e2d6 100755
--- a/install.sh
+++ b/install.sh
@@ -183,10 +183,21 @@ _install_bnb_rocm() {
fi
if [ -n "$_bnb_whl_url" ]; then
substep "installing bitsandbytes for AMD ROCm (pre-release, PR #1887)..."
- if run_install_cmd "$_label (pre-release)" "$_venv_py" -m pip install \
- --force-reinstall --no-cache-dir --no-deps "$_bnb_whl_url"; then
+ _bnb_log=$(mktemp)
+ if "$_venv_py" -m pip install \
+ --disable-pip-version-check \
+ --force-reinstall --no-cache-dir --no-deps \
+ --retries 8 --timeout 90 \
+ "$_bnb_whl_url" >"$_bnb_log" 2>&1; then
+ rm -f "$_bnb_log"
return 0
fi
+ _bnb_rc=$?
+ if _is_verbose; then
+ cat "$_bnb_log" >&2
+ fi
+ rm -f "$_bnb_log"
+ step "warning" "$_label (pre-release) failed (exit code $_bnb_rc)" "$C_WARN" >&2
substep "[WARN] bnb pre-release install failed; falling back to PyPI (4-bit decode broken on ROCm)" "$C_WARN"
fi
run_install_cmd "$_label (pypi fallback)" "$_venv_py" -m pip install \
@@ -245,6 +256,9 @@ _tauri_torch_index_family() {
rocm[0-9]*.[0-9]*) echo "$_diag_family" ;;
*) echo "auto" ;;
esac ;;
+ # AMD arch-specific index (e.g. repo.amd.com/rocm/whl/gfx1151/) --
+ # used for Strix Halo/Point where torch 2.11+rocm7.13 has the real fix.
+ *repo.amd.com/rocm/whl/gfx*|*rocm/whl/gfx*) echo "rocm7.13" ;;
"") echo "none" ;;
*) echo "auto" ;;
esac
@@ -1568,16 +1582,19 @@ _find_no_torch_runtime() {
}
# ── AMD ROCm GPU detection helper ──
-# Returns 0 (true) if an actual AMD GPU is present, 1 (false) otherwise.
-# Checks rocminfo for gfx[1-9]* (excludes gfx000 CPU agent) and
-# amd-smi list for GPU data rows (excludes header-only output).
+# Returns 0 if an AMD GPU is present. Checks rocminfo, amd-smi, then sysfs
+# KFD topology (env-var-independent fallback for when HIP/ROCR_VISIBLE_DEVICES hides devices).
_has_amd_rocm_gpu() {
if command -v rocminfo >/dev/null 2>&1 && \
- rocminfo 2>/dev/null | awk '/Name:[[:space:]]*gfx[0-9]/ && !/Name:[[:space:]]*gfx000/{found=1} END{exit !found}'; then
+ rocminfo 2>/dev/null | awk '/Name:[[:space:]]*gfx[1-9][0-9]/{found=1} END{exit !found}'; then
return 0
elif command -v amd-smi >/dev/null 2>&1 && \
amd-smi list 2>/dev/null | awk '/^GPU[[:space:]]*[:\[][[:space:]]*[0-9]/{ found=1 } END{ exit !found }'; then
return 0
+ elif [ -e /dev/kfd ] && \
+ awk '/gpu_id/{ if ($2+0 > 0) found=1 } END{ exit !found }' \
+ /sys/class/kfd/kfd/topology/nodes/*/properties 2>/dev/null; then
+ return 0
fi
return 1
}
@@ -1656,31 +1673,39 @@ get_torch_index_url() {
if [ -n "$_rocm_tag" ]; then
# Minimum supported: ROCm 6.0 (no PyTorch wheels exist for older)
case "$_rocm_tag" in
- rocm[1-5].*) echo "$_base/cpu"; return ;;
+ rocm[1-5].*)
+ echo "[WARN] ROCm $_rocm_tag detected but PyTorch ROCm wheels require ROCm 6.0+ -- falling back to CPU-only PyTorch" >&2
+ echo "[WARN] Upgrade ROCm: https://rocm.docs.amd.com/en/latest/deploy/linux/index.html" >&2
+ echo "$_base/cpu"; return ;;
esac
- # ROCm 7.2 only has torch 2.11.0 which exceeds current bounds
- # (<2.11.0). Fall back to rocm7.1 index which has torch 2.10.0.
- # Enumerate explicit versions rather than matching rocm6.* so
- # a host on ROCm 6.5 or 6.6 (no PyTorch wheels published) is
- # clipped down to the last supported 6.x (rocm6.4) instead of
- # constructing https://download.pytorch.org/whl/rocm6.5 which
- # returns HTTP 403. PyTorch only ships: rocm5.7, 6.0, 6.1, 6.2,
- # 6.3, 6.4, 7.0, 7.1, 7.2 (and 5.7 is below our minimum).
- # TODO: uncomment rocm7.2 when the torch upper bound is bumped
- # to >=2.11.0.
+ # Supported tags; 6.5+ clips to rocm6.4, 7.3+ caps to rocm7.2.
+ # PyTorch publishes major.minor URLs only (no patch level), so
+ # rocm7.2.1 / rocm6.0.2 / etc. must normalise to rocm7.2 / rocm6.0.
case "$_rocm_tag" in
- rocm6.0|rocm6.0.*|rocm6.1|rocm6.1.*|rocm6.2|rocm6.2.*|rocm6.3|rocm6.3.*|rocm6.4|rocm6.4.*|rocm7.0|rocm7.0.*|rocm7.1|rocm7.1.*)
- echo "$_base/$_rocm_tag" ;;
+ rocm6.0|rocm6.0.*) echo "$_base/rocm6.0" ;;
+ rocm6.1|rocm6.1.*) echo "$_base/rocm6.1" ;;
+ rocm6.2|rocm6.2.*) echo "$_base/rocm6.2" ;;
+ rocm6.3|rocm6.3.*) echo "$_base/rocm6.3" ;;
+ rocm6.4|rocm6.4.*) echo "$_base/rocm6.4" ;;
+ rocm7.0|rocm7.0.*) echo "$_base/rocm7.0" ;;
+ rocm7.1|rocm7.1.*) echo "$_base/rocm7.1" ;;
+ rocm7.2|rocm7.2.*) echo "$_base/rocm7.2" ;;
rocm6.*)
# ROCm 6.5+ (no published PyTorch wheels): clip down
# to the last supported 6.x wheel set.
echo "$_base/rocm6.4" ;;
*)
- # ROCm 7.2+ (including future 10.x+): cap to rocm7.1
- echo "$_base/rocm7.1" ;;
+ # ROCm 7.3+ (future): cap to rocm7.2 (latest known)
+ echo "$_base/rocm7.2" ;;
esac
return
fi
+ # AMD GPU confirmed by rocminfo/amd-smi but ROCm version could not be
+ # read from any source (amd-smi, /opt/rocm/.info/version, hipconfig,
+ # dpkg, rpm). Warn explicitly rather than silently installing CPU PyTorch.
+ echo "[WARN] AMD GPU detected but ROCm version could not be determined -- falling back to CPU-only PyTorch" >&2
+ echo "[WARN] Ensure one of the following is accessible: amd-smi, hipconfig, /opt/rocm/.info/version, rocm-core package" >&2
+ echo "[WARN] To install ROCm: https://rocm.docs.amd.com/en/latest/deploy/linux/index.html" >&2
echo "$_base/cpu"; return
fi
# Parse CUDA version from nvidia-smi output (POSIX-safe, no grep -P).
@@ -1760,9 +1785,9 @@ print('cp{}{}'.format(sys.version_info.major, sys.version_info.minor))
}
_pick_radeon_wheel() {
- # Usage: _pick_radeon_wheel PACKAGE_NAME
+ # Usage: _pick_radeon_wheel PACKAGE_NAME [VERSION_PREFIX]
# Scans $_RADEON_LISTING for the newest wheel whose filename starts exactly
- # with PACKAGE_NAME- and matches _RADEON_PYTAG + linux_x86_64.
+ # with PACKAGE_NAME- (and optionally VERSION_PREFIX) and matches _RADEON_PYTAG + linux_x86_64.
# Prints the full URL (resolving relative hrefs against _RADEON_BASE_URL).
#
# POSIX-compliant pipeline: all href parsing, filtering, and version
@@ -1770,11 +1795,12 @@ _pick_radeon_wheel() {
# for GNU extensions (grep -o, sort -V) that would break under BSD
# or BusyBox coreutils.
_pkg="$1"
+ _ver_prefix="${2:-}"
[ -n "$_RADEON_LISTING" ] || return 1
[ -n "$_RADEON_PYTAG" ] || return 1
_tag="$_RADEON_PYTAG"
_href=$(printf '%s\n' "$_RADEON_LISTING" \
- | awk -v pkg="$_pkg" -v tag="$_tag" '
+ | awk -v pkg="$_pkg" -v tag="$_tag" -v ver_prefix="$_ver_prefix" '
BEGIN { max_pad = ""; max_url = "" }
{
line = $0
@@ -1788,7 +1814,7 @@ _pick_radeon_wheel() {
base = p[n]
sub(/[?#].*/, "", base)
- prefix = pkg "-"
+ prefix = pkg "-" ver_prefix
# Match cpXY-cpXY or cpXY-abi3 with any linux x86_64
# platform tag (linux_x86_64, manylinux_2_28_x86_64,
# manylinux2014_x86_64, etc.)
@@ -1822,6 +1848,12 @@ _pick_radeon_wheel() {
TORCH_INDEX_URL=$(get_torch_index_url)
+# rocm7.2 ships torch 2.11.0 -- adjust the constraint to allow it.
+# All other ROCm tags and CUDA stay within <2.11.0.
+case "$TORCH_INDEX_URL" in
+ */rocm7.2) TORCH_CONSTRAINT="torch>=2.11.0,<2.12.0" ;;
+esac
+
# Auto-detect GPU for AMD ROCm based
# get_torch_index_url must have chosen */rocm*
# (gfx in rocminfo or amd-smi list). Then require rocminfo "Marketing Name:.*Radeon".
@@ -1834,6 +1866,78 @@ case "$TORCH_INDEX_URL" in
fi
;;
esac
+# ── Strix Halo / Strix Point: force rocm7.2 wheels, bypass Radeon repo ───────
+# gfx1151 (Strix Halo) and gfx1150 (Strix Point) have a ROCm 7.1 driver bug
+# that causes a segfault in torch._grouped_mm (moe_utils.py line 167).
+# The Radeon repo now ships cp313 wheels for rocm-rel-7.1, so when
+# _amd_gpu_radeon=true the installer silently lands on the broken combo.
+# Detect these GPUs when TORCH_INDEX_URL is rocm7.1 and override to rocm7.2.
+case "$TORCH_INDEX_URL" in
+ */rocm7.1|*/rocm7.1.*)
+ # Collect every gfx token in rocminfo / amd-smi enumeration order
+ # (skip duplicates), then index by HIP_VISIBLE_DEVICES /
+ # ROCR_VISIBLE_DEVICES so a mixed Strix iGPU + non-Strix dGPU box
+ # where the user selected the dGPU does NOT get rerouted to the
+ # Strix per-gfx index.
+ _gfx_all=""
+ if command -v rocminfo >/dev/null 2>&1; then
+ _gfx_all=$(rocminfo 2>/dev/null | grep -oE 'gfx[1-9][0-9a-z]{2,3}')
+ fi
+ if [ -z "$_gfx_all" ] && command -v amd-smi >/dev/null 2>&1; then
+ _gfx_all=$(amd-smi list 2>/dev/null | grep -oE 'gfx[1-9][0-9a-z]{2,3}')
+ # PowerShell paths also probe `amd-smi static --asic`; mirror it
+ # so a host with hipinfo-less amd-smi reports the gfx target.
+ if [ -z "$_gfx_all" ]; then
+ _gfx_all=$(amd-smi static --asic 2>/dev/null | grep -oE 'gfx[1-9][0-9a-z]{2,3}')
+ fi
+ fi
+ _runtime_gfx=""
+ if [ -n "$_gfx_all" ]; then
+ _vis="${HIP_VISIBLE_DEVICES:-${ROCR_VISIBLE_DEVICES:-}}"
+ _idx=0
+ if [ -n "$_vis" ] && [ "$_vis" != "-1" ]; then
+ _first=${_vis%%,*}
+ case "$_first" in
+ ''|*[!0-9]*) _idx=0 ;;
+ *) _idx=$_first ;;
+ esac
+ fi
+ _runtime_gfx=$(printf '%s\n' "$_gfx_all" | awk -v idx="$_idx" '
+ NF && !seen[$0]++ { vals[n++] = $0 }
+ END {
+ if (idx < 0 || idx >= n) idx = 0
+ if (n > 0) print vals[idx]
+ }')
+ fi
+ _strix_gfx=""
+ case "$_runtime_gfx" in
+ gfx1151|gfx1150) _strix_gfx="$_runtime_gfx" ;;
+ esac
+ if [ -n "$_strix_gfx" ]; then
+ echo "" >&2
+ echo " [WARN] $_strix_gfx (Strix) + ROCm 7.1 detected -- known _grouped_mm segfault" >&2
+ echo " [WARN] ROCm 7.1 wheels are broken for gfx1150/gfx1151 (moe_utils.py:167)" >&2
+ echo " [WARN] Routing to AMD arch-specific index (torch 2.11+rocm7.13 has the real fix)" >&2
+ echo " [WARN] Upgrade ROCm to 7.2+ to use the standard index:" >&2
+ echo " [WARN] https://rocm.docs.amd.com/en/latest/deploy/linux/index.html" >&2
+ echo "" >&2
+ # AMD's arch-specific index serves torch 2.11.0+rocm7.13.0 which has AMD's
+ # actual fix for the gfx1151/gfx1150 _grouped_mm kernel bug -- preferred
+ # 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}"
+ # 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
+ ;;
+esac
_TAURI_TORCH_INDEX_FAMILY=$(_tauri_torch_index_family "$TORCH_INDEX_URL")
if [ "$_amd_gpu_radeon" = true ] && [ "$SKIP_TORCH" = false ]; then
_TAURI_TORCH_INDEX_FAMILY="radeon"
@@ -1841,27 +1945,93 @@ fi
_TAURI_GPU_BRANCH=$(_tauri_gpu_branch "$_TAURI_TORCH_INDEX_FAMILY" "$_amd_gpu_radeon")
tauri_diag_marker "$_TAURI_GPU_BRANCH" "$_TAURI_TORCH_INDEX_FAMILY"
-# ── Print CPU-only hint when no GPU detected ──
+# ── GPU detection summary (mirrors install.ps1 step "gpu" block) ──
+if _has_usable_nvidia_gpu; then
+ step "gpu" "NVIDIA GPU detected"
+elif case "$TORCH_INDEX_URL" in */rocm*|*/gfx*) true ;; *) false ;; esac; then
+ # Probe gfx arch for the display label, honouring HIP_VISIBLE_DEVICES
+ _gpu_disp_gfx_all=""
+ _gpu_disp_mkt=""
+ if command -v rocminfo >/dev/null 2>&1; then
+ _gpu_disp_gfx_all=$(rocminfo 2>/dev/null | grep -oE 'gfx[1-9][0-9a-z]{2,3}' || true)
+ _gpu_disp_mkt=$(rocminfo 2>/dev/null | awk -F': ' \
+ '/Marketing Name:/{gsub(/^[[:space:]]+|[[:space:]]+$/,"", $2); if($2){print $2; exit}}' || true)
+ fi
+ if [ -z "$_gpu_disp_gfx_all" ] && command -v amd-smi >/dev/null 2>&1; then
+ _gpu_disp_gfx_all=$(amd-smi list 2>/dev/null | grep -oE 'gfx[1-9][0-9a-z]{2,3}' || true)
+ [ -z "$_gpu_disp_gfx_all" ] && \
+ _gpu_disp_gfx_all=$(amd-smi static --asic 2>/dev/null | grep -oE 'gfx[1-9][0-9a-z]{2,3}' || true)
+ fi
+ if [ -z "$_gpu_disp_mkt" ] && command -v amd-smi >/dev/null 2>&1; then
+ _gpu_disp_mkt=$(amd-smi static --asic 2>/dev/null | awk -F'[:|]' \
+ '/[Mm]arket.?[Nn]ame/{gsub(/^[[:space:]]+|[[:space:]]+$/,"", $2); if($2){print $2; exit}}' || true)
+ fi
+ _gpu_vis="${HIP_VISIBLE_DEVICES:-${ROCR_VISIBLE_DEVICES:-}}"
+ _gpu_vis_idx=0
+ if [ -n "$_gpu_vis" ] && [ "$_gpu_vis" != "-1" ]; then
+ _gpu_first="${_gpu_vis%%,*}"
+ case "$_gpu_first" in ''|*[!0-9]*) ;; *) _gpu_vis_idx=$_gpu_first ;; esac
+ fi
+ _gpu_disp_gfx=$(printf '%s\n' "$_gpu_disp_gfx_all" | awk -v idx="$_gpu_vis_idx" \
+ 'NF && !seen[$0]++ { a[n++]=$0 } END { if(idx>=n) idx=0; if(n>0) print a[idx] }')
+ # UNSLOTH_ROCM_GFX_ARCH env override (mirrors install.ps1)
+ if [ -n "${UNSLOTH_ROCM_GFX_ARCH:-}" ]; then
+ _gpu_disp_gfx="${UNSLOTH_ROCM_GFX_ARCH}"
+ substep "gfx arch from UNSLOTH_ROCM_GFX_ARCH env override: $_gpu_disp_gfx"
+ # Name-based arch inference when tools don't report gfx (mirrors install.ps1 nameArchTable)
+ elif [ -z "$_gpu_disp_gfx" ] && [ -n "$_gpu_disp_mkt" ]; then
+ case "$_gpu_disp_mkt" in
+ *"9070 XT"*|*9080*) _gpu_disp_gfx="gfx1201" ;; # RDNA 4
+ *9070*|*9060*) _gpu_disp_gfx="gfx1200" ;; # RDNA 4
+ *"8060S"*|*"890M"*|*"Strix Halo"*|*"HX 37"*|*"HX 38"*|*"AI 9 HX"*) _gpu_disp_gfx="gfx1151" ;; # RDNA 3.5 iGPU
+ *"880M"*|*"Strix Point"*|*"AI 9 36"*|*"AI 7 35"*|*"AI 5 34"*) _gpu_disp_gfx="gfx1150" ;; # RDNA 3.5 iGPU
+ *"RX 7900"*|*"RX 7800"*|*"RX 7700"*) _gpu_disp_gfx="gfx1100" ;; # RDNA 3 desktop
+ *"RX 7600"*) _gpu_disp_gfx="gfx1102" ;; # RDNA 3
+ *"780M"*|*"760M"*|*"740M"*|*"Phoenix"*) _gpu_disp_gfx="gfx1103" ;; # RDNA 3 iGPU
+ esac
+ if [ -n "$_gpu_disp_gfx" ]; then
+ substep "gfx arch inferred from GPU name: $_gpu_disp_gfx"
+ substep "Tip: set UNSLOTH_ROCM_GFX_ARCH=$_gpu_disp_gfx to skip inference next time"
+ fi
+ fi
+ # ROCm version via hipconfig, then amd-smi
+ _gpu_rocm_ver=""
+ if command -v hipconfig >/dev/null 2>&1; then
+ _gpu_rocm_ver=$(hipconfig --version 2>/dev/null | awk 'NR==1 && /^[0-9]/{print; exit}' || true)
+ fi
+ if [ -z "$_gpu_rocm_ver" ] && command -v amd-smi >/dev/null 2>&1; then
+ _gpu_rocm_ver=$(amd-smi version 2>/dev/null | awk -F'ROCm version: ' \
+ 'NF>1{gsub(/[[:space:]]/,"", $2); print $2; exit}' || true)
+ fi
+ if [ -n "$_gpu_disp_gfx" ]; then
+ step "gpu" "AMD ROCm ($_gpu_disp_gfx)"
+ else
+ step "gpu" "AMD ROCm"
+ fi
+ _rocm_root="${ROCM_PATH:-${HIP_PATH:-/opt/rocm}}"
+ substep "ROCm: $_rocm_root"
+ [ -n "$_gpu_rocm_ver" ] && substep "hipconfig: $_gpu_rocm_ver"
+ [ -n "$_gpu_disp_mkt" ] && [ -n "$_gpu_disp_gfx" ] && substep "GPU: $_gpu_disp_mkt"
+else
+ step "gpu" "none (CPU-only)" "$C_WARN"
+fi
+
+# ── PyTorch wheel index note ──
case "$TORCH_INDEX_URL" in
*/cpu)
if [ "$SKIP_TORCH" = false ] && [ "$OS" != "macos" ]; then
- echo ""
- echo " NOTE: No GPU detected (nvidia-smi and ROCm not found)."
- echo " Installing CPU-only PyTorch. If you only need GGUF chat/inference,"
- echo " re-run with --no-torch for a faster, lighter install:"
- echo " curl -fsSL https://unsloth.ai/install.sh | sh -s -- --no-torch"
- echo " AMD ROCm users: see https://docs.unsloth.ai/get-started/install-and-update/amd"
- echo ""
+ substep "No GPU detected -- installing CPU-only PyTorch." "$C_WARN"
+ substep "AMD ROCm users: see https://docs.unsloth.ai/get-started/install-and-update/amd"
+ substep "Re-run with --no-torch for GGUF-only (faster, no PyTorch):"
+ substep " curl -fsSL https://unsloth.ai/install.sh | sh -s -- --no-torch"
fi
;;
- */rocm*)
- echo ""
+ */rocm*|*/gfx*)
if [ "$_amd_gpu_radeon" = true ]; then
- echo " AMD Radeon + ROCm detected -- installing PyTorch wheels from repo.radeon.com"
+ substep "wheels: repo.radeon.com (Radeon)"
else
- echo " AMD ROCm detected -- installing ROCm-enabled PyTorch ($TORCH_INDEX_URL)"
+ substep "wheels: $TORCH_INDEX_URL"
fi
- echo ""
;;
esac
@@ -1943,24 +2113,23 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
if [ "$_radeon_listing_ok" = true ]; then
# Require torch, torchvision, torchaudio wheels to all resolve
- # from the Radeon listing. If any is missing for this Python
- # tag, fall through to the standard ROCm index instead of
- # silently mixing Radeon wheels with PyPI defaults.
+ # from the Radeon listing. The repo often publishes multiple
+ # generations simultaneously, so picking the highest-version
+ # for each package independently can assemble a mismatched trio
+ # (e.g. torch 2.10 + torchvision 0.24). To prevent this,
+ # we identify the highest common minor version and downpair
+ # wheels if necessary to ensure a compatible set.
_torch_whl=$(_pick_radeon_wheel "torch" 2>/dev/null) || _torch_whl=""
_tv_whl=$(_pick_radeon_wheel "torchvision" 2>/dev/null) || _tv_whl=""
_ta_whl=$(_pick_radeon_wheel "torchaudio" 2>/dev/null) || _ta_whl=""
_tri_whl=$(_pick_radeon_wheel "triton" 2>/dev/null) || _tri_whl=""
- # Sanity-check torch / torchvision / torchaudio are a
- # matching release. The Radeon repo publishes multiple
- # generations simultaneously, so picking the highest-version
- # wheel for each package independently can assemble a
- # mismatched trio (e.g. torch 2.9.1 + torchvision 0.23.0 +
- # torchaudio 2.9.0 from the current rocm-rel-7.2.1 index).
+
# Check that torch and torchaudio share the same X.Y public
# version prefix, and that torchvision's minor correctly
- # pairs with torch's minor (torchvision = torch.minor - 5
+ # pairs with torch's minor (torchvision = torch.minor + 15
# since torch 2.4 -> torchvision 0.19 -> torch 2.9 ->
# torchvision 0.24).
+ #
# URL-decode each wheel name so %2B -> + before version
# extraction. Real Radeon wheel hrefs are percent-encoded
# (torch-2.10.0%2Brocm7.2.0...), so a plain [+-] terminator
@@ -1968,38 +2137,75 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
# _radeon_versions_match would stay false for every real
# listing, silently forcing a fallback to the generic
# ROCm index.
- _torch_ver=""
- _tv_ver=""
- _ta_ver=""
- if [ -n "$_torch_whl" ]; then
- _torch_name=$(printf '%s' "${_torch_whl##*/}" | sed 's/%2[Bb]/+/g')
- _torch_ver=$(printf '%s\n' "$_torch_name" | sed -n 's|^torch-\([0-9][0-9]*\.[0-9][0-9]*\)\(\.[0-9][0-9]*\)\{0,1\}[+-].*|\1|p')
- fi
- if [ -n "$_tv_whl" ]; then
- _tv_name=$(printf '%s' "${_tv_whl##*/}" | sed 's/%2[Bb]/+/g')
- _tv_ver=$(printf '%s\n' "$_tv_name" | sed -n 's|^torchvision-\([0-9][0-9]*\.[0-9][0-9]*\)\(\.[0-9][0-9]*\)\{0,1\}[+-].*|\1|p')
- fi
- if [ -n "$_ta_whl" ]; then
- _ta_name=$(printf '%s' "${_ta_whl##*/}" | sed 's/%2[Bb]/+/g')
- _ta_ver=$(printf '%s\n' "$_ta_name" | sed -n 's|^torchaudio-\([0-9][0-9]*\.[0-9][0-9]*\)\(\.[0-9][0-9]*\)\{0,1\}[+-].*|\1|p')
- fi
+ _extract_version() {
+ _whl=$1
+ _pkg=$2
+ if [ -n "$_whl" ]; then
+ _name=$(printf '%s' "${_whl##*/}" | sed 's/%2[Bb]/+/g')
+ printf '%s\n' "$_name" | sed -n "s|^${_pkg}-\([0-9][0-9]*\.[0-9][0-9]*\)\(\.[0-9][0-9]*\)\{0,1\}[+-].*|\1|p"
+ fi
+ }
+
+ _torch_ver=$(_extract_version "$_torch_whl" "torch")
+ _tv_ver=$(_extract_version "$_tv_whl" "torchvision")
+ _ta_ver=$(_extract_version "$_ta_whl" "torchaudio")
+
_radeon_versions_match=false
if [ -n "$_torch_ver" ] && [ -n "$_tv_ver" ] && [ -n "$_ta_ver" ]; then
- _torch_major=${_torch_ver%%.*}
_torch_minor=${_torch_ver#*.}
- _ta_major=${_ta_ver%%.*}
_ta_minor=${_ta_ver#*.}
- _tv_major=${_tv_ver%%.*}
_tv_minor=${_tv_ver#*.}
- # torchvision expected minor (e.g. torch 2.9 -> 0.24)
- _expected_tv_minor=$((_torch_minor + 15))
- if [ "$_torch_major" = "$_ta_major" ] && \
- [ "$_torch_minor" = "$_ta_minor" ] && \
- [ "$_tv_major" = "0" ] && \
- [ "$_tv_minor" = "$_expected_tv_minor" ]; then
- _radeon_versions_match=true
- fi
+ _tv_equiv_minor=$((_tv_minor - 15))
+
+ # Determine initial target minor (lowest common denominator)
+ _target_minor=$_torch_minor
+ [ "$_tv_equiv_minor" -lt "$_target_minor" ] && _target_minor=$_tv_equiv_minor
+ [ "$_ta_minor" -lt "$_target_minor" ] && _target_minor=$_ta_minor
+
+ # Loop downwards to find the first complete matching trio.
+ # This avoids aborting if the repo has gaps.
+ _attempts=0
+ while [ "$_attempts" -lt 5 ] && [ "$_target_minor" -ge 0 ]; do
+ _expected_tv_minor=$((_target_minor + 15))
+
+ _curr_torch=$(_pick_radeon_wheel "torch" "2.${_target_minor}." 2>/dev/null) || _curr_torch=""
+ _curr_tv=$(_pick_radeon_wheel "torchvision" "0.${_expected_tv_minor}." 2>/dev/null) || _curr_tv=""
+ _curr_ta=$(_pick_radeon_wheel "torchaudio" "2.${_target_minor}." 2>/dev/null) || _curr_ta=""
+
+ if [ -n "$_curr_torch" ] && [ -n "$_curr_tv" ] && [ -n "$_curr_ta" ]; then
+ # Extract versions from the wheels found in this iteration
+ _c_torch_ver=$(_extract_version "$_curr_torch" "torch")
+ _c_tv_ver=$(_extract_version "$_curr_tv" "torchvision")
+ _c_ta_ver=$(_extract_version "$_curr_ta" "torchaudio")
+
+ # Parse Major.Minor for validation
+ _c_torch_major=${_c_torch_ver%%.*}
+ _c_torch_minor=${_c_torch_ver#*.}
+ _c_ta_major=${_c_ta_ver%%.*}
+ _c_ta_minor=${_c_ta_ver#*.}
+ _c_tv_major=${_c_tv_ver%%.*}
+ _c_tv_minor=${_c_tv_ver#*.}
+
+ # Strict X.Y validation: allow patch versions to differ (e.g. torch 2.9.1 + vision 0.24.0)
+ # as long as the Major and Minor pairing is correct.
+ if [ "$_c_torch_major" = "$_c_ta_major" ] && \
+ [ "$_c_torch_minor" = "$_c_ta_minor" ] && \
+ [ "$_c_tv_major" = "0" ] && \
+ [ "$_c_tv_minor" = "$((_c_torch_minor + 15))" ]; then
+
+ _torch_whl=$_curr_torch
+ _tv_whl=$_curr_tv
+ _ta_whl=$_curr_ta
+ _tri_whl=""
+ _radeon_versions_match=true
+ break
+ fi
+ fi
+ _target_minor=$((_target_minor - 1))
+ _attempts=$((_attempts + 1))
+ done
fi
+
if [ -z "$_torch_whl" ] || [ -z "$_tv_whl" ] || [ -z "$_ta_whl" ] || \
[ "$_radeon_versions_match" != true ]; then
substep "[WARN] Radeon repo lacks a compatible wheel set for this Python; falling back to ROCm index ($TORCH_INDEX_URL)" "$C_WARN"
diff --git a/scripts/verify_import_hoist.py b/scripts/verify_import_hoist.py
new file mode 100644
index 0000000000..606488cc7f
--- /dev/null
+++ b/scripts/verify_import_hoist.py
@@ -0,0 +1,854 @@
+#!/usr/bin/env python3
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
+"""Deterministic, scope-aware verifier for import-hoisting / alias-rename refactors.
+
+The risk when moving `from a import b as _b` (or `import b as _b`) to module top
+and normalizing `_b` -> `b` is twofold:
+
+ 1. DANGLING ALIAS - a `_b` reference is left un-normalized; it now resolves to
+ nothing (NameError) or, worse, to some *other* module-level `_b`.
+ 2. RENAME CLASH - `_b` was an alias on purpose because `b` already meant
+ something else in that scope; normalizing `_b` -> `b` silently re-points the
+ reference at the wrong object (no NameError, no pyflakes warning).
+
+This tool parses BEFORE (a git ref, default origin/main) and AFTER (default HEAD)
+for each file, builds a real LEGB scope model (functions, classes, lambdas,
+comprehensions, global/nonlocal, args, walrus, star-imports), and resolves every
+Name load to its binding. It then compares, PER SCOPE:
+
+ * UNRESOLVED-NEW : loads that resolve to nothing in AFTER but did in BEFORE
+ (or are newly present) -> catches dangling aliases.
+ * TARGET-MISSING : an import *target* (e.g. module `glob`, or
+ `importlib.metadata.version`) that a function resolved to
+ in BEFORE but no longer resolves to in AFTER -> catches a
+ function that lost access to a module it still uses.
+ Robust to alias renames because it compares the *target*,
+ not the local name.
+ * TARGET-CHANGED : a load whose resolved import target differs BEFORE vs
+ AFTER -> catches a rename that re-points to a different
+ module (the clash case).
+ * AMBIGUOUS-BIND : a name bound by BOTH an import and a non-import in the same
+ scope in AFTER (and not in BEFORE) -> the "alias was on
+ purpose / now collides" smell.
+ * MODULE-DUP-IMPORT: a module-level name imported and also defined/assigned at
+ module level (introduced by the change).
+ * NEW-UNUSED-IMPORT: a module-level import added in AFTER that nothing resolves
+ to (informational; re-exports are a known false positive).
+
+Usage:
+ verify_import_hoist.py [--before REF] [--after REF] ... # compare
+ verify_import_hoist.py --self-test # prove it catches bugs
+Exit code 1 if any non-informational finding.
+"""
+
+from __future__ import annotations
+
+import argparse
+import ast
+import builtins
+import re as _re_mod
+import subprocess
+import sys
+from dataclasses import dataclass, field
+
+_BUILTINS = set(dir(builtins)) | {
+ "__file__",
+ "__name__",
+ "__doc__",
+ "__package__",
+ "__spec__",
+ "__loader__",
+ "__builtins__",
+ "__class__",
+ "__annotations__",
+ "__dict__",
+ "__qualname__",
+ "__module__",
+ "__path__",
+ "__debug__",
+ "__import__",
+ "NotImplemented",
+ "Ellipsis",
+ "copyright",
+ "credits",
+ "license",
+ "help",
+ "exit",
+ "quit",
+ "__build_class__",
+ "__cached__",
+ "reveal_type",
+ "reveal_locals",
+}
+
+
+# ---------------------------------------------------------------- scope model
+
+
+@dataclass
+class Binding:
+ kind: str # 'import' | 'importfrom' | 'def' | 'class' | 'other'
+ target: str | None = None # canonical import target id, else None
+
+
+@dataclass
+class Scope:
+ kind: str # 'module' | 'function' | 'class' | 'lambda' | 'comp'
+ qualname: str
+ parent: "Scope | None"
+ bindings: dict[str, list[Binding]] = field(default_factory = dict)
+ globals: set[str] = field(default_factory = set)
+ nonlocals: set[str] = field(default_factory = set)
+ star_import: bool = False
+
+ def add(self, name: str, b: Binding) -> None:
+ self.bindings.setdefault(name, []).append(b)
+
+
+def _import_target(node: ast.AST, alias: ast.alias) -> tuple[str, str]:
+ """Return (bound_name, canonical_target_id) for one import alias."""
+ if isinstance(node, ast.Import):
+ bound = alias.asname or alias.name.split(".")[0]
+ return bound, f"import:{alias.name}"
+ # ImportFrom
+ bound = alias.asname or alias.name
+ mod = ("." * (node.level or 0)) + (node.module or "")
+ return bound, f"from:{mod}:{alias.name}"
+
+
+class _Builder(ast.NodeVisitor):
+ """Builds the scope tree + bindings, and records every (scope, Name-load)."""
+
+ def __init__(self):
+ self.module = Scope("module", "", None)
+ self.uses: list[tuple[Scope, str, int]] = [] # (scope, name, lineno) hard loads
+ self.soft_uses: list[
+ tuple[Scope, str, int]
+ ] = [] # annotations: count as "used"
+ # but never as "unresolved"
+ # (forward refs / string annos)
+
+ def _visit_annotation(self, node, scope: Scope) -> None:
+ """Annotation context: with `from __future__ import annotations` these are
+ never evaluated (strings), and even otherwise they routinely contain forward
+ references. Record contained names as SOFT uses so an import used only in an
+ annotation still counts as used, but a forward-ref name is never 'unresolved'."""
+ if node is None:
+ return
+ for n in ast.walk(node):
+ if isinstance(n, ast.Name) and isinstance(n.ctx, ast.Load):
+ self.soft_uses.append((scope, n.id, n.lineno))
+
+ # -- binding helpers --
+ def _bind_targets(self, scope: Scope, target: ast.AST) -> None:
+ for n in ast.walk(target):
+ if isinstance(n, ast.Name) and isinstance(n.ctx, (ast.Store, ast.Del)):
+ self._bind_name(scope, n.id, Binding("other"))
+ elif isinstance(n, ast.Starred):
+ pass
+
+ def _bind_name(self, scope: Scope, name: str, b: Binding) -> None:
+ if name in scope.globals:
+ self.module.add(name, b)
+ elif name in scope.nonlocals:
+ p = scope.parent
+ while p is not None and p.kind not in ("function", "lambda"):
+ p = p.parent
+ (p or self.module).add(name, b)
+ else:
+ scope.add(name, b)
+
+ # -- generic dispatch within a scope --
+ def _visit_body(self, stmts, scope: Scope) -> None:
+ for s in stmts:
+ self._visit_stmt(s, scope)
+
+ def _visit_stmt(self, node: ast.AST, scope: Scope) -> None:
+ if isinstance(node, (ast.Import, ast.ImportFrom)):
+ star = isinstance(node, ast.ImportFrom) and any(
+ a.name == "*" for a in node.names
+ )
+ if star:
+ scope.star_import = True
+ for alias in node.names:
+ if alias.name == "*":
+ continue
+ bound, target = _import_target(node, alias)
+ kind = "import" if isinstance(node, ast.Import) else "importfrom"
+ self._bind_name(scope, bound, Binding(kind, target))
+ return
+ if isinstance(node, ast.Global):
+ scope.globals.update(node.names)
+ return
+ if isinstance(node, ast.Nonlocal):
+ scope.nonlocals.update(node.names)
+ return
+ if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
+ self._bind_name(scope, node.name, Binding("def"))
+ # decorators / defaults evaluate in the ENCLOSING scope
+ for d in node.decorator_list:
+ self._visit_expr(d, scope)
+ self._visit_arg_defaults(node.args, scope)
+ child = Scope("function", f"{scope.qualname}.{node.name}", scope)
+ self._bind_type_params(node, child)
+ self._bind_args(node.args, child)
+ # arg + return annotations: soft uses (may be strings / forward refs)
+ for a in self._all_args(node.args):
+ self._visit_annotation(a.annotation, child)
+ self._visit_annotation(getattr(node, "returns", None), child)
+ self._visit_body(node.body, child)
+ return
+ if isinstance(node, ast.ClassDef):
+ self._bind_name(scope, node.name, Binding("class"))
+ for d in node.decorator_list:
+ self._visit_expr(d, scope)
+ for b in node.bases:
+ self._visit_expr(b, scope)
+ for kw in node.keywords:
+ self._visit_expr(kw.value, scope)
+ child = Scope("class", f"{scope.qualname}.{node.name}", scope)
+ self._bind_type_params(node, child)
+ self._visit_body(node.body, child)
+ return
+ if isinstance(node, ast.Match):
+ self._visit_expr(node.subject, scope)
+ for case in node.cases:
+ self._bind_pattern(case.pattern, scope)
+ if case.guard is not None:
+ self._visit_expr(case.guard, scope)
+ self._visit_body(case.body, scope)
+ return
+ if isinstance(node, getattr(ast, "TryStar", ())): # py3.11 except*
+ self._visit_body(node.body, scope)
+ for h in node.handlers:
+ if h.type is not None:
+ self._visit_expr(h.type, scope)
+ if h.name:
+ self._bind_name(scope, h.name, Binding("other"))
+ self._visit_body(h.body, scope)
+ self._visit_body(node.orelse, scope)
+ self._visit_body(node.finalbody, scope)
+ return
+ if isinstance(node, getattr(ast, "TypeAlias", ())): # py3.12 `type X = ...`
+ if isinstance(node.name, ast.Name):
+ self._bind_name(scope, node.name.id, Binding("other"))
+ self._visit_annotation(node.value, scope)
+ return
+ if isinstance(node, (ast.Assign, ast.AnnAssign, ast.AugAssign)):
+ targets = node.targets if isinstance(node, ast.Assign) else [node.target]
+ val = node.value
+ if val is not None:
+ self._visit_expr(val, scope)
+ if isinstance(node, ast.AnnAssign) and node.annotation is not None:
+ self._visit_annotation(node.annotation, scope)
+ for t in targets:
+ self._bind_targets(scope, t)
+ # AugAssign target is also a load
+ if isinstance(node, ast.AugAssign):
+ self._record_loads(t, scope)
+ return
+ if isinstance(node, (ast.For, ast.AsyncFor)):
+ self._visit_expr(node.iter, scope)
+ self._bind_targets(scope, node.target)
+ self._visit_body(node.body, scope)
+ self._visit_body(node.orelse, scope)
+ return
+ if isinstance(node, (ast.With, ast.AsyncWith)):
+ for item in node.items:
+ self._visit_expr(item.context_expr, scope)
+ if item.optional_vars is not None:
+ self._bind_targets(scope, item.optional_vars)
+ self._visit_body(node.body, scope)
+ return
+ if isinstance(node, ast.Try):
+ self._visit_body(node.body, scope)
+ for h in node.handlers:
+ if h.type is not None:
+ self._visit_expr(h.type, scope)
+ if h.name:
+ self._bind_name(scope, h.name, Binding("other"))
+ self._visit_body(h.body, scope)
+ self._visit_body(node.orelse, scope)
+ self._visit_body(node.finalbody, scope)
+ return
+ # generic statement: visit all child expressions/stmts in same scope
+ for child in ast.iter_child_nodes(node):
+ if isinstance(child, ast.stmt):
+ self._visit_stmt(child, scope)
+ else:
+ self._visit_expr(child, scope)
+
+ # -- expressions --
+ def _visit_arg_defaults(self, args: ast.arguments, scope: Scope) -> None:
+ for d in list(args.defaults) + [d for d in args.kw_defaults if d is not None]:
+ self._visit_expr(d, scope)
+
+ def _all_args(self, args: ast.arguments) -> list[ast.arg]:
+ out = list(args.posonlyargs) + list(args.args) + list(args.kwonlyargs)
+ if args.vararg:
+ out.append(args.vararg)
+ if args.kwarg:
+ out.append(args.kwarg)
+ return out
+
+ def _bind_args(self, args: ast.arguments, scope: Scope) -> None:
+ for a in self._all_args(args):
+ scope.add(a.arg, Binding("other"))
+
+ def _bind_type_params(self, node, scope: Scope) -> None:
+ for tp in getattr(node, "type_params", []) or []:
+ name = getattr(tp, "name", None)
+ if isinstance(name, str):
+ scope.add(name, Binding("other"))
+ self._visit_annotation(getattr(tp, "bound", None), scope)
+ self._visit_annotation(getattr(tp, "default_value", None), scope)
+
+ def _bind_pattern(self, pat, scope: Scope) -> None:
+ if pat is None:
+ return
+ if isinstance(pat, ast.MatchValue):
+ self._visit_expr(pat.value, scope)
+ elif isinstance(pat, ast.MatchSingleton):
+ pass
+ elif isinstance(pat, ast.MatchSequence):
+ for p in pat.patterns:
+ self._bind_pattern(p, scope)
+ elif isinstance(pat, ast.MatchStar):
+ if pat.name:
+ self._bind_name(scope, pat.name, Binding("other"))
+ elif isinstance(pat, ast.MatchMapping):
+ for k in pat.keys:
+ self._visit_expr(k, scope)
+ for p in pat.patterns:
+ self._bind_pattern(p, scope)
+ if pat.rest:
+ self._bind_name(scope, pat.rest, Binding("other"))
+ elif isinstance(pat, ast.MatchClass):
+ self._visit_expr(pat.cls, scope)
+ for p in pat.patterns:
+ self._bind_pattern(p, scope)
+ for p in pat.kwd_patterns:
+ self._bind_pattern(p, scope)
+ elif isinstance(pat, ast.MatchAs):
+ self._bind_pattern(pat.pattern, scope)
+ if pat.name:
+ self._bind_name(scope, pat.name, Binding("other"))
+ elif isinstance(pat, ast.MatchOr):
+ for p in pat.patterns:
+ self._bind_pattern(p, scope)
+
+ def _record_loads(self, node: ast.AST, scope: Scope) -> None:
+ for n in ast.walk(node):
+ if isinstance(n, ast.Name) and isinstance(n.ctx, ast.Load):
+ self.uses.append((scope, n.id, n.lineno))
+
+ def _visit_expr(self, node: ast.AST, scope: Scope) -> None:
+ if isinstance(node, ast.Name):
+ if isinstance(node.ctx, ast.Load):
+ self.uses.append((scope, node.id, node.lineno))
+ elif isinstance(node.ctx, (ast.Store, ast.Del)):
+ self._bind_name(scope, node.id, Binding("other"))
+ return
+ if isinstance(node, ast.Lambda):
+ self._visit_arg_defaults(node.args, scope)
+ child = Scope("lambda", f"{scope.qualname}.", scope)
+ self._bind_args(node.args, child)
+ self._visit_expr(node.body, child)
+ return
+ if isinstance(
+ node, (ast.ListComp, ast.SetComp, ast.GeneratorExp, ast.DictComp)
+ ):
+ child = Scope("comp", f"{scope.qualname}.", scope)
+ for i, gen in enumerate(node.generators):
+ # first iterable is evaluated in the enclosing scope
+ self._visit_expr(gen.iter, scope if i == 0 else child)
+ self._bind_targets(child, gen.target)
+ for cond in gen.ifs:
+ self._visit_expr(cond, child)
+ if isinstance(node, ast.DictComp):
+ self._visit_expr(node.key, child)
+ self._visit_expr(node.value, child)
+ else:
+ self._visit_expr(node.elt, child)
+ return
+ if isinstance(node, ast.NamedExpr): # walrus binds in enclosing scope
+ self._visit_expr(node.value, scope)
+ if isinstance(node.target, ast.Name):
+ self._bind_name(scope, node.target.id, Binding("other"))
+ return
+ for child in ast.iter_child_nodes(node):
+ if isinstance(child, ast.stmt):
+ self._visit_stmt(child, scope)
+ else:
+ self._visit_expr(child, scope)
+
+ def run(self, tree: ast.Module) -> None:
+ self._visit_body(tree.body, self.module)
+
+
+# ---------------------------------------------------------------- resolution
+
+
+def _any_star(scope: Scope) -> bool:
+ c = scope
+ while c is not None:
+ if c.star_import:
+ return True
+ c = c.parent
+ return False
+
+
+def _resolve(scope: Scope, name: str):
+ """LEGB resolution. Returns (status, bindings) where status in
+ {'local','import','other','builtin','star','unresolved'}."""
+ # global / nonlocal redirection
+ start = scope
+ if name in scope.globals:
+ chain = [_module_of(scope)]
+ elif name in scope.nonlocals:
+ chain = _enclosing_functions(scope)
+ else:
+ chain = _legb_chain(scope)
+ for i, sc in enumerate(chain):
+ if sc is None:
+ continue
+ if name in sc.bindings:
+ binds = sc.bindings[name]
+ if any(b.kind in ("import", "importfrom") for b in binds):
+ return "import", binds
+ return "other", binds
+ if name in _BUILTINS:
+ return "builtin", []
+ if _any_star(start):
+ return "star", []
+ return "unresolved", []
+
+
+def _module_of(scope: Scope) -> Scope:
+ while scope.parent is not None:
+ scope = scope.parent
+ return scope
+
+
+def _enclosing_functions(scope: Scope) -> list[Scope]:
+ out = []
+ p = scope.parent
+ while p is not None:
+ if p.kind in ("function", "lambda"):
+ out.append(p)
+ p = p.parent
+ out.append(_module_of(scope))
+ return out
+
+
+def _legb_chain(scope: Scope) -> list[Scope]:
+ """Immediate scope, then enclosing scopes skipping class scopes, then module."""
+ chain = [scope]
+ p = scope.parent
+ while p is not None:
+ if (
+ p.kind != "class" or p.parent is None
+ ): # module-level class never happens; keep module
+ if p.kind != "class":
+ chain.append(p)
+ p = p.parent
+ return chain
+
+
+# ---------------------------------------------------------------- analysis
+
+
+def _analyze(src: str):
+ tree = ast.parse(src)
+ b = _Builder()
+ b.run(tree)
+ # Per-scope: unresolved load names, and import targets it resolves to.
+ unresolved: dict[str, set[str]] = {}
+ targets_by_scope: dict[str, set[str]] = {}
+ target_by_use: dict[tuple[str, str], set[str]] = {}
+ for scope, name, _ln in b.uses:
+ status, binds = _resolve(scope, name)
+ if status == "unresolved":
+ unresolved.setdefault(scope.qualname, set()).add(name)
+ elif status == "import":
+ tids = {bd.target for bd in binds if bd.target}
+ targets_by_scope.setdefault(scope.qualname, set()).update(tids)
+ target_by_use.setdefault((scope.qualname, name), set()).update(tids)
+ # soft uses (annotations): only contribute to "used", never to "unresolved"
+ for scope, name, _ln in b.soft_uses:
+ status, binds = _resolve(scope, name)
+ if status == "import":
+ tids = {bd.target for bd in binds if bd.target}
+ targets_by_scope.setdefault(scope.qualname, set()).update(tids)
+ # module-level binding info for clash checks
+ module = b.module
+ module_imports = {
+ n: bs
+ for n, bs in module.bindings.items()
+ if any(x.kind in ("import", "importfrom") for x in bs)
+ }
+ module_dup = {
+ n
+ for n, bs in module.bindings.items()
+ if any(x.kind in ("import", "importfrom") for x in bs)
+ and any(x.kind not in ("import", "importfrom") for x in bs)
+ }
+ # ambiguous: any scope where a name is bound by import AND non-import
+ ambiguous: dict[str, set[str]] = {}
+
+ def walk_scopes(scope: Scope):
+ for n, bs in scope.bindings.items():
+ if any(x.kind in ("import", "importfrom") for x in bs) and any(
+ x.kind not in ("import", "importfrom") for x in bs
+ ):
+ ambiguous.setdefault(scope.qualname, set()).add(n)
+ # scope tree isn't stored; rebuild via uses is hard. We approximate with module only.
+
+ walk_scopes(module)
+ return {
+ "unresolved": unresolved,
+ "targets_by_scope": targets_by_scope,
+ "target_by_use": target_by_use,
+ "module_import_targets": {
+ n: {x.target for x in bs if x.target} for n, bs in module_imports.items()
+ },
+ "module_dup": module_dup,
+ "ambiguous": ambiguous,
+ }
+
+
+def _git_show(ref: str, path: str) -> str | None:
+ try:
+ return subprocess.run(
+ ["git", "show", f"{ref}:{path}"], capture_output = True, text = True, check = True
+ ).stdout
+ except subprocess.CalledProcessError:
+ return None
+
+
+def compare(before_src: str, after_src: str, path: str) -> list[tuple[str, str]]:
+ """Return list of (severity, message). severity in BLOCKER/WARN/INFO.
+
+ Blocker signals (precise, no relocation false-positives):
+ UNRESOLVED-NEW - a load became undefined (dangling alias / removed import).
+ NEW-UNUSED-HOIST - a module-level import added by THIS change is resolved by
+ NO load. A correct hoist always wires its new import to a
+ reference; if the alias was left un-normalized OR renamed
+ to the wrong name, the hoisted import ends up unused. This
+ single signal catches BOTH user-described failure modes and
+ does NOT fire for code merely relocated to another file
+ (that removes the import, it doesn't add an unused one).
+ TARGET-CHANGED - the same (scope, name) load resolves to a different import
+ target before vs after (a same-name re-point).
+ """
+ a = _analyze(before_src)
+ b = _analyze(after_src)
+ findings: list[tuple[str, str]] = []
+
+ def used_targets(analysis) -> set[str]:
+ out: set[str] = set()
+ for tids in analysis["targets_by_scope"].values():
+ out |= tids
+ return out
+
+ before_used = used_targets(a)
+ after_used = used_targets(b)
+ before_module_targets: set[str] = set()
+ for tids in a["module_import_targets"].values():
+ before_module_targets |= tids
+ after_module_targets: set[str] = set()
+ for tids in b["module_import_targets"].values():
+ after_module_targets |= tids
+ added_module_targets = after_module_targets - before_module_targets
+
+ # 1. UNRESOLVED-NEW
+ for scope, names in b["unresolved"].items():
+ new = names - a["unresolved"].get(scope, set())
+ for n in sorted(new):
+ findings.append(
+ (
+ "BLOCKER",
+ f"{path}: UNRESOLVED-NEW '{n}' in scope {scope} "
+ f"(undefined after change -> dangling alias / removed import)",
+ )
+ )
+
+ # 2. HOISTED-IMPORT-UNUSED (the core botched-hoist / wrong-rename signal)
+ # A module-level import in AFTER that NO load resolves to, and which was
+ # either newly added by this change OR was actually used before. Excludes:
+ # - relocation (the import is REMOVED, so it's not in after at all)
+ # - stable pre-existing re-exports (unused before AND after, not newly added)
+ for n, tids in b["module_import_targets"].items():
+ if tids & after_used:
+ continue # resolved by something -> fine
+ newly_added = bool(tids - before_module_targets)
+ was_used_before = bool(tids & before_used)
+ if newly_added or was_used_before:
+ why = (
+ "added but unused"
+ if newly_added
+ else "was used before, now unused (references re-pointed)"
+ )
+ findings.append(
+ (
+ "BLOCKER",
+ f"{path}: HOISTED-IMPORT-UNUSED '{n}' ({sorted(tids)}) "
+ f"{why} -> un-normalized alias or wrong rename target?",
+ )
+ )
+
+ # 3. TARGET-CHANGED (same scope+name resolves to a different import target)
+ for key, tafter in b["target_by_use"].items():
+ tbefore = a["target_by_use"].get(key)
+ if tbefore and tbefore != tafter:
+ findings.append(
+ (
+ "BLOCKER",
+ f"{path}: TARGET-CHANGED name '{key[1]}' in {key[0]} "
+ f"{sorted(tbefore)} -> {sorted(tafter)} (rename re-points module)",
+ )
+ )
+
+ # 4. MODULE-DUP-IMPORT introduced
+ for n in sorted(b["module_dup"] - a["module_dup"]):
+ findings.append(
+ (
+ "WARN",
+ f"{path}: MODULE-DUP-IMPORT '{n}' bound by import AND non-import "
+ f"at module level (possible clash)",
+ )
+ )
+
+ # 5. AMBIGUOUS-BIND introduced (module scope)
+ for scope, names in b["ambiguous"].items():
+ new = names - a["ambiguous"].get(scope, set())
+ for n in sorted(new):
+ findings.append(
+ ("WARN", f"{path}: AMBIGUOUS-BIND '{n}' import+non-import in {scope}")
+ )
+
+ # 6. TARGET-MISSING (informational): a scope stopped resolving to an import
+ # target. Real bugs are already covered above; remaining cases are code
+ # relocated to another file (e.g. a moved helper). Shown for transparency.
+ for scope, tbefore in a["targets_by_scope"].items():
+ tafter = b["targets_by_scope"].get(scope, set())
+ for t in sorted(tbefore - tafter):
+ relocated = (
+ ""
+ if t in added_module_targets
+ else " [target not re-added here -> likely relocated/deleted]"
+ )
+ findings.append(
+ ("INFO", f"{path}: TARGET-MISSING {t} in scope {scope}{relocated}")
+ )
+ return findings
+
+
+# ---------------------------------------------------------------- self-test
+
+_SELF_TESTS = {
+ "dangling_alias": (
+ # before: inline aliased import, used as _b
+ "import os\n"
+ "def f():\n"
+ " import glob as _b\n"
+ " return _b.glob('*')\n",
+ # after: hoisted to canonical, but reference NOT normalized -> _b dangles
+ "import os\n" "import glob\n" "def f():\n" " return _b.glob('*')\n",
+ "BLOCKER",
+ ),
+ "rename_clash": (
+ # before: _b is a deliberate alias; `b` already means something else
+ "import re as _b\n" "b = 123\n" "def f():\n" " return _b.compile('x'), b\n",
+ # after: someone normalized _b -> b ; now f().b is the int, re is lost
+ "import re\n" "b = 123\n" "def f():\n" " return b.compile('x'), b\n",
+ "BLOCKER", # TARGET-MISSING from:.. or import:re in f
+ ),
+ "clean_rename": (
+ "def f():\n" " import glob as _g\n" " return _g.glob('*')\n",
+ "import glob\n" "def f():\n" " return glob.glob('*')\n",
+ None, # expect NO blocker
+ ),
+ "clean_dedup_redundant": (
+ "import sys\n" "def f():\n" " import sys\n" " return sys.argv\n",
+ "import sys\n" "def f():\n" " return sys.argv\n",
+ None,
+ ),
+ "from_import_dangling": (
+ # from-import alias left un-normalized
+ "def f():\n"
+ " from importlib.metadata import version as _v\n"
+ " return _v('x')\n",
+ "from importlib.metadata import version\n" "def f():\n" " return _v('x')\n",
+ "BLOCKER",
+ ),
+ "local_var_clash": (
+ # _b renamed to b, but b is a LOCAL variable in f -> import silently unused
+ "def f(b):\n" " import re as _b\n" " return _b.compile(b)\n",
+ "import re\n"
+ "def f(b):\n"
+ " return b.compile(b)\n", # 'b' is the param, not the module
+ "BLOCKER",
+ ),
+ "substring_safe": (
+ # correct _copy->copy rename while a config_copy var exists: NO false positive
+ "def f(config):\n"
+ " import copy as _copy\n"
+ " config_copy = _copy.deepcopy(config)\n"
+ " return config_copy\n",
+ "import copy\n"
+ "def f(config):\n"
+ " config_copy = copy.deepcopy(config)\n"
+ " return config_copy\n",
+ None,
+ ),
+ "attr_access_not_a_use": (
+ # x._b is attribute access, not a use of name _b; removing import _b is fine
+ "import os\n"
+ "def f(x):\n"
+ " import sys as _b\n"
+ " return x._b + _b.argv[0]\n",
+ "import os\n" "import sys\n" "def f(x):\n" " return x._b + sys.argv[0]\n",
+ None,
+ ),
+}
+
+
+def _self_test() -> int:
+ ok = True
+ for name, (before, after, expect) in _SELF_TESTS.items():
+ findings = compare(before, after, f"<{name}>")
+ blockers = [m for sev, m in findings if sev == "BLOCKER"]
+ got = "BLOCKER" if blockers else None
+ passed = got == expect
+ ok = ok and passed
+ print(f"[{'PASS' if passed else 'FAIL'}] {name}: expect={expect} got={got}")
+ for sev, m in findings:
+ print(f" ({sev}) {m}")
+ print("\nSELF-TEST:", "ALL PASS" if ok else "FAILURES")
+ return 0 if ok else 1
+
+
+def _pyflakes_undefined(path: str) -> set[str] | None:
+ """Return the set of names pyflakes reports as 'undefined name' for `path`,
+ or None if pyflakes failed to run/parse the file."""
+ try:
+ proc = subprocess.run(
+ [sys.executable, "-m", "pyflakes", path], capture_output = True, text = True
+ )
+ except Exception:
+ return None
+ if "syntax error" in (proc.stdout + proc.stderr).lower():
+ return None
+ names = set()
+ for line in proc.stdout.splitlines():
+ m = _re_mod.search(r"undefined name '([^']+)'", line)
+ if m:
+ names.add(m.group(1))
+ return names
+
+
+def audit_files(paths: list[str]) -> int:
+ """Single-version robustness audit. For every file: confirm the analyzer does
+ not crash, then cross-check its 'unresolved' names against pyflakes. Any name
+ the resolver flags that pyflakes does NOT call undefined is a tool FALSE
+ POSITIVE (a resolver gap to fix)."""
+ n_files = n_err = n_fp = n_syntax = 0
+ fp_detail: dict[str, set[str]] = {}
+ err_detail: dict[str, str] = {}
+ for path in paths:
+ n_files += 1
+ try:
+ src = open(path, encoding = "utf-8").read()
+ except Exception as e: # unreadable
+ n_err += 1
+ err_detail[path] = f"read: {e}"
+ continue
+ try:
+ res = _analyze(src)
+ except SyntaxError:
+ n_syntax += 1
+ continue
+ except Exception as e: # analyzer crash -> robustness bug
+ n_err += 1
+ err_detail[path] = f"{type(e).__name__}: {e}"
+ continue
+ tool_unresolved = set()
+ for names in res["unresolved"].values():
+ tool_unresolved |= names
+ if not tool_unresolved:
+ continue
+ pf = _pyflakes_undefined(path)
+ if pf is None:
+ continue # pyflakes couldn't adjudicate; skip cross-check
+ false_pos = tool_unresolved - pf
+ if false_pos:
+ n_fp += 1
+ fp_detail[path] = false_pos
+ print(f"audited files : {n_files}")
+ print(f"syntax-skipped : {n_syntax}")
+ print(f"analyzer errors : {n_err}")
+ for p, e in sorted(err_detail.items()):
+ print(f" ERROR {p}: {e}")
+ print(f"false-positive files: {n_fp} (resolver flagged a name pyflakes accepts)")
+ for p, names in sorted(fp_detail.items()):
+ print(f" FP {p}: {sorted(names)}")
+ ok = n_err == 0 and n_fp == 0
+ print(
+ "\nAUDIT:",
+ "ROBUST (no crashes, no false positives vs pyflakes)"
+ if ok
+ else "NEEDS WORK (see above)",
+ )
+ return 0 if ok else 1
+
+
+def main() -> int:
+ ap = argparse.ArgumentParser()
+ ap.add_argument("--before", default = "origin/main")
+ ap.add_argument("--after", default = "HEAD")
+ ap.add_argument("--self-test", action = "store_true")
+ ap.add_argument(
+ "--audit",
+ action = "store_true",
+ help = "single-version robustness audit on filesystem paths",
+ )
+ ap.add_argument("files", nargs = "*")
+ args = ap.parse_args()
+
+ if args.self_test:
+ return _self_test()
+ if args.audit:
+ return audit_files(args.files)
+
+ any_blocker = False
+ for path in args.files:
+ before = _git_show(args.before, path)
+ after = _git_show(args.after, path)
+ if after is None:
+ print(f"SKIP {path}: not found at {args.after}")
+ continue
+ if before is None:
+ before = "" # new file
+ findings = compare(before, after, path)
+ blockers = [f for f in findings if f[0] == "BLOCKER"]
+ warns = [f for f in findings if f[0] == "WARN"]
+ infos = [f for f in findings if f[0] == "INFO"]
+ status = (
+ "CLEAN"
+ if not blockers and not warns
+ else ("BLOCKERS" if blockers else "WARNINGS")
+ )
+ print(f"\n=== {path}: {status} ===")
+ for sev, m in blockers + warns + infos:
+ print(f" [{sev}] {m}")
+ any_blocker = any_blocker or bool(blockers)
+ print(
+ "\nOVERALL:", "FAIL (blockers found)" if any_blocker else "PASS (no blockers)"
+ )
+ return 1 if any_blocker else 0
+
+
+if __name__ == "__main__":
+ sys.exit(main())
diff --git a/studio/backend/core/_torchao_stub.py b/studio/backend/core/_torchao_stub.py
new file mode 100644
index 0000000000..5650a60ee2
--- /dev/null
+++ b/studio/backend/core/_torchao_stub.py
@@ -0,0 +1,142 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""Shared torchao Windows-ROCm import stub.
+
+torchao (pulled in by transformers.quantizers) imports
+torch.distributed._functional_collectives at module level, which imports
+distributed_c10d.py unconditionally — that file crashes on Windows ROCm because
+torch._C._distributed_c10d (the RCCL backend) is absent.
+torch/distributed/__init__.py itself is guarded by `if is_available()` so
+`import torch.distributed` alone is safe; the crash only comes via torchao's
+import chain. Stubbing torchao short-circuits it entirely.
+_StubSubpackageFinder handles any depth of torchao.xxx.yyy imports.
+
+This logic used to be duplicated inline inside run_export_process() and
+run_training_process(); it now lives here so both worker subprocesses call the
+single `install_torchao_windows_rocm_stub()` entrypoint before importing
+transformers / unsloth_zoo.
+"""
+
+from __future__ import annotations
+
+import sys
+import types
+import importlib.abc
+import importlib.machinery
+
+_STUB_SENTINEL = object()
+
+
+# Metaclass for stub types so that isinstance(x, StubClass) returns False
+# instead of raising TypeError ("arg 2 must be a type").
+# peft/tuners/lora/torchao.py does:
+# from torchao.dtypes import AffineQuantizedTensor, LinearActivationQuantizedTensor
+# isinstance(weight, (AffineQuantizedTensor, LinearActivationQuantizedTensor))
+# If those names resolve to stub modules rather than types, isinstance() raises.
+class _StubTypeMeta(type):
+ def __instancecheck__(cls, instance):
+ return False
+
+ def __subclasscheck__(cls, subclass):
+ return False
+
+ def __getattr__(cls, attr):
+ if attr.startswith("__"):
+ raise AttributeError(attr)
+ child = _StubTypeMeta(attr, (), {})
+ setattr(cls, attr, child)
+ return child
+
+ def __call__(cls, *args, **kwargs):
+ return None
+
+
+def _make_stub_type(name):
+ """Stub class: accepted by isinstance() (always False), supports attr access."""
+ return _StubTypeMeta(name, (), {})
+
+
+def _make_mod_stub(mod_name):
+ m = types.ModuleType(mod_name)
+ m.__path__ = []
+ m.__package__ = mod_name
+ m._unsloth_stub = _STUB_SENTINEL
+ m.__spec__ = importlib.machinery.ModuleSpec(mod_name, loader = None, is_package = True)
+
+ def _ga(attr, _m = m, _n = mod_name):
+ if attr.startswith("__"):
+ raise AttributeError(attr)
+ # Return a stub CLASS (not a module) so that isinstance(x, attr)
+ # works and returns False instead of raising TypeError.
+ child = _make_stub_type(f"{_n}.{attr}")
+ setattr(_m, attr, child)
+ return child
+
+ m.__getattr__ = _ga
+ return m
+
+
+class _StubSubpackageLoader(importlib.abc.Loader):
+ def __init__(self, mod_name):
+ self._mod_name = mod_name
+
+ def create_module(self, spec):
+ return _make_mod_stub(self._mod_name)
+
+ def exec_module(self, module):
+ pass
+
+
+class _StubSubpackageFinder(importlib.abc.MetaPathFinder):
+ def find_spec(self, fullname, path, target = None):
+ if "." not in fullname:
+ return None
+ parent = sys.modules.get(fullname.rsplit(".", 1)[0])
+ if parent is None:
+ return None
+ if getattr(parent, "_unsloth_stub", None) is not _STUB_SENTINEL:
+ return None
+ return importlib.machinery.ModuleSpec(
+ fullname, _StubSubpackageLoader(fullname), is_package = True
+ )
+
+
+def install_torchao_windows_rocm_stub() -> None:
+ """Pre-stub torchao on Windows ROCm so transformers/peft imports don't crash.
+
+ No-op on every other platform (Windows CUDA included — there torchao is real
+ and shadowing it would break torchao-based quantization paths). Must run
+ before any import of transformers / unsloth_zoo. Safe to call once per worker
+ process.
+ """
+ # Gate on the active torch runtime, not env-var presence -- HIP_PATH /
+ # ROCM_PATH stay set after a user installs the HIP SDK and reverts to a
+ # CUDA torch wheel. AMD SDK / Radeon ROCm wheels may not set torch.version.hip
+ # but still encode "rocm" in torch.__version__, so accept either.
+ _is_win32_rocm = False
+ if sys.platform == "win32":
+ try:
+ import torch as _torch_probe
+
+ _is_win32_rocm = bool(
+ getattr(getattr(_torch_probe, "version", None), "hip", None)
+ or "rocm" in getattr(_torch_probe, "__version__", "").lower()
+ )
+ del _torch_probe
+ except Exception:
+ pass
+ if _is_win32_rocm:
+ # Register the finder only on Windows ROCm -- on other platforms there
+ # are no stub modules seeded, so appending is a pure accumulation.
+ sys.meta_path.append(_StubSubpackageFinder())
+ # Seed torchao top-level + key submodules; the finder handles the rest.
+ for _tao_name in (
+ "torchao",
+ "torchao.quantization",
+ "torchao.dtypes",
+ "torchao.float8",
+ "torchao.utils",
+ ):
+ if _tao_name not in sys.modules:
+ sys.modules[_tao_name] = _make_mod_stub(_tao_name)
diff --git a/studio/backend/core/export/worker.py b/studio/backend/core/export/worker.py
index f77b1966c4..defcff924b 100644
--- a/studio/backend/core/export/worker.py
+++ b/studio/backend/core/export/worker.py
@@ -439,6 +439,15 @@ def run_export_process(
'Install for better performance: pip install "triton-windows<3.7"'
)
+ # ── 1c. Stub torchao on Windows ROCm ──
+ # Shared with the training worker; see core/_torchao_stub.py for the full
+ # rationale (torchao -> torch.distributed._functional_collectives crashes on
+ # Windows ROCm because the RCCL backend is absent). No-op off Windows ROCm.
+ # Must run before any import of transformers / unsloth_zoo.
+ from core._torchao_stub import install_torchao_windows_rocm_stub
+
+ install_torchao_windows_rocm_stub()
+
# ── 2. Import ML libraries (fresh in this clean process) ──
try:
_send_response(
diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py
index f1331219cc..7e26882e95 100644
--- a/studio/backend/core/inference/llama_cpp.py
+++ b/studio/backend/core/inference/llama_cpp.py
@@ -17,6 +17,7 @@ import struct
import structlog
from loggers import get_logger
import shutil
+import signal
import socket
import subprocess
import sys
@@ -978,9 +979,6 @@ class LlamaCppBackend:
7. llama-server on PATH (system install)
8. ./bin/llama-server (legacy: extracted binary)
"""
- import os
- import sys
-
binary_name = "llama-server.exe" if sys.platform == "win32" else "llama-server"
# 1. Env var — direct path to binary
@@ -1251,6 +1249,33 @@ class LlamaCppBackend:
return total
+ @staticmethod
+ def _amd_apu_wants_unified_memory() -> bool:
+ """True only for AMD unified-memory APUs (gfx1150/gfx1151), where
+ GGML_CUDA_ENABLE_UNIFIED_MEMORY lets llama.cpp use shared system RAM.
+ False for discrete AMD, NVIDIA, CPU and macOS (the env hurts discrete
+ GPUs). ROCm reuses torch.cuda.*; the gcnArchName suffix is stripped."""
+ try:
+ import torch
+
+ if getattr(torch.version, "hip", None) is None:
+ return False
+ if not (hasattr(torch, "cuda") and torch.cuda.is_available()):
+ return False
+ for _i in range(torch.cuda.device_count()):
+ try:
+ _arch = (
+ getattr(torch.cuda.get_device_properties(_i), "gcnArchName", "")
+ or ""
+ )
+ except Exception:
+ continue
+ if _arch.split(":")[0].strip().lower() in {"gfx1150", "gfx1151"}:
+ return True
+ except Exception:
+ return False
+ return False
+
@staticmethod
def _get_gpu_free_memory() -> list[tuple[int, int]]:
"""Query free memory per GPU.
@@ -1268,8 +1293,6 @@ class LlamaCppBackend:
Returns list of (gpu_index, free_mib) sorted by index. Empty
list if no supported GPU is reachable.
"""
- import os
-
# ── NVIDIA via nvidia-smi ────────────────────────────────────
try:
result = subprocess.run(
@@ -3171,6 +3194,14 @@ class LlamaCppBackend:
env = child_env_without_native_path_secret()
binary_dir = str(Path(binary).parent)
+ # AMD unified-memory APUs (gfx1150/gfx1151): let llama.cpp use
+ # shared system RAM. setdefault so a user value wins.
+ if self._amd_apu_wants_unified_memory():
+ env.setdefault("GGML_CUDA_ENABLE_UNIFIED_MEMORY", "1")
+ logger.info(
+ "AMD unified-memory APU: set GGML_CUDA_ENABLE_UNIFIED_MEMORY=1"
+ )
+
if sys.platform == "win32":
# See _build_windows_path_dirs for ordering. #5106.
path_dirs = self._build_windows_path_dirs(
@@ -3180,6 +3211,24 @@ class LlamaCppBackend:
)
existing_path = env.get("PATH", "")
env["PATH"] = ";".join(path_dirs) + ";" + existing_path
+
+ # ROCm: the llama.cpp prebuilt bundles its own rocblas.dll
+ # but NOT the Tensile kernel library files it needs
+ # (rocblas/library/TensileLibrary*.dat + *.hsaco). The
+ # bundled DLL searches relative to its own location by
+ # default (i.e. /rocblas/library/) which does
+ # not exist, causing a silent crash on the first GEMM.
+ # ROCBLAS_TENSILE_LIBPATH overrides that search to point at
+ # the ROCm installation where the kernel files actually are.
+ _hip_path = os.environ.get(
+ "HIP_PATH", os.environ.get("ROCM_PATH", "")
+ )
+ if _hip_path:
+ _rocblas_lib = os.path.join(
+ _hip_path, "bin", "rocblas", "library"
+ )
+ if os.path.isdir(_rocblas_lib):
+ env.setdefault("ROCBLAS_TENSILE_LIBPATH", _rocblas_lib)
else:
# Linux: set LD_LIBRARY_PATH for shared libs next to the binary
# and CUDA runtime libs (libcudart, libcublas, etc.)
@@ -3888,10 +3937,6 @@ class LlamaCppBackend:
Falls back to pgrep + /proc//exe on Linux when psutil is
not installed.
"""
- import os
- import signal
- import sys
-
try:
# -- Build the ownership allowlist --------------------------------
# Two kinds of matches:
diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py
index b9643cac6a..0365b3ffd8 100644
--- a/studio/backend/core/training/trainer.py
+++ b/studio/backend/core/training/trainer.py
@@ -6,8 +6,10 @@ Unsloth Training Backend
Integrates Unsloth training capabilities with the FastAPI backend
"""
+import gc
import os
import sys
+import types
# Prevent tokenizer parallelism deadlocks when datasets uses multiprocessing fork
os.environ["TOKENIZERS_PARALLELISM"] = "false"
@@ -42,7 +44,10 @@ from utils.hardware import (
get_visible_gpu_count,
)
-torch._dynamo.config.recompile_limit = 64
+# recompile_limit was removed in some ROCm torch builds (e.g. pytorch.org/whl/rocm6.2).
+# Guard so training doesn't crash on RDNA2/RDNA3 with older ROCm torch wheels.
+if hasattr(torch._dynamo.config, "recompile_limit"):
+ torch._dynamo.config.recompile_limit = 64
from unsloth import FastLanguageModel, FastVisionModel, is_bfloat16_supported
from unsloth.chat_templates import get_chat_template
@@ -417,8 +422,6 @@ class UnslothTrainer:
in sys.modules. When the next training run calls dataset.map(num_proc=N),
forked child processes inherit this stale state and deadlock.
"""
- import sys as _sys
-
# Remove cloned audio repo paths from sys.path
base_dir = os.path.dirname(os.path.abspath(__file__))
audio_paths = [
@@ -433,15 +436,15 @@ class UnslothTrainer:
removed_paths = []
for path in audio_paths:
- if path in _sys.path:
- _sys.path.remove(path)
+ if path in sys.path:
+ sys.path.remove(path)
removed_paths.append(path)
# Remove stale audio modules from sys.modules
prefixes = ("snac", "whisper", "sparktts", "outetts")
- removed_modules = [key for key in _sys.modules if key.startswith(prefixes)]
+ removed_modules = [key for key in sys.modules if key.startswith(prefixes)]
for key in removed_modules:
- del _sys.modules[key]
+ del sys.modules[key]
if removed_paths or removed_modules:
logger.info(
@@ -538,10 +541,9 @@ class UnslothTrainer:
# clear_unsloth_compiled_cache() deletes the disk cache, but the flag
# prevents re-compilation — leaving missing cache files. Reloading
# restores original class definitions so Unsloth can re-compile cleanly.
- import sys as _sys
import importlib
- for _key, _mod in list(_sys.modules.items()):
+ for _key, _mod in list(sys.modules.items()):
if "transformers.models." in _key and ".modeling_" in _key:
if hasattr(_mod, "__UNSLOTH_PATCHED__"):
try:
@@ -657,6 +659,23 @@ class UnslothTrainer:
f"Using device_map='{device_map}' ({get_visible_gpu_count()} GPU(s) visible)"
)
+ # AMD ROCm hardware without native bfloat16 (e.g. RDNA2 / gfx103x)
+ # crashes with an LLVM error at the first bf16 kernel dispatch if
+ # dtype=None lets unsloth auto-pick bf16. Force float16 there so that
+ # path is never reached. NVIDIA keeps dtype=None so unsloth's own
+ # bf16/fp16/float32 auto-detection (including FORCE_FLOAT32 models) is
+ # honored -- older NVIDIA without bf16 (T4/V100) must NOT be coerced to
+ # float16 here, which the previous unconditional branch did wrongly.
+ # Derive ROCm inline (not hardware.IS_ROCM) because that flag is unset
+ # until detect_hardware() runs, which isn't guaranteed in this subprocess.
+ _is_rocm = (
+ bool(getattr(torch.version, "hip", None))
+ or "rocm" in torch.__version__.lower()
+ )
+ _auto_dtype = (
+ torch.float16 if (_is_rocm and not is_bfloat16_supported()) else None
+ )
+
# Branch based on model type
if self._audio_type == "csm":
# CSM: FastModel + auto_model=CsmForConditionalGeneration + load_in_4bit=False
@@ -666,7 +685,7 @@ class UnslothTrainer:
self.model, self.tokenizer = FastModel.from_pretrained(
model_name = model_name,
max_seq_length = max_seq_length,
- dtype = None,
+ dtype = _auto_dtype,
auto_model = CsmForConditionalGeneration,
load_in_4bit = False,
device_map = device_map,
@@ -683,7 +702,7 @@ class UnslothTrainer:
self.model, self.tokenizer = FastModel.from_pretrained(
model_name = model_name,
- dtype = None,
+ dtype = _auto_dtype,
load_in_4bit = False,
device_map = device_map,
full_finetuning = full_finetuning,
@@ -705,7 +724,7 @@ class UnslothTrainer:
self.model, self.tokenizer = FastLanguageModel.from_pretrained(
model_name = model_name,
max_seq_length = max_seq_length,
- dtype = None,
+ dtype = _auto_dtype,
load_in_4bit = load_in_4bit,
device_map = device_map,
full_finetuning = full_finetuning,
@@ -777,7 +796,7 @@ class UnslothTrainer:
self.model, self.tokenizer = FastModel.from_pretrained(
model_name = model_name,
max_seq_length = max_seq_length,
- dtype = None,
+ dtype = _auto_dtype,
load_in_4bit = load_in_4bit,
device_map = device_map,
full_finetuning = full_finetuning,
@@ -791,7 +810,7 @@ class UnslothTrainer:
self.model, self.tokenizer = FastVisionModel.from_pretrained(
model_name = model_name,
max_seq_length = max_seq_length,
- dtype = None, # Auto-detect
+ dtype = _auto_dtype,
load_in_4bit = load_in_4bit,
device_map = device_map,
full_finetuning = full_finetuning,
@@ -824,7 +843,7 @@ class UnslothTrainer:
self.model, self.tokenizer = FastLanguageModel.from_pretrained(
model_name = model_name,
max_seq_length = max_seq_length,
- dtype = None, # Auto-detect
+ dtype = _auto_dtype,
load_in_4bit = load_in_4bit,
device_map = device_map,
full_finetuning = full_finetuning,
@@ -1188,7 +1207,6 @@ class UnslothTrainer:
We patch at both instance AND class level for maximum reliability,
and strip non-TransformersKwargs params that Unsloth/PEFT inject.
"""
- import types
import torch
import torch.nn as nn
from transformers.models.csm.modeling_csm import (
@@ -1730,7 +1748,6 @@ class UnslothTrainer:
logger.info("Freeing SNAC codec model from GPU...\n")
snac_model.to("cpu")
del snac_model
- import gc
gc.collect()
torch.cuda.empty_cache()
@@ -1754,13 +1771,10 @@ class UnslothTrainer:
Mirrors Spark_TTS_(0_5B).ipynb: encode audio with BiCodec (semantic + global tokens),
format as special-token text strings for SFTTrainer with dataset_text_field="text".
"""
- import sys
import torch
import numpy as np
import torchaudio.transforms as T
- import subprocess
-
device = "cuda" if torch.cuda.is_available() else "cpu"
# The sparktts Python package lives in the SparkAudio/Spark-TTS GitHub repo,
@@ -1960,7 +1974,6 @@ class UnslothTrainer:
audio_tokenizer.model.cpu()
audio_tokenizer.feature_extractor.cpu()
del audio_tokenizer
- import gc
gc.collect()
torch.cuda.empty_cache()
@@ -1989,7 +2002,6 @@ class UnslothTrainer:
OuteTTS AudioProcessor for speaker representations, PromptProcessor for
training prompts. Outputs text strings for SFTTrainer with dataset_text_field="text".
"""
- import sys
import io
import tempfile
import torch
@@ -2173,7 +2185,6 @@ class UnslothTrainer:
del whisper_model
del audio_processor
del prompt_processor
- import gc
gc.collect()
torch.cuda.empty_cache()
diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py
index 632b38d75a..a825321597 100644
--- a/studio/backend/core/training/worker.py
+++ b/studio/backend/core/training/worker.py
@@ -21,6 +21,9 @@ import shutil
import sys
import time
import traceback
+import gc
+import re
+import types
import subprocess as _sp
from pathlib import Path
from typing import Any, Callable
@@ -70,6 +73,58 @@ _TILELANG_INSTALL_TIMEOUT_S = 600
_TVM_FFI_BROKEN_VERSIONS = ("0.1.10", "0.1.11")
_FAST_PATH_HOOKS_SKIP_ENV = "UNSLOTH_STUDIO_SKIP_FAST_PATH_HOOKS"
+# Module-level handle so the torch.library.Library registration survives past
+# run_training_process() and is not garbage collected mid-run.
+_WINDOWS_ROCM_GROUPED_MM_LIB = None
+
+# Worker subprocesses inherit the parent env but not the parent's
+# os.add_dll_directory registrations. Replicate main.py's Windows ROCm DLL
+# setup at module load so the first `import torch` can find amdhip64.dll even
+# when HIP_PATH\bin is not on the system PATH. Handles retained at module
+# scope so they are not garbage collected.
+_ROCM_DLL_HANDLES: list = []
+if sys.platform == "win32":
+
+ def _add_rocm_dll_dirs_worker() -> None:
+ _candidates: list[str] = []
+ for _var in ("HIP_PATH", "ROCM_PATH"):
+ _val = os.environ.get(_var)
+ if _val:
+ _candidates.append(os.path.join(_val, "bin"))
+ _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), key = _ver_key, reverse = True
+ ):
+ _bin = os.path.join(_default_root, _ver, "bin")
+ if os.path.isdir(_bin):
+ _candidates.append(_bin)
+ except OSError:
+ pass
+ for _d in _candidates:
+ if os.path.isdir(_d):
+ try:
+ _ROCM_DLL_HANDLES.append(os.add_dll_directory(_d))
+ except (OSError, AttributeError):
+ pass
+
+ _add_rocm_dll_dirs_worker()
+ del _add_rocm_dll_dirs_worker
+
def _model_wants_causal_conv1d(model_name: str) -> bool:
name = model_name.lower()
@@ -320,11 +375,21 @@ def _install_package_wheel_first(
f"{snippet}",
)
else:
- logger.error(
- "Failed to install %s from PyPI:\n%s",
- display_name,
- result.stdout,
- )
+ if sys.platform == "win32":
+ # No prebuilt wheel and no source build toolchain on Windows --
+ # this is expected for packages like causal-conv1d. Log at
+ # info so users aren't alarmed by what looks like an error.
+ logger.info(
+ "%s is not available on Windows (no prebuilt wheel); skipping",
+ display_name,
+ )
+ logger.debug("Install output:\n%s", result.stdout)
+ else:
+ logger.error(
+ "Failed to install %s from PyPI:\n%s",
+ display_name,
+ result.stdout,
+ )
return False
if is_hip:
@@ -337,6 +402,9 @@ def _install_package_wheel_first(
def _ensure_causal_conv1d_fast_path(event_queue: Any, model_name: str) -> None:
if not _model_wants_causal_conv1d(model_name):
return
+ if sys.platform == "win32":
+ logger.info("causal-conv1d: no prebuilt wheel for Windows; skipping")
+ return
_install_package_wheel_first(
event_queue = event_queue,
@@ -404,6 +472,11 @@ def _ensure_flash_linear_attention_unconditional(event_queue: Any) -> bool:
"""Install pinned FLA + fla-core with --no-deps. Returns True iff importable post-call."""
if os.getenv(_FLA_SKIP_ENV) == "1":
return False
+ if sys.platform == "win32":
+ logger.info(
+ "Skipping flash-linear-attention install: no prebuilt wheel for Windows"
+ )
+ return False
if sys.version_info < _FLA_MIN_PYTHON:
logger.info(
"Skipping flash-linear-attention install: requires Python >= %d.%d, have %s",
@@ -483,10 +556,17 @@ def _ensure_flash_linear_attention_unconditional(event_queue: Any) -> bool:
return False
if result.returncode != 0:
- logger.warning(
- "flash-linear-attention install failed (continuing on torch fallback):\n%s",
- result.stdout,
- )
+ if sys.platform == "win32":
+ logger.info(
+ "flash-linear-attention not available on Windows (no prebuilt wheel); "
+ "continuing on torch fallback"
+ )
+ logger.debug("Install output:\n%s", result.stdout)
+ else:
+ logger.warning(
+ "flash-linear-attention install failed (continuing on torch fallback):\n%s",
+ result.stdout,
+ )
_send_status(
event_queue,
"flash-linear-attention install failed; continuing without it",
@@ -607,15 +687,61 @@ def _tilelang_importable() -> bool:
def _torch_has_hip() -> bool:
- """True iff torch is a ROCm build; `torch.version.hip` is the only reliable signal on x86_64 ROCm."""
+ """True iff torch is a ROCm build.
+
+ `torch.version.hip` covers official PyTorch ROCm wheels; AMD SDK / Radeon
+ wheels can leave it unset but still encode "rocm" in `torch.__version__`.
+ """
try:
import torch as _torch
- return getattr(_torch.version, "hip", None) is not None
+ return bool(
+ getattr(_torch.version, "hip", None)
+ or "rocm" in getattr(_torch, "__version__", "").lower()
+ )
except Exception:
return False
+def _rocm_classify_unified_memory(props: Any) -> tuple[str, bool]:
+ """Classify a ROCm device as unified-memory (APU) or discrete.
+
+ Returns ``(gcn_arch, is_unified)`` where:
+ - ``gcn_arch`` is the canonical arch string (e.g. ``"gfx1151"``) when a
+ known attribute is present, or ``""`` when all arch attrs are absent.
+ - ``is_unified`` is ``True`` for AMD APUs with a shared GPU/system-RAM pool
+ (gfx1150 Strix Point, gfx1151 Strix Halo) — these need a lower
+ ``set_per_process_memory_fraction`` cap to leave headroom for the OS.
+
+ Classification priority:
+ 1. ``gcnArchName`` / variant spellings (stable, naming-independent).
+ 2. Device-name substring match as a last-resort fallback when all arch
+ attrs are absent (AMD SDK / Radeon wheels may not populate them):
+ - gfx1150 Strix Point: ``Radeon 890M``, ``Radeon 880M``
+ - gfx1151 Strix Halo: ``Radeon 8060S`` (Ryzen AI MAX+ 395),
+ ``Radeon 8050S`` (cut-down SKU)
+ """
+ gcn_arch = ""
+ for _attr in ("gcnArchName", "gcn_arch_name", "arch_name", "gfx_arch_name"):
+ _v = (getattr(props, _attr, "") or "").split(":")[0].strip()
+ if _v:
+ gcn_arch = _v
+ break
+
+ if gcn_arch:
+ return gcn_arch, gcn_arch in {"gfx1150", "gfx1151"}
+
+ # Arch attrs absent — fall back to device-name matching.
+ dev_lower = (getattr(props, "name", "") or "").lower()
+ is_unified = (
+ "890m" in dev_lower
+ or "880m" in dev_lower
+ or "8060s" in dev_lower
+ or "8050s" in dev_lower
+ )
+ return gcn_arch, is_unified
+
+
def _tilelang_platform_supported() -> bool:
"""True iff a tilelang 0.1.8 wheel will load: Linux x86_64/aarch64, non-HIP torch.
@@ -881,6 +1007,9 @@ def _install_fast_path_hooks(event_queue: Any, model_name: str) -> None:
_ensure_tilelang_backend_unconditional(eq)
def _causal_conv1d_install(eq: Any) -> bool:
+ if sys.platform == "win32":
+ logger.info("causal-conv1d: no prebuilt wheel for Windows; skipping")
+ return False
ok = _install_package_wheel_first(
event_queue = eq,
import_name = "causal_conv1d",
@@ -1133,7 +1262,6 @@ def _run_mlx_training(event_queue, stop_queue, config):
Mirrors the event_queue protocol so the parent process pump works unchanged.
"""
import time
- import gc
import math
import threading
import queue as _queue
@@ -1893,6 +2021,340 @@ def run_training_process(
'Install for better performance: pip install "triton-windows<3.7"'
)
+ # ── 1d. Stub torchao on Windows ROCm ──
+ # Shared with the export worker; see core/_torchao_stub.py for the full
+ # rationale (torchao -> torch.distributed._functional_collectives crashes on
+ # Windows ROCm because the RCCL backend is absent). No-op off Windows ROCm.
+ # Must run before any import of transformers / unsloth_zoo.
+ from core._torchao_stub import install_torchao_windows_rocm_stub
+
+ install_torchao_windows_rocm_stub()
+
+ # ── 1e. Ensure torch.distributed helper attrs are present ──
+ # Single-GPU training never initialises the process group, so these helpers
+ # are never called — but transformers/trl import them unconditionally.
+ _td_stubs = {
+ "is_initialized": lambda: False,
+ "is_available": lambda: False,
+ "is_torchelastic_launched": lambda: False,
+ "get_rank": lambda: 0,
+ "get_world_size": lambda: 1,
+ "barrier": lambda: None,
+ }
+
+ try:
+ import torch.distributed as _td
+
+ for _name, _stub in _td_stubs.items():
+ if not hasattr(_td, _name):
+ setattr(_td, _name, _stub)
+ except Exception:
+ _td_mock = types.ModuleType("torch.distributed")
+ for _name, _stub in _td_stubs.items():
+ setattr(_td_mock, _name, _stub)
+ sys.modules["torch.distributed"] = _td_mock
+ try:
+ import torch as _torch
+
+ _torch.distributed = _td_mock
+ except Exception:
+ pass
+
+ # ── 1f. Windows ROCm runtime patches ──
+ # torch._grouped_mm has a null HIP kernel on gfx1200 (ROCm ≤ 7.12 Windows),
+ # causing 0xC0000005 (access violation) during training.
+ #
+ # Root cause: the JitDecomp autograd decomposition system (NOT torch.compile)
+ # dispatches _grouped_mm → _fused_adagrad_ → _grouped_mm HIP → null crash.
+ # TORCHDYNAMO_DISABLE=1 stops the compiler frontend but does NOT stop
+ # JitDecomp, so we must also override the CUDA dispatch key for _grouped_mm
+ # with a safe Python fallback.
+ #
+ # Fixed in AMD's wheel: torch==2.11.0+rocm7.13.0 — the 3-D batch and grouped
+ # (with offs) variants of _grouped_mm now have working HIP kernels on gfx1200.
+ # We gate the dispatch override on HIP < 7.13 so users on the fixed wheel get
+ # the real GPU kernel rather than our Python fallback.
+ #
+ # Verified: null on torch==2.10.0+rocm7.12.0; fixed on torch==2.11.0+rocm7.13.0.
+ #
+ # Schema: _grouped_mm(Tensor self, Tensor mat2, Tensor? offs=None,
+ # Tensor? bias=None, ScalarType? out_dtype=None) -> Tensor
+ # offs: optional group-split offsets (MoE-style variable-size batches)
+ #
+ # torch is already in sys.modules from section 1e's `import torch.distributed`.
+ # Module-level _WINDOWS_ROCM_GROUPED_MM_LIB keeps the registration alive past
+ # function return / mid-run GC.
+ global _WINDOWS_ROCM_GROUPED_MM_LIB
+ if sys.platform == "win32":
+ _torch_for_rocm = sys.modules.get("torch")
+ # Broad check: torch.version.hip OR "rocm" in torch.__version__.
+ # AMD SDK / Radeon Windows wheels do not always populate
+ # torch.version.hip; without the broad check the BNB version pin,
+ # dynamo-disable, and _grouped_mm fallback below silently skip
+ # (matches the torchao stub gate above and main.py).
+ _build_version_for_rocm = (
+ getattr(_torch_for_rocm, "__version__", "").lower()
+ if _torch_for_rocm is not None
+ else ""
+ )
+ _is_win_rocm_torch = bool(
+ _torch_for_rocm is not None
+ and (
+ getattr(getattr(_torch_for_rocm, "version", None), "hip", None)
+ or "rocm" in _build_version_for_rocm
+ )
+ )
+ if _is_win_rocm_torch:
+ # Disable dynamo (belt-and-suspenders; JitDecomp patch below is the
+ # real fix, but keeping dynamo off avoids any other compile paths).
+ if "TORCHDYNAMO_DISABLE" not in os.environ:
+ os.environ["TORCHDYNAMO_DISABLE"] = "1"
+ logger.info("Windows ROCm: torch.compile (dynamo) disabled")
+
+ # BNB auto-detects the HIP version from torch.version.hip and uses
+ # it to choose which DLL to load (e.g. "7.13" → rocm713.dll).
+ # AMD's Windows BNB prerelease wheel ships only one rocm DLL, and its
+ # version suffix does not always match the torch HIP version (e.g.
+ # torch==2.11.0+rocm7.13.0 ships HIP 7.13, but the BNB wheel still
+ # ships rocm72.dll). We detect the actual DLL name from the installed
+ # package and override BNB's auto-detection. "72" is a safe fallback
+ # if detection fails. Callers may override by pre-setting the var.
+ if "BNB_ROCM_VERSION" not in os.environ:
+ _bnb_rocm_ver = None
+ try:
+ import glob as _glob
+ import importlib.util as _ilu
+ import re as _re
+
+ _bnb_spec = _ilu.find_spec("bitsandbytes")
+ if _bnb_spec and _bnb_spec.submodule_search_locations:
+ _all_vers: list[str] = []
+ for _pkg_dir in _bnb_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.
+ # Filesystem glob order is not guaranteed, so always
+ # sort rather than stopping at the first match.
+ if _all_vers:
+ _bnb_rocm_ver = max(_all_vers, key = lambda v: int(v))
+ except Exception:
+ pass
+ _bnb_rocm_ver = _bnb_rocm_ver or "72"
+ os.environ["BNB_ROCM_VERSION"] = _bnb_rocm_ver
+ logger.info(
+ "Windows ROCm: set BNB_ROCM_VERSION=%s "
+ "(detected from installed BNB wheel; "
+ "overrides torch.version.hip auto-detection)",
+ _bnb_rocm_ver,
+ )
+
+ # Parse HIP version for the kernel-fix gate below.
+ # torch.version.hip can be "7.13.99004", "7.2.0", etc.
+ # AMD SDK / Radeon wheels may leave torch.version.hip unset and
+ # encode the ROCm version in torch.__version__ instead
+ # (e.g. "2.11.0+rocm7.13.0" or "2.9.0+rocmsdk20251116"); fall back
+ # to that string when version.hip is missing.
+ def _hip_ver_at_least(major: int, minor: int) -> bool:
+ _hip_str = getattr(
+ getattr(_torch_for_rocm, "version", None), "hip", None
+ )
+ if not _hip_str:
+ # Try the standard "+rocmX.Y.Z" embedded version first
+ # (e.g. "2.11.0+rocm7.13.0").
+ _ver_match = re.search(r"rocm(\d+)\.(\d+)", _build_version_for_rocm)
+ if _ver_match:
+ return (
+ int(_ver_match.group(1)),
+ int(_ver_match.group(2)),
+ ) >= (major, minor)
+ # AMD SDK / Radeon Windows wheels encode the build as
+ # "+rocmsdk" (e.g. "2.9.0+rocmsdk20251116") with no
+ # explicit rocmX.Y component. The rocmsdk format was
+ # introduced after the gfx120X null-kernel fix landed in
+ # ROCm 7.13, so any wheel with this suffix is new enough to
+ # have working HIP kernels. Treat as >= 7.13 rather than
+ # falling back to False and installing the Python workaround
+ # on a wheel that doesn't need it.
+ if "rocmsdk" in _build_version_for_rocm:
+ logger.debug(
+ "Windows ROCm: AMD SDK wheel detected (%r); "
+ "assuming HIP >= %d.%d (rocmsdk wheels post-date "
+ "the gfx120X null-kernel fix)",
+ _build_version_for_rocm,
+ major,
+ minor,
+ )
+ return True
+ return False
+ try:
+ _parts = [int(x) for x in str(_hip_str).split(".")[:2]]
+ if len(_parts) < 2:
+ logger.warning(
+ "Windows ROCm: torch.version.hip %r has fewer than "
+ "two components; cannot compare against %d.%d",
+ _hip_str,
+ major,
+ minor,
+ )
+ return False
+ return (_parts[0], _parts[1]) >= (major, minor)
+ except ValueError:
+ logger.warning(
+ "Windows ROCm: could not parse torch.version.hip %r as "
+ "a version number; assuming HIP < %d.%d",
+ _hip_str,
+ major,
+ minor,
+ )
+ return False
+
+ # _grouped_mm HIP kernel was null on gfx1200 in ROCm ≤ 7.12,
+ # causing 0xC0000005. AMD fixed it in ROCm 7.13 (torch 2.11+).
+ # Only install the Python fallback on the affected versions so users
+ # on 7.13+ get the real GPU kernel for MoE workloads.
+ if not _hip_ver_at_least(7, 13):
+ try:
+ import warnings as _warnings
+
+ _gm_lib = _torch_for_rocm.library.Library("aten", "IMPL")
+
+ def _grouped_mm_safe_impl(
+ self, mat2, offs = None, bias = None, out_dtype = None
+ ):
+ """Python mm/bmm fallback for _grouped_mm on gfx1200 (null HIP kernel, ROCm ≤ 7.12)."""
+ _t = _torch_for_rocm
+ if offs is None:
+ # No offsets: behave like the real op, which
+ # accepts either (M, K) x (K, N) -> mm, or 3-D
+ # batched inputs -> bmm. Picking torch.mm
+ # unconditionally previously raised "self must be
+ # a matrix" on 3-D MoE workloads.
+ if self.dim() == 3 and mat2.dim() == 3:
+ result = _t.bmm(self.contiguous(), mat2.contiguous())
+ elif self.dim() == 3 and mat2.dim() == 2:
+ # Broadcast 2-D mat2 across the batch dim.
+ result = _t.matmul(self.contiguous(), mat2.contiguous())
+ elif self.dim() == 2 and mat2.dim() == 3:
+ # Broadcast 2-D self across batch via matmul semantics.
+ result = _t.matmul(self.contiguous(), mat2.contiguous())
+ else:
+ result = _t.mm(self.contiguous(), mat2.contiguous())
+ else:
+ # Grouped case: offs[i] is the exclusive end-row of
+ # group i in `self`; mat2 may be 3-D or 2-D.
+ offs_list = offs.tolist()
+ pieces = []
+ prev = 0
+ for idx, end in enumerate(offs_list):
+ end = int(end)
+ a_part = self[prev:end].contiguous()
+ if mat2.dim() == 3:
+ b_part = mat2[idx].contiguous()
+ else:
+ b_part = mat2.contiguous()
+ pieces.append(_t.mm(a_part, b_part))
+ prev = end
+ # Include any trailing rows not covered by offs
+ if prev < self.shape[0]:
+ a_tail = self[prev:].contiguous()
+ b_tail = (
+ mat2[-1].contiguous()
+ if mat2.dim() == 3
+ else mat2.contiguous()
+ )
+ pieces.append(_t.mm(a_tail, b_tail))
+ result = (
+ _t.cat(pieces, dim = 0)
+ if pieces
+ else _t.zeros(
+ 0,
+ mat2.shape[-1],
+ device = self.device,
+ dtype = self.dtype,
+ )
+ )
+ if bias is not None:
+ result = result + bias
+ if out_dtype is not None:
+ result = result.to(out_dtype)
+ elif result.dtype != self.dtype:
+ result = result.to(self.dtype)
+ return result
+
+ with _warnings.catch_warnings():
+ _warnings.simplefilter("ignore")
+ _gm_lib.impl("_grouped_mm", _grouped_mm_safe_impl, "CUDA")
+
+ _WINDOWS_ROCM_GROUPED_MM_LIB = _gm_lib # prevent GC
+ logger.info(
+ "Windows ROCm: patched _grouped_mm CUDA dispatch "
+ "(null HIP kernel on gfx1200, ROCm ≤ 7.12 — "
+ "bypassed with Python mm fallback)"
+ )
+ except Exception as _patch_exc:
+ logger.warning(
+ "Windows ROCm: could not patch _grouped_mm — "
+ "training may crash with 0xC0000005: %s",
+ _patch_exc,
+ )
+ else:
+ logger.info(
+ "Windows ROCm: HIP >= 7.13 — _grouped_mm kernel is functional, "
+ "skipping Python fallback (AMD fixed gfx1200 null kernel in ROCm 7.13)"
+ )
+
+ # ── 1g. ROCm OOM guard ──
+ # On RDNA 4 (gfx1200/gfx1201) and other ROCm GPUs, exhausting VRAM can
+ # cause a HIP driver hang that freezes the entire system rather than
+ # raising a Python exception. set_per_process_memory_fraction caps the
+ # HIP allocator so PyTorch raises OutOfMemoryError before hitting the
+ # hardware limit, giving the UI a clean error instead of a system freeze.
+ # Only applied on ROCm -- NVIDIA CUDA has a graceful OOM path and does
+ # not need this cap.
+ # Unified-memory APUs (gfx1150 Strix Point / gfx1151 Strix Halo) share GPU
+ # and system RAM in one pool: 0.90 of 128 GB starves the OS. Use 0.80 there.
+ # Primary classifier: gcnArchName from device properties — stable within a
+ # product family and naming-independent. AMD SDK / Radeon wheels may omit
+ # gcnArchName or expose it under a variant spelling, so we try several attr
+ # names then fall back to known device-name markers as a last resort.
+ # Non-fatal: silently skipped if torch is not importable.
+ if _hw.IS_ROCM:
+ try:
+ import torch as _torch_mem
+
+ if _torch_mem.cuda.is_available():
+ # Classify unified vs discrete via _rocm_classify_unified_memory.
+ # See that function's docstring for classification priority.
+ _props = _torch_mem.cuda.get_device_properties(0)
+ _dev_name = _props.name
+ _gcn_arch, _is_unified = _rocm_classify_unified_memory(_props)
+ if _is_unified and not _gcn_arch:
+ logger.debug(
+ "ROCm OOM guard: gcnArchName absent -- inferred "
+ "unified memory from device name %r; applying 0.80 cap",
+ _dev_name,
+ )
+ _mem_fraction = 0.80 if _is_unified else 0.90
+ _torch_mem.cuda.set_per_process_memory_fraction(_mem_fraction)
+ logger.info(
+ "ROCm OOM guard: set_per_process_memory_fraction(%.2f) — "
+ "%s memory host (%s, %s)",
+ _mem_fraction,
+ "unified" if _is_unified else "discrete",
+ _dev_name,
+ _gcn_arch or "unknown arch",
+ )
+ except Exception as _oom_guard_err:
+ logger.debug("Could not set GPU memory fraction: %s", _oom_guard_err)
+
# ── 2. Now import ML libraries (fresh in this clean process) ──
try:
_send_status(event_queue, "Importing Unsloth...")
@@ -2347,14 +2809,38 @@ def run_training_process(
)
except Exception as exc:
- event_queue.put(
- {
- "type": "error",
- "error": str(exc),
- "stack": traceback.format_exc(limit = 20),
- "ts": time.time(),
- }
+ _exc_str = str(exc).lower()
+ _is_oom = (
+ "out of memory" in _exc_str
+ or "hip out of memory" in _exc_str
+ or "cuda out of memory" in _exc_str
+ or type(exc).__name__ == "OutOfMemoryError"
)
+ if _is_oom:
+ _oom_msg = (
+ "GPU ran out of VRAM during training.\n"
+ "To fix: reduce max_seq_length (e.g. 2048–4096), enable "
+ "gradient_checkpointing=True, lower per_device_train_batch_size, "
+ "or use a smaller model / higher quantization."
+ )
+ logger.error("Training stopped: GPU OOM — %s", exc)
+ event_queue.put(
+ {
+ "type": "error",
+ "error": _oom_msg,
+ "stack": traceback.format_exc(limit = 20),
+ "ts": time.time(),
+ }
+ )
+ else:
+ event_queue.put(
+ {
+ "type": "error",
+ "error": str(exc),
+ "stack": traceback.format_exc(limit = 20),
+ "ts": time.time(),
+ }
+ )
def _send_status(event_queue: Any, message: str) -> None:
diff --git a/studio/backend/main.py b/studio/backend/main.py
index bd007d2e37..884a5e5841 100644
--- a/studio/backend/main.py
+++ b/studio/backend/main.py
@@ -12,6 +12,110 @@ from pathlib import Path as _Path
# Suppress annoying C-level dependency warnings globally
os.environ["PYTHONWARNINGS"] = "ignore"
+# ── Windows AMD ROCm DLL injection ──────────────────────────────────────────
+# Python 3.8+ ignores PATH for extension modules; register ROCm bin dirs with
+# os.add_dll_directory() so amdhip64.dll etc. are found before any torch import.
+if sys.platform == "win32":
+ # Retained at module scope -- os.add_dll_directory returns a handle that
+ # removes the search-path entry when garbage collected.
+ _ROCM_DLL_HANDLES: list = []
+
+ def _add_rocm_dll_dirs() -> None:
+ candidates = []
+ # 1. HIP_PATH / ROCM_PATH -- set by the AMD HIP SDK installer
+ for _var in ("HIP_PATH", "ROCM_PATH"):
+ _val = os.environ.get(_var)
+ if _val:
+ candidates.append(os.path.join(_val, "bin"))
+ # 2. Standard AMD installer location: C:\Program Files\AMD\ROCm\\bin
+ # Scan all installed versions, newest first.
+ _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), key = _ver_key, reverse = True
+ ):
+ _bin = os.path.join(_default_root, _ver, "bin")
+ if os.path.isdir(_bin):
+ candidates.append(_bin)
+ except OSError:
+ pass
+ for _d in candidates:
+ if os.path.isdir(_d):
+ try:
+ _ROCM_DLL_HANDLES.append(os.add_dll_directory(_d))
+ except (OSError, AttributeError):
+ pass
+
+ _add_rocm_dll_dirs()
+ del _add_rocm_dll_dirs
+
+ # ── Windows AMD ROCm: set BNB_ROCM_VERSION before any bitsandbytes import ─
+ # bitsandbytes on Windows ROCm tries to load libbitsandbytes_rocm.dll
+ # where comes from torch.version.hip (e.g. "7.13..." → "713").
+ # The installed BNB wheel ships rocm72.dll (not rocm713.dll), so without
+ # this the server process crashes with "Configured ROCm binary not found".
+ # Detect the available DLL, fall back to "72", and set BNB_ROCM_VERSION
+ # before any import that pulls in bitsandbytes (mirrors worker.py logic).
+ # Gate on the rocm bnb DLL (the exact file this configures) or HIP_PATH/
+ # ROCM_PATH, not on torch.version.hip: that needed importing torch on every
+ # Windows host (NVIDIA/CPU included), adding seconds to startup. Radeon
+ # wheels without HIP_PATH still ship the rocm bnb DLL, so they are covered.
+ if "BNB_ROCM_VERSION" not in os.environ:
+ import glob as _glob
+ import logging as _logging
+
+ _hip_env = bool(os.environ.get("HIP_PATH") or os.environ.get("ROCM_PATH"))
+ _bnb_rocm_ver = None
+ _found_rocm_bnb = False
+ try:
+ import importlib.util as _ilu
+
+ _bnb_spec = _ilu.find_spec("bitsandbytes")
+ # submodule_search_locations (not spec.origin) handles editable installs.
+ if _bnb_spec and _bnb_spec.submodule_search_locations:
+ import re as _re_bnb
+
+ _all_vers_main: list[str] = []
+ for _pkg_dir in _bnb_spec.submodule_search_locations:
+ for _dll in _glob.glob(
+ os.path.join(_pkg_dir, "libbitsandbytes_rocm*.dll")
+ ):
+ _found_rocm_bnb = True
+ _km = _re_bnb.search(
+ r"libbitsandbytes_rocm(\d+)\.dll", os.path.basename(_dll)
+ )
+ if _km:
+ _all_vers_main.append(_km.group(1))
+ if _all_vers_main:
+ _bnb_rocm_ver = max(_all_vers_main, key = lambda v: int(v))
+ except Exception as _e:
+ _logging.getLogger(__name__).warning(
+ "Windows ROCm: BNB DLL detection failed (%s); falling back to version '72'",
+ _e,
+ )
+ # rocm bnb DLL present, or HIP_PATH/ROCM_PATH set (DLL unparsable -> "72").
+ if _found_rocm_bnb or _hip_env:
+ _bnb_rocm_ver_final = _bnb_rocm_ver or "72"
+ os.environ["BNB_ROCM_VERSION"] = _bnb_rocm_ver_final
+ _logging.getLogger(__name__).info(
+ "Windows ROCm: set BNB_ROCM_VERSION=%s (from installed BNB wheel)",
+ _bnb_rocm_ver_final,
+ )
+
# Ensure backend dir is on sys.path so _platform_compat is importable when
# main.py is launched directly (e.g. `uvicorn main:app`).
_backend_dir = str(_Path(__file__).parent)
@@ -764,8 +868,6 @@ def _strip_crossorigin(html_bytes: bytes) -> bytes:
@font-face downloads to fail silently. Stripping the attribute
makes them regular same-origin fetches that work on any protocol.
"""
- import re as _re
-
html = html_bytes.decode("utf-8")
html = _re.sub(r'\s+crossorigin(?:="[^"]*")?', "", html)
return html.encode("utf-8")
diff --git a/studio/backend/tests/test_amd_apu_unified_memory.py b/studio/backend/tests/test_amd_apu_unified_memory.py
new file mode 100644
index 0000000000..e0b819d54b
--- /dev/null
+++ b/studio/backend/tests/test_amd_apu_unified_memory.py
@@ -0,0 +1,52 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""GGML_CUDA_ENABLE_UNIFIED_MEMORY must be set only for AMD unified-memory APUs
+(gfx1150/gfx1151), never for discrete AMD, NVIDIA, CPU or macOS."""
+
+from __future__ import annotations
+
+import sys
+import types
+
+import pytest
+
+from core.inference.llama_cpp import LlamaCppBackend
+
+
+def _fake_torch(hip, archs, *, cuda_ok = True):
+ t = types.ModuleType("torch")
+ t.version = types.SimpleNamespace(hip = hip)
+ t.cuda = types.SimpleNamespace(
+ is_available = lambda: cuda_ok,
+ device_count = lambda: len(archs),
+ get_device_properties = lambda i: types.SimpleNamespace(gcnArchName = archs[i]),
+ )
+ return t
+
+
+@pytest.mark.parametrize(
+ "hip,archs,expected",
+ [
+ ("6.2.0", ["gfx1151:xnack-"], True), # Strix Halo APU (suffix stripped)
+ ("6.2.0", ["gfx1150"], True), # Strix Point APU
+ ("6.2.0", ["gfx1100"], False), # discrete RDNA3
+ ("6.2.0", ["gfx1201"], False), # discrete RDNA4
+ ("6.2.0", ["gfx942"], False), # MI300X (data center)
+ (None, ["sm_90"], False), # NVIDIA (no torch.version.hip)
+ ("6.2.0", ["gfx1100", "gfx1151"], True), # mixed dGPU + APU
+ ],
+)
+def test_apu_unified_memory_gating(monkeypatch, hip, archs, expected):
+ monkeypatch.setitem(sys.modules, "torch", _fake_torch(hip, archs))
+ assert LlamaCppBackend._amd_apu_wants_unified_memory() is expected
+
+
+def test_cpu_no_cuda_returns_false(monkeypatch):
+ monkeypatch.setitem(sys.modules, "torch", _fake_torch("6.2.0", [], cuda_ok = False))
+ assert LlamaCppBackend._amd_apu_wants_unified_memory() is False
+
+
+def test_missing_torch_returns_false(monkeypatch):
+ monkeypatch.setitem(sys.modules, "torch", None)
+ assert LlamaCppBackend._amd_apu_wants_unified_memory() is False
diff --git a/studio/backend/tests/test_gguf_routing.py b/studio/backend/tests/test_gguf_routing.py
new file mode 100644
index 0000000000..1b299cba19
--- /dev/null
+++ b/studio/backend/tests/test_gguf_routing.py
@@ -0,0 +1,102 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""
+Tests for GGUF routing in detect_gguf_model.
+
+Regression test for the bug where a .gguf file temporarily appears
+inaccessible on Windows during llama-server process teardown, causing
+is_file() to return False and the model to be routed to the transformers
+backend instead of llama-server.
+"""
+
+import sys
+import os
+import types
+from pathlib import Path
+from unittest.mock import patch
+
+# Stub structlog before importing backend modules (mirrors other tests in this suite)
+if "structlog" not in sys.modules:
+
+ class _DummyLogger:
+ def __getattr__(self, _):
+ return lambda *a, **k: None
+
+ sys.modules["structlog"] = types.SimpleNamespace(
+ get_logger = lambda *a, **k: _DummyLogger(),
+ BoundLogger = _DummyLogger,
+ )
+
+sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))
+
+from utils.models.model_config import detect_gguf_model
+
+
+def test_detects_gguf_file_normally(tmp_path):
+ """Normal case: .gguf file exists and is accessible."""
+ gguf = tmp_path / "gpt-oss-20b-MXFP4.gguf"
+ gguf.write_bytes(b"")
+ result = detect_gguf_model(str(gguf))
+ assert result is not None
+ assert result.endswith("gpt-oss-20b-MXFP4.gguf")
+
+
+def test_detects_gguf_when_stat_raises_oserror(tmp_path):
+ """
+ Regression: on Windows, both is_file() and exists() call stat() internally.
+ During the brief lock window after llama-server is killed, stat() raises
+ OSError, causing both to return False. detect_gguf_model must still route
+ to llama-server based on the file extension alone.
+ """
+ gguf = tmp_path / "gpt-oss-20b-MXFP4.gguf"
+ gguf.write_bytes(b"")
+
+ original_stat = Path.stat
+
+ def flaky_stat(self, **kwargs):
+ if self.suffix.lower() == ".gguf":
+ raise OSError("file temporarily inaccessible (Windows lock window)")
+ return original_stat(self, **kwargs)
+
+ with patch.object(Path, "stat", flaky_stat):
+ result = detect_gguf_model(str(gguf))
+
+ assert result is not None, (
+ "detect_gguf_model returned None when stat() raised OSError. "
+ "This causes the model to fall through to the transformers backend."
+ )
+
+
+def test_does_not_detect_mmproj_as_main_model(tmp_path):
+ """mmproj files must never be returned as the primary model."""
+ mmproj = tmp_path / "mmproj-model-f16.gguf"
+ mmproj.write_bytes(b"")
+ result = detect_gguf_model(str(mmproj))
+ assert result is None
+
+
+def test_detects_gguf_in_directory(tmp_path):
+ """Directory containing a .gguf file is resolved to that file."""
+ gguf = tmp_path / "model-Q4_K_M.gguf"
+ gguf.write_bytes(b"")
+ result = detect_gguf_model(str(tmp_path))
+ assert result is not None
+ assert result.endswith("model-Q4_K_M.gguf")
+
+
+def test_directory_named_like_gguf_scans_inside(tmp_path):
+ """A directory named *.gguf resolves the real .gguf inside, not itself."""
+ gguf_dir = tmp_path / "mymodel.gguf"
+ gguf_dir.mkdir()
+ inner = gguf_dir / "model-Q4_K_M.gguf"
+ inner.write_bytes(b"")
+ result = detect_gguf_model(str(gguf_dir))
+ assert result is not None
+ assert result.endswith("model-Q4_K_M.gguf")
+
+
+def test_returns_none_for_non_gguf_path(tmp_path):
+ """Non-.gguf paths with no .gguf files inside return None."""
+ result = detect_gguf_model(str(tmp_path))
+ assert result is None
diff --git a/studio/backend/tests/test_lemonade_llamacpp_rocm_bins_mock.py b/studio/backend/tests/test_lemonade_llamacpp_rocm_bins_mock.py
new file mode 100644
index 0000000000..5d2d672890
--- /dev/null
+++ b/studio/backend/tests/test_lemonade_llamacpp_rocm_bins_mock.py
@@ -0,0 +1,430 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""Validates that the installer correctly resolves lemonade ROCm prebuilt assets.
+
+Uses a faked HostInfo so no AMD GPU is needed. Network calls to the lemonade
+GitHub API are stubbed out so the suite runs without internet access and is
+not subject to rate limits.
+"""
+
+from __future__ import annotations
+
+import importlib
+import sys
+from pathlib import Path
+from unittest.mock import patch
+
+import pytest
+
+_studio = Path(__file__).resolve().parent.parent.parent
+if str(_studio) not in sys.path:
+ sys.path.insert(0, str(_studio))
+
+_mod = importlib.import_module("install_llama_prebuilt")
+HostInfo = _mod.HostInfo
+resolve_lemonade_rocm_choice = getattr(_mod, "resolve_lemonade_rocm_choice", None)
+_LEMONADE_GFX_FAMILIES = getattr(_mod, "_LEMONADE_GFX_FAMILIES", None)
+
+if resolve_lemonade_rocm_choice is None or _LEMONADE_GFX_FAMILIES is None:
+ pytest.skip("PR symbols not present - check branch", allow_module_level = True)
+
+
+@pytest.fixture(autouse = True)
+def _clear_lemonade_release_cache():
+ """Prevent cross-test pollution of the lemonade release lru_cache when
+ future tests vary the fetch_json mock return value."""
+ _cache = getattr(_mod, "_fetch_lemonade_release_cached", None)
+ if _cache is not None and hasattr(_cache, "cache_clear"):
+ _cache.cache_clear()
+ yield
+ if _cache is not None and hasattr(_cache, "cache_clear"):
+ _cache.cache_clear()
+
+
+_STUB_TAG = "b1262"
+_STUB_OS_PREFIXES = ("ubuntu", "windows")
+_STUB_FAMILIES = ("gfx1151", "gfx1150", "gfx120X", "gfx110X", "gfx103X")
+
+
+def _stub_lemonade_release() -> dict:
+ """Minimal lemonade release payload covering all supported GPU/OS combinations."""
+ assets = [
+ {
+ "name": f"llama-{_STUB_TAG}-{prefix}-rocm-{family}-x64.zip",
+ "browser_download_url": (
+ f"https://github.com/lemonade-sdk/llamacpp-rocm/releases/download/"
+ f"{_STUB_TAG}/llama-{_STUB_TAG}-{prefix}-rocm-{family}-x64.zip"
+ ),
+ }
+ for prefix in _STUB_OS_PREFIXES
+ for family in _STUB_FAMILIES
+ ]
+ return {"tag_name": _STUB_TAG, "assets": assets}
+
+
+def _make_rocm_host(gfx_target: str, *, windows: bool = False) -> HostInfo:
+ return HostInfo(
+ system = "Windows" if windows else "Linux",
+ machine = "amd64" if windows else "x86_64",
+ is_windows = windows,
+ is_linux = not windows,
+ is_macos = False,
+ is_x86_64 = True,
+ is_arm64 = False,
+ nvidia_smi = None,
+ driver_cuda_version = None,
+ compute_caps = [],
+ visible_cuda_devices = None,
+ has_physical_nvidia = False,
+ has_usable_nvidia = False,
+ has_rocm = True,
+ rocm_gfx_target = gfx_target,
+ )
+
+
+def _lookup_family(gfx: str) -> str | None:
+ for prefix, family in _LEMONADE_GFX_FAMILIES:
+ if gfx.startswith(prefix):
+ return family
+ return None
+
+
+# ---------------------------------------------------------------------------
+# GPU family mapping
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.parametrize(
+ "gfx,expected_family",
+ [
+ ("gfx1151", "gfx1151"),
+ ("gfx1150", "gfx1150"),
+ ("gfx1201", "gfx120X"),
+ ("gfx1200", "gfx120X"),
+ ("gfx1100", "gfx110X"),
+ ("gfx1030", "gfx103X"),
+ ],
+)
+def test_gpu_family_mapping(gfx, expected_family):
+ assert _lookup_family(gfx) == expected_family
+
+
+def test_unknown_gpu_not_in_families():
+ assert _lookup_family("gfx999") is None
+
+
+# ---------------------------------------------------------------------------
+# Asset resolution - hits real lemonade GitHub API
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.parametrize(
+ "gfx,os_prefix,windows",
+ [
+ ("gfx1151", "ubuntu", False),
+ ("gfx1150", "ubuntu", False),
+ ("gfx1201", "ubuntu", False),
+ ("gfx1100", "ubuntu", False),
+ ("gfx1030", "ubuntu", False),
+ ("gfx1151", "windows", True),
+ ("gfx1100", "windows", True),
+ ],
+)
+def test_asset_resolves_for_known_gpu(gfx, os_prefix, windows):
+ host = _make_rocm_host(gfx, windows = windows)
+ with patch.object(_mod, "fetch_json", return_value = _stub_lemonade_release()):
+ result = resolve_lemonade_rocm_choice(
+ host, os_prefix, "default", llama_tag = "latest"
+ )
+ assert (
+ result is not None
+ ), f"Installer will NOT fetch lemonade binary for {gfx} ({os_prefix})"
+ assert _lookup_family(gfx) in result.name
+ assert result.url.startswith("https://github.com/lemonade-sdk/llamacpp-rocm")
+
+
+def test_unknown_gpu_falls_through_to_upstream():
+ host = _make_rocm_host("gfx999")
+ result = resolve_lemonade_rocm_choice(host, "ubuntu", "default", llama_tag = "latest")
+ assert result is None
+
+
+# ---------------------------------------------------------------------------
+# Simple-policy dispatcher must plan a lemonade ROCm attempt for AMD-only hosts.
+# This is the path setup.sh actually invokes (via --simple-policy), so the
+# lemonade integration is useless if it isn't wired in here.
+# ---------------------------------------------------------------------------
+
+direct_linux_release_plan = getattr(_mod, "direct_linux_release_plan", None)
+direct_upstream_release_plan = getattr(_mod, "direct_upstream_release_plan", None)
+
+
+def _stub_unsloth_release(release_tag: str = "b9022") -> dict:
+ # Minimal payload that parse_direct_linux_release_bundle accepts. It
+ # requires at least one `app-{label}-linux-x64*.tar.gz` asset for the
+ # bundle to be recognised; we ship a bare CPU one so the planner has a
+ # baseline non-ROCm attempt to fall through to.
+ asset_name = f"app-{release_tag}-linux-x64.tar.gz"
+ return {
+ "tag_name": release_tag,
+ "name": release_tag,
+ "assets": [
+ {
+ "name": asset_name,
+ "browser_download_url": f"https://example.invalid/{asset_name}",
+ },
+ ],
+ }
+
+
+@pytest.mark.skipif(
+ direct_linux_release_plan is None,
+ reason = "simple-policy dispatcher not present on this branch",
+)
+def test_simple_policy_plans_lemonade_for_rocm_host():
+ host = _make_rocm_host("gfx1151")
+ with patch.object(_mod, "fetch_json", return_value = _stub_lemonade_release()):
+ plan = direct_linux_release_plan(
+ _stub_unsloth_release(),
+ host,
+ "unslothai/llama.cpp",
+ "latest",
+ )
+ assert plan is not None, "ROCm host should not be skipped by simple-policy planner"
+ kinds = [a.install_kind for a in plan.attempts]
+ assert (
+ "linux-rocm" in kinds
+ ), f"simple-policy planner did not include a lemonade ROCm attempt; got {kinds}"
+ rocm_attempt = next(a for a in plan.attempts if a.install_kind == "linux-rocm")
+ assert rocm_attempt.source_label == "lemonade"
+ assert "gfx1151" in rocm_attempt.name
+
+
+@pytest.mark.skipif(
+ direct_upstream_release_plan is None,
+ reason = "simple-policy dispatcher not present on this branch",
+)
+def test_simple_policy_plans_lemonade_for_windows_hip_host():
+ host = _make_rocm_host("gfx1151", windows = True)
+ release = {
+ "tag_name": "b9022",
+ "name": "b9022",
+ "assets": [],
+ }
+ with patch.object(_mod, "fetch_json", return_value = _stub_lemonade_release()):
+ plan = direct_upstream_release_plan(
+ release, host, "ggml-org/llama.cpp", "latest"
+ )
+ assert plan is not None, "Windows ROCm host should plan a lemonade HIP attempt"
+ kinds = [a.install_kind for a in plan.attempts]
+ assert (
+ "windows-hip" in kinds
+ ), f"simple-policy planner did not include a lemonade HIP attempt; got {kinds}"
+
+
+@pytest.mark.skipif(
+ direct_upstream_release_plan is None,
+ reason = "simple-policy dispatcher not present on this branch",
+)
+def test_simple_policy_windows_hip_falls_back_to_upstream_when_lemonade_unavailable():
+ """If lemonade returns None (e.g. gfx999 or transient API failure), the planner
+ must still include the upstream HIP asset rather than silently downgrading to CPU."""
+ host = _make_rocm_host("gfx999", windows = True)
+ hip_asset = "llama-b9022-bin-win-hip-radeon-x64.zip"
+ release = {
+ "tag_name": "b9022",
+ "name": "b9022",
+ "assets": [
+ {
+ "name": hip_asset,
+ "browser_download_url": f"https://example.invalid/{hip_asset}",
+ },
+ ],
+ }
+ plan = direct_upstream_release_plan(release, host, "ggml-org/llama.cpp", "latest")
+ assert plan is not None
+ kinds = [a.install_kind for a in plan.attempts]
+ assert (
+ "windows-hip" in kinds
+ ), f"upstream HIP asset not included as fallback; got {kinds}"
+ hip_attempt = next(a for a in plan.attempts if a.install_kind == "windows-hip")
+ assert hip_attempt.source_label == "upstream"
+
+
+# ── Follow-up: pinned-tag URL helper, URL trust pinning, opt-out env, autouse cache clear ──
+
+
+def test_lemonade_release_api_url_pinned_tag():
+ """A pinned llama_tag must produce the /releases/tags/ URL."""
+ assert _mod._lemonade_release_api_for("b1262").endswith("/releases/tags/b1262")
+ assert _mod._lemonade_release_api_for("latest").endswith("/releases/latest")
+ assert _mod._lemonade_release_api_for("").endswith("/releases/latest")
+
+
+def test_lemonade_release_api_url_encodes_tag():
+ """Unexpected slashes / hashes in the tag must be URL-encoded so the URL
+ cannot be reshaped (defence in depth -- tags should already be sanitised
+ upstream)."""
+ url = _mod._lemonade_release_api_for("b1260/../latest")
+ assert "/releases/tags/b1260%2F..%2Flatest" in url
+ assert "//latest" not in url.split("/releases/tags/", 1)[1]
+
+
+def test_lemonade_resolver_skipped_by_opt_out_env(monkeypatch):
+ """UNSLOTH_DISABLE_LEMONADE_ROCM=1 must short-circuit the resolver."""
+ monkeypatch.setenv("UNSLOTH_DISABLE_LEMONADE_ROCM", "1")
+ host = _make_rocm_host("gfx1151")
+ res = resolve_lemonade_rocm_choice(host, "ubuntu", "linux-rocm", llama_tag = "latest")
+ assert res is None
+
+
+def test_lemonade_resolver_rejects_non_github_url(monkeypatch):
+ """If the GitHub API response somehow contained an off-host download URL,
+ the resolver must refuse to use it (lemonade assets are not in the
+ approved-hash manifest)."""
+ bad_release = {
+ "tag_name": _STUB_TAG,
+ "assets": [
+ {
+ "name": f"llama-{_STUB_TAG}-ubuntu-rocm-gfx1151-x64.zip",
+ "browser_download_url": "https://attacker.invalid/llama.zip",
+ },
+ ],
+ }
+ host = _make_rocm_host("gfx1151")
+ with patch.object(_mod, "fetch_json", return_value = bad_release):
+ res = resolve_lemonade_rocm_choice(
+ host, "ubuntu", "linux-rocm", llama_tag = "latest"
+ )
+ assert res is None
+
+
+def test_lemonade_resolver_rejects_http_scheme():
+ assert not _mod._is_trusted_github_release_url(
+ "http://github.com/lemonade-sdk/llamacpp-rocm/releases/download/x/y.zip",
+ "lemonade-sdk/llamacpp-rocm",
+ )
+
+
+def test_lemonade_resolver_accepts_github_cdn():
+ # Real GitHub release CDN URLs carry the /github-production-release-asset- prefix.
+ assert _mod._is_trusted_github_release_url(
+ "https://objects.githubusercontent.com/github-production-release-asset-abc123/456/789?token=x",
+ "lemonade-sdk/llamacpp-rocm",
+ )
+
+
+def test_lemonade_resolver_rejects_arbitrary_cdn_path():
+ # A CDN URL without the release-asset path prefix must be rejected.
+ assert not _mod._is_trusted_github_release_url(
+ "https://objects.githubusercontent.com/abc/def",
+ "lemonade-sdk/llamacpp-rocm",
+ )
+
+
+def test_lemonade_resolver_accepts_release_path():
+ url = "https://github.com/lemonade-sdk/llamacpp-rocm/releases/download/b1262/llama-b1262-ubuntu-rocm-gfx1151-x64.zip"
+ assert _mod._is_trusted_github_release_url(url, "lemonade-sdk/llamacpp-rocm")
+
+
+def test_lemonade_resolver_rejects_wrong_repo():
+ """A github.com release URL for a different repo must be rejected."""
+ assert not _mod._is_trusted_github_release_url(
+ "https://github.com/attacker/llamacpp-rocm/releases/download/x/y.zip",
+ "lemonade-sdk/llamacpp-rocm",
+ )
+
+
+def test_lemonade_resolver_rejects_empty_browser_download_url():
+ """An asset entry with an empty browser_download_url must fall through."""
+ release = {
+ "tag_name": _STUB_TAG,
+ "assets": [
+ {
+ "name": f"llama-{_STUB_TAG}-ubuntu-rocm-gfx1151-x64.zip",
+ "browser_download_url": "",
+ },
+ ],
+ }
+ host = _make_rocm_host("gfx1151")
+ with patch.object(_mod, "fetch_json", return_value = release):
+ res = resolve_lemonade_rocm_choice(
+ host, "ubuntu", "linux-rocm", llama_tag = "latest"
+ )
+ assert res is None
+
+
+def test_lemonade_runtime_patterns_include_hip_runtime():
+ """linux-rocm overlay must use a broad lib glob to catch all bundled .so files.
+
+ Lemonade ZIPs carry transitive deps (libamd_comgr, libLLVM, libclang-cpp,
+ ...) whose names change across ROCm releases. A broad ``lib*.so*`` glob
+ avoids having to enumerate every transitive dependency by name.
+ """
+ from install_llama_prebuilt import runtime_patterns_for_choice, AssetChoice
+
+ choice = AssetChoice(
+ repo = "lemonade-sdk/llamacpp-rocm",
+ tag = "b1262",
+ name = "llama-b1262-ubuntu-rocm-gfx1151-x64.zip",
+ url = "https://github.com/lemonade-sdk/llamacpp-rocm/releases/download/b1262/x.zip",
+ source_label = "lemonade",
+ install_kind = "linux-rocm",
+ )
+ pats = runtime_patterns_for_choice(choice)
+ # The broad glob must be present so every .so in the lemonade bundle
+ # (including transitive deps added in future ROCm releases) gets overlaid.
+ assert "lib*.so*" in pats, f"'lib*.so*' missing from linux-rocm patterns: {pats}"
+
+
+_pick_rocm_gfx_target = getattr(_mod, "_pick_rocm_gfx_target", None)
+
+
+@pytest.mark.skipif(
+ _pick_rocm_gfx_target is None,
+ reason = "_pick_rocm_gfx_target not present on this branch",
+)
+def test_pick_rocm_gfx_target_honors_cuda_visible_devices(monkeypatch):
+ """AMD HIP honours CUDA_VISIBLE_DEVICES identically to HIP_VISIBLE_DEVICES;
+ on a gfx1151 + gfx1100 mixed host, CUDA_VISIBLE_DEVICES=1 must select gfx1100."""
+ # Two GPUs; rocminfo reports each token twice (as in the real tool output).
+ probe_out = "gfx1151\ngfx1151\ngfx1100\ngfx1100"
+ monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
+ monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
+ monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "1")
+ assert _pick_rocm_gfx_target(probe_out) == "gfx1100"
+
+
+@pytest.mark.skipif(
+ _pick_rocm_gfx_target is None,
+ reason = "_pick_rocm_gfx_target not present on this branch",
+)
+def test_pick_rocm_gfx_target_cuda_visible_devices_minus_one_returns_none(monkeypatch):
+ """CUDA_VISIBLE_DEVICES=-1 means no GPU visible; resolver must return None."""
+ probe_out = "gfx1151\ngfx1100"
+ monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising = False)
+ monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
+ monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "-1")
+ assert _pick_rocm_gfx_target(probe_out) is None
+
+
+@pytest.mark.skipif(
+ _pick_rocm_gfx_target is None,
+ reason = "_pick_rocm_gfx_target not present on this branch",
+)
+def test_pick_rocm_gfx_target_same_arch_multi_gpu(monkeypatch):
+ """Regression: [gfx1100, gfx1100, gfx1151] with HIP_VISIBLE_DEVICES=2 must
+ return gfx1151, not fall back to GPU 0 due to dict.fromkeys collapsing the
+ two gfx1100 entries into one and making index 2 out of range."""
+ # Simulate rocminfo output for 3 GPUs (2x gfx1100 dGPU + 1x gfx1151 APU).
+ # Each GPU gets its own Agent section with a few token mentions.
+ probe_out = (
+ "***\nAgent 1\n***\n gfx1100 some info\n gfx1100\n"
+ "***\nAgent 2\n***\n gfx1100 some info\n gfx1100\n"
+ "***\nAgent 3\n***\n gfx1151 some info\n gfx1151\n"
+ )
+ monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
+ monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising = False)
+ monkeypatch.setenv("HIP_VISIBLE_DEVICES", "2")
+ assert _pick_rocm_gfx_target(probe_out) == "gfx1151"
diff --git a/studio/backend/tests/test_log_filter_no_truncation.py b/studio/backend/tests/test_log_filter_no_truncation.py
index d78643f5b9..d9a6e2bc4a 100644
--- a/studio/backend/tests/test_log_filter_no_truncation.py
+++ b/studio/backend/tests/test_log_filter_no_truncation.py
@@ -2,27 +2,11 @@
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""
-Regression tests for studio.backend.loggers.handlers.filter_sensitive_data.
+Regression tests for loggers.handlers.filter_sensitive_data.
-Context: filter_sensitive_data was originally written with a base64-detection
-heuristic that truncated any string >100 chars containing ',' or '/' down to
-20 chars + '...'. The block was dormant until PR #5246 wired the processor
-into the structlog chain to redact native-path leases. Once active, the
-heuristic ate normal log lines emitted by llama_cpp_backend (GGUF size
-summary, mmproj selection, the full llama-server command line) and any
-exception traceback that happened to contain a file path.
-
-These tests pin two properties:
-
-1. Long, comma- or slash-bearing log messages flow through filter_sensitive_data
- unchanged. The exact strings exercised match the call sites at
- studio/backend/core/inference/llama_cpp.py:2117, :2283, and :2312 that
- were truncated in the original bug report.
-
-2. PR #5246's native-path lease redaction still fires for both the inline
- ``native_path_lease=...`` regex form and the ``nativePathLease`` dict-key
- form. This guards against future regressions that strip redaction along
- with the truncation block.
+Pins two properties: (1) long strings with commas/slashes pass through
+unchanged (the base64-truncation heuristic from PR #5246 was too aggressive),
+and (2) native-path lease redaction still fires for both inline and dict-key forms.
"""
from loggers.handlers import filter_sensitive_data
diff --git a/studio/backend/tests/test_rocm_oom_guard.py b/studio/backend/tests/test_rocm_oom_guard.py
new file mode 100644
index 0000000000..2ce9b55789
--- /dev/null
+++ b/studio/backend/tests/test_rocm_oom_guard.py
@@ -0,0 +1,176 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""Unit tests for _rocm_classify_unified_memory (ROCm OOM-guard classifier).
+
+Covers the three classification paths:
+ Path 1 – canonical gcnArchName attribute present.
+ Path 2 – gcnArchName absent, alternate-spelling attribute present.
+ Path 3 – ALL arch attrs absent; falls back to device-name substring match.
+
+Regression for: Strix Halo (gfx1151) misclassified as discrete on AMD SDK /
+Radeon wheels that populate props.name = "Radeon 8060S Graphics" but do NOT
+set any gcnArchName attribute. Without the 8060s/8050s name patterns the
+fallback returned is_unified=False, applying the 0.90 fraction instead of
+0.80 and leaving only ~12.8 GiB OS headroom on a 128 GiB unified-memory pool.
+"""
+
+from __future__ import annotations
+
+from types import SimpleNamespace
+
+import pytest
+
+from core.training.worker import _rocm_classify_unified_memory
+
+
+# ── helpers ──────────────────────────────────────────────────────────────────
+
+
+def _props(**kwargs) -> SimpleNamespace:
+ """Build a fake device-properties object with the given attributes."""
+ return SimpleNamespace(**kwargs)
+
+
+# ── Path 1: canonical gcnArchName ────────────────────────────────────────────
+
+
+class TestCanonicalGcnArchName:
+ """gcnArchName is present and populated."""
+
+ @pytest.mark.parametrize(
+ "arch, expected_unified",
+ [
+ ("gfx1150", True), # Strix Point
+ ("gfx1151", True), # Strix Halo
+ ("gfx1100", False), # Navi 31 (RX 7900 XTX) — discrete
+ ("gfx906", False), # MI50 — discrete server GPU
+ ("gfx1201", False), # RX 9070 XT — discrete
+ ],
+ )
+ def test_canonical_attr(self, arch: str, expected_unified: bool) -> None:
+ props = _props(gcnArchName = arch, name = "irrelevant")
+ gcn, is_unified = _rocm_classify_unified_memory(props)
+ assert gcn == arch
+ assert is_unified is expected_unified
+
+ def test_arch_with_colon_suffix_stripped(self) -> None:
+ """gcnArchName can carry xnack/sramecc suffix; only the base is kept."""
+ props = _props(gcnArchName = "gfx1151:xnack-", name = "irrelevant")
+ gcn, is_unified = _rocm_classify_unified_memory(props)
+ assert gcn == "gfx1151"
+ assert is_unified is True
+
+ def test_canonical_attr_wins_over_name(self) -> None:
+ """Arch attr takes priority; device name should be ignored."""
+ # Discrete arch, but name looks like a unified SKU — arch must win.
+ props = _props(gcnArchName = "gfx1100", name = "Radeon 890M")
+ gcn, is_unified = _rocm_classify_unified_memory(props)
+ assert gcn == "gfx1100"
+ assert is_unified is False
+
+
+# ── Path 2: alternate-spelling fallback ──────────────────────────────────────
+
+
+class TestAlternateSpellingFallback:
+ """gcnArchName is missing but an alternate attr spelling is present."""
+
+ @pytest.mark.parametrize(
+ "attr_name",
+ ["gcn_arch_name", "arch_name", "gfx_arch_name"],
+ )
+ def test_alternate_attr_unified(self, attr_name: str) -> None:
+ props = _props(**{attr_name: "gfx1151"}, name = "Radeon 8060S Graphics")
+ gcn, is_unified = _rocm_classify_unified_memory(props)
+ assert gcn == "gfx1151"
+ assert is_unified is True
+
+ @pytest.mark.parametrize(
+ "attr_name",
+ ["gcn_arch_name", "arch_name", "gfx_arch_name"],
+ )
+ def test_alternate_attr_discrete(self, attr_name: str) -> None:
+ props = _props(**{attr_name: "gfx1201"}, name = "Radeon RX 9070 XT")
+ gcn, is_unified = _rocm_classify_unified_memory(props)
+ assert gcn == "gfx1201"
+ assert is_unified is False
+
+ def test_first_non_empty_attr_wins(self) -> None:
+ """When multiple alternate attrs are present the first non-empty one wins."""
+ props = _props(gcn_arch_name = "gfx1151", arch_name = "gfx1100", name = "irrelevant")
+ gcn, is_unified = _rocm_classify_unified_memory(props)
+ assert gcn == "gfx1151"
+ assert is_unified is True
+
+
+# ── Path 3: device-name fallback ─────────────────────────────────────────────
+
+
+class TestDeviceNameFallback:
+ """ALL arch attrs absent — classifier must rely solely on device name."""
+
+ # --- unified-memory devices that MUST be detected ---
+
+ @pytest.mark.parametrize(
+ "device_name",
+ [
+ # gfx1150 Strix Point
+ "Radeon 890M",
+ "AMD Radeon 890M Graphics",
+ "RADEON 890M", # case-insensitive
+ "Radeon 880M",
+ "AMD Radeon 880M Graphics",
+ # gfx1151 Strix Halo — the regression case from the review
+ "Radeon 8060S Graphics", # Ryzen AI MAX+ 395 (as returned by torch)
+ "AMD Radeon 8060S",
+ "Radeon 8050S Graphics", # cut-down Strix Halo SKU
+ "AMD Radeon 8050S",
+ # case variants
+ "RADEON 8060S GRAPHICS",
+ "radeon 8050s",
+ ],
+ )
+ def test_unified_memory_detected(self, device_name: str) -> None:
+ props = _props(name = device_name)
+ gcn, is_unified = _rocm_classify_unified_memory(props)
+ assert gcn == "", f"expected empty gcn_arch, got {gcn!r}"
+ assert (
+ is_unified is True
+ ), f"device {device_name!r} should be classified as unified-memory"
+
+ # --- discrete devices that must NOT be mis-classified ---
+
+ @pytest.mark.parametrize(
+ "device_name",
+ [
+ "Radeon RX 9070 XT",
+ "AMD Radeon RX 7900 XTX",
+ "Radeon RX 6900 XT",
+ "Radeon Pro W7900",
+ "AMD Instinct MI300X",
+ # Names that contain superficially similar substrings but are discrete
+ "Radeon RX 580",
+ "Radeon VII",
+ ],
+ )
+ def test_discrete_not_misclassified(self, device_name: str) -> None:
+ props = _props(name = device_name)
+ gcn, is_unified = _rocm_classify_unified_memory(props)
+ assert gcn == ""
+ assert (
+ is_unified is False
+ ), f"discrete device {device_name!r} should NOT be classified as unified-memory"
+
+ def test_empty_name_returns_false(self) -> None:
+ """Completely absent name must not crash and must default to discrete."""
+ props = _props() # no 'name' attr at all
+ gcn, is_unified = _rocm_classify_unified_memory(props)
+ assert gcn == ""
+ assert is_unified is False
+
+ def test_none_name_returns_false(self) -> None:
+ props = _props(name = None)
+ gcn, is_unified = _rocm_classify_unified_memory(props)
+ assert gcn == ""
+ assert is_unified is False
diff --git a/studio/backend/utils/hardware/amd.py b/studio/backend/utils/hardware/amd.py
index fdb1ab4520..48d5890399 100644
--- a/studio/backend/utils/hardware/amd.py
+++ b/studio/backend/utils/hardware/amd.py
@@ -11,18 +11,35 @@ nvidia.py counterparts.
import json
import math
import os
+import platform
import re
import subprocess
+import sys
from typing import Any, Optional
from loggers import get_logger
from utils.native_path_leases import child_env_without_native_path_secret
+from utils.subprocess_compat import windows_hidden_subprocess_kwargs
logger = get_logger(__name__)
+# amd-smi on Windows must initialise the full ROCm runtime on first call, which
+# can take 15-25 s on cold hardware. Linux is consistently < 2 s.
+_AMD_SMI_DEFAULT_TIMEOUT = 30 if platform.system() == "Windows" else 10
-def _run_amd_smi(*args: str, timeout: int = 5) -> Optional[Any]:
+# Circuit breaker: stop calling amd-smi after this many consecutive failures.
+# On Windows, each failed call spawns a process that may show a UAC/DiskPart
+# elevation prompt. Once we know amd-smi doesn't work we stop polling it.
+_AMD_SMI_FAILURE_LIMIT = 3
+_amd_smi_consecutive_failures = 0
+_amd_smi_disabled = False
+
+
+def _run_amd_smi(*args: str, timeout: int = _AMD_SMI_DEFAULT_TIMEOUT) -> Optional[Any]:
"""Run amd-smi with the given arguments and return parsed JSON, or None."""
+ global _amd_smi_consecutive_failures, _amd_smi_disabled
+ if _amd_smi_disabled:
+ return None
try:
result = subprocess.run(
["amd-smi", *args, "--json"],
@@ -30,13 +47,40 @@ def _run_amd_smi(*args: str, timeout: int = 5) -> Optional[Any]:
text = True,
timeout = timeout,
env = child_env_without_native_path_secret(),
+ **windows_hidden_subprocess_kwargs(),
)
except (OSError, subprocess.TimeoutExpired) as e:
- logger.warning("amd-smi query failed: %s", e)
+ if isinstance(e, FileNotFoundError):
+ # amd-smi ships with Adrenalin, not the HIP SDK -- absence is
+ # expected on HIP SDK-only Windows setups. Log at debug only.
+ logger.debug("amd-smi not found (not in PATH): %s", e)
+ else:
+ logger.warning("amd-smi query failed: %s", e)
+ _amd_smi_consecutive_failures += 1
+ if _amd_smi_consecutive_failures >= _AMD_SMI_FAILURE_LIMIT:
+ logger.info(
+ "amd-smi not available (not installed; expected on HIP SDK-only systems); "
+ "GPU VRAM polling disabled"
+ )
+ _amd_smi_disabled = True
return None
- if result.returncode != 0 or not result.stdout.strip():
+ if result.returncode != 0:
logger.warning("amd-smi returned code %d", result.returncode)
+ _amd_smi_consecutive_failures += 1
+ if _amd_smi_consecutive_failures >= _AMD_SMI_FAILURE_LIMIT:
+ logger.info(
+ "amd-smi not available (not installed; expected on HIP SDK-only systems); "
+ "GPU VRAM polling disabled"
+ )
+ _amd_smi_disabled = True
return None
+ if not result.stdout.strip():
+ # amd-smi exited successfully but produced no output (e.g. no GPUs
+ # visible on this query, or a version that emits nothing for --json).
+ # This is not a tool failure, so don't count against the circuit breaker.
+ logger.debug("amd-smi exited 0 but returned no output")
+ return None
+ _amd_smi_consecutive_failures = 0 # reset on success
try:
return json.loads(result.stdout)
except json.JSONDecodeError:
@@ -352,7 +396,7 @@ def get_visible_gpu_utilization(
)
parsed_id = _parse_numeric(raw_id)
if parsed_id is None:
- logger.debug(
+ logger.warning(
"amd-smi GPU id %r could not be parsed; falling back to "
"enumeration index %d",
raw_id,
@@ -360,7 +404,15 @@ def get_visible_gpu_utilization(
)
idx = fallback_idx
else:
- idx = int(parsed_id)
+ rounded = round(parsed_id)
+ if rounded != parsed_id:
+ logger.warning(
+ "amd-smi GPU id %r parsed as non-integer %r; truncating to %d",
+ raw_id,
+ parsed_id,
+ rounded,
+ )
+ idx = int(rounded)
if idx not in visible_set:
continue
metrics = _extract_gpu_metrics(gpu_data)
diff --git a/studio/backend/utils/hardware/hardware.py b/studio/backend/utils/hardware/hardware.py
index ede37e2953..180fde8f13 100644
--- a/studio/backend/utils/hardware/hardware.py
+++ b/studio/backend/utils/hardware/hardware.py
@@ -16,8 +16,16 @@ Usage:
...
"""
+import copy
+import gc
+import glob
import os
import platform
+import re
+import subprocess
+import sys
+import types
+from importlib.metadata import PackageNotFoundError, version as pkg_version
import structlog
from loggers import get_logger
from enum import Enum
@@ -120,11 +128,13 @@ def detect_hardware() -> DeviceType:
# Distinguish AMD ROCm (HIP) from NVIDIA CUDA for display purposes.
# DeviceType stays CUDA since torch.cuda.* works on ROCm via HIP.
- if getattr(torch.version, "hip", None) is not None:
+ # AMD's repo.radeon.com SDK wheels (e.g. 2.9.0+rocmsdk20251116) do
+ # not set torch.version.hip, so fall back to checking __version__.
+ _hip_ver = getattr(torch.version, "hip", None)
+ if _hip_ver is not None or "rocm" in torch.__version__.lower():
IS_ROCM = True
- print(
- f"Hardware detected: ROCm (HIP {torch.version.hip}) -- {device_name}"
- )
+ _hip_label = _hip_ver or torch.__version__
+ print(f"Hardware detected: ROCm (HIP {_hip_label}) -- {device_name}")
else:
print(f"Hardware detected: CUDA -- {device_name}")
return DEVICE
@@ -176,8 +186,6 @@ def clear_gpu_cache():
Clear GPU memory cache for the current device.
Safe to call on any platform — no-ops gracefully.
"""
- import gc
-
gc.collect()
device = get_device()
@@ -359,8 +367,6 @@ def get_package_versions() -> Dict[str, Optional[str]]:
Returns dict with keys: unsloth, torch, transformers, cuda.
Missing packages yield None.
"""
- from importlib.metadata import version as pkg_version, PackageNotFoundError
-
packages = ("unsloth", "torch", "transformers")
versions: Dict[str, Optional[str]] = {}
@@ -466,7 +472,7 @@ def _smi_query(func_name: str, *args, **kwargs) -> Optional[Dict[str, Any]]:
try:
func = getattr(_backend, func_name)
result = func(*args, **kwargs)
- if result.get("available"):
+ if isinstance(result, dict) and result.get("available"):
return result
except Exception as e:
logger.warning("%s %s query failed: %s", backend_name, func_name, e)
@@ -479,9 +485,6 @@ def _read_apple_gpu_stats() -> Dict[str, Any]:
Returns dict with utilization_pct, vram_used_bytes (system-wide GPU memory).
Returns empty dict on failure.
"""
- import subprocess
- import re
-
try:
result = subprocess.run(
["ioreg", "-r", "-c", "AGXAccelerator"],
@@ -506,6 +509,133 @@ def _read_apple_gpu_stats() -> Dict[str, Any]:
}
+def _rocm_linux_sysfs_gpu_busy_pct() -> Optional[float]:
+ """Query AMD GPU compute utilization via Linux DRM sysfs gpu_busy_percent."""
+ if platform.system() != "Linux":
+ return None
+ try:
+ files = glob.glob("/sys/class/drm/card*/device/gpu_busy_percent")
+ if not files:
+ return None
+ values = [int(open(f).read().strip()) for f in files]
+ return round(sum(values) / len(values), 1)
+ except Exception:
+ return None
+
+
+def _rocm_linux_sysfs_temp_c() -> Optional[float]:
+ """Query AMD GPU edge temperature via Linux DRM hwmon sysfs (temp1_input, millidegrees C)."""
+ if platform.system() != "Linux":
+ return None
+ try:
+ files = glob.glob("/sys/class/drm/card*/device/hwmon/hwmon*/temp1_input")
+ if not files:
+ return None
+ temps = [int(open(f).read().strip()) / 1000.0 for f in files]
+ return round(max(temps), 1)
+ except Exception:
+ return None
+
+
+def _rocm_linux_sysfs_power_w() -> Optional[float]:
+ """Query AMD GPU average power draw via Linux DRM hwmon sysfs (microwatts)."""
+ if platform.system() != "Linux":
+ return None
+ try:
+ for pattern in (
+ "/sys/class/drm/card*/device/hwmon/hwmon*/power1_average",
+ "/sys/class/drm/card*/device/hwmon/hwmon*/power1_input",
+ ):
+ files = glob.glob(pattern)
+ if files:
+ watts = sum(int(open(f).read().strip()) / 1_000_000.0 for f in files)
+ return round(watts, 1)
+ return None
+ except Exception:
+ return None
+
+
+def _rocm_windows_perf_counter_gpu_util_pct() -> Optional[float]:
+ """Query AMD GPU compute utilization via Windows Performance Counters (3D engine nodes)."""
+ if platform.system() != "Windows":
+ return None
+ try:
+ ps = (
+ "$s=(Get-Counter '\\GPU Engine(*engtype_3D*)\\Utilization Percentage'"
+ " -ErrorAction SilentlyContinue).CounterSamples;"
+ "if($s){[math]::Min(($s|Measure-Object CookedValue -Sum).Sum,100)}else{-1}"
+ )
+ r = subprocess.run(
+ ["powershell", "-NoProfile", "-NonInteractive", "-Command", ps],
+ capture_output = True,
+ text = True,
+ timeout = 5,
+ )
+ if r.returncode != 0 or not r.stdout.strip():
+ return None
+ val = float(r.stdout.strip())
+ return round(val, 1) if val >= 0 else None
+ except Exception:
+ return None
+
+
+def _rocm_linux_sysfs_vram_gb() -> tuple[Optional[float], Optional[float]]:
+ """Query system-wide AMD GPU VRAM via Linux DRM sysfs.
+
+ Reads /sys/class/drm/card*/device/mem_info_vram_* which the kernel
+ updates in real-time across all processes. No tools required.
+ Returns (used_gb, total_gb) or (None, None) on failure.
+ """
+ if platform.system() != "Linux":
+ return None, None
+ try:
+ used_files = glob.glob("/sys/class/drm/card*/device/mem_info_vram_used")
+ total_files = glob.glob("/sys/class/drm/card*/device/mem_info_vram_total")
+ if not used_files or not total_files:
+ return None, None
+ used_bytes = sum(int(open(f).read().strip()) for f in used_files)
+ total_bytes = sum(int(open(f).read().strip()) for f in total_files)
+ if total_bytes == 0:
+ return None, None
+ return round(used_bytes / (1024**3), 2), round(total_bytes / (1024**3), 2)
+ except Exception:
+ return None, None
+
+
+def _rocm_windows_perf_counter_vram_gb() -> tuple[Optional[float], Optional[float]]:
+ """Query system-wide dedicated GPU VRAM via Windows Performance Counters.
+
+ Uses the same data source as Task Manager so it reflects cross-process
+ usage accurately. Works for any GPU vendor without amd-smi or nvidia-smi.
+ Returns (used_gb, total_gb) or (None, None) on failure.
+ """
+ if platform.system() != "Windows":
+ return None, None
+ try:
+ ps = (
+ "$s=(Get-Counter '\\GPU Adapter Memory(*)\\Dedicated Usage'"
+ " -ErrorAction SilentlyContinue).CounterSamples;"
+ "if($s){($s|Measure-Object CookedValue -Sum).Sum}else{-1}"
+ )
+ r = subprocess.run(
+ ["powershell", "-NoProfile", "-NonInteractive", "-Command", ps],
+ capture_output = True,
+ text = True,
+ timeout = 5,
+ )
+ if r.returncode != 0 or not r.stdout.strip():
+ return None, None
+ used_bytes = float(r.stdout.strip())
+ if used_bytes < 0:
+ return None, None
+ import torch as _torch
+
+ total_bytes = _torch.cuda.get_device_properties(0).total_memory
+ return round(used_bytes / (1024**3), 2), round(total_bytes / (1024**3), 2)
+ except Exception:
+ return None, None
+
+
def get_gpu_utilization() -> Dict[str, Any]:
"""Return a live snapshot of device utilization information."""
device = get_device()
@@ -514,7 +644,78 @@ def get_gpu_utilization() -> Dict[str, Any]:
result = _smi_query("get_primary_gpu_utilization")
if result is not None:
result["backend"] = _backend_label(device)
+ if IS_ROCM:
+ # Fix unified-memory VRAM on AMD iGPUs (Strix Halo etc.)
+ _reconcile_primary_rocm_unified_memory(
+ result, _get_parent_visible_gpu_spec()
+ )
return result
+ # SMI tool unavailable or returned no usable data. On Windows, query
+ # the Performance Counter API (same source as Task Manager) for
+ # system-wide dedicated VRAM — covers cross-process usage that
+ # torch.cuda.mem_get_info cannot see from the Studio server process.
+ if IS_ROCM and platform.system() == "Windows":
+ _win_used, _win_total = _rocm_windows_perf_counter_vram_gb()
+ if _win_used is not None and _win_total is not None:
+ _win_util = _rocm_windows_perf_counter_gpu_util_pct()
+ return {
+ "available": True,
+ "backend": _backend_label(device),
+ "gpu_utilization_pct": _win_util,
+ "temperature_c": None,
+ "vram_used_gb": _win_used,
+ "vram_total_gb": _win_total,
+ "vram_utilization_pct": round((_win_used / _win_total) * 100, 1)
+ if _win_total > 0
+ else None,
+ "power_draw_w": None,
+ "power_limit_w": None,
+ "power_utilization_pct": None,
+ }
+ # Linux: DRM sysfs gives system-wide VRAM across all processes, no tools needed.
+ if IS_ROCM and platform.system() == "Linux":
+ _linux_used, _linux_total = _rocm_linux_sysfs_vram_gb()
+ if _linux_used is not None and _linux_total is not None:
+ _linux_util = _rocm_linux_sysfs_gpu_busy_pct()
+ _linux_temp = _rocm_linux_sysfs_temp_c()
+ _linux_power = _rocm_linux_sysfs_power_w()
+ return {
+ "available": True,
+ "backend": _backend_label(device),
+ "gpu_utilization_pct": _linux_util,
+ "temperature_c": _linux_temp,
+ "vram_used_gb": _linux_used,
+ "vram_total_gb": _linux_total,
+ "vram_utilization_pct": round((_linux_used / _linux_total) * 100, 1)
+ if _linux_total > 0
+ else None,
+ "power_draw_w": _linux_power,
+ "power_limit_w": None,
+ "power_utilization_pct": None,
+ }
+ # Last resort: torch mem_get_info (process-local).
+ _visible_spec = _get_parent_visible_gpu_spec()
+ _numeric_ids = _visible_spec.get("numeric_ids") or [0]
+ _primary_idx = [_numeric_ids[0]] if _numeric_ids else [0]
+ _torch_devices = _torch_get_per_device_info(_primary_idx)
+ if _torch_devices:
+ _td = _torch_devices[0]
+ _total = _td["total_gb"]
+ _used = _td["used_gb"]
+ return {
+ "available": True,
+ "backend": _backend_label(device),
+ "gpu_utilization_pct": None,
+ "temperature_c": None,
+ "vram_used_gb": _used,
+ "vram_total_gb": _total,
+ "vram_utilization_pct": round((_used / _total) * 100, 1)
+ if _total > 0
+ else None,
+ "power_draw_w": None,
+ "power_limit_w": None,
+ "power_utilization_pct": None,
+ }
# MLX path: single _read_apple_gpu_stats() call carries both VRAM-used
# bytes and GPU utilization %. psutil for unified-memory total is cheap.
@@ -578,6 +779,77 @@ def get_gpu_utilization() -> Dict[str, Any]:
return {"available": False, "backend": _backend_label(device)}
+def _apply_unified_memory_correction(
+ device_metrics: Dict[str, Any], torch_info: Dict[str, Any]
+) -> None:
+ """Per-device reconciliation: when torch reports a larger memory total
+ than amd-smi, overwrite the smi VRAM fields in place.
+
+ Used by both the multi-device and primary-device reconciliation helpers
+ so the two endpoints stay in sync on AMD iGPUs with unified memory.
+ """
+ torch_total_gb = torch_info["total_gb"]
+ smi_total_gb = device_metrics.get("vram_total_gb") or 0.0
+ if torch_total_gb > smi_total_gb:
+ torch_used_gb = torch_info["used_gb"]
+ device_metrics["vram_total_gb"] = torch_total_gb
+ device_metrics["vram_used_gb"] = torch_used_gb
+ device_metrics["vram_utilization_pct"] = (
+ round((torch_used_gb / torch_total_gb) * 100, 1)
+ if torch_total_gb > 0
+ else None
+ )
+ logger.debug(
+ "ROCm unified memory: replaced amd-smi VRAM (%.2f GB) with "
+ "torch mem_get_info total (%.2f GB) for device %s",
+ smi_total_gb,
+ torch_total_gb,
+ torch_info.get("index"),
+ )
+
+
+def _reconcile_rocm_unified_memory(
+ utilization: Dict[str, Any], device_indices: list[int]
+) -> None:
+ """Fix amd-smi VRAM for ROCm unified-memory GPUs (e.g. Strix Halo).
+
+ amd-smi reports only the dedicated slice (~512 MB); torch sees the full
+ GTT pool (~128 GB). When torch total > smi total, overwrite per-device
+ VRAM fields so GPU selection uses the real available memory.
+ """
+ torch_devices = _torch_get_per_device_info(device_indices)
+ if not torch_devices:
+ return
+ torch_by_index = {td["index"]: td for td in torch_devices}
+ for dev in utilization.get("devices", []):
+ td = torch_by_index.get(dev.get("index"))
+ if td is None:
+ continue
+ _apply_unified_memory_correction(dev, td)
+
+
+def _reconcile_primary_rocm_unified_memory(
+ utilization: Dict[str, Any], parent_visible_spec: Dict[str, Any]
+) -> None:
+ """Same fix as _reconcile_rocm_unified_memory for the flat primary-GPU dict."""
+ numeric_ids = parent_visible_spec.get("numeric_ids")
+ if numeric_ids is None:
+ # No visibility env var set: torch ordinal 0 is the primary device.
+ primary_idx = [0]
+ elif len(numeric_ids) == 0:
+ # Empty mask (HIP_VISIBLE_DEVICES="" or "-1"): no GPU is visible to
+ # this process. Querying torch device 0 would raise a RuntimeError or
+ # return stale/wrong data, so bail out rather than writing bad values
+ # into the utilization dict.
+ return
+ else:
+ primary_idx = [int(numeric_ids[0])]
+ torch_devices = _torch_get_per_device_info(primary_idx)
+ if not torch_devices:
+ return
+ _apply_unified_memory_correction(utilization, torch_devices[0])
+
+
def get_visible_gpu_utilization() -> Dict[str, Any]:
device = get_device()
@@ -590,6 +862,10 @@ def get_visible_gpu_utilization() -> Dict[str, Any]:
)
if result is not None:
result["backend"] = _backend_label(device)
+ numeric_ids = parent_visible_spec.get("numeric_ids")
+ if IS_ROCM and numeric_ids is not None:
+ # Fix unified-memory VRAM on AMD iGPUs (Strix Halo etc.)
+ _reconcile_rocm_unified_memory(result, numeric_ids)
return result
# Torch-based fallback for CUDA (nvidia-smi unavailable, AMD ROCm) and XPU (Intel)
@@ -689,7 +965,15 @@ def _get_parent_visible_gpu_spec() -> Dict[str, Any]:
# Use explicit None checks (not `or`) so empty string "" is honoured
# as "no visible GPUs" rather than falling through to CUDA_VISIBLE_DEVICES.
cuda_visible = None
- if IS_ROCM:
+ # Prefer ROCm masks only on a ROCm host, or when no CUDA mask is set, so a
+ # stale HIP_VISIBLE_DEVICES on an NVIDIA host can't override CUDA_VISIBLE_DEVICES.
+ _is_rocm_spec = IS_ROCM or (
+ "CUDA_VISIBLE_DEVICES" not in os.environ
+ and (
+ "HIP_VISIBLE_DEVICES" in os.environ or "ROCR_VISIBLE_DEVICES" in os.environ
+ )
+ )
+ if _is_rocm_spec:
hip_vis = os.environ.get("HIP_VISIBLE_DEVICES")
rocr_vis = os.environ.get("ROCR_VISIBLE_DEVICES")
if hip_vis is not None:
@@ -865,7 +1149,57 @@ def _load_config_for_gpu_estimate(model_name: str, hf_token: Optional[str] = Non
def _determine_attention_impl_for_gpu_estimate(config) -> str:
- import copy as _copy
+ # torch.distributed is incomplete on Windows ROCm — torch._C is a C
+ # extension (not a package), so Python cannot import the submodule
+ # torch._C._distributed_c10d that torch.distributed depends on.
+ # Inject an empty stub into sys.modules BEFORE importing torch.distributed
+ # so the import succeeds, then patch the missing process-group helpers.
+ if sys.platform == "win32" and IS_ROCM:
+ # Dummy class for any name torch.distributed tries to import from these stubs
+ class _Dummy:
+ pass
+
+ for _c10d_name in (
+ "torch._C._distributed_c10d",
+ "torch._C._distributed_autograd",
+ "torch._C._distributed_rpc",
+ ):
+ if _c10d_name not in sys.modules:
+ _stub = types.ModuleType(_c10d_name)
+ # torch.distributed imports these names from _distributed_c10d;
+ # provide no-op dummies so the import doesn't raise AttributeError.
+ for _sym in (
+ "FakeProcessGroup",
+ "ProcessGroup",
+ "Work",
+ "Store",
+ "PrefixStore",
+ "FileStore",
+ "TCPStore",
+ "HashStore",
+ "Reducer",
+ "Logger",
+ "DistributedDebugLevel",
+ "GradBucket",
+ "BuiltinCommHookType",
+ ):
+ setattr(_stub, _sym, _Dummy)
+ sys.modules[_c10d_name] = _stub
+
+ try:
+ import torch.distributed as _td
+
+ for _attr, _stub in (
+ ("is_initialized", lambda: False),
+ ("is_available", lambda: False),
+ ("get_rank", lambda: 0),
+ ("get_world_size", lambda: 1),
+ ("is_torchelastic_launched", lambda: False),
+ ):
+ if not hasattr(_td, _attr):
+ setattr(_td, _attr, _stub)
+ except ImportError:
+ pass
from unsloth.models._utils import resolve_attention_implementation
from transformers import AutoModel, AutoModelForCausalLM
@@ -875,7 +1209,7 @@ def _determine_attention_impl_for_gpu_estimate(config) -> str:
# `sub_configs` and propagates to nested text_config / sub-configs, so a
# shallow copy still mutates those shared inner objects on the cached
# config returned by _load_config_for_gpu_estimate. Deepcopy isolates them.
- config_copy = _copy.deepcopy(config)
+ config_copy = copy.deepcopy(config)
model_class = None
for auto_model in (AutoModelForCausalLM, AutoModel):
@@ -1062,7 +1396,10 @@ def estimate_required_model_memory_gb(
_determine_attention_impl_for_gpu_estimate(config)
)
except Exception as e:
- logger.warning(
+ # Log at debug: on Windows ROCm the torch.distributed stub does
+ # not implement Store, so this fires on every estimate call.
+ # It is expected and non-actionable -- eager is the safe fallback.
+ logger.debug(
"Could not resolve attention implementation for '%s': %s",
estimate_model,
e,
@@ -1552,14 +1889,35 @@ def apply_gpu_ids(gpu_ids) -> None:
# parent process already set a ROCm visibility variable -- that
# way a downstream ROCm process inherits the narrowed mask even
# before Studio's hardware detection has classified the host.
+ # Final fallback: probe torch.version.hip so AMD workers without
+ # HIP_VISIBLE_DEVICES still get the correct ROCm visibility mask.
_inherits_rocm_visibility = (
"HIP_VISIBLE_DEVICES" in os.environ or "ROCR_VISIBLE_DEVICES" in os.environ
)
- if IS_ROCM or _inherits_rocm_visibility:
+ _is_rocm = IS_ROCM or _inherits_rocm_visibility
+ if not _is_rocm:
+ # torch.version.hip is a non-empty string on ROCm, None on CUDA.
+ # AMD SDK / Radeon ROCm wheels can leave torch.version.hip unset but
+ # still encode "rocm" in torch.__version__, matching detect_hardware().
+ # Broad except: a probe failure must never crash a training worker.
+ try:
+ import torch as _torch
+
+ _is_rocm = (
+ getattr(_torch.version, "hip", None) is not None
+ or "rocm" in getattr(_torch, "__version__", "").lower()
+ )
+ except Exception as e:
+ logger.debug(
+ "apply_gpu_ids: torch ROCm probe skipped (%s: %s)",
+ type(e).__name__,
+ e,
+ )
+ if _is_rocm:
os.environ["HIP_VISIBLE_DEVICES"] = value
os.environ["ROCR_VISIBLE_DEVICES"] = value
_visible_gpu_count = None
- if IS_ROCM or _inherits_rocm_visibility:
+ if _is_rocm:
logger.info("Applied gpu_ids: CUDA_VISIBLE_DEVICES='%s' (rocm)", value)
else:
logger.info("Applied gpu_ids: CUDA_VISIBLE_DEVICES='%s'", value)
@@ -1652,8 +2010,6 @@ def safe_num_proc(desired: Optional[int] = None) -> int:
Returns:
A safe integer ≥ 1.
"""
- import sys
-
# Windows and macOS use 'spawn' for multiprocessing -- the overhead of
# re-importing torch/transformers/unsloth per worker is typically slower
# than single-process.
@@ -1704,8 +2060,6 @@ def dataset_map_num_proc(desired: Optional[int] = None) -> Optional[int]:
``datasets`` treats ``num_proc=1`` as multiprocessing (creates ``Pool(1)``).
Only ``num_proc=None`` guarantees in-process execution.
"""
- import sys
-
if sys.platform in ("win32", "darwin"):
return None
return safe_num_proc(desired)
diff --git a/studio/backend/utils/models/model_config.py b/studio/backend/utils/models/model_config.py
index 993995ee57..dc34444ccb 100644
--- a/studio/backend/utils/models/model_config.py
+++ b/studio/backend/utils/models/model_config.py
@@ -1156,13 +1156,18 @@ def detect_gguf_model(path: str) -> Optional[str]:
p = Path(path)
# Case 1: direct .gguf file
- if p.suffix.lower() == ".gguf" and p.is_file():
+ if p.suffix.lower() == ".gguf":
if _is_mmproj(p.name):
return None
- # Use absolute (not resolve) to preserve symlink names -- e.g.
- # Ollama .studio_links/model.gguf -> blobs/sha256-... should
- # keep the readable symlink name, not the opaque blob hash.
- return str(p.absolute())
+ # Extension is authoritative: don't gate on is_file()/exists(), which
+ # can fail in the Windows lock window after llama-server is killed.
+ try:
+ is_dir = p.is_dir()
+ except OSError:
+ is_dir = False # stat() unavailable in the lock window
+ if not is_dir:
+ return str(p.absolute()) # absolute() keeps symlink names readable
+ # Directory named "*.gguf": fall through to the dir scan below.
# Case 2: directory containing .gguf files (skip mmproj)
if p.is_dir():
diff --git a/studio/backend/utils/wheel_utils.py b/studio/backend/utils/wheel_utils.py
index 5c42e890d1..e0ce02261b 100644
--- a/studio/backend/utils/wheel_utils.py
+++ b/studio/backend/utils/wheel_utils.py
@@ -15,6 +15,7 @@ import urllib.request
from typing import Callable
from utils.native_path_leases import child_env_without_native_path_secret
+from utils.subprocess_compat import windows_hidden_subprocess_kwargs
_logger = logging.getLogger(__name__)
@@ -106,6 +107,7 @@ def probe_torch_wheel_env(*, timeout: int | None = None) -> dict[str, str] | Non
text = True,
timeout = timeout,
env = child_env_without_native_path_secret(),
+ **windows_hidden_subprocess_kwargs(),
)
except subprocess.TimeoutExpired:
return None
diff --git a/studio/frontend/src/components/app-sidebar.tsx b/studio/frontend/src/components/app-sidebar.tsx
index 849e017ea8..4e4a6130cd 100644
--- a/studio/frontend/src/components/app-sidebar.tsx
+++ b/studio/frontend/src/components/app-sidebar.tsx
@@ -49,7 +49,7 @@ import {
Edit03Icon,
Globe02Icon,
HelpCircleIcon,
- Logout01Icon,
+ Logout05Icon,
Search01Icon,
PowerIcon,
PencilEdit02Icon,
@@ -796,7 +796,7 @@ export function AppSidebar() {
void navigate({ to: "/login" });
}}
>
-
+
{t("shell.navigation.logOut")}
setShutdownOpen(true)}>
diff --git a/studio/install_llama_prebuilt.py b/studio/install_llama_prebuilt.py
index b9d63ecee9..38fa31fff8 100644
--- a/studio/install_llama_prebuilt.py
+++ b/studio/install_llama_prebuilt.py
@@ -9,6 +9,7 @@ from __future__ import annotations
import argparse
import errno
import fnmatch
+import functools
import hashlib
import json
import os
@@ -18,6 +19,7 @@ import re
import shutil
import site
import socket
+import struct
import subprocess
import sys
import tarfile
@@ -29,7 +31,7 @@ import urllib.parse
import urllib.request
import zipfile
from contextlib import contextmanager
-from dataclasses import dataclass, field
+from dataclasses import dataclass, field, replace as dataclasses_replace
try:
from filelock import FileLock, Timeout as FileLockTimeout
@@ -100,6 +102,39 @@ DEFAULT_PUBLISHED_SHA256_ASSET = os.environ.get(
)
UPSTREAM_REPO = "ggml-org/llama.cpp"
UPSTREAM_RELEASES_API = f"https://api.github.com/repos/{UPSTREAM_REPO}/releases/latest"
+
+LEMONADE_ROCM_REPO = "lemonade-sdk/llamacpp-rocm"
+LEMONADE_ROCM_RELEASES_API = (
+ f"https://api.github.com/repos/{LEMONADE_ROCM_REPO}/releases/latest"
+)
+
+
+def _lemonade_release_api_for(llama_tag: str) -> str:
+ """Return the GitHub API URL for the lemonade release that matches a
+ requested llama.cpp tag.
+
+ When llama_tag is unset or "latest", point at /releases/latest. When the
+ caller has pinned a specific tag (e.g. "b1260"), point at the same tag in
+ lemonade. Lemonade tracks `ggml-org/llama.cpp` build tags (e.g. "b1260")
+ but is NOT guaranteed to publish every upstream build -- lemonade may be
+ several builds behind ggml-org. Pinning to a specific tag that lemonade
+ skipped will produce a 404 and the caller falls through to the upstream
+ tarball; that is intentional so pinned installs stay reproducible.
+ Do NOT pass a `unslothai/llama.cpp` fork tag -- the fork uses its own
+ namespace and will always 404 against lemonade.
+
+ The tag is URL-encoded with `safe=""` so an unexpected slash / hash / query
+ character cannot reshape the URL.
+ """
+ normalized = (llama_tag or "").strip()
+ if not normalized or normalized.lower() == "latest":
+ return LEMONADE_ROCM_RELEASES_API
+ return (
+ f"https://api.github.com/repos/{LEMONADE_ROCM_REPO}/releases/tags/"
+ f"{urllib.parse.quote(normalized, safe = '')}"
+ )
+
+
TEST_MODEL_URL = (
"https://huggingface.co/ggml-org/models/resolve/main/tinyllamas/stories260K.gguf"
)
@@ -126,6 +161,14 @@ DEFAULT_MAX_PREBUILT_RELEASE_FALLBACKS = env_int(
2,
minimum = 1,
)
+# Deeper macOS-only walk-back: upstream can ship a run of prebuilts built for a
+# newer macOS than the host, only caught at validate time, so an older host must
+# skip the whole run. Free on new hosts (first plan validates, extras unused).
+DEFAULT_MAX_MACOS_RELEASE_FALLBACKS = env_int(
+ "UNSLOTH_LLAMA_MAX_MACOS_RELEASE_FALLBACKS",
+ 16,
+ minimum = 1,
+)
FORCE_COMPILE_DEFAULT_REF = os.environ.get("UNSLOTH_LLAMA_FORCE_COMPILE_REF", "master")
DIRECT_LINUX_BUNDLE_PROFILES: dict[str, dict[str, Any]] = {
@@ -196,6 +239,10 @@ class HostInfo:
has_physical_nvidia: bool
has_usable_nvidia: bool
has_rocm: bool = False
+ rocm_gfx_target: str | None = None
+ # (major, minor) from platform.mac_ver(); None off macOS or if unparseable.
+ # Skips a macos prebuilt whose minimum-OS exceeds this host.
+ macos_version: tuple[int, int] | None = None
@dataclass
@@ -1268,9 +1315,26 @@ def direct_linux_release_plan(
selection = linux_cuda_choice_from_release(host, bundle)
if selection is not None:
attempts.extend(selection.attempts)
- cpu_choice = published_asset_choice_for_kind(bundle, "linux-cpu")
- if cpu_choice is not None:
- attempts.append(cpu_choice)
+ if host.has_rocm and not host.has_usable_nvidia:
+ # Per-GPU lemonade prebuilts ship the ROCm runtime libs alongside
+ # llama.cpp, so they install cleanly even on hosts (e.g. gfx1151
+ # Strix Halo) that the upstream combined-ROCm tarball doesn't cover.
+ # The "ubuntu" label is lemonade's asset naming convention only --
+ # the binary is a manylinux-style glibc build that runs on Arch,
+ # Fedora, openSUSE, etc. as long as the host glibc is recent enough.
+ # Do NOT append the CPU asset for ROCm-only hosts: if lemonade fails
+ # validation we want validate_prebuilt_attempts to raise PrebuiltFallback
+ # so the caller triggers the HIP source build, not silently install a
+ # CPU-only binary.
+ lemonade_choice = resolve_lemonade_rocm_choice(
+ host, "ubuntu", "linux-rocm", llama_tag = requested_tag
+ )
+ if lemonade_choice is not None:
+ attempts.append(lemonade_choice)
+ else:
+ cpu_choice = published_asset_choice_for_kind(bundle, "linux-cpu")
+ if cpu_choice is not None:
+ attempts.append(cpu_choice)
if not attempts:
raise PrebuiltFallback("no compatible Linux prebuilt asset was found")
approved_checksums = synthetic_checksums_for_release(
@@ -1336,6 +1400,25 @@ def direct_upstream_release_plan(
torch_preference.selection_log,
)
)
+ elif host.has_rocm:
+ lemonade_choice = resolve_lemonade_rocm_choice(
+ host, "windows", "windows-hip", llama_tag = requested_tag
+ )
+ if lemonade_choice is not None:
+ attempts.append(lemonade_choice)
+ hip_asset = f"llama-{release_tag}-bin-win-hip-radeon-x64.zip"
+ hip_url = assets.get(hip_asset)
+ if hip_url:
+ attempts.append(
+ AssetChoice(
+ repo = repo,
+ tag = release_tag,
+ name = hip_asset,
+ url = hip_url,
+ source_label = "upstream",
+ install_kind = "windows-hip",
+ )
+ )
cpu_asset = f"llama-{release_tag}-bin-win-cpu-x64.zip"
cpu_url = assets.get(cpu_asset)
if cpu_url:
@@ -1457,6 +1540,14 @@ def resolve_simple_install_release_plans(
requested_tag == "latest" and not published_release_tag
)
release_limit = max(1, max_release_fallbacks)
+ # macOS may need to walk past a run of too-new prebuilts. Only when the host
+ # version is known; otherwise keep the default (cannot tell up front).
+ if (
+ host.is_macos
+ and allow_older_release_fallback
+ and host.macos_version is not None
+ ):
+ release_limit = max(release_limit, DEFAULT_MAX_MACOS_RELEASE_FALLBACKS)
plans: list[InstallReleasePlan] = []
last_error: PrebuiltFallback | None = None
@@ -2604,6 +2695,72 @@ def run_capture(
return result
+def _pick_rocm_gfx_target(out: str) -> str | None:
+ """Choose the gfx target rocminfo / hipinfo report for the active GPU.
+
+ A bare first-match picked the wrong device on mixed APU + dGPU hosts
+ (e.g. Strix Halo gfx1151 + discrete RX 7900 gfx1100). Respect
+ HIP_VISIBLE_DEVICES / ROCR_VISIBLE_DEVICES / CUDA_VISIBLE_DEVICES so the
+ asset matches what HIP actually runs on. Falls back to the first GPU when
+ no env var is set.
+
+ rocminfo / hipinfo print the same gfx token multiple times per GPU (Name,
+ ISA, marketing-name). We first try to split the output on per-GPU section
+ headers (rocminfo: "Agent N" blocks, hipinfo: "device#N" entries) and take
+ exactly one gfx token per section. This gives the correct per-GPU list even
+ on same-arch multi-GPU hosts (e.g. two RX 7900 XTX cards) where global
+ dict.fromkeys dedup would collapse both cards to a single entry and make
+ HIP_VISIBLE_DEVICES=1 point out of range.
+
+ Falls back to insertion-order dedup when the output has no recognisable
+ section markers (flat gfx-string inputs, unit-test stubs, etc.).
+
+ Empty / "-1" env values mean no AMD GPU is visible to HIP: return None.
+ """
+ # Try to build a per-GPU token list by splitting on section boundaries.
+ # rocminfo sections are introduced by "Agent N" lines (optionally between
+ # rows of asterisks). hipinfo sections start with "device#N".
+ _sections = re.split(
+ r"(?mi)^\s*\*+\s*$\s*agent\s+\d+\s*$|\bdevice\s*#\s*\d+\b",
+ out,
+ )
+ if len(_sections) > 1:
+ # Section-based: one gfx token per GPU section preserves physical order.
+ _tokens: list[str] = []
+ for _sec in _sections[1:]:
+ _m = re.search(r"gfx[1-9][0-9a-z]{2,3}", _sec.lower())
+ if _m:
+ _tokens.append(_m.group(0))
+ else:
+ # Fallback: insertion-order dedup (handles flat strings / unknown formats).
+ _raw = re.findall(r"gfx[1-9][0-9a-z]{2,3}", out.lower())
+ _tokens = list(dict.fromkeys(_raw))
+
+ if not _tokens:
+ return None
+
+ _vis_raw = None
+ # AMD's HIP runtime honours all three env vars with identical semantics.
+ for _env in ("HIP_VISIBLE_DEVICES", "ROCR_VISIBLE_DEVICES", "CUDA_VISIBLE_DEVICES"):
+ _val = os.environ.get(_env)
+ if _val is not None:
+ _vis_raw = _val
+ break
+ if _vis_raw is not None:
+ _vis = _vis_raw.strip()
+ # Empty or "-1" means "no AMD GPU visible" (matches the rest of Studio).
+ if _vis == "" or _vis == "-1":
+ return None
+ _first = _vis.split(",")[0].strip()
+ try:
+ _idx = int(_first)
+ if 0 <= _idx < len(_tokens):
+ return _tokens[_idx]
+ except ValueError:
+ pass
+ return _tokens[0]
+
+
def detect_host() -> HostInfo:
system = platform.system()
machine = platform.machine().lower()
@@ -2613,6 +2770,8 @@ def detect_host() -> HostInfo:
is_x86_64 = machine in {"x86_64", "amd64"}
is_arm64 = machine in {"arm64", "aarch64"}
+ macos_version = parse_macos_version(platform.mac_ver()[0]) if is_macos else None
+
nvidia_smi = shutil.which("nvidia-smi")
driver_cuda_version = None
compute_caps: list[str] = []
@@ -2707,6 +2866,7 @@ def detect_host() -> HostInfo:
return bool(re.search(r"(?im)^gpu\s*[:\[]\s*\d", stdout))
has_rocm = False
+ rocm_gfx_target: str | None = None
if is_linux:
for _cmd, _check in (
# rocminfo: look for a real gfx GPU id (3-4 chars, nonzero first digit).
@@ -2729,14 +2889,33 @@ def detect_host() -> HostInfo:
if _result.returncode == 0 and _result.stdout.strip():
if _check(_result.stdout):
has_rocm = True
+ rocm_gfx_target = _pick_rocm_gfx_target(_result.stdout)
break
elif is_windows:
- # Windows: prefer active probes that validate GPU presence
+ # Windows: prefer active probes that validate GPU presence.
+ # hipinfo / amd-smi are often NOT on PATH -- the HIP SDK installer
+ # sets HIP_PATH / ROCM_PATH but does not always add the bin dir to
+ # the system PATH. Mirror setup.ps1's fallback: check the env-var
+ # bin dirs before giving up so that `has_rocm` is not silently False
+ # on machines where the PATH is not yet updated.
+ def _resolve_exe(name: str) -> str | None:
+ """Return full path to `name`, checking PATH then HIP_PATH/ROCM_PATH bin."""
+ found = shutil.which(name)
+ if found:
+ return found
+ for _env in ("HIP_PATH", "ROCM_PATH"):
+ _root = os.environ.get(_env)
+ if _root:
+ _candidate = os.path.join(_root, "bin", f"{name}.exe")
+ if os.path.isfile(_candidate):
+ return _candidate
+ return None
+
for _cmd, _check in (
(["hipinfo"], lambda out: "gcnarchname" in out.lower()),
(["amd-smi", "list"], _amd_smi_has_gpu),
):
- _exe = shutil.which(_cmd[0])
+ _exe = _resolve_exe(_cmd[0])
if not _exe:
continue
try:
@@ -2746,6 +2925,8 @@ def detect_host() -> HostInfo:
if _result.returncode == 0 and _result.stdout.strip():
if _check(_result.stdout):
has_rocm = True
+ # hipinfo reports "gcnArchName: gfx1100" -- extract if present
+ rocm_gfx_target = _pick_rocm_gfx_target(_result.stdout)
break
# Note: amdhip64.dll presence alone is NOT treated as GPU evidence
# since the HIP SDK can be installed without an AMD GPU.
@@ -2765,6 +2946,8 @@ def detect_host() -> HostInfo:
has_physical_nvidia = has_physical_nvidia,
has_usable_nvidia = has_usable_nvidia,
has_rocm = has_rocm,
+ rocm_gfx_target = rocm_gfx_target,
+ macos_version = macos_version,
)
@@ -3263,6 +3446,185 @@ def _detect_host_rocm_version() -> tuple[int, int] | None:
return None
+# Map detected gfx IDs to lemonade-sdk asset family suffixes.
+# More-specific prefixes must come before shorter ones (e.g. gfx1151 before gfx110).
+_LEMONADE_GFX_FAMILIES: list[tuple[str, str]] = [
+ ("gfx1151", "gfx1151"),
+ ("gfx1150", "gfx1150"),
+ ("gfx120", "gfx120X"),
+ ("gfx110", "gfx110X"),
+ ("gfx103", "gfx103X"),
+]
+
+
+def _lemonade_gfx_family(gfx_id: str) -> str | None:
+ gfx_id = gfx_id.lower().strip()
+ for prefix, family in _LEMONADE_GFX_FAMILIES:
+ if gfx_id.startswith(prefix):
+ return family
+ return None
+
+
+def _is_trusted_github_release_url(url: str, expected_repo: str) -> bool:
+ """Validate a release asset URL points at GitHub's expected hosts.
+
+ Accepts:
+ https://github.com/{expected_repo}/releases/download/...
+ https://objects.githubusercontent.com/... (GitHub's release CDN)
+ Anything else (including http://, raw.githubusercontent.com, gist, etc.)
+ is rejected so a malicious API response cannot redirect downloads to an
+ attacker-chosen host.
+ """
+ if not isinstance(url, str) or not url:
+ return False
+ try:
+ parsed = urllib.parse.urlparse(url)
+ except Exception:
+ return False
+ if parsed.scheme != "https":
+ return False
+ host = (parsed.netloc or "").lower()
+ if host == "objects.githubusercontent.com":
+ # GitHub's release CDN. Restrict to release-asset paths so a tampered
+ # API response pointing at an arbitrary CDN object is still rejected.
+ # Real release asset URLs carry the "/github-production-release-asset-"
+ # prefix; gist / raw / avatar CDN paths do not.
+ return parsed.path.startswith("/github-production-release-asset-")
+ if host == "github.com":
+ return parsed.path.startswith(f"/{expected_repo}/releases/download/")
+ return False
+
+
+@functools.lru_cache(maxsize = 8)
+def _fetch_lemonade_release_cached(api_url: str, llama_tag: str) -> "dict | None":
+ """Cached wrapper around fetch_json for lemonade release lookups.
+
+ resolve_lemonade_rocm_choice() is called twice per install (once from the
+ direct planner, once from resolve_upstream_asset_choice) with identical
+ arguments. Without memoisation, each install hits api.github.com twice,
+ doubling the rate-limit failure surface on busy CI runners. Cache is
+ process-scoped; tests that need to vary fetch_json's return value across
+ invocations should call cache_clear().
+ """
+ try:
+ return fetch_json(api_url)
+ except Exception as exc:
+ normalized = (llama_tag or "").strip().lower()
+ if normalized and normalized != "latest":
+ log(
+ f"Could not fetch {LEMONADE_ROCM_REPO} release for "
+ f"llama_tag={llama_tag!r} ({exc}); skipping lemonade prebuilt"
+ )
+ else:
+ log(f"Could not fetch {LEMONADE_ROCM_REPO} latest release: {exc}")
+ return None
+
+
+def resolve_lemonade_rocm_choice(
+ host: HostInfo,
+ os_prefix: str,
+ install_kind: str,
+ llama_tag: str = "latest",
+) -> "AssetChoice | None":
+ """Return an AssetChoice from lemonade-sdk/llamacpp-rocm for the detected GPU, or None.
+
+ os_prefix: lemonade's asset filename label, NOT a host-distro filter.
+ Pass "ubuntu" for any Linux host (Arch, Fedora, openSUSE,
+ Debian, ...) -- lemonade only publishes one Linux variant
+ and it is a manylinux-style glibc build that runs on any
+ distro with a recent-enough glibc. Pass "windows" for
+ Windows hosts.
+ install_kind: "linux-rocm" or "windows-hip"
+ llama_tag: the requested upstream llama.cpp tag ("latest" or a pinned
+ release like "b1260"). When pinned, the resolver fetches
+ the matching lemonade release. When the pinned tag is not
+ published by lemonade we skip silently (and the caller
+ falls through to upstream) rather than drift to whatever
+ lemonade ships as latest.
+ """
+ if not host.rocm_gfx_target:
+ return None
+ # Opt-out for users who want the upstream HIP build path only -- lemonade
+ # binaries are downloaded without entries in the approved-hash manifest, so
+ # the integrity gate is functional validation only.
+ if os.environ.get("UNSLOTH_DISABLE_LEMONADE_ROCM", "").strip().lower() in (
+ "1",
+ "true",
+ "yes",
+ ):
+ log("UNSLOTH_DISABLE_LEMONADE_ROCM is set; skipping lemonade-sdk prebuilt")
+ return None
+ gfx_family = _lemonade_gfx_family(host.rocm_gfx_target)
+ if gfx_family is None:
+ log(
+ f"AMD GPU {host.rocm_gfx_target!r} is not covered by lemonade-sdk ROCm prebuilts; "
+ "skipping lemonade prebuilt"
+ )
+ return None
+ api_url = _lemonade_release_api_for(llama_tag)
+ release = _fetch_lemonade_release_cached(api_url, llama_tag)
+ if release is None:
+ return None
+ release_tag = release.get("tag_name") if isinstance(release, dict) else None
+ if not isinstance(release_tag, str) or not release_tag:
+ log(
+ f"Unexpected {LEMONADE_ROCM_REPO} release payload; skipping lemonade prebuilt"
+ )
+ return None
+ assets = release_asset_map(release)
+ asset_name = f"llama-{release_tag}-{os_prefix}-rocm-{gfx_family}-x64.zip"
+ if asset_name not in assets:
+ log(
+ f"{LEMONADE_ROCM_REPO}@{release_tag} has no asset {asset_name!r}; "
+ "skipping lemonade prebuilt"
+ )
+ return None
+ asset_url = assets[asset_name]
+ if not asset_url:
+ # release_asset_map defaults to "" when an asset row is missing
+ # browser_download_url; skip cleanly instead of letting
+ # download_file("") raise a less obvious error downstream.
+ log(
+ f"{LEMONADE_ROCM_REPO}@{release_tag} asset {asset_name!r} has no "
+ "browser_download_url; skipping lemonade prebuilt"
+ )
+ return None
+ # Defence in depth: lemonade browser_download_url should be on github.com
+ # or githubusercontent.com. A compromised GitHub API response that
+ # redirects to an attacker-chosen host would otherwise be honoured
+ # silently (lemonade assets are not in the approved-hash manifest).
+ if not _is_trusted_github_release_url(asset_url, LEMONADE_ROCM_REPO):
+ log(
+ f"{LEMONADE_ROCM_REPO}@{release_tag} asset {asset_name!r} points "
+ f"to an unexpected host ({asset_url!r}); refusing to download "
+ "lemonade prebuilt"
+ )
+ return None
+ # Note: lemonade tags Linux assets with "ubuntu" but the binary is a
+ # generic glibc build that runs on any distro (Arch, Fedora, ...), so
+ # this attempt is selected for all Linux ROCm hosts, not just Ubuntu.
+ log(
+ f"AMD GPU {host.rocm_gfx_target!r} ({gfx_family}) -- "
+ f"trying lemonade-sdk ROCm prebuilt {asset_name} "
+ f"(works on any glibc Linux, not just Ubuntu)"
+ )
+ log(
+ f"NOTE: lemonade-sdk/llamacpp-rocm releases are not covered by the "
+ f"Unsloth approved-hash manifest; download integrity relies on "
+ f"functional validation (llama-bench / llama-server smoke tests) "
+ f"after extraction. Set UNSLOTH_DISABLE_LEMONADE_ROCM=1 to skip "
+ f"lemonade and fall back to the upstream HIP build path."
+ )
+ return AssetChoice(
+ repo = LEMONADE_ROCM_REPO,
+ tag = release_tag,
+ name = asset_name,
+ url = asset_url,
+ source_label = "lemonade",
+ install_kind = install_kind,
+ )
+
+
def resolve_upstream_asset_choice(host: HostInfo, llama_tag: str) -> AssetChoice:
upstream_assets = github_release_assets(UPSTREAM_REPO, llama_tag)
if host.is_linux and host.is_x86_64:
@@ -3271,6 +3633,15 @@ def resolve_upstream_asset_choice(host: HostInfo, llama_tag: str) -> AssetChoice
# the exact GPU target via rocminfo, which is more reliable for consumer
# GPUs (e.g. gfx1151) that may not be in the prebuilt.
if host.has_rocm and not host.has_usable_nvidia:
+ # Try lemonade-sdk per-GPU prebuilt first: these are built against
+ # specific gfx targets and bundle all required ROCm runtime libs.
+ lemonade_choice = resolve_lemonade_rocm_choice(
+ host, "ubuntu", "linux-rocm", llama_tag = llama_tag
+ )
+ if lemonade_choice is not None:
+ return lemonade_choice
+
+ # Fall back to upstream combined ROCm tarball.
# Scan upstream assets for any rocm- prebuilt. When the
# host ROCm runtime version is known, pick the newest candidate
# whose major.minor is <= host version -- otherwise a ROCm 6.4
@@ -3350,8 +3721,14 @@ def resolve_upstream_asset_choice(host: HostInfo, llama_tag: str) -> AssetChoice
return attempts[0]
raise PrebuiltFallback("no compatible Windows CUDA asset was found")
- # AMD ROCm on Windows: try HIP prebuilt
+ # AMD ROCm on Windows: try lemonade per-GPU prebuilt first, then upstream HIP
if host.has_rocm:
+ lemonade_choice = resolve_lemonade_rocm_choice(
+ host, "windows", "windows-hip", llama_tag = llama_tag
+ )
+ if lemonade_choice is not None:
+ return lemonade_choice
+
hip_name = f"llama-{llama_tag}-bin-win-hip-radeon-x64.zip"
if hip_name in upstream_assets:
log(
@@ -3892,6 +4269,22 @@ def runtime_patterns_for_choice(choice: AssetChoice) -> list[str]:
)
+def runtime_subdirs_for_choice(choice: AssetChoice) -> list[str]:
+ """Subdirectory names within the archive root that must be copied into
+ the overlay directory alongside the flat shared libraries.
+
+ hipBLASLt and rocBLAS expect their Tensile kernel catalog trees
+ (hipblaslt/library// and rocblas/library//) to sit next to
+ their shared libraries at runtime. These trees are multi-level and
+ cannot be handled by copy_globs (filename-only matching, flat copy)."""
+ if choice.source_label == "lemonade" and choice.install_kind in {
+ "linux-rocm",
+ "windows-hip",
+ }:
+ return ["hipblaslt", "rocblas"]
+ return []
+
+
def metadata_patterns_for_choice(choice: AssetChoice) -> list[str]:
patterns = ["BUILD_INFO.txt", "THIRD_PARTY_LICENSES.txt"]
if choice.install_kind.startswith("windows"):
@@ -4216,6 +4609,10 @@ def install_from_archives(
copy_globs(
source_dir, overlay_dir, runtime_patterns_for_choice(choice), required = True
)
+ for _subdir in runtime_subdirs_for_choice(choice):
+ _src_subdir = source_dir / _subdir
+ if _src_subdir.is_dir():
+ shutil.copytree(_src_subdir, overlay_dir / _subdir, dirs_exist_ok = True)
if runtime_extract_dir is not None:
# The runtime archive only contributes the CUDA DLLs.
# Restrict the overlay to the cudart bundle's known
@@ -4498,6 +4895,189 @@ def linux_runtime_dirs(binary_path: Path) -> list[str]:
return linux_runtime_dirs_for_required_libraries(missing)
+# macOS prebuilt compatibility. Upstream macos prebuilts built on a newer macOS
+# (e.g. minos=26, referencing Metal-4 symbols) fail dyld load on macOS 14/15. We
+# read the host macOS version and each binary's minimum-OS so selection can skip
+# a too-new prebuilt and walk back to the newest release that runs on this host.
+# Mach-O constants (Apple mach-o/fat.h, mach-o/loader.h, mach/machine.h).
+_MACHO_FAT_MAGICS = {0xCAFEBABE, 0xCAFEBABF} # universal binary (32/64-bit fat)
+_LC_VERSION_MIN_MACOSX = 0x24 # legacy min-macOS load command
+_LC_BUILD_VERSION = 0x32 # modern platform+minos+sdk load command
+_MACHO_PLATFORM_MACOS = 1 # LC_BUILD_VERSION platform id for macOS (iOS=2, ...)
+# CPU types (base | ABI64); used to pick the host slice in a fat binary.
+_CPU_TYPE_X86_64 = 0x01000007
+_CPU_TYPE_ARM64 = 0x0100000C
+
+
+def parse_macos_version(value: str | None) -> tuple[int, int] | None:
+ """Parse a macOS product version string into (major, minor).
+
+ Handles "14.7.1", "15.5", "26.0" and bare "26". Returns None when the
+ value is empty or cannot be parsed (callers then defer to runtime
+ validation rather than rejecting every prebuilt)."""
+ if not value:
+ return None
+ match = re.match(r"\s*(\d+)(?:\.(\d+))?", str(value))
+ if not match:
+ return None
+ return int(match.group(1)), int(match.group(2) or 0)
+
+
+def host_supports_macos_minos(host: HostInfo, minos: tuple[int, int] | None) -> bool:
+ """True if a prebuilt requiring `minos` can load on this host. Unknown host
+ version or unknown minos -> True: let runtime validation decide instead of
+ rejecting a binary we cannot reason about."""
+ if minos is None or host.macos_version is None:
+ return True
+ return host.macos_version >= minos
+
+
+def _macho_slice_minos(data: bytes, offset: int) -> tuple[int, int] | None:
+ """Minimum macOS for a single thin Mach-O at `offset`, via LC_BUILD_VERSION
+ (platform macOS) or the legacy LC_VERSION_MIN_MACOSX. None if absent."""
+ if offset + 4 > len(data):
+ return None
+ magic = struct.unpack_from(">I", data, offset)[0]
+ if magic in (0xFEEDFACE, 0xFEEDFACF):
+ endian, is64 = ">", magic == 0xFEEDFACF
+ elif magic in (0xCEFAEDFE, 0xCFFAEDFE):
+ endian, is64 = "<", magic == 0xCFFAEDFE
+ else:
+ return None
+ header_size = 32 if is64 else 28
+ if offset + header_size > len(data):
+ return None
+ ncmds = struct.unpack_from(endian + "I", data, offset + 16)[0]
+ cursor = offset + header_size
+ for _ in range(ncmds):
+ if cursor + 8 > len(data):
+ break
+ cmd, cmdsize = struct.unpack_from(endian + "II", data, cursor)
+ if cmdsize < 8:
+ break
+ if cmd == _LC_BUILD_VERSION and cursor + 16 <= len(data):
+ platform_id, minos = struct.unpack_from(endian + "II", data, cursor + 8)
+ if platform_id == _MACHO_PLATFORM_MACOS:
+ return (minos >> 16) & 0xFFFF, (minos >> 8) & 0xFF
+ elif cmd == _LC_VERSION_MIN_MACOSX and cursor + 12 <= len(data):
+ version = struct.unpack_from(endian + "I", data, cursor + 8)[0]
+ return (version >> 16) & 0xFFFF, (version >> 8) & 0xFF
+ cursor += cmdsize
+ return None
+
+
+def macho_minimum_macos(
+ path: Path, host: HostInfo | None = None
+) -> tuple[int, int] | None:
+ """Minimum macOS (major, minor) a Mach-O binary or dylib requires.
+
+ Pure-Python so it works on consumer Macs without the Xcode command line
+ tools (otool/vtool). For universal binaries it prefers the host-arch slice,
+ else the highest minos found. Returns None for non-Mach-O files or when no
+ version load command is present."""
+ try:
+ data = path.read_bytes()
+ except Exception:
+ return None
+ if len(data) < 8:
+ return None
+ magic = struct.unpack_from(">I", data, 0)[0]
+ if magic in _MACHO_FAT_MAGICS:
+ is64 = magic == 0xCAFEBABF
+ nfat = struct.unpack_from(">I", data, 4)[0]
+ entry = 8
+ slices: list[tuple[int, tuple[int, int]]] = []
+ for _ in range(nfat):
+ if is64:
+ if entry + 32 > len(data):
+ break
+ cputype = struct.unpack_from(">I", data, entry)[0]
+ slice_offset = struct.unpack_from(">Q", data, entry + 8)[0]
+ entry += 32
+ else:
+ if entry + 20 > len(data):
+ break
+ cputype = struct.unpack_from(">I", data, entry)[0]
+ slice_offset = struct.unpack_from(">I", data, entry + 8)[0]
+ entry += 20
+ minos = _macho_slice_minos(data, slice_offset)
+ if minos is not None:
+ slices.append((cputype, minos))
+ if not slices:
+ return None
+ if host is not None:
+ want = (
+ _CPU_TYPE_ARM64
+ if host.is_arm64
+ else (_CPU_TYPE_X86_64 if host.is_x86_64 else None)
+ )
+ for cputype, minos in slices:
+ if cputype == want:
+ return minos
+ return max(minos for _cputype, minos in slices)
+ return _macho_slice_minos(data, 0)
+
+
+def looks_like_macos_incompatibility(text: str) -> bool:
+ """True when dyld output means a prebuilt needs a newer macOS than the host
+ (the runtime backstop for cases the static minos scan cannot read)."""
+ if not text:
+ return False
+ if "built for macOS" in text and "newer than running OS" in text:
+ return True
+ return "Symbol not found" in text and "MTLResidency" in text
+
+
+def macos_binary_minos_issues(
+ binaries: Iterable[Path],
+ install_dir: Path,
+ host: HostInfo,
+) -> list[str]:
+ """Issue strings for every installed Mach-O whose minimum macOS exceeds the
+ host. Scans the given executables plus every bundled .dylib next to them --
+ the dyld failure originates in libggml-metal.dylib, not the executable."""
+ candidates: list[Path] = list(binaries)
+ bin_dir = install_dir / "build" / "bin"
+ if bin_dir.is_dir():
+ candidates.extend(sorted(bin_dir.rglob("*.dylib")))
+
+ issues: list[str] = []
+ seen: set[Path] = set()
+ for path in candidates:
+ try:
+ resolved = path.resolve()
+ except Exception:
+ resolved = path
+ if resolved in seen or not path.exists():
+ continue
+ seen.add(resolved)
+ minos = macho_minimum_macos(path, host)
+ if minos is not None and not host_supports_macos_minos(host, minos):
+ issues.append(
+ f"{path.name}: built for macOS {minos[0]}.{minos[1]} > "
+ f"host macOS {host.macos_version[0]}.{host.macos_version[1]}"
+ )
+ return issues
+
+
+def preflight_macos_installed_binaries(
+ binaries: Iterable[Path],
+ install_dir: Path,
+ host: HostInfo,
+) -> None:
+ """Reject a macos prebuilt whose minimum-OS is newer than the host so the
+ release walk-back advances to the newest compatible release. No-op when the
+ host macOS version is unknown (runtime validation remains the backstop)."""
+ if not host.is_macos or host.macos_version is None:
+ return
+ issues = macos_binary_minos_issues(binaries, install_dir, host)
+ if issues:
+ raise PrebuiltFallback(
+ "macos prebuilt requires a newer macOS than this host:\n"
+ + "\n".join(issues)
+ )
+
+
def preflight_linux_installed_binaries(
binaries: Iterable[Path],
install_dir: Path,
@@ -4656,10 +5236,17 @@ def validate_quantize(
or not quantized_path.exists()
or quantized_path.stat().st_size == 0
):
+ combined = result.stdout + ("\n" + result.stderr if result.stderr else "")
+ # Backstop for prebuilts the static minos scan could not read: a dyld
+ # "built for macOS N" / missing Metal symbol failure means this binary
+ # needs a newer macOS than the host, so fall back to an older release.
+ prefix = (
+ "macos prebuilt requires a newer macOS than this host: "
+ if looks_like_macos_incompatibility(combined)
+ else ""
+ )
raise PrebuiltFallback(
- "llama-quantize validation failed:\n"
- + result.stdout
- + ("\n" + result.stderr if result.stderr else "")
+ prefix + "llama-quantize validation failed:\n" + combined
)
@@ -4964,6 +5551,15 @@ def apply_approved_hashes(
approved_attempts: list[AssetChoice] = []
missing_assets: list[str] = []
for attempt in attempts:
+ # External prebuilts (e.g. lemonade-sdk) are not listed in the
+ # approved-hash manifest; they are explicitly documented as relying
+ # on functional validation only (llama-bench / smoke tests).
+ # Passing them through here lets the caller include both a lemonade
+ # attempt and a hash-approved upstream fallback in the same list
+ # without apply_approved_hashes discarding the lemonade entry.
+ if attempt.source_label == "lemonade":
+ approved_attempts.append(attempt)
+ continue
approved = approved_hash_for_attempt(attempt)
if approved is None:
missing_assets.append(attempt.name)
@@ -5068,6 +5664,14 @@ def resolve_install_release_plans(
requested_tag == "latest" and not published_release_tag
)
release_limit = max(1, max_release_fallbacks)
+ # macOS may need to walk past a run of too-new prebuilts. Only when the host
+ # version is known; otherwise keep the default (cannot tell up front).
+ if (
+ host.is_macos
+ and allow_older_release_fallback
+ and host.macos_version is not None
+ ):
+ release_limit = max(release_limit, DEFAULT_MAX_MACOS_RELEASE_FALLBACKS)
plans: list[InstallReleasePlan] = []
last_error: PrebuiltFallback | None = None
@@ -5164,6 +5768,12 @@ def write_prebuilt_metadata(
"asset": choice.name,
"asset_sha256": choice.expected_sha256,
"source": choice.source_label,
+ # Binary-side repo/tag for non-upstream sources (e.g. lemonade).
+ # published_repo/release_tag always refer to the unsloth source tree;
+ # these capture where the actual binaries came from so the install
+ # summary can show both (e.g. "unslothai/llama.cpp@b9334 + lemonade@b1280").
+ "binary_repo": choice.repo,
+ "binary_release_tag": choice.tag,
"source_asset": source_asset_name,
"source_sha256": source_sha256,
"source_commit": approved_checksums.source_commit,
@@ -5242,7 +5852,7 @@ def runtime_payload_health_groups(choice: AssetChoice) -> list[list[str]]:
["libllama.so*"],
["libggml.so*"],
["libggml-base.so*"],
- ["libggml-cpu-*.so*"],
+ ["libggml-cpu*.so*"],
["libmtmd.so*"],
]
if choice.install_kind == "linux-cuda":
@@ -5251,7 +5861,7 @@ def runtime_payload_health_groups(choice: AssetChoice) -> list[list[str]]:
["libllama.so*"],
["libggml.so*"],
["libggml-base.so*"],
- ["libggml-cpu-*.so*"],
+ ["libggml-cpu*.so*"],
["libmtmd.so*"],
["libggml-cuda.so*"],
]
@@ -5267,7 +5877,7 @@ def runtime_payload_health_groups(choice: AssetChoice) -> list[list[str]]:
["libllama.so*"],
["libggml.so*"],
["libggml-base.so*"],
- ["libggml-cpu-*.so*"],
+ ["libggml-cpu*.so*"],
["libmtmd.so*"],
["libggml-hip.so*"],
]
@@ -5446,6 +6056,7 @@ def validate_prebuilt_choice(
choice, host, install_dir, work_dir
)
preflight_linux_installed_binaries((server_path, quantize_path), install_dir, host)
+ preflight_macos_installed_binaries((server_path, quantize_path), install_dir, host)
ensure_repo_shape(install_dir)
write_prebuilt_metadata(
install_dir,
@@ -5565,8 +6176,11 @@ def install_prebuilt(
published_release_tag: str,
*,
simple_policy: bool = False,
+ override_has_rocm: bool = False,
) -> None:
host = detect_host()
+ if override_has_rocm and not host.has_rocm:
+ host = dataclasses_replace(host, has_rocm = True)
choice: AssetChoice | None = None
try:
with install_lock(install_lock_path(install_dir)):
@@ -5700,6 +6314,17 @@ def parse_args() -> argparse.Namespace:
action = "store_true",
help = "Use the simplified platform-specific prebuilt selection policy.",
)
+ parser.add_argument(
+ "--has-rocm",
+ action = "store_true",
+ default = False,
+ help = (
+ "Assert that an AMD ROCm GPU is present. When set, skips the internal "
+ "hipinfo/amd-smi probe and forces has_rocm=True in the host profile. "
+ "Used by setup.ps1/setup.sh to forward their own ROCm detection result "
+ "so the HIP llama.cpp prebuilt is selected even when hipinfo is not on PATH."
+ ),
+ )
resolve_group = parser.add_mutually_exclusive_group()
resolve_group.add_argument(
"--resolve-llama-tag",
@@ -5820,6 +6445,7 @@ def main() -> int:
published_repo = args.published_repo,
published_release_tag = args.published_release_tag or "",
simple_policy = args.simple_policy,
+ override_has_rocm = args.has_rocm,
)
return EXIT_SUCCESS
diff --git a/studio/install_python_stack.py b/studio/install_python_stack.py
index 158218062d..c0a78751b7 100644
--- a/studio/install_python_stack.py
+++ b/studio/install_python_stack.py
@@ -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",
)
diff --git a/studio/setup.ps1 b/studio/setup.ps1
index afa00409fb..6c5f4171c4 100644
--- a/studio/setup.ps1
+++ b/studio/setup.ps1
@@ -685,13 +685,228 @@ if (-not $HasNvidiaSmi) {
}
}
}
+# ── AMD ROCm detection (Windows): probe hipinfo/amd-smi for actual GPU ──
+$HasROCm = $false
+$HipSdkInstalled = $false # HIP SDK binary found (independent of device accessibility)
+$ROCmGpuLabel = $null
+$script:ROCmGfxArch = $null
if (-not $HasNvidiaSmi) {
+ # hipinfo: PATH first, then HIP_PATH/ROCM_PATH bin fallback (mirrors NVIDIA smi path resolution).
+ # AMD HIP SDK sets HIP_PATH but may not add the bin dir to PATH depending on install type.
+ $hipinfoExe = Get-Command hipinfo -ErrorAction SilentlyContinue
+ if (-not $hipinfoExe) {
+ $hipRoot = if ($env:HIP_PATH) { $env:HIP_PATH } elseif ($env:ROCM_PATH) { $env:ROCM_PATH } else { $null }
+ $hipEnvLabel = if ($env:HIP_PATH) { "HIP_PATH" } else { "ROCM_PATH" }
+ if ($hipRoot) {
+ $hipinfoCandidate = Join-Path $hipRoot "bin\hipinfo.exe"
+ if (Test-Path $hipinfoCandidate) {
+ substep "[WARN] hipinfo not on PATH -- located via ${hipEnvLabel}: $hipinfoCandidate" "Yellow"
+ substep " Add '$(Join-Path $hipRoot 'bin')' to your PATH to suppress this warning" "Yellow"
+ substep " Quick fix: [Environment]::SetEnvironmentVariable('PATH',`$env:PATH+';$(Join-Path $hipRoot 'bin')','User')" "Yellow"
+ $hipinfoExe = [PSCustomObject]@{ Source = $hipinfoCandidate }
+ } else {
+ substep "[WARN] ${hipEnvLabel}=$hipRoot is set but hipinfo.exe not found at $hipinfoCandidate" "Yellow"
+ substep " HIP SDK install may be incomplete -- re-install from:" "Yellow"
+ substep " https://rocm.docs.amd.com/en/latest/deploy/windows/index.html" "Yellow"
+ }
+ }
+ }
+ if ($hipinfoExe) {
+ $HipSdkInstalled = $true # binary found → SDK is installed regardless of device state
+ try {
+ $hipOut = & $hipinfoExe.Source 2>&1 | Out-String
+ if ($LASTEXITCODE -eq 0 -and $hipOut -match "(?i)gcnArchName") {
+ $HasROCm = $true
+ $_hipAllArches = @([regex]::Matches($hipOut, "(?im)^\s*gcnArchName\s*:\s*(\S+)") | ForEach-Object { ($_.Groups[1].Value -split ':')[0].Trim().ToLower() })
+ $_hipVisIdx = if ($env:HIP_VISIBLE_DEVICES -match '^\d') { [int]($env:HIP_VISIBLE_DEVICES -split ',')[0] } elseif ($env:ROCR_VISIBLE_DEVICES -match '^\d') { [int]($env:ROCR_VISIBLE_DEVICES -split ',')[0] } else { 0 }
+ if ($_hipAllArches.Count -gt 0) {
+ $script:ROCmGfxArch = if ($_hipVisIdx -lt $_hipAllArches.Count) { $_hipAllArches[$_hipVisIdx] } else { $_hipAllArches[0] }
+ $ROCmGpuLabel = "AMD ROCm ($script:ROCmGfxArch)"
+ } else {
+ $ROCmGpuLabel = "AMD ROCm"
+ }
+ } elseif ($LASTEXITCODE -ne 0) {
+ # hipinfo ran but returned a HIP runtime error (e.g. "no ROCm-capable device detected")
+ $firstLine = ($hipOut -split '\r?\n' | Where-Object { $_.Trim() } | Select-Object -First 1)
+ substep "[WARN] hipinfo returned a HIP runtime error (exit $LASTEXITCODE)" "Yellow"
+ substep " $firstLine" "Yellow"
+ substep " Ensure ROCm drivers are installed: https://rocm.docs.amd.com/en/latest/deploy/windows/index.html" "Yellow"
+ }
+ } catch {}
+ }
+ # amd-smi fallback: HIP runtime present but hipinfo unavailable (no full HIP SDK).
+ # Confirms GPU visibility via 'list', then attempts 'static --asic' to extract
+ # the gfx arch that hipinfo would have provided. Critical for Strix Halo
+ # (gfx1151) and other iGPUs where only the HIP runtime is installed.
+ if (-not $HasROCm) {
+ $amdSmiExe = Get-Command "amd-smi" -ErrorAction SilentlyContinue
+ if ($amdSmiExe) {
+ try {
+ $smiOut = & $amdSmiExe.Source list 2>&1 | Out-String
+ if ($LASTEXITCODE -eq 0 -and $smiOut -match "(?im)^GPU\s*[:\[]\s*\d") {
+ $HasROCm = $true
+ # Attempt 1: newer amd-smi versions embed the gfx arch in list output.
+ # Collect ALL gfx tokens in output order so that on mixed-arch systems
+ # we can honour HIP_VISIBLE_DEVICES / ROCR_VISIBLE_DEVICES and pick the
+ # arch for the *runtime-visible* GPU rather than always the first one.
+ # Do NOT deduplicate: a dual same-arch system (e.g. two gfx1151 APUs)
+ # must produce a 2-element array so HIP_VISIBLE_DEVICES=1 selects the
+ # second GPU rather than triggering a false out-of-range warning.
+ # Note: this mapping assumes amd-smi lists GPUs in the same order as
+ # HIP enumerates them (both follow PCI bus order in practice); it may
+ # give the wrong arch when GPU indices are non-contiguous (very rare).
+ $allGfxArches = @([regex]::Matches($smiOut, '(?i)\b(gfx\d+[a-z]?)\b') |
+ ForEach-Object { $_.Groups[1].Value.ToLower() })
+ if ($allGfxArches.Count -gt 0) {
+ # Resolve which GPU index is runtime-visible. When a single
+ # integer index is set, use it; fall back to index 0 otherwise
+ # (comma-separated lists or unset → first GPU, same as before).
+ $visGpu = if ($env:HIP_VISIBLE_DEVICES) { $env:HIP_VISIBLE_DEVICES }
+ elseif ($env:ROCR_VISIBLE_DEVICES) { $env:ROCR_VISIBLE_DEVICES }
+ else { $null }
+ $gpuIdx = 0
+ if ($visGpu -match '^\s*(\d+)\s*$') { $gpuIdx = [int]$Matches[1] }
+ if ($gpuIdx -ge $allGfxArches.Count) {
+ substep "[WARN] HIP/ROCR_VISIBLE_DEVICES index $gpuIdx is out of range ($($allGfxArches.Count) GPU(s) detected); defaulting to GPU 0 for arch selection" "Yellow"
+ $gpuIdx = 0
+ }
+ $script:ROCmGfxArch = $allGfxArches[$gpuIdx]
+ $ROCmGpuLabel = "AMD ROCm ($script:ROCmGfxArch)"
+ } else {
+ # Attempt 2: 'static --asic' exposes ASIC details on ROCm 6+,
+ # including the GFX target needed for wheel index selection.
+ $smiAsicOut = ""
+ try { $smiAsicOut = & $amdSmiExe.Source static --asic 2>&1 | Out-String } catch {}
+ if ($smiAsicOut -match "(?i)\b(gfx\d+[a-z]?)\b") {
+ $script:ROCmGfxArch = $Matches[1].ToLower()
+ $ROCmGpuLabel = "AMD ROCm ($script:ROCmGfxArch)"
+ } elseif ($smiAsicOut -match "(?im)Market.?Name\s*[:\|]\s*([^\r\n]+)") {
+ $ROCmGpuLabel = "AMD ROCm ($($Matches[1].Trim()))"
+ } else {
+ $ROCmGpuLabel = "AMD ROCm"
+ }
+ }
+ }
+ } catch {}
+ }
+ }
+ # WMI fallback: AMD GPU in device list but no HIP SDK → guide the user.
+ # WMI gives a marketing name (e.g. "AMD Radeon 890M") but never a gfx arch.
+ # $HasROCm is intentionally NOT set here — we cannot confirm ROCm runtime
+ # support without hipinfo or amd-smi. The name is saved to $ROCmGpuLabel
+ # so the name-based inference below can still attempt an arch lookup.
+ if (-not $HasROCm) {
+ try {
+ $wmiGpu = Get-WmiObject Win32_VideoController -ErrorAction SilentlyContinue |
+ Where-Object { $_.Name -match "AMD|Radeon" } |
+ Select-Object -First 1
+ if ($wmiGpu) { $ROCmGpuLabel = $wmiGpu.Name }
+ } catch {}
+ }
+ # ── Arch resolution: env-var override → name inference ──────────────────
+ # Runs after all probe methods. Covers users whose amd-smi version is too
+ # old to report the GFX target and who don't have hipinfo (HIP-runtime-only
+ # installs, common on Strix Halo / iGPU systems).
+ if ($HasROCm -and -not $script:ROCmGfxArch) {
+ # 1. Manual override: set UNSLOTH_ROCM_GFX_ARCH=gfx1151 before running.
+ if ($env:UNSLOTH_ROCM_GFX_ARCH) {
+ $script:ROCmGfxArch = $env:UNSLOTH_ROCM_GFX_ARCH.Trim().ToLower()
+ $ROCmGpuLabel = "AMD ROCm ($script:ROCmGfxArch)"
+ substep "gfx arch from UNSLOTH_ROCM_GFX_ARCH env override: $script:ROCmGfxArch" "Cyan"
+ }
+ # 2. Best-effort name → arch lookup from marketing name (amd-smi / WMI).
+ # Ordered most-specific first; first match wins.
+ elseif ($ROCmGpuLabel) {
+ $nameArchTable = @(
+ @{ P = "9070 XT|9080"; A = "gfx1201" } # RDNA 4
+ @{ P = "9070|9060"; A = "gfx1200" } # RDNA 4
+ @{ P = "8060S|890M|Strix Halo|HX 37[05]|HX 38[05]|AI 9 HX"; A = "gfx1151" } # RDNA 3.5 iGPU (Strix Halo / Radeon 8060S retail)
+ @{ P = "880M|Strix Point|AI 9 36[05]|AI 7 35[05]|AI 5 34[05]"; A = "gfx1150" } # RDNA 3.5 iGPU (Strix Point)
+ @{ P = "RX 7900|RX 7800|RX 7700(?! S)"; A = "gfx1100" } # RDNA 3 desktop
+ @{ P = "RX 7600"; A = "gfx1102" } # RDNA 3
+ @{ P = "780M|760M|740M|Phoenix"; A = "gfx1103" } # RDNA 3 iGPU (Phoenix)
+ )
+ foreach ($row in $nameArchTable) {
+ if ($ROCmGpuLabel -match $row.P) {
+ $script:ROCmGfxArch = $row.A
+ $ROCmGpuLabel = "AMD ROCm ($script:ROCmGfxArch)"
+ substep "gfx arch inferred from GPU name: $script:ROCmGfxArch" "Cyan"
+ substep "Tip: set UNSLOTH_ROCM_GFX_ARCH=$script:ROCmGfxArch to skip inference next time" "Cyan"
+ break
+ }
+ }
+ }
+ }
+ # Capture ROCm version early for display and wheel selection.
+ # Run whenever the HIP SDK binary is present, not just when the device is accessible --
+ # hipconfig --version works even when hipinfo reports no ROCm device (driver issue).
+ if ($HasROCm -or $HipSdkInstalled) {
+ $script:ROCmVersion = $null
+ $hipConfigExe = Get-Command hipconfig -ErrorAction SilentlyContinue
+ if (-not $hipConfigExe) {
+ $hipRoot = if ($env:HIP_PATH) { $env:HIP_PATH } elseif ($env:ROCM_PATH) { $env:ROCM_PATH } else { $null }
+ if ($hipRoot) {
+ $hipConfigCandidate = Join-Path $hipRoot "bin\hipconfig.exe"
+ if (Test-Path $hipConfigCandidate) {
+ $hipConfigEnvLabel = if ($env:HIP_PATH) { "HIP_PATH" } else { "ROCM_PATH" }
+ substep "[WARN] hipconfig not on PATH -- located via ${hipConfigEnvLabel}: $hipConfigCandidate" "Yellow"
+ $hipConfigExe = [PSCustomObject]@{ Source = $hipConfigCandidate }
+ }
+ }
+ }
+ if ($hipConfigExe) {
+ try {
+ $hipVerOut = & $hipConfigExe.Source --version 2>&1 | Out-String
+ if ($LASTEXITCODE -eq 0) {
+ $hipVerLine = ($hipVerOut -split '\r?\n' | Where-Object { $_.Trim() } | Select-Object -First 1).Trim()
+ if ($hipVerLine -match '(\d+\.\d+)') {
+ $script:ROCmVersion = $Matches[1]
+ $script:ROCmVersionFull = $hipVerLine
+ }
+ }
+ } catch {}
+ }
+ if (-not $script:ROCmVersion) {
+ $amdSmiVer = Get-Command "amd-smi" -ErrorAction SilentlyContinue
+ if ($amdSmiVer) {
+ try {
+ $smiVerOut = & $amdSmiVer.Source version 2>&1 | Out-String
+ if ($LASTEXITCODE -eq 0 -and $smiVerOut -match 'ROCm version:\s*(\d+\.\d+)') { $script:ROCmVersion = $Matches[1] }
+ } catch {}
+ }
+ }
+ }
+}
+
+if ($HasNvidiaSmi) {
+ step "gpu" "NVIDIA GPU detected"
+} elseif ($HasROCm) {
+ step "gpu" $ROCmGpuLabel
+ $hipSdkPath = if ($env:HIP_PATH) { $env:HIP_PATH } elseif ($env:ROCM_PATH) { $env:ROCM_PATH } else { "on system PATH" }
+ substep "HIP SDK: $hipSdkPath"
+ if ($script:ROCmVersionFull) { substep "hipconfig: $script:ROCmVersionFull" }
+} elseif ($HipSdkInstalled -and $ROCmGpuLabel) {
+ # HIP SDK is installed but ROCm can't see the device (driver issue, not SDK issue)
+ $sdkVer = if ($script:ROCmVersionFull) { " (HIP $script:ROCmVersionFull)" } else { "" }
Write-Host ""
- step "gpu" "none (chat-only / GGUF)" "Yellow"
- substep "Training and GPU inference require an NVIDIA GPU with drivers installed." "Yellow"
+ step "gpu" "AMD GPU detected -- not ROCm-accessible$sdkVer" "Yellow"
+ substep "Detected: $ROCmGpuLabel" "Yellow"
+ substep "[WARN] HIP SDK is installed but hipinfo reports no ROCm-capable device." "Yellow"
+ substep " This is a driver issue, not an SDK issue." "Yellow"
+ substep " Ensure the ROCm compute driver is installed alongside the display driver:" "Yellow"
+ substep " https://rocm.docs.amd.com/en/latest/deploy/windows/index.html" "Yellow"
+} elseif ($ROCmGpuLabel) {
+ Write-Host ""
+ step "gpu" "AMD GPU detected -- HIP SDK not found" "Yellow"
+ substep "Detected: $ROCmGpuLabel" "Yellow"
+ substep "Install the HIP SDK for ROCm GPU inference:" "Yellow"
+ substep "https://rocm.docs.amd.com/en/latest/deploy/windows/index.html" "Yellow"
Write-Host ""
} else {
- step "gpu" "NVIDIA GPU detected"
+ Write-Host ""
+ step "gpu" "none (chat-only / GGUF)" "Yellow"
+ substep "Training and GPU inference require an NVIDIA or AMD ROCm GPU." "Yellow"
+ Write-Host ""
}
# ============================================
@@ -1102,6 +1317,13 @@ if (-not $CudaArch) {
step "cuda" "skipped (no NVIDIA GPU detected)" "Yellow"
}
+if ($HasROCm) {
+ $rocmVerLabel = if ($script:ROCmVersionFull) { "ROCm $script:ROCmVersionFull" } elseif ($script:ROCmVersion) { "ROCm $script:ROCmVersion" } else { "ROCm (version unknown)" }
+ step "rocm" $rocmVerLabel
+} elseif ($ROCmGpuLabel) {
+ step "rocm" "HIP SDK not found -- GPU-accelerated training unavailable" "Yellow"
+}
+
# ============================================
# 1f. Node.js / npm (skip if pip-installed or Tauri -- only needed for frontend build)
# ============================================
@@ -1516,7 +1738,7 @@ if (-not $PythonCmd) {
exit 1
}
-substep "Using $PythonCmd ($(& $PythonCmd --version 2>&1))"
+substep "Python found: $PythonCmd"
# The venv must already exist (created by install.ps1); this script only
# updates packages. UNSLOTH_STUDIO_HOME (or STUDIO_HOME alias) overrides the
@@ -1684,6 +1906,13 @@ if (-not (Test-Path -LiteralPath $VenvDir)) {
exit 1
} else {
substep "reusing existing virtual environment at $VenvDir"
+ $_venvPyExe = Join-Path $VenvDir "Scripts\python.exe"
+ if (Test-Path -LiteralPath $_venvPyExe) {
+ try {
+ $_venvPyVer = (& $_venvPyExe --version 2>&1 | Out-String).Trim()
+ if ($_venvPyVer) { substep $_venvPyVer }
+ } catch {}
+ }
}
# pip and python write to stderr even on success (progress bars, warnings).
@@ -1795,9 +2024,91 @@ if ($HasNvidiaSmi) {
$CuTag = "cpu"
}
+# ── GPU arch → newest compatible Windows ROCm wheel release ──
+# Wheels bundle their own ROCm runtime; the installed HIP SDK version does
+# not constrain which release to use. Always picks the newest release that
+# supports the GPU architecture.
+# ── AMD Windows ROCm torch override ──────────────────────────────────────────
+# Uses AMD's arch-specific pip index (repo.amd.com/rocm/whl/{arch}/).
+# Wheels bundle their own ROCm runtime; HIP SDK version is irrelevant.
+$ROCmGfxArch = $script:ROCmGfxArch
+$ROCmIndexUrl = $null
+if ($HasROCm -and $CuTag -eq "cpu") {
+ $amdIndexBase = if ($env:UNSLOTH_ROCM_WINDOWS_MIRROR) { $env:UNSLOTH_ROCM_WINDOWS_MIRROR.TrimEnd('/') } else { "https://repo.amd.com/rocm/whl" }
+ $archFamilyMap = @{
+ "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
+ }
+ # gfx120X and Strix have a null _grouped_mm kernel on torch <2.11.0.
+ # Mirrors the $torchFloorMap in install.ps1 so both installers enforce
+ # the same floor and ceiling when pulling from AMD's per-arch index.
+ $torchFloorMap = @{
+ "gfx1201" = "torch>=2.11.0,<2.12.0"; "gfx1200" = "torch>=2.11.0,<2.12.0"
+ "gfx1151" = "torch>=2.11.0,<2.12.0"; "gfx1150" = "torch>=2.11.0,<2.12.0"
+ }
+ # Companion ranges for torchvision/torchaudio -- must stay in sync with the
+ # torch ceiling so pip can always find a consistent trio on AMD's per-arch
+ # index. AMD publishes each package independently and may add a newer
+ # torchvision (e.g. 0.27 for torch 2.12) before removing 0.26, which would
+ # cause pip to resolve an ABI-incompatible set if these are left bare.
+ # Matches _ROCM_TORCH_PKG_SPECS["rocm7.2"] in install_python_stack.py.
+ # Bump all three ceilings together when torch 2.12.x is validated.
+ $torchvisionFloorMap = @{
+ "gfx1201" = "torchvision>=0.26.0,<0.27.0"; "gfx1200" = "torchvision>=0.26.0,<0.27.0"
+ "gfx1151" = "torchvision>=0.26.0,<0.27.0"; "gfx1150" = "torchvision>=0.26.0,<0.27.0"
+ }
+ $torchaudioFloorMap = @{
+ "gfx1201" = "torchaudio>=2.11.0,<2.12.0"; "gfx1200" = "torchaudio>=2.11.0,<2.12.0"
+ "gfx1151" = "torchaudio>=2.11.0,<2.12.0"; "gfx1150" = "torchaudio>=2.11.0,<2.12.0"
+ }
+ $archFamily = if ($ROCmGfxArch -and $archFamilyMap.ContainsKey($ROCmGfxArch)) { $archFamilyMap[$ROCmGfxArch] } else { $null }
+ $ROCmTorchSpec = if ($ROCmGfxArch -and $torchFloorMap.ContainsKey($ROCmGfxArch)) { $torchFloorMap[$ROCmGfxArch] } else { "torch" }
+ $ROCmVisionSpec = if ($ROCmGfxArch -and $torchvisionFloorMap.ContainsKey($ROCmGfxArch)) { $torchvisionFloorMap[$ROCmGfxArch] } else { "torchvision" }
+ $ROCmAudioSpec = if ($ROCmGfxArch -and $torchaudioFloorMap.ContainsKey($ROCmGfxArch)) { $torchaudioFloorMap[$ROCmGfxArch] } else { "torchaudio" }
+ if ($archFamily) {
+ $ROCmIndexUrl = "$amdIndexBase/$archFamily/"
+ } elseif ($ROCmGfxArch) {
+ # GPU arch detected but not in the supported wheel map — warn explicitly
+ # so the user knows why they are getting CPU PyTorch instead of ROCm.
+ substep "[WARN] AMD GPU ($ROCmGfxArch) not in supported arch list -- falling back to CPU-only PyTorch" "Yellow"
+ substep " Supported: gfx1200/1201 (RDNA 4), gfx1150/1151 (RDNA 3.5), gfx1100-1103 (RDNA 3), gfx90a, gfx908" "Yellow"
+ } else {
+ # HIP SDK present ($HasROCm=true via amd-smi) but gcnArchName was not
+ # readable — warn rather than silently falling back to CPU PyTorch.
+ substep "[WARN] AMD GPU detected (HIP SDK present) but GPU arch could not be read -- falling back to CPU-only PyTorch" "Yellow"
+ substep " Arch detection requires hipinfo to report gcnArchName. Re-install the HIP SDK if this is unexpected." "Yellow"
+ }
+}
+
$PyTorchWhlBase = if ($env:UNSLOTH_PYTORCH_MIRROR) { $env:UNSLOTH_PYTORCH_MIRROR.TrimEnd('/') } else { "https://download.pytorch.org/whl" }
-if ($CuTag -eq "cpu") {
+if ($ROCmIndexUrl) {
+ substep "installing PyTorch (AMD ROCm, $ROCmGfxArch)..."
+ if ($ROCmTorchSpec -ne "torch") {
+ substep " enforcing $ROCmTorchSpec $ROCmVisionSpec $ROCmAudioSpec (known _grouped_mm bug in older wheels)" "Cyan"
+ }
+ if ($script:UnslothVerbose) {
+ Fast-Install $ROCmTorchSpec $ROCmVisionSpec $ROCmAudioSpec --force-reinstall --index-url $ROCmIndexUrl
+ $torchInstallExit = $LASTEXITCODE
+ $output = ""
+ } else {
+ $output = Fast-Install $ROCmTorchSpec $ROCmVisionSpec $ROCmAudioSpec --force-reinstall --index-url $ROCmIndexUrl | Out-String
+ $torchInstallExit = $LASTEXITCODE
+ }
+ if ($torchInstallExit -ne 0) {
+ Write-Host "[WARN] AMD ROCm PyTorch install failed -- falling back to CPU" -ForegroundColor Yellow
+ Write-Host $output -ForegroundColor Yellow
+ $ROCmIndexUrl = $null
+ } else {
+ # Tell install_python_stack.py to skip probe + suppress manual-install warning.
+ $env:UNSLOTH_ROCM_TORCH_INSTALLED = "1"
+ }
+}
+
+if (-not $ROCmIndexUrl -and $CuTag -eq "cpu") {
substep "installing PyTorch (CPU-only)..."
if ($script:UnslothVerbose) {
Fast-Install torch torchvision torchaudio --index-url "$PyTorchWhlBase/cpu"
@@ -1812,7 +2123,7 @@ if ($CuTag -eq "cpu") {
Write-Host $output -ForegroundColor Red
exit 1
}
-} else {
+} elseif (-not $ROCmIndexUrl) {
substep "installing PyTorch with CUDA support ($CuTag)..."
substep "(This download is ~2.8 GB -- may take a few minutes)"
if ($script:UnslothVerbose) {
@@ -2115,6 +2426,23 @@ if ($env:UNSLOTH_LLAMA_FORCE_COMPILE -eq "1") {
substep "installing prebuilt llama.cpp bundle (preferred path)..."
if (Test-Path -LiteralPath $LlamaCppDir) {
substep "Existing llama.cpp install detected -- validating staged prebuilt update before replacement"
+ # If the existing install is the wrong kind (e.g. windows-cpu on a ROCm
+ # machine that should have windows-hip), remove it so the installer is
+ # forced to download the correct variant rather than skipping on tag match.
+ $existingMetaPath = Join-Path $LlamaCppDir "UNSLOTH_PREBUILT_INFO.json"
+ if (Test-Path $existingMetaPath) {
+ try {
+ $existingMeta = Get-Content $existingMetaPath -Raw | ConvertFrom-Json
+ $existingKind = $existingMeta.install_kind
+ $expectedKind = if ($HasROCm) { "windows-hip" } elseif ($HasNvidiaSmi) { "windows-cuda" } else { "windows-cpu" }
+ if ($existingKind -and $existingKind -ne $expectedKind) {
+ substep "Removing mismatched llama.cpp install (found '$existingKind', need '$expectedKind')..."
+ Remove-Item -Recurse -Force -LiteralPath $LlamaCppDir -ErrorAction SilentlyContinue
+ }
+ } catch {
+ # unreadable metadata -- let the installer handle it
+ }
+ }
}
# why: install_llama_prebuilt.py uses os.replace(), which would displace
# an unrelated $env:UNSLOTH_STUDIO_HOME\llama.cpp before the source-build
@@ -2129,6 +2457,9 @@ if ($env:UNSLOTH_LLAMA_FORCE_COMPILE -eq "1") {
"--published-repo", $HelperReleaseRepo,
"--simple-policy"
)
+ if ($HasROCm) {
+ $prebuiltArgs += "--has-rocm"
+ }
if ($env:UNSLOTH_LLAMA_RELEASE_TAG) {
$prebuiltArgs += @("--published-release-tag", $env:UNSLOTH_LLAMA_RELEASE_TAG)
}
diff --git a/studio/setup.sh b/studio/setup.sh
index f2395255c0..3a149d465f 100755
--- a/studio/setup.sh
+++ b/studio/setup.sh
@@ -181,12 +181,21 @@ if not isinstance(payload, dict):
repo = str(payload.get("published_repo") or "").strip()
release_tag = str(payload.get("release_tag") or "").strip()
llama_tag = str(payload.get("tag") or "").strip()
+source = str(payload.get("source") or "").strip()
+binary_repo = str(payload.get("binary_repo") or "").strip()
+binary_tag = str(payload.get("binary_release_tag") or "").strip()
if not repo or not release_tag:
raise SystemExit(0)
-message = f"installed release: {repo}@{release_tag}"
-if llama_tag and llama_tag != release_tag:
- message += f" (tag {llama_tag})"
+# For non-upstream sources (e.g. lemonade) the published_repo/release_tag
+# refer to the unsloth source tree while the actual binaries came from a
+# different repo. Show both so the log is unambiguous.
+if source and source != "upstream" and binary_repo and binary_tag and binary_repo != repo:
+ message = f"installed release: {repo}@{release_tag} + {source}@{binary_tag}"
+else:
+ message = f"installed release: {repo}@{release_tag}"
+ if llama_tag and llama_tag != release_tag:
+ message += f" (tag {llama_tag})"
print(message)
PY
}
@@ -659,6 +668,81 @@ if [ "$_NEED_T5_INSTALL" = true ]; then
fi
fi
+# ── GPU detection summary (mirrors setup.ps1 step "gpu" block) ──
+_setup_amd_detected=false
+_setup_gfx_all=""
+_setup_mkt=""
+if command -v rocminfo >/dev/null 2>&1 && \
+ rocminfo 2>/dev/null | awk '/Name:[[:space:]]*gfx[1-9][0-9]/{found=1} END{exit !found}'; then
+ _setup_amd_detected=true
+ _setup_gfx_all=$(rocminfo 2>/dev/null | grep -oE 'gfx[1-9][0-9a-z]{2,3}' || true)
+ _setup_mkt=$(rocminfo 2>/dev/null | awk -F': ' \
+ '/Marketing Name:/{gsub(/^[[:space:]]+|[[:space:]]+$/,"", $2); if($2){print $2; exit}}' || true)
+elif command -v amd-smi >/dev/null 2>&1 && \
+ amd-smi list 2>/dev/null | awk '/^GPU[[:space:]]*[:\[][[:space:]]*[0-9]/{ found=1 } END{ exit !found }'; then
+ _setup_amd_detected=true
+ _setup_gfx_all=$(amd-smi list 2>/dev/null | grep -oE 'gfx[1-9][0-9a-z]{2,3}' || true)
+ [ -z "$_setup_gfx_all" ] && \
+ _setup_gfx_all=$(amd-smi static --asic 2>/dev/null | grep -oE 'gfx[1-9][0-9a-z]{2,3}' || true)
+ _setup_mkt=$(amd-smi static --asic 2>/dev/null | awk -F'[:|]' \
+ '/[Mm]arket.?[Nn]ame/{gsub(/^[[:space:]]+|[[:space:]]+$/,"", $2); if($2){print $2; exit}}' || true)
+fi
+
+if command -v nvidia-smi >/dev/null 2>&1 && \
+ nvidia-smi -L 2>/dev/null | awk '/^GPU[[:space:]]+[0-9]+:/{found=1} END{exit !found}'; then
+ step "gpu" "NVIDIA GPU detected"
+elif [ "$_setup_amd_detected" = true ]; then
+ _setup_vis="${HIP_VISIBLE_DEVICES:-${ROCR_VISIBLE_DEVICES:-}}"
+ _setup_vis_idx=0
+ if [ -n "$_setup_vis" ] && [ "$_setup_vis" != "-1" ]; then
+ _setup_first="${_setup_vis%%,*}"
+ case "$_setup_first" in ''|*[!0-9]*) ;; *) _setup_vis_idx=$_setup_first ;; esac
+ fi
+ _setup_gfx=$(printf '%s\n' "$_setup_gfx_all" | awk -v idx="$_setup_vis_idx" \
+ 'NF && !seen[$0]++ { a[n++]=$0 } END { if(idx>=n) idx=0; if(n>0) print a[idx] }')
+ # UNSLOTH_ROCM_GFX_ARCH env override (mirrors setup.ps1)
+ if [ -n "${UNSLOTH_ROCM_GFX_ARCH:-}" ]; then
+ _setup_gfx="${UNSLOTH_ROCM_GFX_ARCH}"
+ substep "gfx arch from UNSLOTH_ROCM_GFX_ARCH env override: $_setup_gfx"
+ # Name-based arch inference when tools don't report gfx (mirrors setup.ps1 nameArchTable)
+ elif [ -z "$_setup_gfx" ] && [ -n "$_setup_mkt" ]; then
+ case "$_setup_mkt" in
+ *"9070 XT"*|*9080*) _setup_gfx="gfx1201" ;; # RDNA 4
+ *9070*|*9060*) _setup_gfx="gfx1200" ;; # RDNA 4
+ *"8060S"*|*"890M"*|*"Strix Halo"*|*"HX 37"*|*"HX 38"*|*"AI 9 HX"*) _setup_gfx="gfx1151" ;; # RDNA 3.5 iGPU
+ *"880M"*|*"Strix Point"*|*"AI 9 36"*|*"AI 7 35"*|*"AI 5 34"*) _setup_gfx="gfx1150" ;; # RDNA 3.5 iGPU
+ *"RX 7900"*|*"RX 7800"*|*"RX 7700"*) _setup_gfx="gfx1100" ;; # RDNA 3 desktop
+ *"RX 7600"*) _setup_gfx="gfx1102" ;; # RDNA 3
+ *"780M"*|*"760M"*|*"740M"*|*"Phoenix"*) _setup_gfx="gfx1103" ;; # RDNA 3 iGPU
+ esac
+ if [ -n "$_setup_gfx" ]; then
+ substep "gfx arch inferred from GPU name: $_setup_gfx"
+ substep "Tip: set UNSLOTH_ROCM_GFX_ARCH=$_setup_gfx to skip inference next time"
+ fi
+ fi
+ # ROCm version via hipconfig, then amd-smi
+ _setup_rocm_ver=""
+ if command -v hipconfig >/dev/null 2>&1; then
+ _setup_rocm_ver=$(hipconfig --version 2>/dev/null | awk 'NR==1 && /^[0-9]/{print; exit}' || true)
+ fi
+ if [ -z "$_setup_rocm_ver" ] && command -v amd-smi >/dev/null 2>&1; then
+ _setup_rocm_ver=$(amd-smi version 2>/dev/null | awk -F'ROCm version: ' \
+ 'NF>1{gsub(/[[:space:]]/,"", $2); print $2; exit}' || true)
+ fi
+ if [ -n "$_setup_gfx" ]; then
+ step "gpu" "AMD ROCm ($_setup_gfx)"
+ else
+ step "gpu" "AMD ROCm"
+ fi
+ _setup_rocm_root="${ROCM_PATH:-${HIP_PATH:-/opt/rocm}}"
+ substep "ROCm: $_setup_rocm_root"
+ [ -n "$_setup_rocm_ver" ] && substep "hipconfig: $_setup_rocm_ver"
+ [ -n "$_setup_mkt" ] && [ -n "$_setup_gfx" ] && substep "GPU: $_setup_mkt"
+else
+ step "gpu" "none (chat-only / GGUF)" "$C_WARN"
+ substep "Training and GPU inference require an NVIDIA or AMD ROCm GPU."
+fi
+
# ── 7. Prefer prebuilt llama.cpp bundles before any source build path ──
# Nest llama.cpp under $STUDIO_HOME only for real env-overrides; legacy
# default keeps ~/.unsloth/llama.cpp so pre-PR builds are still discovered.
@@ -817,6 +901,22 @@ else
fi
fi
+# Source-built llama.cpp installs do not have the prebuilt metadata used above
+# for exact release matching. Reuse a complete local source build unless the
+# caller explicitly requested a rebuild or a PR-specific llama.cpp checkout.
+if [ "$_NEED_LLAMA_SOURCE_BUILD" = true ] && \
+ [ "$_LLAMA_FORCE_COMPILE" != "1" ] && \
+ [ -z "$_LLAMA_PR" ] && \
+ [ -x "$LLAMA_CPP_DIR/build/bin/llama-server" ] && \
+ [ -x "$LLAMA_CPP_DIR/build/bin/llama-quantize" ]; then
+ step "llama.cpp" "existing source build found; skipping rebuild"
+ ln -sf build/bin/llama-quantize "$LLAMA_CPP_DIR/llama-quantize"
+ if [ "$_STUDIO_HOME_IS_CUSTOM" = true ]; then
+ : > "$LLAMA_CPP_DIR/$_STUDIO_OWNED_MARKER" 2>/dev/null || true
+ fi
+ _NEED_LLAMA_SOURCE_BUILD=false
+fi
+
# ── 8. WSL: pre-install GGUF build dependencies for fallback source builds ──
# On WSL, sudo requires a password and can't be entered during GGUF export
# (runs in a non-interactive subprocess). Install build deps here instead.
@@ -974,7 +1074,8 @@ else
fi
if [ "$BUILD_OK" = true ]; then
- CMAKE_ARGS="-DLLAMA_BUILD_TESTS=OFF -DLLAMA_BUILD_EXAMPLES=OFF -DLLAMA_BUILD_SERVER=ON -DGGML_NATIVE=ON"
+ # Set Release explicitly (llama.cpp only defaults to it on non-MSVC/Xcode).
+ CMAKE_ARGS="-DCMAKE_BUILD_TYPE=Release -DLLAMA_BUILD_TESTS=OFF -DLLAMA_BUILD_EXAMPLES=OFF -DLLAMA_BUILD_SERVER=ON -DGGML_NATIVE=ON"
_TRY_METAL_CPU_FALLBACK=false
_HOST_SYSTEM="$(uname -s 2>/dev/null || true)"
_HOST_MACHINE="$(uname -m 2>/dev/null || true)"
@@ -983,6 +1084,15 @@ else
_IS_MACOS_ARM64=true
fi
+ # macOS: pin a low deployment target so the source build loads on
+ # older macOS too (else a macOS 26 host stamps minos=26). Set before
+ # CPU_FALLBACK_CMAKE_ARGS copies CMAKE_ARGS so both paths inherit it.
+ if [ "$_HOST_SYSTEM" = "Darwin" ]; then
+ _MACOS_DEPLOYMENT_TARGET="${UNSLOTH_MACOS_DEPLOYMENT_TARGET:-13.3}"
+ CMAKE_ARGS="$CMAKE_ARGS -DCMAKE_OSX_DEPLOYMENT_TARGET=${_MACOS_DEPLOYMENT_TARGET}"
+ export MACOSX_DEPLOYMENT_TARGET="${_MACOS_DEPLOYMENT_TARGET}"
+ fi
+
if command -v ccache &>/dev/null; then
CMAKE_ARGS="$CMAKE_ARGS -DCMAKE_C_COMPILER_LAUNCHER=ccache -DCMAKE_CXX_COMPILER_LAUNCHER=ccache -DCMAKE_CUDA_COMPILER_LAUNCHER=ccache"
fi
@@ -1088,6 +1198,32 @@ else
_BUILD_DESC="building (ROCm)"
CMAKE_ARGS="$CMAKE_ARGS -DGGML_HIP=ON"
+
+ # ROCm 7.x ships clang-20 which on Ubuntu 24.04+ defaults to the
+ # highest-numbered gcc lib dir (/usr/lib/gcc/x86_64-linux-gnu/14/)
+ # which contains runtime objects but NOT C++ headers, causing:
+ # fatal error: 'cstdlib' file not found
+ # Find the newest gcc install dir that actually has both the
+ # runtime dir AND /usr/include/c++/ headers, then pass it
+ # to clang via --gcc-install-dir so HIP builds succeed.
+ _GCC_INSTALL_DIR=""
+ _gcc_pm="$(gcc -print-multiarch 2>/dev/null)"
+ case "$_gcc_pm" in
+ *-linux-gnu*) _GCC_MULTIARCH="$_gcc_pm" ;;
+ *) _GCC_MULTIARCH="$(uname -m)-linux-gnu" ;;
+ esac
+ for _gcc_ver in 14 13 12 11; do
+ if [ -d "/usr/lib/gcc/$_GCC_MULTIARCH/$_gcc_ver/include" ] && \
+ [ -d "/usr/include/c++/$_gcc_ver" ]; then
+ _GCC_INSTALL_DIR="/usr/lib/gcc/$_GCC_MULTIARCH/$_gcc_ver"
+ break
+ fi
+ done
+ if [ -n "$_GCC_INSTALL_DIR" ]; then
+ CMAKE_ARGS="$CMAKE_ARGS -DCMAKE_HIP_FLAGS=--gcc-install-dir=\"$_GCC_INSTALL_DIR\""
+ substep "ROCm HIP gcc install dir: $_GCC_INSTALL_DIR"
+ fi
+
export ROCM_PATH="$ROCM_ROOT"
export HIP_PATH="$ROCM_ROOT"
diff --git a/studio/src-tauri/Cargo.lock b/studio/src-tauri/Cargo.lock
index c73d7cc0c5..df350e4011 100644
--- a/studio/src-tauri/Cargo.lock
+++ b/studio/src-tauri/Cargo.lock
@@ -2731,15 +2731,14 @@ dependencies = [
[[package]]
name = "openssl"
-version = "0.10.76"
+version = "0.10.80"
source = "registry+https://github.com/rust-lang/crates.io-index"
-checksum = "951c002c75e16ea2c65b8c7e4d3d51d5530d8dfa7d060b4776828c88cfb18ecf"
+checksum = "a45fa2aa886c42762255da344f0a0d313e254066c46aad76f300c3d3da62d967"
dependencies = [
"bitflags 2.11.0",
"cfg-if",
"foreign-types 0.3.2",
"libc",
- "once_cell",
"openssl-macros",
"openssl-sys",
]
@@ -2763,9 +2762,9 @@ checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe"
[[package]]
name = "openssl-sys"
-version = "0.9.112"
+version = "0.9.116"
source = "registry+https://github.com/rust-lang/crates.io-index"
-checksum = "57d55af3b3e226502be1526dfdba67ab0e9c96fc293004e79576b2b9edb0dbdb"
+checksum = "f28a22dc7140cda5f096e5e7724a6962ca81a7f8bfd2979f9b18c11af56318c4"
dependencies = [
"cc",
"libc",
@@ -2978,7 +2977,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5d5285893bb5eb82e6aaf5d59ee909a06a16737a8970984dd7746ba9283498d6"
dependencies = [
"phf_shared 0.10.0",
- "rand 0.8.5",
+ "rand 0.8.6",
]
[[package]]
@@ -2988,7 +2987,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3c80231409c20246a13fddb31776fb942c38553c51e871f8cbd687a4cfb5843d"
dependencies = [
"phf_shared 0.11.3",
- "rand 0.8.5",
+ "rand 0.8.6",
]
[[package]]
@@ -3361,9 +3360,9 @@ dependencies = [
[[package]]
name = "rand"
-version = "0.8.5"
+version = "0.8.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
-checksum = "34af8d1a0e25924bc5b7c43c079c942339d8f0a8b57c39049bef581b46327404"
+checksum = "5ca0ecfa931c29007047d1bc58e623ab12e5590e8c7cc53200d5202b69266d8a"
dependencies = [
"libc",
"rand_chacha 0.3.1",
@@ -3382,9 +3381,9 @@ dependencies = [
[[package]]
name = "rand"
-version = "0.10.0"
+version = "0.10.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
-checksum = "bc266eb313df6c5c09c1c7b1fbe2510961e5bcd3add930c1e31f7ed9da0feff8"
+checksum = "d2e8e8bcc7961af1fdac401278c6a831614941f6164ee3bf4ce61b7edb162207"
dependencies = [
"chacha20",
"getrandom 0.4.2",
@@ -3765,9 +3764,9 @@ checksum = "f87165f0995f63a9fbeea62b64d10b4d9d8e78ec6d7d51fb2125fda7bb36788f"
[[package]]
name = "rustls-webpki"
-version = "0.103.10"
+version = "0.103.13"
source = "registry+https://github.com/rust-lang/crates.io-index"
-checksum = "df33b2b81ac578cabaf06b89b0631153a3f416b0a886e8a7a1707fb51abbd1ef"
+checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e"
dependencies = [
"ring",
"rustls-pki-types",
@@ -4450,9 +4449,9 @@ dependencies = [
[[package]]
name = "tar"
-version = "0.4.45"
+version = "0.4.46"
source = "registry+https://github.com/rust-lang/crates.io-index"
-checksum = "22692a6476a21fa75fdfc11d452fda482af402c008cdbaf3476414e122040973"
+checksum = "3f6221d9a6003c78398e3b239969f352578258df48c8eb051caadae0015bc840"
dependencies = [
"filetime",
"libc",
@@ -5380,7 +5379,7 @@ dependencies = [
"log",
"open",
"process-wrap",
- "rand 0.10.0",
+ "rand 0.10.1",
"regex",
"reqwest 0.12.28",
"serde",
diff --git a/tests/sh/test_get_torch_index_url.sh b/tests/sh/test_get_torch_index_url.sh
index 89ec32fba5..e20fd1ca86 100755
--- a/tests/sh/test_get_torch_index_url.sh
+++ b/tests/sh/test_get_torch_index_url.sh
@@ -193,10 +193,10 @@ _result=$(run_func "$_dir")
assert_eq "ROCm 7.1 -> rocm7.1" "https://download.pytorch.org/whl/rocm7.1" "$_result"
rm -rf "$_dir"
-# 11) ROCm 7.2 (no nvidia-smi) -> rocm7.1 (capped due to torch <2.11.0)
+# 11) ROCm 7.2 (no nvidia-smi) -> rocm7.2
_dir=$(make_mock_amd_smi "7.2")
_result=$(run_func "$_dir")
-assert_eq "ROCm 7.2 -> rocm7.1 (capped)" "https://download.pytorch.org/whl/rocm7.1" "$_result"
+assert_eq "ROCm 7.2 -> rocm7.2" "https://download.pytorch.org/whl/rocm7.2" "$_result"
rm -rf "$_dir"
# 12) Both nvidia-smi and amd-smi present -> CUDA takes precedence
@@ -231,10 +231,10 @@ _result=$(run_func "$_dir")
assert_eq "ROCm 7.0 -> rocm7.0" "https://download.pytorch.org/whl/rocm7.0" "$_result"
rm -rf "$_dir"
-# 17) ROCm 8.0 (future, no nvidia-smi) -> rocm7.1 (capped)
+# 17) ROCm 8.0 (future, no nvidia-smi) -> rocm7.2 (capped to latest known)
_dir=$(make_mock_amd_smi "8.0")
_result=$(run_func "$_dir")
-assert_eq "ROCm 8.0 -> rocm7.1 (capped)" "https://download.pytorch.org/whl/rocm7.1" "$_result"
+assert_eq "ROCm 8.0 -> rocm7.2 (capped)" "https://download.pytorch.org/whl/rocm7.2" "$_result"
rm -rf "$_dir"
# 18) Malformed amd-smi output (empty version field) -> cpu
diff --git a/tests/studio/install/conftest.py b/tests/studio/install/conftest.py
new file mode 100644
index 0000000000..8738ef2319
--- /dev/null
+++ b/tests/studio/install/conftest.py
@@ -0,0 +1,20 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
+
+"""Pytest configuration for studio/install tests.
+
+install_python_stack.py does ``from backend.utils.wheel_utils import ...``
+which requires the ``studio/`` directory to be on sys.path. When tests are
+run from the repo root (the normal case), the studio package is not
+automatically importable, so we add it here.
+"""
+
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+
+# /studio → makes `backend` importable as a package
+_STUDIO_DIR = Path(__file__).resolve().parents[3] / "studio"
+if str(_STUDIO_DIR) not in sys.path:
+ sys.path.insert(0, str(_STUDIO_DIR))
diff --git a/tests/studio/install/test_macos_version_compat.py b/tests/studio/install/test_macos_version_compat.py
new file mode 100644
index 0000000000..0adc91a622
--- /dev/null
+++ b/tests/studio/install/test_macos_version_compat.py
@@ -0,0 +1,322 @@
+"""Tests for the host-macOS-version-aware llama.cpp prebuilt selection added
+for the Mac "Failing CI" fix.
+
+Covers: parse_macos_version, host_supports_macos_minos, the pure-Python Mach-O
+minimum-OS parser (macho_minimum_macos), the dyld-incompatibility classifier,
+the install preflight that rejects a too-new prebuilt, and the deeper macOS
+release walk-back in resolve_simple_install_release_plans.
+
+No GPU, no network, no torch, no real Mach-O toolchain required -- the Mach-O
+samples are synthesized in-process and all I/O is monkeypatched.
+"""
+
+import importlib.util
+import struct
+import sys
+from pathlib import Path
+
+import pytest
+
+
+PACKAGE_ROOT = Path(__file__).resolve().parents[3]
+MODULE_PATH = PACKAGE_ROOT / "studio" / "install_llama_prebuilt.py"
+SPEC = importlib.util.spec_from_file_location(
+ "studio_install_llama_prebuilt_macos", MODULE_PATH
+)
+assert SPEC is not None and SPEC.loader is not None
+ILP = importlib.util.module_from_spec(SPEC)
+sys.modules[SPEC.name] = ILP
+SPEC.loader.exec_module(ILP)
+
+HostInfo = ILP.HostInfo
+PrebuiltFallback = ILP.PrebuiltFallback
+
+_CPU_TYPE_ARM64 = 0x0100000C
+_CPU_TYPE_X86_64 = 0x01000007
+
+
+def make_macos_host(macos_version, *, arm64 = True):
+ return HostInfo(
+ system = "Darwin",
+ machine = "arm64" if arm64 else "x86_64",
+ is_windows = False,
+ is_linux = False,
+ is_macos = True,
+ is_x86_64 = not arm64,
+ is_arm64 = arm64,
+ nvidia_smi = None,
+ driver_cuda_version = None,
+ compute_caps = [],
+ visible_cuda_devices = None,
+ has_physical_nvidia = False,
+ has_usable_nvidia = False,
+ macos_version = macos_version,
+ )
+
+
+def thin_macho(minos = (14, 0), *, cputype = _CPU_TYPE_ARM64, build_version = True):
+ """Synthesize a minimal little-endian 64-bit Mach-O carrying a macOS
+ minimum-version load command."""
+ encoded = (minos[0] << 16) | (minos[1] << 8)
+ if build_version:
+ # LC_BUILD_VERSION: cmd, cmdsize, platform(=1 macOS), minos, sdk, ntools
+ load_command = struct.pack("<6I", 0x32, 24, 1, encoded, encoded, 0)
+ else:
+ # LC_VERSION_MIN_MACOSX: cmd, cmdsize, version, sdk
+ load_command = struct.pack("<4I", 0x24, 16, encoded, encoded)
+ header = struct.pack("<8I", 0xFEEDFACF, cputype, 0, 0x2, 1, len(load_command), 0, 0)
+ return header + load_command
+
+
+def fat_macho(slices):
+ """Synthesize a big-endian universal binary from (cputype, thin_bytes)."""
+ header = struct.pack(">2I", 0xCAFEBABE, len(slices))
+ data_offset = 8 + 20 * len(slices)
+ arch_entries = b""
+ body = b""
+ for cputype, thin in slices:
+ offset = data_offset + len(body)
+ arch_entries += struct.pack(">5I", cputype, 0, offset, len(thin), 0)
+ body += thin
+ return header + arch_entries + body
+
+
+class TestParseMacosVersion:
+ @pytest.mark.parametrize(
+ "value, expected",
+ [
+ ("14.7.1", (14, 7)),
+ ("15.5", (15, 5)),
+ ("26.0", (26, 0)),
+ ("26", (26, 0)),
+ ("13", (13, 0)),
+ ("", None),
+ (None, None),
+ ("not-a-version", None),
+ ],
+ )
+ def test_parse(self, value, expected):
+ assert ILP.parse_macos_version(value) == expected
+
+
+class TestHostSupportsMacosMinos:
+ def test_older_host_rejects_newer_prebuilt(self):
+ assert not ILP.host_supports_macos_minos(make_macos_host((14, 0)), (26, 0))
+
+ def test_same_version_supported(self):
+ assert ILP.host_supports_macos_minos(make_macos_host((26, 0)), (26, 0))
+
+ def test_newer_host_supports_older_prebuilt(self):
+ assert ILP.host_supports_macos_minos(make_macos_host((15, 5)), (14, 0))
+
+ def test_unknown_host_defers_to_runtime(self):
+ assert ILP.host_supports_macos_minos(make_macos_host(None), (26, 0))
+
+ def test_unknown_minos_defers_to_runtime(self):
+ assert ILP.host_supports_macos_minos(make_macos_host((14, 0)), None)
+
+
+class TestMachoMinimumMacos:
+ def test_build_version_thin(self, tmp_path):
+ path = tmp_path / "lib.dylib"
+ path.write_bytes(thin_macho((26, 0)))
+ assert ILP.macho_minimum_macos(path) == (26, 0)
+
+ def test_legacy_version_min_thin(self, tmp_path):
+ path = tmp_path / "lib.dylib"
+ path.write_bytes(thin_macho((14, 0), build_version = False))
+ assert ILP.macho_minimum_macos(path) == (14, 0)
+
+ def test_universal_prefers_host_arch_slice(self, tmp_path):
+ # arm64 slice needs macOS 14, x86_64 slice needs macOS 26.
+ path = tmp_path / "fat"
+ path.write_bytes(
+ fat_macho(
+ [
+ (_CPU_TYPE_ARM64, thin_macho((14, 0), cputype = _CPU_TYPE_ARM64)),
+ (_CPU_TYPE_X86_64, thin_macho((26, 0), cputype = _CPU_TYPE_X86_64)),
+ ]
+ )
+ )
+ assert ILP.macho_minimum_macos(path, make_macos_host((14, 0))) == (14, 0)
+ assert ILP.macho_minimum_macos(path, make_macos_host((26, 0), arm64 = False)) == (
+ 26,
+ 0,
+ )
+
+ def test_non_macho_returns_none(self, tmp_path):
+ path = tmp_path / "script.sh"
+ path.write_bytes(b'#!/bin/sh\nexec real "$@"\n')
+ assert ILP.macho_minimum_macos(path) is None
+
+ def test_missing_file_returns_none(self, tmp_path):
+ assert ILP.macho_minimum_macos(tmp_path / "nope") is None
+
+
+class TestLooksLikeMacosIncompatibility:
+ def test_built_for_newer_os(self):
+ assert ILP.looks_like_macos_incompatibility(
+ "dyld: ... (built for macOS 26.0 which is newer than running OS)"
+ )
+
+ def test_metal_residency_symbol(self):
+ assert ILP.looks_like_macos_incompatibility(
+ "Symbol not found: _OBJC_CLASS_$_MTLResidencySetDescriptor"
+ )
+
+ def test_benign_error(self):
+ assert not ILP.looks_like_macos_incompatibility("some unrelated failure")
+
+ def test_empty(self):
+ assert not ILP.looks_like_macos_incompatibility("")
+
+
+class TestPreflightMacosInstalledBinaries:
+ def _install_dir(self, tmp_path, dylib_minos):
+ bin_dir = tmp_path / "build" / "bin"
+ bin_dir.mkdir(parents = True)
+ (bin_dir / "libggml-metal.dylib").write_bytes(thin_macho(dylib_minos))
+ server = tmp_path / "llama-server"
+ server.write_bytes(thin_macho(dylib_minos))
+ quantize = tmp_path / "llama-quantize"
+ quantize.write_bytes(thin_macho(dylib_minos))
+ return tmp_path, (server, quantize)
+
+ def test_rejects_too_new_dylib(self, tmp_path):
+ install_dir, binaries = self._install_dir(tmp_path, (26, 0))
+ with pytest.raises(PrebuiltFallback, match = "newer macOS"):
+ ILP.preflight_macos_installed_binaries(
+ binaries, install_dir, make_macos_host((14, 0))
+ )
+
+ def test_accepts_compatible_prebuilt(self, tmp_path):
+ install_dir, binaries = self._install_dir(tmp_path, (14, 0))
+ # Must not raise on a macOS 15 host.
+ ILP.preflight_macos_installed_binaries(
+ binaries, install_dir, make_macos_host((15, 5))
+ )
+
+ def test_skips_when_host_version_unknown(self, tmp_path):
+ install_dir, binaries = self._install_dir(tmp_path, (26, 0))
+ # Unknown host version -> defer to runtime validation, do not raise.
+ ILP.preflight_macos_installed_binaries(
+ binaries, install_dir, make_macos_host(None)
+ )
+
+ def test_noop_on_non_macos_host(self, tmp_path):
+ install_dir, binaries = self._install_dir(tmp_path, (26, 0))
+ linux_host = HostInfo(
+ system = "Linux",
+ machine = "x86_64",
+ is_windows = False,
+ is_linux = True,
+ is_macos = False,
+ is_x86_64 = True,
+ is_arm64 = False,
+ nvidia_smi = None,
+ driver_cuda_version = None,
+ compute_caps = [],
+ visible_cuda_devices = None,
+ has_physical_nvidia = False,
+ has_usable_nvidia = False,
+ )
+ ILP.preflight_macos_installed_binaries(binaries, install_dir, linux_host)
+
+
+def _fake_macos_releases(tags):
+ return [
+ {
+ "tag_name": tag,
+ "assets": [
+ {
+ "name": f"llama-{tag}-bin-macos-arm64.tar.gz",
+ "browser_download_url": f"https://example.com/{tag}.tar.gz",
+ }
+ ],
+ }
+ for tag in tags
+ ]
+
+
+class TestMacosReleaseWalkback:
+ """A known-version macOS host must generate enough older-release plans to
+ walk back past a run of too-new prebuilts; unknown-version and non-macOS
+ hosts keep the conservative 2-release default."""
+
+ TAGS = [f"b{n}" for n in range(9437, 9400, -1)] # 37 newest-first releases
+
+ def _patch_releases(self, monkeypatch):
+ monkeypatch.setattr(
+ ILP,
+ "iter_release_payloads_by_time",
+ lambda repo, published_release_tag, requested_tag: _fake_macos_releases(
+ self.TAGS
+ ),
+ )
+
+ def test_known_macos_host_walks_back_deeper(self, monkeypatch):
+ self._patch_releases(monkeypatch)
+ _tag, plans = ILP.resolve_simple_install_release_plans(
+ "latest",
+ make_macos_host((14, 0)),
+ "ggml-org/llama.cpp",
+ "",
+ )
+ assert len(plans) == ILP.DEFAULT_MAX_MACOS_RELEASE_FALLBACKS
+ assert len(plans) > ILP.DEFAULT_MAX_PREBUILT_RELEASE_FALLBACKS
+
+ def test_unknown_macos_host_uses_default(self, monkeypatch):
+ self._patch_releases(monkeypatch)
+ _tag, plans = ILP.resolve_simple_install_release_plans(
+ "latest",
+ make_macos_host(None),
+ "ggml-org/llama.cpp",
+ "",
+ )
+ assert len(plans) == ILP.DEFAULT_MAX_PREBUILT_RELEASE_FALLBACKS
+
+
+class TestForwardsBackwardsCompat:
+ """The gate is host >= prebuilt minos with no hardcoded version, so it holds
+ for older and future macOS alike. Emulate the walk-back over a release set
+ spanning several minos tiers and assert each host takes the newest release
+ it can load."""
+
+ # Newest first: future 27 builds, current 26 builds, an old 14 tier, a 13.
+ RELEASES = [
+ ("b9600", (27, 0)),
+ ("b9450", (26, 0)),
+ ("b9415", (14, 0)),
+ ("b8300", (13, 0)),
+ ]
+
+ def _select(self, tmp_path, host_version):
+ for tag, minos in self.RELEASES:
+ bin_dir = tmp_path / tag / "build" / "bin"
+ bin_dir.mkdir(parents = True)
+ (bin_dir / "libggml-metal.dylib").write_bytes(thin_macho(minos))
+ try:
+ ILP.preflight_macos_installed_binaries(
+ (), tmp_path / tag, make_macos_host(host_version)
+ )
+ return tag
+ except PrebuiltFallback:
+ continue
+ return None
+
+ @pytest.mark.parametrize(
+ "host_version, expected",
+ [
+ ((13, 0), "b8300"), # older host takes the older prebuilt
+ ((14, 7), "b9415"), # backwards: skip 26/27, take newest that loads
+ ((15, 5), "b9415"),
+ ((26, 0), "b9450"), # unchanged: newest <= host
+ ((27, 1), "b9600"), # forwards: future host takes the future build
+ ],
+ )
+ def test_selects_newest_loadable(self, tmp_path, host_version, expected):
+ assert self._select(tmp_path, host_version) == expected
+
+ def test_host_below_prebuilt_floor_falls_through(self, tmp_path):
+ # macOS 12 is below every prebuilt -> nothing matches -> source build.
+ assert self._select(tmp_path, (12, 0)) is None
diff --git a/tests/studio/install/test_rocm_support.py b/tests/studio/install/test_rocm_support.py
index e6f1ae1c65..ec99e9a8d9 100644
--- a/tests/studio/install/test_rocm_support.py
+++ b/tests/studio/install/test_rocm_support.py
@@ -51,6 +51,9 @@ _ensure_rocm_torch = stack_mod._ensure_rocm_torch
_has_rocm_gpu = stack_mod._has_rocm_gpu
_has_usable_nvidia_gpu = stack_mod._has_usable_nvidia_gpu
_ROCM_TORCH_INDEX = stack_mod._ROCM_TORCH_INDEX
+_windows_rocm_index_url = stack_mod._windows_rocm_index_url
+_detect_windows_gfx_arch = stack_mod._detect_windows_gfx_arch
+_install_bnb_windows_rocm = stack_mod._install_bnb_windows_rocm
def _extract_sh_function_body(source: str, name: str) -> str:
@@ -596,6 +599,7 @@ class TestEnsureRocmTorch:
_ensure_rocm_torch()
mock_pip.assert_not_called()
+ @patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@@ -616,7 +620,9 @@ class TestEnsureRocmTorch:
assert "rocm7.1" in str(mock_pip.call_args_list[0])
assert mock_pip_try.call_count >= 1
assert "bitsandbytes" in str(mock_pip_try.call_args_list[0])
+ assert mock_pip_try.call_args.kwargs["force_pip"] is True
+ @patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@@ -648,6 +654,7 @@ class TestEnsureRocmTorch:
_ensure_rocm_torch()
mock_pip.assert_not_called()
+ @patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@@ -662,12 +669,13 @@ class TestEnsureRocmTorch:
captured = capsys.readouterr()
assert "unreadable" in captured.out
+ @patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
- def test_rocm_72_selects_71_tag(self, mock_ver, mock_gpu, mock_nvidia, mock_pip):
- """ROCm 7.2 should select rocm7.1 tag (capped, not in mapping)."""
+ def test_rocm_72_selects_72_tag(self, mock_ver, mock_gpu, mock_nvidia, mock_pip):
+ """ROCm 7.2 should select rocm7.2 tag (now in mapping with torch 2.11.0)."""
mock_probe = MagicMock()
mock_probe.returncode = 0
mock_probe.stdout = b"\n"
@@ -675,8 +683,9 @@ class TestEnsureRocmTorch:
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
torch_call = mock_pip.call_args_list[0]
- assert "rocm7.1" in str(torch_call)
+ assert "rocm7.2" in str(torch_call)
+ @patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@@ -696,6 +705,7 @@ class TestEnsureRocmTorch:
assert mock_pip.call_count == 1
assert "rocm7.1" in str(mock_pip.call_args_list[0])
assert mock_pip_try.call_count >= 1
+ assert mock_pip_try.call_args.kwargs["force_pip"] is True
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@@ -720,9 +730,10 @@ class TestRocmTorchIndex:
keys = list(_ROCM_TORCH_INDEX.keys())
assert keys == sorted(keys, reverse = True)
- def test_rocm_72_not_in_mapping(self):
- """ROCm 7.2 should NOT be in the active mapping (torch 2.11.0 exceeds bound)."""
- assert (7, 2) not in _ROCM_TORCH_INDEX
+ def test_rocm_72_in_mapping(self):
+ """ROCm 7.2 should be in the active mapping (torch 2.11.0 now supported)."""
+ assert (7, 2) in _ROCM_TORCH_INDEX
+ assert _ROCM_TORCH_INDEX[(7, 2)] == "rocm7.2"
def test_rocm_71_maps_correctly(self):
assert _ROCM_TORCH_INDEX[(7, 1)] == "rocm7.1"
@@ -740,7 +751,7 @@ class TestRocmTorchIndex:
assert "radeon" not in tag
def test_newer_rocm_selects_best_match(self):
- """ROCm 7.2 (not in map) should select rocm7.1 via >= comparison."""
+ """ROCm 7.2 (now in map) should select rocm7.2 directly."""
ver = (7, 2)
tag = next(
(
@@ -750,7 +761,7 @@ class TestRocmTorchIndex:
),
None,
)
- assert tag == "rocm7.1"
+ assert tag == "rocm7.2"
def test_rocm_64_selects_64(self):
ver = (6, 4)
@@ -778,7 +789,7 @@ class TestHardwareRocmFlag:
hw_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
)
- source = hw_path.read_text()
+ source = hw_path.read_text(encoding = "utf-8")
assert "IS_ROCM: bool" in source and "False" in source
def test_hardware_py_sets_is_rocm_on_hip(self):
@@ -786,7 +797,7 @@ class TestHardwareRocmFlag:
hw_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
)
- source = hw_path.read_text()
+ source = hw_path.read_text(encoding = "utf-8")
assert 'torch.version, "hip"' in source or "torch.version.hip" in source
def test_hardware_py_still_returns_cuda_for_rocm(self):
@@ -794,7 +805,7 @@ class TestHardwareRocmFlag:
hw_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
)
- source = hw_path.read_text()
+ source = hw_path.read_text(encoding = "utf-8")
# Ensure ROCM is NOT a DeviceType member
enum_section = source.split("class DeviceType")[1].split("\n\n")[0]
assert "ROCM" not in enum_section
@@ -804,7 +815,7 @@ class TestHardwareRocmFlag:
hw_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
)
- source = hw_path.read_text()
+ source = hw_path.read_text(encoding = "utf-8")
assert '"rocm"' in source
def test_hardware_py_device_type_cuda_references_intact(self):
@@ -812,7 +823,7 @@ class TestHardwareRocmFlag:
hw_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
)
- source = hw_path.read_text()
+ source = hw_path.read_text(encoding = "utf-8")
# Key functions that must still reference DeviceType.CUDA
assert "DeviceType.CUDA" in source
assert "DEVICE = DeviceType.CUDA" in source
@@ -822,7 +833,7 @@ class TestHardwareRocmFlag:
init_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "__init__.py"
)
- source = init_path.read_text()
+ source = init_path.read_text(encoding = "utf-8")
assert "IS_ROCM" in source
def test_is_rocm_in_all_list(self):
@@ -830,7 +841,7 @@ class TestHardwareRocmFlag:
init_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "__init__.py"
)
- source = init_path.read_text()
+ source = init_path.read_text(encoding = "utf-8")
# Extract __all__ section
assert '"IS_ROCM"' in source
@@ -839,13 +850,37 @@ class TestHardwareRocmFlag:
hw_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
)
- source = hw_path.read_text()
+ source = hw_path.read_text(encoding = "utf-8")
# Find the get_package_versions function body
func_start = source.find("def get_package_versions")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert '"cuda"' in func_body
assert '"rocm"' in func_body
+ def test_distributed_stubs_cover_is_torchelastic_launched(self):
+ """_determine_attention_impl_for_gpu_estimate must stub is_torchelastic_launched.
+
+ resolve_attention_implementation calls is_torchelastic_launched() on
+ Windows ROCm where torch.distributed ships without that helper, causing
+ a warning: 'module torch.distributed has no attribute is_torchelastic_launched'.
+ """
+ hw_path = (
+ PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
+ )
+ source = hw_path.read_text(encoding = "utf-8")
+ assert "is_torchelastic_launched" in source
+
+ def test_distributed_stubs_cover_core_helpers(self):
+ """_determine_attention_impl_for_gpu_estimate must stub the four core distributed helpers."""
+ hw_path = (
+ PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
+ )
+ source = hw_path.read_text(encoding = "utf-8")
+ for attr in ("is_initialized", "is_available", "get_rank", "get_world_size"):
+ assert (
+ attr in source
+ ), f"distributed stub for '{attr}' missing from hardware.py"
+
# =============================================================================
# TEST: tokenizer_utils.py -- error message
@@ -858,13 +893,13 @@ class TestTokenizerErrorMessage:
def test_no_old_amd_message(self):
"""Old 'We do not support AMD' message should be gone."""
tu_path = PACKAGE_ROOT / "unsloth" / "tokenizer_utils.py"
- source = tu_path.read_text()
+ source = tu_path.read_text(encoding = "utf-8")
assert "We do not support AMD" not in source
def test_new_message_has_docs_link(self):
"""New message should point to Unsloth AMD docs."""
tu_path = PACKAGE_ROOT / "unsloth" / "tokenizer_utils.py"
- source = tu_path.read_text()
+ source = tu_path.read_text(encoding = "utf-8")
assert "docs.unsloth.ai" in source or "No GPU detected" in source
@@ -879,7 +914,7 @@ class TestInstallShStructure:
def test_no_here_strings(self):
"""install.sh must not use <<< (not POSIX)."""
sh_path = PACKAGE_ROOT / "install.sh"
- source = sh_path.read_text()
+ source = sh_path.read_text(encoding = "utf-8")
# <<< is bash-only; breaks dash
for i, line in enumerate(source.splitlines(), 1):
stripped = line.lstrip()
@@ -890,7 +925,7 @@ class TestInstallShStructure:
def test_rocm_detection_present(self):
"""install.sh should have ROCm detection in get_torch_index_url."""
sh_path = PACKAGE_ROOT / "install.sh"
- source = sh_path.read_text()
+ source = sh_path.read_text(encoding = "utf-8")
assert "amd-smi" in source
assert "rocm" in source.lower()
@@ -905,7 +940,7 @@ class TestInstallShStructure:
block.
"""
sh_path = PACKAGE_ROOT / "install.sh"
- source = sh_path.read_text()
+ source = sh_path.read_text(encoding = "utf-8")
body = _extract_sh_function_body(source, "get_torch_index_url")
nvidia_call = body.find("_has_usable_nvidia_gpu")
no_nvidia_branch = body.find('if [ -z "$_smi" ]')
@@ -926,44 +961,45 @@ class TestInstallShStructure:
def test_bitsandbytes_amd_install(self):
"""install.sh should install bitsandbytes for AMD when ROCm detected."""
sh_path = PACKAGE_ROOT / "install.sh"
- source = sh_path.read_text()
+ source = sh_path.read_text(encoding = "utf-8")
assert "bitsandbytes" in source
assert "rocm*)" in source # case pattern for ROCm URLs
def test_cpu_hint_mentions_amd(self):
"""CPU-only hint should mention AMD ROCm."""
sh_path = PACKAGE_ROOT / "install.sh"
- source = sh_path.read_text()
+ source = sh_path.read_text(encoding = "utf-8")
assert "ROCm" in source
- def test_rocm72_capped_to_71(self):
- """ROCm 7.2+ should fall back to rocm7.1 index."""
+ def test_rocm72_supported_future_capped(self):
+ """ROCm 7.2 should pass through directly; 7.3+ falls back to rocm7.2."""
sh_path = PACKAGE_ROOT / "install.sh"
- source = sh_path.read_text()
- assert 'echo "$_base/rocm7.1"' in source # fallback for unknown versions
+ source = sh_path.read_text(encoding = "utf-8")
+ assert 'echo "$_base/rocm7.2"' in source # fallback for unknown future versions
# Allowlisted versions should pass through directly
assert "rocm6.*" in source
assert "rocm7.0" in source
assert "rocm7.1" in source
+ assert "rocm7.2" in source
def test_rocm_tag_validation_guard_exists(self):
"""install.sh should validate _rocm_tag with a case guard."""
sh_path = PACKAGE_ROOT / "install.sh"
- source = sh_path.read_text()
+ source = sh_path.read_text(encoding = "utf-8")
assert "rocm[1-9]*.[0-9]*)" in source
assert '_rocm_tag=""' in source # rejection path
def test_dpkg_epoch_handling(self):
"""install.sh should strip Debian epoch prefix from dpkg-query output."""
sh_path = PACKAGE_ROOT / "install.sh"
- source = sh_path.read_text()
+ source = sh_path.read_text(encoding = "utf-8")
assert "sed 's/^[0-9]*://' " in source or "sed 's/^[0-9]*://'" in source
def test_no_double_bracket_in_rocm_block(self):
"""ROCm detection block should not use [[ ]] (bash-only, not POSIX).
Note: [[:space:]], [[:digit:]] etc. are valid POSIX character classes, not bash [[ ]]."""
sh_path = PACKAGE_ROOT / "install.sh"
- source = sh_path.read_text()
+ source = sh_path.read_text(encoding = "utf-8")
func_start = source.find("get_torch_index_url()")
func_end = source.find("\n}", func_start)
func_body = source[func_start:func_end]
@@ -982,7 +1018,7 @@ class TestInstallShStructure:
def test_no_arithmetic_expansion_in_rocm_block(self):
"""ROCm detection block should not use (( )) (bash-only)."""
sh_path = PACKAGE_ROOT / "install.sh"
- source = sh_path.read_text()
+ source = sh_path.read_text(encoding = "utf-8")
func_start = source.find("get_torch_index_url()")
func_end = source.find("\n}", func_start)
func_body = source[func_start:func_end]
@@ -997,7 +1033,7 @@ class TestInstallShStructure:
def test_macos_returns_cpu_before_rocm_check(self):
"""macOS should return CPU immediately (before any ROCm check)."""
sh_path = PACKAGE_ROOT / "install.sh"
- source = sh_path.read_text()
+ source = sh_path.read_text(encoding = "utf-8")
func_start = source.find("get_torch_index_url()")
func_body = source[func_start:]
darwin_pos = func_body.find("Darwin")
@@ -1060,6 +1096,12 @@ class TestLiveRegression:
# Load worker.py module
_WORKER_PATH = PACKAGE_ROOT / "studio" / "backend" / "core" / "training" / "worker.py"
+_EXPORT_WORKER_PATH = (
+ PACKAGE_ROOT / "studio" / "backend" / "core" / "export" / "worker.py"
+)
+# The torchao Windows-ROCm stub was de-duplicated out of the export/training
+# workers into a shared module; both workers now call into it.
+_TORCHAO_STUB_PATH = PACKAGE_ROOT / "studio" / "backend" / "core" / "_torchao_stub.py"
# The wheel-probe subprocess was hoisted out of worker.py into wheel_utils
# during the wheel-resolver refactor; the probe script literal lives there.
_WHEEL_UTILS_PATH = PACKAGE_ROOT / "studio" / "backend" / "utils" / "wheel_utils.py"
@@ -1071,12 +1113,12 @@ class TestWorkerRocmMambaSsm:
def test_probe_returns_hip_version_field(self):
"""The wheel probe should include hip_version, and worker.py should
consume it."""
- assert "hip_version" in _WHEEL_UTILS_PATH.read_text()
- assert "hip_version" in _WORKER_PATH.read_text()
+ assert "hip_version" in _WHEEL_UTILS_PATH.read_text(encoding = "utf-8")
+ assert "hip_version" in _WORKER_PATH.read_text(encoding = "utf-8")
def test_probe_script_has_getattr_hip(self):
"""Probe script should use getattr for torch.version.hip (safe on CUDA)."""
- source = _WHEEL_UTILS_PATH.read_text()
+ source = _WHEEL_UTILS_PATH.read_text(encoding = "utf-8")
assert "getattr(torch.version, 'hip', None)" in source
def test_direct_wheel_url_returns_none_without_cuda_major(self):
@@ -1118,22 +1160,22 @@ class TestWorkerRocmMambaSsm:
def test_hipcc_check_exists_in_source(self):
"""worker.py should check for hipcc before ROCm source builds."""
- source = _WORKER_PATH.read_text()
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "hipcc" in source
def test_rocm_source_build_status_message(self):
"""worker.py should send a specific status for ROCm source compilation."""
- source = _WORKER_PATH.read_text()
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "Compiling" in source and "from source for ROCm" in source
def test_rocm_build_failure_message(self):
"""worker.py should send a clear error on ROCm build failure."""
- source = _WORKER_PATH.read_text()
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "Failed to compile" in source and "for ROCm" in source
def test_timeout_on_install(self):
"""worker.py should have a timeout on pip install subprocess."""
- source = _WORKER_PATH.read_text()
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
assert "TimeoutExpired" in source
assert "timeout" in source
@@ -1154,7 +1196,7 @@ class TestAmdGpuMonitoring:
def test_amd_py_has_required_functions(self):
"""amd.py should export the same function signatures as nvidia.py."""
amd_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "amd.py"
- source = amd_path.read_text()
+ source = amd_path.read_text(encoding = "utf-8")
assert "def get_physical_gpu_count" in source
assert "def get_primary_gpu_utilization" in source
assert "def get_visible_gpu_utilization" in source
@@ -1299,7 +1341,7 @@ class TestHardwareAmdBranching:
hw_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
)
- source = hw_path.read_text()
+ source = hw_path.read_text(encoding = "utf-8")
assert "from . import amd" in source
def test_hardware_branches_on_is_rocm_for_utilization(self):
@@ -1309,7 +1351,7 @@ class TestHardwareAmdBranching:
hw_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
)
- source = hw_path.read_text()
+ source = hw_path.read_text(encoding = "utf-8")
func_start = source.find("def get_gpu_utilization")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert '_smi_query("get_primary_gpu_utilization"' in func_body
@@ -1327,7 +1369,7 @@ class TestHardwareAmdBranching:
hw_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
)
- source = hw_path.read_text()
+ source = hw_path.read_text(encoding = "utf-8")
func_start = source.find("def get_visible_gpu_utilization")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
# The dispatcher call may wrap onto multiple lines; allow whitespace
@@ -1348,13 +1390,55 @@ class TestHardwareAmdBranching:
hw_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
)
- source = hw_path.read_text()
+ source = hw_path.read_text(encoding = "utf-8")
func_start = source.find("def get_physical_gpu_count")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "IS_ROCM" in func_body
assert "from . import amd" in func_body
+# =============================================================================
+# TEST: hardware.py -- apply_gpu_ids ROCm fallback (issue #5180)
+# =============================================================================
+
+
+class TestApplyGpuIdsRocmFallback:
+ """Verify apply_gpu_ids sets HIP_VISIBLE_DEVICES on ROCm hosts even when
+ IS_ROCM is still False (worker subprocess before detect_hardware runs)."""
+
+ def test_apply_gpu_ids_falls_back_to_torch_version_hip(self):
+ """apply_gpu_ids should probe torch.version.hip when IS_ROCM is False and no ROCm env vars are set."""
+ hw_path = (
+ PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
+ )
+ source = hw_path.read_text(encoding = "utf-8")
+ func_start = source.find("def apply_gpu_ids")
+ func_body = source[func_start : source.find("\ndef ", func_start + 1)]
+ assert 'getattr(_torch.version, "hip", None)' in func_body
+
+ def test_apply_gpu_ids_sets_hip_and_rocr_visible_devices(self):
+ """apply_gpu_ids should set both HIP_VISIBLE_DEVICES and ROCR_VISIBLE_DEVICES on ROCm."""
+ hw_path = (
+ PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
+ )
+ source = hw_path.read_text(encoding = "utf-8")
+ func_start = source.find("def apply_gpu_ids")
+ func_body = source[func_start : source.find("\ndef ", func_start + 1)]
+ assert 'os.environ["HIP_VISIBLE_DEVICES"] = value' in func_body
+ assert 'os.environ["ROCR_VISIBLE_DEVICES"] = value' in func_body
+
+ def test_apply_gpu_ids_rocm_fallback_is_guarded_by_try_except(self):
+ """torch import in apply_gpu_ids must be wrapped in try/except so a missing torch never crashes."""
+ hw_path = (
+ PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
+ )
+ source = hw_path.read_text(encoding = "utf-8")
+ func_start = source.find("def apply_gpu_ids")
+ func_body = source[func_start : source.find("\ndef ", func_start + 1)]
+ assert "import torch as _torch" in func_body
+ assert "except Exception" in func_body
+
+
# =============================================================================
# TEST: install_python_stack.py -- Windows AMD warning
# =============================================================================
@@ -1365,18 +1449,18 @@ class TestWindowsRocmWarning:
def test_windows_amd_warning_in_source(self):
"""install_python_stack.py should warn Windows AMD users."""
- source = _STACK_PATH.read_text()
- assert "AMD GPU detected on Windows" in source
+ source = _STACK_PATH.read_text(encoding = "utf-8")
+ assert "AMD GPU detected" in source
def test_windows_amd_warning_checks_hipinfo_or_amdsmi(self):
"""Warning should check for hipinfo or amd-smi."""
- source = _STACK_PATH.read_text()
+ source = _STACK_PATH.read_text(encoding = "utf-8")
assert "hipinfo" in source
assert "amd-smi" in source
def test_windows_amd_warning_has_docs_link(self):
"""Warning should include AMD docs link."""
- source = _STACK_PATH.read_text()
+ source = _STACK_PATH.read_text(encoding = "utf-8")
assert "docs.unsloth.ai/get-started/install-and-update/amd" in source
@@ -1391,7 +1475,7 @@ class TestIsRdnaExpansion:
def test_is_rdna_source_has_rdna2(self):
"""is_rdna() should include RDNA2 architectures."""
utils_path = PACKAGE_ROOT / "unsloth" / "kernels" / "utils.py"
- source = utils_path.read_text()
+ source = utils_path.read_text(encoding = "utf-8")
func_start = source.find("def is_rdna()")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "gfx1030" in func_body
@@ -1405,7 +1489,7 @@ class TestIsRdnaExpansion:
def test_is_rdna_source_has_rdna3(self):
"""is_rdna() should include RDNA3 architectures."""
utils_path = PACKAGE_ROOT / "unsloth" / "kernels" / "utils.py"
- source = utils_path.read_text()
+ source = utils_path.read_text(encoding = "utf-8")
func_start = source.find("def is_rdna()")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "gfx1100" in func_body
@@ -1416,7 +1500,7 @@ class TestIsRdnaExpansion:
def test_is_rdna_source_has_rdna35(self):
"""is_rdna() should include RDNA3.5 architectures."""
utils_path = PACKAGE_ROOT / "unsloth" / "kernels" / "utils.py"
- source = utils_path.read_text()
+ source = utils_path.read_text(encoding = "utf-8")
func_start = source.find("def is_rdna()")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "gfx1150" in func_body
@@ -1426,7 +1510,7 @@ class TestIsRdnaExpansion:
def test_is_rdna_source_has_rdna4(self):
"""is_rdna() should include RDNA4 architectures."""
utils_path = PACKAGE_ROOT / "unsloth" / "kernels" / "utils.py"
- source = utils_path.read_text()
+ source = utils_path.read_text(encoding = "utf-8")
func_start = source.find("def is_rdna()")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "gfx1200" in func_body
@@ -1435,7 +1519,7 @@ class TestIsRdnaExpansion:
def test_is_cdna_not_changed(self):
"""is_cdna() should remain unchanged (no RDNA architectures added)."""
utils_path = PACKAGE_ROOT / "unsloth" / "kernels" / "utils.py"
- source = utils_path.read_text()
+ source = utils_path.read_text(encoding = "utf-8")
func_start = source.find("def is_cdna()")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "gfx940" in func_body
@@ -1447,5 +1531,1072 @@ class TestIsRdnaExpansion:
assert "gfx1100" not in func_body
+# =============================================================================
+# TEST: install_python_stack.py -- _windows_rocm_index_url arch mapping
+# =============================================================================
+
+
+class TestWindowsRocmIndexUrl:
+ """Verify GPU arch → AMD pip index URL mapping."""
+
+ def test_gfx1200_maps_to_gfx120x_all(self):
+ url = stack_mod._windows_rocm_index_url("gfx1200")
+ assert url is not None
+ assert "gfx120X-all" in url
+
+ def test_gfx1201_maps_to_gfx120x_all(self):
+ url = stack_mod._windows_rocm_index_url("gfx1201")
+ assert url is not None
+ assert "gfx120X-all" in url
+
+ def test_gfx1151_maps_to_gfx1151(self):
+ url = stack_mod._windows_rocm_index_url("gfx1151")
+ assert url is not None
+ assert "gfx1151" in url
+
+ def test_gfx1150_maps_to_gfx1150(self):
+ url = stack_mod._windows_rocm_index_url("gfx1150")
+ assert url is not None
+ assert "gfx1150" in url
+
+ def test_gfx1100_maps_to_gfx110x_all(self):
+ url = stack_mod._windows_rocm_index_url("gfx1100")
+ assert url is not None
+ assert "gfx110X-all" in url
+
+ def test_unknown_arch_returns_none(self):
+ assert stack_mod._windows_rocm_index_url("gfx9999") is None
+
+ def test_none_arch_returns_none(self):
+ assert stack_mod._windows_rocm_index_url(None) is None
+
+ def test_url_ends_with_slash(self):
+ """AMD pip index URLs must end with / for --index-url compatibility."""
+ url = stack_mod._windows_rocm_index_url("gfx1200")
+ assert url is not None
+ assert url.endswith("/")
+
+ def test_base_url_uses_repo_amd_com_by_default(self):
+ url = stack_mod._windows_rocm_index_url("gfx1200")
+ assert url is not None
+ assert "repo.amd.com" in url
+
+ def test_mirror_env_var_overrides_base(self, monkeypatch):
+ monkeypatch.setenv(
+ "UNSLOTH_ROCM_WINDOWS_MIRROR", "https://my-mirror.example.com/rocm/whl"
+ )
+ # Reload module-level constant by calling helper directly
+ url = stack_mod._windows_rocm_index_url("gfx1200")
+ # The env var is read at module load time for _ROCM_WINDOWS_INDEX_BASE,
+ # so just verify the helper itself doesn't error.
+ assert url is not None
+
+
+# =============================================================================
+# TEST: install_python_stack.py -- _detect_windows_gfx_arch
+# =============================================================================
+
+
+class TestDetectWindowsGfxArch:
+ """Verify hipinfo parsing for GPU arch detection on Windows."""
+
+ def test_returns_none_when_hipinfo_not_on_path(self):
+ with patch("shutil.which", return_value = None):
+ result = stack_mod._detect_windows_gfx_arch()
+ assert result is None
+
+ def test_parses_gcnarchname_from_hipinfo_output(self):
+ mock_result = MagicMock()
+ mock_result.returncode = 0
+ mock_result.stdout = b"gcnArchName : gfx1200\nsome other line\n"
+ with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
+ with patch("subprocess.run", return_value = mock_result):
+ result = stack_mod._detect_windows_gfx_arch()
+ assert result == "gfx1200"
+
+ def test_returns_none_on_nonzero_returncode(self):
+ mock_result = MagicMock()
+ mock_result.returncode = 1
+ mock_result.stdout = b"gcnArchName : gfx1200\n"
+ with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
+ with patch("subprocess.run", return_value = mock_result):
+ result = stack_mod._detect_windows_gfx_arch()
+ assert result is None
+
+ def test_returns_none_when_no_gcnarchname_in_output(self):
+ mock_result = MagicMock()
+ mock_result.returncode = 0
+ mock_result.stdout = b"deviceName : Radeon RX 9060 XT\n"
+ with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
+ with patch("subprocess.run", return_value = mock_result):
+ result = stack_mod._detect_windows_gfx_arch()
+ assert result is None
+
+ def test_returns_none_on_timeout(self):
+ with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
+ with patch(
+ "subprocess.run",
+ side_effect = subprocess.TimeoutExpired("hipinfo", 10),
+ ):
+ result = stack_mod._detect_windows_gfx_arch()
+ assert result is None
+
+ def test_strips_whitespace_from_arch(self):
+ mock_result = MagicMock()
+ mock_result.returncode = 0
+ mock_result.stdout = b" gcnArchName : gfx1201 \n"
+ with patch("shutil.which", return_value = "/usr/bin/hipinfo"):
+ with patch("subprocess.run", return_value = mock_result):
+ result = stack_mod._detect_windows_gfx_arch()
+ assert result == "gfx1201"
+
+
+# =============================================================================
+# TEST: install_python_stack.py -- _install_bnb_windows_rocm
+# =============================================================================
+
+
+class TestInstallBnbWindowsRocm:
+ """Verify AMD Windows BNB wheel install helper."""
+
+ def test_calls_pip_install_try_with_win_amd64_url(self):
+ """Should call pip_install_try with the win_amd64 wheel URL via plain pip."""
+ with patch.object(stack_mod, "pip_install_try", return_value = True) as mock_pip:
+ stack_mod._install_bnb_windows_rocm()
+ assert mock_pip.call_count == 1
+ call_args = str(mock_pip.call_args_list[0])
+ assert "bitsandbytes" in call_args
+ assert "win_amd64" in call_args
+ # Must force plain pip (uv mangles the bitsandbytes wheel) -- see
+ # https://unsloth.ai/docs/get-started/install/amd/amd-hackathon
+ assert mock_pip.call_args.kwargs.get("force_pip") is True
+
+ def test_forces_plain_pip_not_uv(self):
+ """The bnb wheel must be installed with plain pip, never uv."""
+ with patch.object(stack_mod, "pip_install_try", return_value = True) as mock_pip:
+ stack_mod._install_bnb_windows_rocm()
+ assert mock_pip.call_args.kwargs.get("force_pip") is True
+
+ def test_does_not_touch_uv_skip_env_var(self):
+ """The UV_SKIP_WHEEL_FILENAME_CHECK hack is gone; the env must be untouched."""
+ observed = {}
+
+ def _capture(*args, **kwargs):
+ observed["during"] = os.environ.get("UV_SKIP_WHEEL_FILENAME_CHECK")
+ return True
+
+ with patch.dict(os.environ, {}, clear = False):
+ os.environ.pop("UV_SKIP_WHEEL_FILENAME_CHECK", None)
+ with patch.object(stack_mod, "pip_install_try", side_effect = _capture):
+ stack_mod._install_bnb_windows_rocm()
+ assert observed.get("during") is None
+ assert "UV_SKIP_WHEEL_FILENAME_CHECK" not in os.environ
+
+ def test_returns_false_on_pip_failure(self):
+ """A failed pip_install_try must surface as a False return, not BNB_ROCM_VERSION."""
+ with patch.dict(os.environ, {}, clear = False):
+ os.environ.pop("BNB_ROCM_VERSION", None)
+ with patch.object(stack_mod, "pip_install_try", return_value = False):
+ result = stack_mod._install_bnb_windows_rocm()
+ assert result is False
+ assert "BNB_ROCM_VERSION" not in os.environ
+
+ def test_no_op_when_win_amd64_url_missing(self):
+ """Should be silent no-op if win_amd64 key absent from _BNB_ROCM_PRERELEASE_URLS."""
+ with patch.object(stack_mod, "_BNB_ROCM_PRERELEASE_URLS", {}):
+ with patch.object(stack_mod, "pip_install_try") as mock_pip:
+ stack_mod._install_bnb_windows_rocm()
+ mock_pip.assert_not_called()
+
+ def test_sets_bnb_rocm_version_from_detected_dll(self):
+ """BNB_ROCM_VERSION is set from the DLL detected after install."""
+ with patch.dict(os.environ, {}, clear = False):
+ os.environ.pop("BNB_ROCM_VERSION", None)
+ with patch.object(stack_mod, "pip_install_try", return_value = True):
+ with patch.object(
+ stack_mod, "_detect_bnb_rocm_dll_ver", return_value = "72"
+ ):
+ stack_mod._install_bnb_windows_rocm()
+ assert os.environ.get("BNB_ROCM_VERSION") == "72"
+
+ def test_sets_bnb_rocm_version_from_newer_dll(self):
+ """If AMD ships a newer DLL (e.g. rocm713.dll), that version is used."""
+ with patch.dict(os.environ, {}, clear = False):
+ os.environ.pop("BNB_ROCM_VERSION", None)
+ with patch.object(stack_mod, "pip_install_try", return_value = True):
+ with patch.object(
+ stack_mod, "_detect_bnb_rocm_dll_ver", return_value = "713"
+ ):
+ stack_mod._install_bnb_windows_rocm()
+ assert os.environ.get("BNB_ROCM_VERSION") == "713"
+
+ def test_falls_back_to_72_when_detection_fails(self):
+ """Falls back to '72' when DLL detection returns None."""
+ with patch.dict(os.environ, {}, clear = False):
+ os.environ.pop("BNB_ROCM_VERSION", None)
+ with patch.object(stack_mod, "pip_install_try", return_value = True):
+ with patch.object(
+ stack_mod, "_detect_bnb_rocm_dll_ver", return_value = None
+ ):
+ stack_mod._install_bnb_windows_rocm()
+ assert os.environ.get("BNB_ROCM_VERSION") == "72"
+
+ def test_does_not_override_existing_bnb_rocm_version(self):
+ """An explicit BNB_ROCM_VERSION in the caller's env must not be clobbered."""
+ with patch.dict(os.environ, {"BNB_ROCM_VERSION": "60"}):
+ with patch.object(stack_mod, "pip_install_try", return_value = True):
+ stack_mod._install_bnb_windows_rocm()
+ assert os.environ.get("BNB_ROCM_VERSION") == "60"
+
+
+class TestDetectBnbRocmDllVer:
+ """Unit tests for _detect_bnb_rocm_dll_ver()."""
+
+ def test_returns_none_when_bnb_not_installed(self):
+ """Returns None if bitsandbytes is not importable."""
+ import importlib.util
+
+ with patch.object(importlib.util, "find_spec", return_value = None):
+ assert stack_mod._detect_bnb_rocm_dll_ver() is None
+
+ def test_detects_rocm72_dll(self, tmp_path):
+ """Returns '72' when libbitsandbytes_rocm72.dll is present."""
+ (tmp_path / "libbitsandbytes_rocm72.dll").write_text("")
+ mock_spec = MagicMock()
+ mock_spec.submodule_search_locations = [str(tmp_path)]
+ import importlib.util
+
+ with patch.object(importlib.util, "find_spec", return_value = mock_spec):
+ assert stack_mod._detect_bnb_rocm_dll_ver() == "72"
+
+ def test_detects_rocm713_dll(self, tmp_path):
+ """Returns '713' when libbitsandbytes_rocm713.dll is present."""
+ (tmp_path / "libbitsandbytes_rocm713.dll").write_text("")
+ mock_spec = MagicMock()
+ mock_spec.submodule_search_locations = [str(tmp_path)]
+ import importlib.util
+
+ with patch.object(importlib.util, "find_spec", return_value = mock_spec):
+ assert stack_mod._detect_bnb_rocm_dll_ver() == "713"
+
+ def test_returns_none_when_only_cuda_dlls(self, tmp_path):
+ """Returns None when only CUDA DLLs are present (no ROCm DLL)."""
+ (tmp_path / "libbitsandbytes_cuda121.dll").write_text("")
+ mock_spec = MagicMock()
+ mock_spec.submodule_search_locations = [str(tmp_path)]
+ import importlib.util
+
+ with patch.object(importlib.util, "find_spec", return_value = mock_spec):
+ assert stack_mod._detect_bnb_rocm_dll_ver() is None
+
+ def test_picks_highest_suffix_when_multiple_dlls(self, tmp_path):
+ """Returns the highest numeric suffix when multiple ROCm DLL variants exist.
+
+ Filesystem glob order is not guaranteed, so the function must not stop
+ at the first match — it must always return the highest one.
+ """
+ (tmp_path / "libbitsandbytes_rocm72.dll").write_text("")
+ (tmp_path / "libbitsandbytes_rocm713.dll").write_text("")
+ mock_spec = MagicMock()
+ mock_spec.submodule_search_locations = [str(tmp_path)]
+ import importlib.util
+
+ with patch.object(importlib.util, "find_spec", return_value = mock_spec):
+ assert stack_mod._detect_bnb_rocm_dll_ver() == "713"
+
+
+# =============================================================================
+# TEST: install_python_stack.py -- UNSLOTH_ROCM_TORCH_INSTALLED early-return path
+# =============================================================================
+
+
+class TestRocmTorchInstalledEnvVar:
+ """Verify UNSLOTH_ROCM_TORCH_INSTALLED=1 skips main install but still installs BNB."""
+
+ @staticmethod
+ def _ok_torch_probe(*a, **kw):
+ # subprocess.run probe returns 0 when torch imports as ROCm
+ rv = MagicMock()
+ rv.returncode = 0
+ return rv
+
+ @patch.object(stack_mod, "_install_bnb_windows_rocm")
+ @patch.object(stack_mod, "pip_install")
+ def test_env_var_skips_main_pip_install(self, mock_pip, mock_bnb):
+ """UNSLOTH_ROCM_TORCH_INSTALLED=1 should not trigger torch pip_install."""
+ with (
+ patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}),
+ patch.object(stack_mod.subprocess, "run", side_effect = self._ok_torch_probe),
+ ):
+ stack_mod._ensure_rocm_torch()
+ mock_pip.assert_not_called()
+
+ @patch.object(stack_mod, "_install_bnb_windows_rocm")
+ @patch.object(stack_mod, "pip_install")
+ def test_env_var_calls_bnb_install(self, mock_pip, mock_bnb):
+ """UNSLOTH_ROCM_TORCH_INSTALLED=1 should still call _install_bnb_windows_rocm."""
+ with (
+ patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}),
+ patch.object(stack_mod.subprocess, "run", side_effect = self._ok_torch_probe),
+ ):
+ stack_mod._ensure_rocm_torch()
+ mock_bnb.assert_called_once()
+
+ @patch.object(stack_mod, "_install_bnb_windows_rocm")
+ @patch.object(stack_mod, "pip_install")
+ def test_env_var_sets_rocm_windows_flag(self, mock_pip, mock_bnb):
+ """UNSLOTH_ROCM_TORCH_INSTALLED=1 should set _rocm_windows_torch_installed."""
+ stack_mod._rocm_windows_torch_installed = False
+ with (
+ patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}),
+ patch.object(stack_mod.subprocess, "run", side_effect = self._ok_torch_probe),
+ ):
+ stack_mod._ensure_rocm_torch()
+ assert stack_mod._rocm_windows_torch_installed is True
+
+ @patch.object(stack_mod, "_install_bnb_windows_rocm")
+ @patch.object(stack_mod, "pip_install")
+ def test_env_var_falls_through_when_torch_missing(self, mock_pip, mock_bnb):
+ """If the venv was wiped between runs, the stale env-var must not suppress reinstall."""
+ stack_mod._rocm_windows_torch_installed = False
+
+ def _bad_probe(*a, **kw):
+ rv = MagicMock()
+ rv.returncode = 1
+ return rv
+
+ with (
+ patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}),
+ patch.object(stack_mod.subprocess, "run", side_effect = _bad_probe),
+ patch.object(stack_mod, "IS_WINDOWS", False),
+ patch.object(stack_mod, "IS_MACOS", True),
+ ):
+ stack_mod._ensure_rocm_torch()
+ # macOS branch is the next exit -- but the point is the early-return did NOT fire.
+ mock_bnb.assert_not_called()
+
+
+# =============================================================================
+# TEST: worker.py -- Windows ROCm patches (source-level checks)
+# =============================================================================
+
+
+class TestWorkerWindowsRocmPatches:
+ """Verify worker.py contains the required Windows ROCm runtime patches."""
+
+ def test_grouped_mm_dispatch_patch_present(self):
+ """worker.py must register a _grouped_mm CUDA dispatch override."""
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ assert '_gm_lib.impl("_grouped_mm"' in source
+
+ def test_grouped_mm_patch_targets_cuda_dispatch_key(self):
+ """The dispatch override must target the CUDA key (not CompositeImplicitAutograd)."""
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ assert '"_grouped_mm", _grouped_mm_safe_impl, "CUDA"' in source
+
+ def test_grouped_mm_lib_kept_alive(self):
+ """The Library object must be stored to prevent GC clearing the registration."""
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ assert "_WINDOWS_ROCM_GROUPED_MM_LIB" in source
+
+ def test_grouped_mm_handles_offs_grouped_case(self):
+ """_grouped_mm fallback must handle the grouped (offs!=None) variant."""
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ assert "offs_list" in source
+ assert "offs.tolist()" in source
+
+ def test_worker_calls_shared_torchao_stub(self):
+ """worker.py must invoke the shared torchao stub entrypoint."""
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ assert "install_torchao_windows_rocm_stub()" in source
+
+ def test_export_worker_calls_shared_torchao_stub(self):
+ """export/worker.py must invoke the same shared torchao stub entrypoint."""
+ source = _EXPORT_WORKER_PATH.read_text(encoding = "utf-8")
+ assert "install_torchao_windows_rocm_stub()" in source
+
+ def test_torchao_stub_uses_stub_type_meta(self):
+ """Torchao stub must use _StubTypeMeta so isinstance() returns False not TypeError."""
+ source = _TORCHAO_STUB_PATH.read_text(encoding = "utf-8")
+ assert "_StubTypeMeta" in source
+
+ def test_stub_type_meta_has_instancecheck(self):
+ """_StubTypeMeta must define __instancecheck__ returning False."""
+ source = _TORCHAO_STUB_PATH.read_text(encoding = "utf-8")
+ assert "__instancecheck__" in source
+
+ def test_stub_subpackage_finder_registered(self):
+ """_StubSubpackageFinder must be appended to sys.meta_path."""
+ source = _TORCHAO_STUB_PATH.read_text(encoding = "utf-8")
+ assert "sys.meta_path.append(_StubSubpackageFinder())" in source
+
+ def test_torchao_key_submodules_pre_stubbed(self):
+ """Key torchao submodules (dtypes, quantization) must be pre-stubbed."""
+ source = _TORCHAO_STUB_PATH.read_text(encoding = "utf-8")
+ assert "torchao.dtypes" in source
+ assert "torchao.quantization" in source
+
+ def test_torchdynamo_disabled_on_windows_rocm(self):
+ """worker.py should disable dynamo on Windows ROCm as belt-and-suspenders."""
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ assert "TORCHDYNAMO_DISABLE" in source
+
+ def test_bnb_rocm_version_set_on_windows_rocm(self):
+ """worker.py must set BNB_ROCM_VERSION in the Windows ROCm section.
+
+ BNB auto-detects HIP version from torch.version.hip, which can mismatch
+ the DLL suffix in the AMD prerelease wheel. The worker must detect the
+ actual DLL suffix and override BNB's auto-detection before ML imports.
+ """
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ # Env var must be set
+ assert "BNB_ROCM_VERSION" in source
+ # Detection helper must be used
+ assert "_detect_bnb_rocm_dll_ver" in source or "libbitsandbytes_rocm" in source
+ # "72" must appear as the safe fallback
+ assert '"72"' in source or "'72'" in source
+
+ def test_bnb_rocm_version_set_before_ml_imports(self):
+ """BNB_ROCM_VERSION must appear in section 1f, before section 2 ML imports."""
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ idx_bnb = source.find("BNB_ROCM_VERSION")
+ # Use the specific section-2 marker that appears in the worker process
+ # entry-point function (not the trainer helper which has its own "# ── 2.").
+ idx_sec2 = source.find("# ── 2. Now import ML libraries")
+ assert idx_bnb != -1, "BNB_ROCM_VERSION not found in worker.py"
+ assert (
+ idx_sec2 != -1
+ ), "'# ── 2. Now import ML libraries' marker not found in worker.py"
+ assert idx_bnb < idx_sec2, (
+ "BNB_ROCM_VERSION must be set before section 2 ML imports "
+ f"(found at {idx_bnb}, section 2 at {idx_sec2})"
+ )
+
+ def test_grouped_mm_patch_guarded_by_windows_and_hip_check(self):
+ """_grouped_mm patch must only apply on Windows + HIP torch."""
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ # Must check sys.platform == "win32"
+ assert 'sys.platform == "win32"' in source
+ # Must gate on HIP version — code uses getattr chain: "version" and "hip"
+ assert '"version"' in source and '"hip"' in source
+
+ def test_hip_ver_at_least_helper_defined(self):
+ """_hip_ver_at_least helper must be defined inside the Windows ROCm block."""
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ assert "def _hip_ver_at_least(major: int, minor: int)" in source
+
+ def test_grouped_mm_patch_gated_on_hip_lt_713(self):
+ """_grouped_mm patch must be skipped on HIP >= 7.13 (AMD fixed the bug in ROCm 7.13)."""
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ # The guard must call _hip_ver_at_least with exactly (7, 13)
+ assert "_hip_ver_at_least(7, 13)" in source
+ # The patch must be inside the `if not` branch (negated guard)
+ assert "if not _hip_ver_at_least(7, 13):" in source
+
+ def test_grouped_mm_hip_713_skip_message_present(self):
+ """worker.py must log a message when skipping the patch on HIP >= 7.13."""
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ assert "HIP >= 7.13" in source
+ assert "7.13" in source
+
+ def test_grouped_mm_patch_else_branch_present(self):
+ """An else branch must follow the _hip_ver_at_least gate (skip path for 7.13+)."""
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ # There must be an else: after the if not _hip_ver_at_least(7, 13): block
+ gate_idx = source.find("if not _hip_ver_at_least(7, 13):")
+ assert gate_idx != -1, "Version gate not found in worker.py"
+ # The else: branch must appear after the gate
+ else_idx = source.find("else:", gate_idx)
+ assert else_idx != -1, "else: branch after _hip_ver_at_least gate not found"
+
+ def test_hip_ver_at_least_handles_amd_version_format(self):
+ """_hip_ver_at_least must split on '.' and compare only major.minor (handles '7.13.99004')."""
+ source = _WORKER_PATH.read_text(encoding = "utf-8")
+ # Must split the version string and take the first two parts
+ assert 'split(".")[:2]' in source or ".split('.')[:2]" in source
+
+
+# =============================================================================
+# TEST: install_python_stack.py -- _ROCM_TORCH_PKG_SPECS mapping
+# =============================================================================
+
+
+class TestRocmTorchPkgSpecs:
+ """Verify per-tag torch version specs are correct."""
+
+ def test_rocm72_has_torch_211(self):
+ """rocm7.2 should specify torch 2.11.x."""
+ specs = stack_mod._ROCM_TORCH_PKG_SPECS.get("rocm7.2")
+ assert specs is not None
+ torch_spec = specs[0]
+ assert "2.11" in torch_spec
+
+ def test_default_caps_below_211(self):
+ """Default spec (rocm7.1 and earlier) should cap below 2.11."""
+ specs = stack_mod._ROCM_TORCH_PKG_SPECS.get("_default")
+ assert specs is not None
+ torch_spec = specs[0]
+ assert "<2.11" in torch_spec
+
+ def test_specs_have_torch_vision_audio(self):
+ """Each entry should be a 3-tuple: torch, torchvision, torchaudio."""
+ for tag, specs in stack_mod._ROCM_TORCH_PKG_SPECS.items():
+ assert len(specs) == 3, f"{tag}: expected (torch, torchvision, torchaudio)"
+ assert "torch" in specs[0]
+ assert "torchvision" in specs[1]
+ assert "torchaudio" in specs[2]
+
+ def test_gfx_to_amd_index_covers_rdna4(self):
+ """_GFX_TO_AMD_INDEX_ARCH must cover gfx1200 and gfx1201 (RDNA 4)."""
+ mapping = stack_mod._GFX_TO_AMD_INDEX_ARCH
+ assert mapping.get("gfx1200") == "gfx120X-all"
+ assert mapping.get("gfx1201") == "gfx120X-all"
+
+ def test_gfx_to_amd_index_covers_strix_halo(self):
+ """_GFX_TO_AMD_INDEX_ARCH must cover gfx1151 and gfx1150 (RDNA 3.5)."""
+ mapping = stack_mod._GFX_TO_AMD_INDEX_ARCH
+ assert mapping.get("gfx1151") == "gfx1151"
+ assert mapping.get("gfx1150") == "gfx1150"
+
+ def test_gfx_to_amd_index_covers_rdna3(self):
+ """_GFX_TO_AMD_INDEX_ARCH must cover gfx1100-gfx1103 (RDNA 3)."""
+ mapping = stack_mod._GFX_TO_AMD_INDEX_ARCH
+ for arch in ("gfx1100", "gfx1101", "gfx1102", "gfx1103"):
+ assert mapping.get(arch) == "gfx110X-all", f"{arch} missing from mapping"
+
+
+# =============================================================================
+# TEST: setup.ps1 / install.ps1 -- Strix Halo gfx arch detection
+# =============================================================================
+
+_SETUP_PS1_PATH = PACKAGE_ROOT / "studio" / "setup.ps1"
+_INSTALL_PS1_PATH = PACKAGE_ROOT / "install.ps1"
+
+
+class TestStrixHaloGfxArchDetection:
+ """Verify that setup.ps1 and install.ps1 have robust gfx arch detection
+ for Strix Halo / iGPU users who only have the HIP runtime (no hipinfo)."""
+
+ def test_amd_smi_static_asic_attempted_in_setup(self):
+ """setup.ps1 must try 'amd-smi static --asic' when list output lacks gfx arch."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "static --asic" in source
+
+ def test_amd_smi_static_asic_attempted_in_install(self):
+ """install.ps1 must try 'amd-smi static --asic' when list output lacks gfx arch."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "static --asic" in source
+
+ def test_env_var_override_in_setup(self):
+ """setup.ps1 must honour UNSLOTH_ROCM_GFX_ARCH as a manual arch override."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "UNSLOTH_ROCM_GFX_ARCH" in source
+
+ def test_env_var_override_in_install(self):
+ """install.ps1 must honour UNSLOTH_ROCM_GFX_ARCH as a manual arch override."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "UNSLOTH_ROCM_GFX_ARCH" in source
+
+ def test_name_arch_table_covers_strix_halo_in_setup(self):
+ """setup.ps1 name→arch table must map 890M / Strix Halo to gfx1151."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "gfx1151" in source
+ assert "890M" in source or "Strix Halo" in source
+
+ def test_name_arch_table_covers_strix_halo_in_install(self):
+ """install.ps1 name→arch table must map 890M / Strix Halo to gfx1151."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "gfx1151" in source
+ assert "890M" in source or "Strix Halo" in source
+
+ def test_name_arch_table_covers_strix_point_in_setup(self):
+ """setup.ps1 name→arch table must map 880M / Strix Point to gfx1150."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "gfx1150" in source
+ assert "880M" in source or "Strix Point" in source
+
+ def test_name_arch_table_covers_strix_point_in_install(self):
+ """install.ps1 name→arch table must map 880M / Strix Point to gfx1150."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "gfx1150" in source
+ assert "880M" in source or "Strix Point" in source
+
+ def test_name_arch_table_covers_rdna3_phoenix_in_setup(self):
+ """setup.ps1 name→arch table must map 780M / Phoenix to gfx1103."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "gfx1103" in source
+ assert "780M" in source or "Phoenix" in source
+
+ def test_wmi_does_not_set_hasrocm_in_setup(self):
+ """WMI block in setup.ps1 must NOT set $HasROCm = $true (no runtime confirmation)."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ # Find the WMI block and confirm HasROCm is not set inside it
+ wmi_idx = source.find("Win32_VideoController")
+ assert wmi_idx != -1, "WMI block not found in setup.ps1"
+ # The nearest HasROCm = $true must not appear between the WMI block
+ # and the closing brace of that if-block. We check by confirming
+ # $HasROCm = $true does NOT appear within 300 chars of the WMI call.
+ wmi_context = source[wmi_idx : wmi_idx + 300]
+ assert "$HasROCm = $true" not in wmi_context
+
+ def test_gfx_arch_regex_parses_from_amd_smi_output(self):
+ """Both files must use the gfx\\d+[a-z]? regex to parse arch from amd-smi output."""
+ for path in (_SETUP_PS1_PATH, _INSTALL_PS1_PATH):
+ source = path.read_text(encoding = "utf-8")
+ # The regex pattern used to match gfx arches
+ assert (
+ "gfx\\d+" in source or r"gfx\d+" in source
+ ), f"gfx arch regex not found in {path.name}"
+
+
+# =============================================================================
+# TEST: HIP SDK tool path resolution via HIP_PATH / ROCM_PATH env vars
+# =============================================================================
+
+
+class TestHipSdkEnvPathResolution:
+ """Verify that both install scripts resolve hipinfo/hipconfig via HIP_PATH
+ and ROCM_PATH when the tools are not on $PATH, and emit explicit warnings."""
+
+ # ── hipinfo resolution ────────────────────────────────────────────────────
+
+ def test_setup_checks_hip_path_for_hipinfo(self):
+ """setup.ps1 must reference HIP_PATH when resolving hipinfo."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "HIP_PATH" in source
+ assert "hipinfo" in source
+
+ def test_install_checks_hip_path_for_hipinfo(self):
+ """install.ps1 must reference HIP_PATH when resolving hipinfo."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "HIP_PATH" in source
+ assert "hipinfo" in source
+
+ def test_setup_checks_rocm_path_as_hipinfo_fallback(self):
+ """setup.ps1 must also check ROCM_PATH as a secondary hipinfo fallback."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "ROCM_PATH" in source
+ # Confirm the fallback pattern: HIP_PATH ?? ROCM_PATH (or equivalent elseif)
+ assert "ROCM_PATH" in source and "HIP_PATH" in source
+
+ def test_install_checks_rocm_path_as_hipinfo_fallback(self):
+ """install.ps1 must also check ROCM_PATH as a secondary hipinfo fallback."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "ROCM_PATH" in source
+ assert "ROCM_PATH" in source and "HIP_PATH" in source
+
+ def test_setup_resolves_hipinfo_via_bin_subdir(self):
+ """setup.ps1 must join the env var root with 'bin\\hipinfo.exe'."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert r"bin\hipinfo.exe" in source
+
+ def test_install_resolves_hipinfo_via_bin_subdir(self):
+ """install.ps1 must join the env var root with 'bin\\hipinfo.exe'."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert r"bin\hipinfo.exe" in source
+
+ # ── hipinfo not-on-PATH warning ───────────────────────────────────────────
+
+ def test_setup_warns_when_hipinfo_not_on_path(self):
+ """setup.ps1 must warn when hipinfo is found via env var but not on PATH."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "hipinfo not on PATH" in source
+
+ def test_install_warns_when_hipinfo_not_on_path(self):
+ """install.ps1 must warn when hipinfo is found via env var but not on PATH."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "hipinfo not on PATH" in source
+
+ # ── warn when HIP_PATH set but exe missing ────────────────────────────────
+
+ def test_setup_warns_when_hip_path_set_but_exe_missing(self):
+ """setup.ps1 must warn when HIP_PATH is set but hipinfo.exe is not present."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ # The warning must mention that the SDK install may be incomplete
+ assert "incomplete" in source or "not found at" in source
+
+ def test_install_warns_when_hip_path_set_but_exe_missing(self):
+ """install.ps1 must warn when HIP_PATH is set but hipinfo.exe is not present."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "incomplete" in source or "not found at" in source
+
+ # ── hipinfo runtime error warning ─────────────────────────────────────────
+
+ def test_setup_warns_on_hipinfo_nonzero_exit(self):
+ """setup.ps1 must warn when hipinfo runs but returns a non-zero exit code."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "HIP runtime error" in source or "runtime error" in source.lower()
+
+ def test_install_warns_on_hipinfo_nonzero_exit(self):
+ """install.ps1 must warn when hipinfo runs but returns a non-zero exit code."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "HIP runtime error" in source or "runtime error" in source.lower()
+
+ # ── hipconfig resolution ──────────────────────────────────────────────────
+
+ def test_setup_resolves_hipconfig_via_bin_subdir(self):
+ """setup.ps1 must also fall back to HIP_PATH/bin/hipconfig.exe for version detection."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert r"bin\hipconfig.exe" in source
+
+ def test_install_resolves_hipconfig_via_bin_subdir(self):
+ """install.ps1 must also fall back to HIP_PATH/bin/hipconfig.exe for version detection."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert r"bin\hipconfig.exe" in source
+
+ def test_setup_warns_when_hipconfig_not_on_path(self):
+ """setup.ps1 must warn when hipconfig is found via env var but not on PATH."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "hipconfig not on PATH" in source
+
+ def test_install_warns_when_hipconfig_not_on_path(self):
+ """install.ps1 must warn when hipconfig is found via env var but not on PATH."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "hipconfig not on PATH" in source
+
+ # ── PATH fix hint ─────────────────────────────────────────────────────────
+
+ def test_setup_provides_path_fix_hint(self):
+ """setup.ps1 must tell the user how to add the HIP bin dir to PATH."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ # Should mention adding to PATH or SetEnvironmentVariable
+ assert "PATH" in source and (
+ "SetEnvironmentVariable" in source or "Add" in source
+ )
+
+ def test_install_provides_path_fix_hint(self):
+ """install.ps1 must tell the user how to add the HIP bin dir to PATH."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "PATH" in source and (
+ "SetEnvironmentVariable" in source or "Add" in source
+ )
+
+
+# =============================================================================
+# TEST: HIP SDK detected substep -- path + hipconfig version shown in terminal
+# =============================================================================
+
+
+class TestHipSdkDetectedSubstep:
+ """Verify that both scripts print HIP SDK path and full hipconfig version
+ as substeps under the gpu step when AMD ROCm is successfully detected."""
+
+ def test_setup_prints_hip_sdk_path_substep(self):
+ """setup.ps1 must print an 'HIP SDK:' substep showing the resolved path."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "HIP SDK:" in source
+
+ def test_install_prints_hip_sdk_path_substep(self):
+ """install.ps1 must print an 'HIP SDK:' substep showing the resolved path."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "HIP SDK:" in source
+
+ def test_setup_shows_hipconfig_full_version(self):
+ """setup.ps1 must capture and display the full hipconfig version string."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "ROCmVersionFull" in source or "hipconfig:" in source
+
+ def test_install_shows_hipconfig_full_version(self):
+ """install.ps1 must capture and display the full hipconfig version string."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "ROCmVersionFull" in source or "hipconfig:" in source
+
+ def test_setup_captures_full_version_not_just_major_minor(self):
+ """setup.ps1 must store the raw hipconfig output line, not just major.minor."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "ROCmVersionFull" in source
+
+ def test_install_captures_full_version_not_just_major_minor(self):
+ """install.ps1 must store the raw hipconfig output line, not just major.minor."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "ROCmVersionFull" in source
+
+ def test_setup_uses_hip_path_or_rocm_path_for_sdk_display(self):
+ """setup.ps1 HIP SDK path substep must check HIP_PATH then ROCM_PATH."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "HIP_PATH" in source and "ROCM_PATH" in source
+
+ def test_install_uses_hip_path_or_rocm_path_for_sdk_display(self):
+ """install.ps1 HIP SDK path substep must check HIP_PATH then ROCM_PATH."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "HIP_PATH" in source and "ROCM_PATH" in source
+
+ def test_setup_rocm_step_uses_full_version(self):
+ """setup.ps1 'rocm' step label must prefer the full version string."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "ROCmVersionFull" in source and "rocm" in source
+
+
+# =============================================================================
+# TEST: install.sh -- Strix Halo rocm7.1 → rocm7.2 override
+# =============================================================================
+
+_INSTALL_SH_PATH = PACKAGE_ROOT / "install.sh"
+_SETUP_SH_PATH = PACKAGE_ROOT / "studio" / "setup.sh"
+
+
+class TestStrixRocm71Override:
+ """Verify install.sh skips Radeon repo and routes to AMD arch-specific index
+ for gfx1151/gfx1150 when ROCm 7.1 would otherwise be selected (known _grouped_mm segfault).
+ AMD's repo.amd.com/rocm/whl/gfx1151/ serves torch 2.11+rocm7.13 which has the real fix."""
+
+ def test_strix_gfx_detection_in_install_sh(self):
+ """install.sh must detect gfx1151 and gfx1150 for the override."""
+ source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
+ assert "gfx1151" in source and "gfx1150" in source
+
+ def test_rocm71_override_to_amd_arch_index_in_install_sh(self):
+ """install.sh must override TORCH_INDEX_URL to AMD arch-specific index for Strix."""
+ source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
+ # The override must route to AMD's arch-specific index (repo.amd.com/rocm/whl)
+ assert "repo.amd.com/rocm/whl" in source
+ assert "_strix_gfx" in source
+ # 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 + 500]
+ assert "_strix_gfx" in ctx
+
+ def test_radeon_repo_bypassed_for_strix_in_install_sh(self):
+ """install.sh must set _amd_gpu_radeon=false when Strix + ROCm 7.1 detected."""
+ source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
+ assert "_amd_gpu_radeon=false" in source
+
+ def test_strix_override_warns_with_moe_utils_reference(self):
+ """install.sh must emit a [WARN] mentioning the moe_utils segfault."""
+ source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
+ assert "moe_utils" in source or "_grouped_mm" in source
+
+ def test_strix_override_only_fires_on_rocm71(self):
+ """install.sh must scope the Strix override to rocm7.1 only (not rocm7.2+)."""
+ source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
+ strix_idx = source.find("_strix_gfx")
+ assert strix_idx != -1
+ # Look back for the rocm7.1 pattern within 600 chars before _strix_gfx
+ context_before = source[max(0, strix_idx - 2400) : strix_idx]
+ assert "rocm7.1" in context_before
+
+ def test_torch_constraint_updated_for_strix_amd_index(self):
+ """install.sh must set TORCH_CONSTRAINT>=2.11 when routing Strix to AMD index."""
+ source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
+ assert "TORCH_CONSTRAINT" in source and "2.11" in source
+
+ def test_amd_rocm_mirror_env_var_respected(self):
+ """install.sh must honour UNSLOTH_AMD_ROCM_MIRROR for air-gapped installs."""
+ source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
+ assert "UNSLOTH_AMD_ROCM_MIRROR" in source
+
+ def test_tauri_family_recognises_amd_arch_url(self):
+ """_tauri_torch_index_family must return a rocm* family for AMD arch-specific URLs."""
+ source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
+ # The function must have a case branch for repo.amd.com/rocm/whl/gfx* URLs
+ assert "rocm/whl/gfx" in source
+
+
+# =============================================================================
+# TEST: setup.sh -- gcc-install-dir fix for Ubuntu 24.04 + ROCm 7.x clang-20
+# =============================================================================
+
+
+class TestSetupShGccInstallDir:
+ """Verify setup.sh applies the --gcc-install-dir flag when building llama.cpp
+ with HIP on Ubuntu 24.04+ to work around ROCm 7.x clang-20 header path bug."""
+
+ def test_gcc_install_dir_search_loop_present(self):
+ """setup.sh must iterate gcc versions 14→11 to find one with C++ headers."""
+ source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
+ assert "_GCC_INSTALL_DIR" in source
+ assert "/usr/lib/gcc/x86_64-linux-gnu" in source
+
+ def test_gcc_install_dir_checks_include_dir(self):
+ """setup.sh must check that the gcc dir has an 'include' subdirectory."""
+ source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
+ assert "include" in source and "_GCC_INSTALL_DIR" in source
+
+ def test_gcc_install_dir_appended_to_cmake_hip_flags(self):
+ """setup.sh must pass --gcc-install-dir via CMAKE_HIP_FLAGS."""
+ source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
+ assert "CMAKE_HIP_FLAGS" in source
+ assert "gcc-install-dir" in source
+
+ def test_gcc_install_dir_only_applied_in_hip_build_block(self):
+ """The --gcc-install-dir fix must only apply in the HIP/ROCm build branch."""
+ source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
+ # GGML_HIP=ON must appear before gcc-install-dir in the source
+ hip_idx = source.find("GGML_HIP=ON")
+ gcc_idx = source.find("gcc-install-dir")
+ assert hip_idx != -1 and gcc_idx != -1
+ assert hip_idx < gcc_idx
+
+ def test_gcc_install_dir_logs_substep(self):
+ """setup.sh must print a substep when the gcc install dir is resolved."""
+ source = _SETUP_SH_PATH.read_text(encoding = "utf-8")
+ assert "gcc install dir" in source or "GCC_INSTALL_DIR" in source
+
+
+# =============================================================================
+# TEST: main.py -- BNB_ROCM_VERSION server startup + distributed stubs
+# =============================================================================
+
+_MAIN_PY_PATH = PACKAGE_ROOT / "studio" / "backend" / "main.py"
+_HARDWARE_PY_PATH = (
+ PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
+)
+
+
+class TestServerStartupRocmFixes:
+ """Verify main.py sets BNB_ROCM_VERSION before any bitsandbytes import and
+ hardware.py injects torch._C._distributed_c10d stubs before torch.distributed."""
+
+ # ── BNB_ROCM_VERSION in server process ────────────────────────────────────
+
+ def test_main_py_sets_bnb_rocm_version(self):
+ """main.py must set BNB_ROCM_VERSION in the server process before imports."""
+ source = _MAIN_PY_PATH.read_text(encoding = "utf-8")
+ assert "BNB_ROCM_VERSION" in source
+
+ def test_main_py_bnb_detection_scoped_to_win32(self):
+ """main.py BNB_ROCM_VERSION logic must be inside the win32 platform guard."""
+ source = _MAIN_PY_PATH.read_text(encoding = "utf-8")
+ win32_idx = source.find('sys.platform == "win32"')
+ bnb_idx = source.find("BNB_ROCM_VERSION")
+ assert win32_idx != -1 and bnb_idx != -1
+ assert win32_idx < bnb_idx
+
+ def test_main_py_bnb_dll_detection_uses_glob(self):
+ """main.py must scan for libbitsandbytes_rocm*.dll to find the right version."""
+ source = _MAIN_PY_PATH.read_text(encoding = "utf-8")
+ assert "libbitsandbytes_rocm" in source
+
+ def test_main_py_bnb_falls_back_to_72(self):
+ """main.py must fall back to BNB_ROCM_VERSION='72' when no DLL is found."""
+ source = _MAIN_PY_PATH.read_text(encoding = "utf-8")
+ assert '"72"' in source or "'72'" in source
+
+ def test_main_py_bnb_only_set_when_not_already_in_env(self):
+ """main.py must not override an existing BNB_ROCM_VERSION env var."""
+ source = _MAIN_PY_PATH.read_text(encoding = "utf-8")
+ assert '"BNB_ROCM_VERSION" not in os.environ' in source
+
+ # ── torch._C._distributed_c10d stubs in hardware.py ──────────────────────
+
+ def test_hardware_py_injects_distributed_c10d_stub(self):
+ """hardware.py must inject torch._C._distributed_c10d into sys.modules."""
+ source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
+ assert "_distributed_c10d" in source
+
+ def test_hardware_py_stub_injected_before_distributed_import(self):
+ """The sys.modules stub must be injected BEFORE import torch.distributed."""
+ source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
+ c10d_idx = source.find("_distributed_c10d")
+ dist_idx = source.find("import torch.distributed")
+ assert c10d_idx != -1 and dist_idx != -1
+ assert c10d_idx < dist_idx
+
+ def test_hardware_py_stub_uses_types_moduletype(self):
+ """hardware.py must create the stub with types.ModuleType."""
+ source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
+ assert "ModuleType" in source
+
+ def test_hardware_py_stub_scoped_to_win32(self):
+ """hardware.py distributed stub injection must be gated on win32."""
+ source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
+ assert 'platform == "win32"' in source or "win32" in source
+
+ def test_hardware_py_stub_exposes_fake_process_group(self):
+ """hardware.py stub must set FakeProcessGroup so torch.distributed doesn't raise AttributeError."""
+ source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
+ assert "FakeProcessGroup" in source
+
+ def test_hardware_py_stub_exposes_process_group(self):
+ """hardware.py stub must set ProcessGroup on the c10d stub."""
+ source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
+ assert "ProcessGroup" in source
+
+ def test_hardware_py_stub_uses_setattr_for_symbols(self):
+ """hardware.py must use setattr to populate stub symbols dynamically."""
+ source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
+ assert "setattr" in source
+
+ def test_hardware_py_stub_all_c10d_siblings_covered(self):
+ """hardware.py must stub all three torch._C._distributed_* submodules."""
+ source = _HARDWARE_PY_PATH.read_text(encoding = "utf-8")
+ assert "_distributed_c10d" in source
+ assert "_distributed_autograd" in source
+ assert "_distributed_rpc" in source
+
+
+# =============================================================================
+# TEST: install.ps1 / setup.ps1 -- HipSdkInstalled flag (SDK found, device inaccessible)
+# =============================================================================
+
+
+class TestHipSdkInstalledButDeviceInaccessible:
+ """Verify that when hipinfo is found but exits non-zero (device not ROCm-accessible),
+ both scripts distinguish this from 'HIP SDK not found' and emit the correct message."""
+
+ def test_install_ps1_has_hip_sdk_installed_flag(self):
+ """install.ps1 must track HipSdkInstalled separately from HasROCm."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "HipSdkInstalled" in source
+
+ def test_setup_ps1_has_hip_sdk_installed_flag(self):
+ """setup.ps1 must track HipSdkInstalled separately from HasROCm."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "HipSdkInstalled" in source
+
+ def test_install_ps1_sets_flag_when_hipinfo_binary_found(self):
+ """install.ps1 must set HipSdkInstalled=true inside the 'if ($hipinfoExe)' block."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ # HipSdkInstalled must be assigned inside the hipinfoExe block
+ hipinfo_block_idx = source.find("if ($hipinfoExe)")
+ sdk_flag_idx = source.find("$HipSdkInstalled = $true", hipinfo_block_idx)
+ assert hipinfo_block_idx != -1 and sdk_flag_idx != -1
+ assert sdk_flag_idx > hipinfo_block_idx
+
+ def test_setup_ps1_sets_flag_when_hipinfo_binary_found(self):
+ """setup.ps1 must set HipSdkInstalled=true inside the 'if ($hipinfoExe)' block."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ hipinfo_block_idx = source.find("if ($hipinfoExe)")
+ sdk_flag_idx = source.find("$HipSdkInstalled = $true", hipinfo_block_idx)
+ assert hipinfo_block_idx != -1 and sdk_flag_idx != -1
+ assert sdk_flag_idx > hipinfo_block_idx
+
+ def test_install_ps1_version_capture_runs_when_sdk_installed(self):
+ """install.ps1 must capture hipconfig version when HipSdkInstalled even if HasROCm is false."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "HasROCm -or $HipSdkInstalled" in source or "$HipSdkInstalled" in source
+
+ def test_setup_ps1_version_capture_runs_when_sdk_installed(self):
+ """setup.ps1 must capture hipconfig version when HipSdkInstalled even if HasROCm is false."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "HasROCm -or $HipSdkInstalled" in source or "$HipSdkInstalled" in source
+
+ def test_install_ps1_distinct_message_for_sdk_found_but_device_inaccessible(self):
+ """install.ps1 must show 'not ROCm-accessible' message (not 'HIP SDK not found') when SDK present."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "not ROCm-accessible" in source
+
+ def test_setup_ps1_distinct_message_for_sdk_found_but_device_inaccessible(self):
+ """setup.ps1 must show 'not ROCm-accessible' message (not 'HIP SDK not found') when SDK present."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "not ROCm-accessible" in source
+
+ def test_install_ps1_driver_guidance_in_sdk_found_branch(self):
+ """install.ps1 must tell user this is a driver issue, not an SDK issue."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "driver issue" in source
+
+ def test_setup_ps1_driver_guidance_in_sdk_found_branch(self):
+ """setup.ps1 must tell user this is a driver issue, not an SDK issue."""
+ source = _SETUP_PS1_PATH.read_text(encoding = "utf-8")
+ assert "driver issue" in source
+
+ def test_install_ps1_cpu_hint_distinguishes_driver_vs_no_sdk(self):
+ """install.ps1 CPU-only hint must say 'GPU not ROCm-accessible' not 'require the HIP SDK' when SDK found."""
+ source = _INSTALL_PS1_PATH.read_text(encoding = "utf-8")
+ assert "GPU not ROCm-accessible" in source
+
+
if __name__ == "__main__":
pytest.main([__file__, "-v"])
diff --git a/unsloth/models/loader_utils.py b/unsloth/models/loader_utils.py
index 99da5f799e..c0c6455a2e 100644
--- a/unsloth/models/loader_utils.py
+++ b/unsloth/models/loader_utils.py
@@ -68,7 +68,10 @@ def _get_env_int(keys):
def _infer_distributed_ranks():
- if torch.distributed.is_available() and torch.distributed.is_initialized():
+ if (
+ torch.distributed.is_available()
+ and getattr(torch.distributed, "is_initialized", lambda: False)()
+ ):
try:
return torch.distributed.get_rank(), torch.distributed.get_world_size()
except Exception:
]