diff --git a/install.sh b/install.sh index 8e30264df0..e3fdf10cdd 100755 --- a/install.sh +++ b/install.sh @@ -1994,11 +1994,18 @@ get_torch_index_url() { # is the convenience form (cpu, cu124, cu126, cu128, cu130, rocm6.4, ...) # appended to the mirror base so UNSLOTH_PYTORCH_MIRROR is still honoured. if [ -n "${UNSLOTH_TORCH_INDEX_URL:-}" ]; then - echo "${UNSLOTH_TORCH_INDEX_URL%/}"; return + # Strip ALL trailing slashes (match the Python side's .rstrip("/") and the + # Strix mirror handling below) -- a double/triple-slash URL 404s on strict + # pip proxies (artifactory, sonatype). + _url="${UNSLOTH_TORCH_INDEX_URL}" + while [ "${_url%/}" != "$_url" ]; do _url="${_url%/}"; done + echo "$_url"; return fi if [ -n "${UNSLOTH_TORCH_INDEX_FAMILY:-}" ]; then - _family="${UNSLOTH_TORCH_INDEX_FAMILY#/}" - echo "$_base/${_family%/}"; return + _family="${UNSLOTH_TORCH_INDEX_FAMILY}" + while [ "${_family#/}" != "$_family" ]; do _family="${_family#/}"; done + while [ "${_family%/}" != "$_family" ]; do _family="${_family%/}"; done + echo "$_base/$_family"; return fi # macOS: always CPU (no CUDA support) case "$(uname -s)" in Darwin) echo "$_base/cpu"; return ;; esac @@ -2396,7 +2403,16 @@ _maybe_bootstrap_rocm_wsl() { [ -n "$_rw_tmp" ] && rm -f "$_rw_tmp" return 0 } -_maybe_bootstrap_rocm_wsl || true +# When the caller pins the wheel index (UNSLOTH_TORCH_INDEX_URL / _FAMILY), +# honour it everywhere downstream: skip the WSL ROCm bootstrap (which can run +# sudo + large downloads after probing /dev/dxg) and the Radeon/Strix rerouting +# below (which would re-probe the GPU and overwrite the pinned URL). A headless / +# container / CI build must get exactly the index it asked for. +_torch_index_pinned=false +if [ -n "${UNSLOTH_TORCH_INDEX_URL:-}" ] || [ -n "${UNSLOTH_TORCH_INDEX_FAMILY:-}" ]; then + _torch_index_pinned=true +fi +[ "$_torch_index_pinned" = true ] || _maybe_bootstrap_rocm_wsl || true TORCH_INDEX_URL=$(get_torch_index_url) @@ -2424,7 +2440,11 @@ 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". +# Skipped entirely when the index is pinned: an explicit override (even a ROCm +# one like UNSLOTH_TORCH_INDEX_FAMILY=rocm6.4) must not be rerouted to the +# Radeon/Strix repos by GPU probing. _amd_gpu_radeon=false +if [ "$_torch_index_pinned" = false ]; then case "$TORCH_INDEX_URL" in */rocm*) if _has_amd_rocm_gpu && command -v rocminfo >/dev/null 2>&1 && \ @@ -2505,6 +2525,7 @@ case "$TORCH_INDEX_URL" in fi ;; esac +fi # _torch_index_pinned guard (Radeon + Strix reroute) _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" diff --git a/studio/install_python_stack.py b/studio/install_python_stack.py index 222cfacb72..a2c1e44556 100644 --- a/studio/install_python_stack.py +++ b/studio/install_python_stack.py @@ -1429,6 +1429,23 @@ NO_TORCH = _infer_no_torch() # GPU detection. Values: "cuda", "rocm", or "cpu". Empty means unknown # (standalone `unsloth studio update` runs, where we re-detect normally). _TORCH_BACKEND: str = os.environ.get("UNSLOTH_TORCH_BACKEND", "").lower() +# When install.sh did not run (standalone `unsloth studio update`) but the caller +# pinned the wheel index explicitly, derive the backend from that override so the +# CUDA/ROCm repair helpers honour it instead of re-probing the GPU and possibly +# reinstalling a different family. Classify on the final URL/family segment, +# mirroring install.sh's UNSLOTH_TORCH_BACKEND case. +if not _TORCH_BACKEND: + _idx_override = ( + os.environ.get("UNSLOTH_TORCH_INDEX_URL", "").strip() + or os.environ.get("UNSLOTH_TORCH_INDEX_FAMILY", "").strip() + ) + _idx_leaf = _idx_override.rstrip("/").rsplit("/", 1)[-1].lower() + if _idx_leaf.startswith(("rocm", "gfx")): + _TORCH_BACKEND = "rocm" + elif _idx_leaf == "cpu": + _TORCH_BACKEND = "cpu" + elif _idx_leaf.startswith("cu"): + _TORCH_BACKEND = "cuda" def _torch_step_label(suffix: str) -> str: diff --git a/tests/sh/test_get_torch_index_url.sh b/tests/sh/test_get_torch_index_url.sh index b9d0c5e5a1..0c538f0a12 100755 --- a/tests/sh/test_get_torch_index_url.sh +++ b/tests/sh/test_get_torch_index_url.sh @@ -415,6 +415,14 @@ assert_eq "url override beats family override -> url" "https://mirror.example.co _result=$(UNSLOTH_TORCH_INDEX_FAMILY="" UNSLOTH_TORCH_INDEX_URL="" run_func "none") assert_eq "empty overrides ignored -> detected cpu" "https://download.pytorch.org/whl/cpu" "$_result" +# 46) ALL trailing slashes are stripped from a URL override (not just one). +_result=$(UNSLOTH_TORCH_INDEX_URL="https://mirror.example.com/whl/cu128///" run_func "none") +assert_eq "url override double slash stripped" "https://mirror.example.com/whl/cu128" "$_result" + +# 47) Leading and trailing slashes stripped from a family override. +_result=$(UNSLOTH_TORCH_INDEX_FAMILY="//cu128//" run_func "none") +assert_eq "family override slashes stripped" "https://download.pytorch.org/whl/cu128" "$_result" + rm -f "$_FUNC_FILE" rm -rf "$_FAKE_SMI_DIR" rm -rf "$_TOOLS_DIR"