From 5f1f30ec82d097d92d1093d1d557427dcbf079e6 Mon Sep 17 00:00:00 2001
From: oobabooga
Date: Sun, 19 Jul 2026 09:46:22 -0300
Subject: [PATCH 01/41] Studio: GPU memory configuration for GGUF models
(#6414)
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
* Studio: GPU memory dropdown — llama.cpp --fit on and manual gpu-layers/cpu-moe
* Studio: simplify GPU memory changes (reuse ParamSlider, GPU_LAYERS_ALL, loadedGpuMemoryFields helper)
* Studio: GPU picker — choose which GPUs a GGUF model loads on (gpu_ids)
* Studio: simplify GPU picker (share /api/system fetch, validate gpu_ids)
* Studio: GPU picker review fixes (gate relative indices, no cross-model leak, validate, types)
* Studio: group GPU controls under a collapsible GPU section
* Studio: GPU feature review fixes (fix fit-ctx test, behavior-test the floor, comment accuracy)
* Studio: make GPU a top-level settings section (not nested under Model)
* Studio: flatten GPU controls into the Model section, group by GPU/context/generation
* Studio: move GPU Memory to the bottom of Model with its dependent controls beneath it
* Studio: move GPU Memory below Tensor Parallelism and GPUs below GPU Memory
* Studio: tighten GPU Memory and GPU Layers tooltip copy
* Studio: fix fit-mode context slider track-click, restore GPU Memory tooltip, shorten fit dropdown label
* Studio: GPU Memory tooltip one mode per line, briefer
* Studio: note HIP_VISIBLE_DEVICES (ROCm) in the GPUs picker tooltip
* Studio: narrow the GPU Memory dropdown to fit the shortened label
* Studio: use 'llama.cpp --fit' in the GPU Memory tooltip for consistency
* Studio: allow Tensor Parallelism in Manual GPU mode
* Studio: graduated MoE-on-CPU offload (--n-cpu-moe) replacing the all-or-nothing toggle
* Studio: size the MoE-offload slider for staged (deferred-load) models
* Studio: share one GGUF header walk for the context-length and MoE-count readers
* Studio: size the GPU Layers slider for staged models (one staged-header read)
* Studio: move Tensor Parallelism below the GPUs picker
* Studio: GPU split (--tensor-split) per-GPU model share in Manual mode
* Studio: tolerate whitespace in GPU split input, move it below GPU Layers
* Studio: rename the GPU split control to "Split ratio"
* Studio: Split ratio sends explicit even input; fix blank=free-VRAM (not even) copy
* Studio: tighten llama.cpp --fit VRAM margin with --fit-target 512
* Studio: GPU memory review fixes (rollback re-baseline, single-GPU TP gate, accurate copy)
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Studio: move Split ratio below MoE Layers on CPU
* Studio: address PR review (fix GPU-info hydration race, share fit context-length across load paths)
* Studio: address codex review (manual single-GPU TP guard, GPU-aware spec defaults in fit/manual, GGUF-only context/preference)
* Studio: address codex review round 2 (gpu_present seed, single-GPU tensor-split guard, staged manual-knob reset, strip inherited offload flags)
* Studio: address codex review round 3 (strip inherited --n-cpu-moe, CPU-fallback warning in Manual mode)
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Studio: address codex review round 4 (preserve pinned fit context across a later Apply)
* Studio: address codex review round 5 (honor GPU picker for diffusion GGUFs, clear fit pin on cross-model switch)
* Studio: preserve the pending GPU Memory mode when staging a model
* Studio: pin diffusion GPU device order and reset GPU-memory state for diffusion loads
* Studio: address codex review round 6 (fit-Auto rollback context, preserve manual non-tensor split modes, persist GPU mode on load not select)
* Studio: persist the applied GPU Memory mode, not the requested one (skip diffusion loads)
* Studio: replace Manual-mode split-ratio field with per-GPU layer sliders
* Studio: clarify per-GPU layer split hint for tensor-parallel mode
* Studio: address codex review round 7 (allow GGUF gpu_ids past the legacy guard, replay GPU-memory fields on respawn)
* Studio: address codex review round 8 (size the validate preflight like the load in fit mode, across both load paths)
* Studio: skip the training-OOM guard for llama.cpp --fit GGUF loads (they spill to RAM)
* Studio: drop the now-redundant compare-path validate sizing (the --fit guard skip makes it moot)
* Studio: address codex review round 9 (keep the training guard for fit loads, forward gpu_ids to validate, strip inherited manual tensor-split)
* Studio: address codex review round 10 (gate GPU-memory adoption on is_gguf, record manual knobs only in Manual mode)
* Studio: handle diffusion GGUFs symmetrically in the GPU Memory controls (preserve the standing mode preference, hide the inapplicable mode/TP controls)
* Studio: remember the GPU Memory settings per model
* Studio: consolidate --fit mode and Manual mode into a single Manual mode
* Studio: preserve the per-GPU layer split across GPU Layers changes
* Studio: trim overly long GPU Memory comments
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* address GPU memory config review comments
* trim redundant GPU memory tests
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Reconcile manual-mode TP drops with the #6659 drop-site invariants
* Preserve quantized KV in manual --fit, charge GGUF companions in full, reconcile GPU pick on load
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Clear stale GPU baseline on non-GGUF loads so it can't read as dirty
* Fix no-context-shift test for the conditional -c flag
* Credit manual GPU-layer offload for cached HF GGUFs
* Reset per-model load knobs on GGUF quant switch
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Strip inherited tensor-split when manual ratio is cleared
* Match auto-load validation to safetensors placement
* Reset editable manual knobs after Auto GGUF loads
* Record a single device for diffusion GPU picks
* Reset per-model GPU knobs before applying saved settings
* Address review comments
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Guard manual tensor splits and keep remembered context on auto-load
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Snapshot compare knobs, seed splits from free VRAM, flag zero-offload loads
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Exempt CPU-only loads from the guard floor and harden compare and reseed paths
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Reach full offload from the layers slider and charge extras drafters in the guard
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Warm the GPU device cache before pick reconciles and disable staged GPU controls
* Align the training guard with inherited extras, spec mode, and compare targets
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Hide GPUs from companion-less zero-offload loads
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Size diffusion picks per device, own manual offload flags, reject XPU picks
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Drop tensor flags at zero layers and exempt CPU-pinned drafters
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Allowlist the zero-layer tensor parallel drop site
* Keep validate and load guards on the same extras and refresh stale baselines
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Drop mismatched manual tensor splits before launch
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Gate XPU picks on the real backend field and harden split and hydration paths
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Weight full GPUs as zero, clamp split shares, and refine the zero-layer mask gate
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Carry fit context across mode changes and align drafter and picker gates
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Catch variant switches, uncached diffusion repos, and text-only mmproj skips
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Check companions on the first device and size native and remote zero-layer loads
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Replace the training guard's precise VRAM modeling with a conservative bound
* Baseline context pins on non-GGUF hydration and reprobe list-seeded staged GGUFs
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Size manual splits by their largest share and preserve resolved context from Default
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Default-deny unsized required companions and price KV at the effective cache dtype
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Reserve MTP draft KV and MLA target-copy in the training guard
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Size tensor-parallel loads per device and show GPU controls for native GGUFs
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Reserve MTP overhead for uncached remote GGUFs and the mmproj runtime factor
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Drop the training-coexistence VRAM estimation this PR added
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Gate remembered load settings to GGUF picks
* Lock the remaining load-time controls during a staged load
* Clear the stale native-path token on compare loads
* Drop a stale guard reference from the zero-offload masking comment
* Seed GPU baselines from the rollback response and drop never-emitted offload flags
* Match validate's training guard to load and keep the native reload token
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Trim verbose GPU-memory comments
* Thread the variants header walk off the event loop, honor device pins on zero-offload, and hold staged GPU edits
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Honor manual placement and classify pinned zero-offload loads
* Close diffusion admission and status hydration gaps
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Check the actual diffusion GPU during training
* Align staged baselines and manual reload dedupe
* Fix GGUF placement and rollback state
* Harden manual GGUF placement boundaries
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Remove unused resolve_tensor_parallel import in llama_cpp.py
The name is used only in llama_server_args.py, routes/inference.py, and tests,
not in llama_cpp.py; the unused hoisted import trips the import-hoist verifier
in the source-lint CI job.
* Fix diffusion GPU dedup and training guard for non-numeric device tokens
The diffusion runner drives only its single lowest device and the backend
records that one device (self._gpu_ids = [sorted(gpu_ids)[0]]), but the reload
dedupe compared it against the full requested list, so a multi-GPU pick that
resolves to the same device forced a needless reload. Normalize the request the
same way for a loaded diffusion model in both _already_in_target_state and the
route _request_matches_loaded_settings.
The chat-during-training coexistence guard called int() on the single-device
token and hard-rejected when it could not parse. A non-numeric token (a CUDA
UUID / MIG handle) now sizes against the whole visible pool like the GGUF guard
instead of falsely blocking the load, and an empty token (a CPU-only runner such
as a CPU diffusion GGUF) is allowed outright since it uses no GPU VRAM.
* Tighten comments added by the GPU memory config changes
* Harden GGUF placement from independent review: VRAM sizing, diffusion TP reset, tensor_split validation
- Training coexistence guard: a single-device runner pinned through an
unresolvable UUID/MIG token was sized against the aggregate visible-VRAM pool,
so a load could pass on capacity it cannot use and then OOM active training.
Size against the worst-case visible device (min free) instead, keeping the
guard's documented default-deny contract. The empty-token (CPU-only runner)
allow path is unchanged.
- Diffusion startup: _start_diffusion_server now resets self._tensor_parallel to
False alongside the other placement resets. A prior tensor-parallel chat load
(process killed but not fully unload-reset) otherwise left /status misreporting
tensor parallelism and made an identical diffusion re-Apply reload against the
stale state.
- tensor_split: reject negative / non-finite / all-zero splits up front. They
were dropped at launch but still compared raw in the reload dedupe, so an
identical Apply reloaded indefinitely.
- Tests: the shared httpx stub was incomplete and, installed via setdefault
before real httpx loaded, broke a combined pytest run (collection errors on
httpx.Response). Import the real installed httpx instead.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
---------
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: danielhanchen
Co-authored-by: danielhanchen
---
studio/backend/core/inference/llama_cpp.py | 715 +++++++++++++-
.../core/inference/llama_server_args.py | 27 +-
studio/backend/main.py | 14 +
studio/backend/models/inference.py | 144 ++-
studio/backend/routes/inference.py | 518 ++++++++---
studio/backend/routes/models.py | 6 +-
studio/backend/routes/training_vram.py | 42 +-
.../tests/test_chat_load_during_training.py | 333 ++++++-
studio/backend/tests/test_gguf_metadata.py | 73 ++
studio/backend/tests/test_gpu_memory_mode.py | 879 ++++++++++++++++++
studio/backend/tests/test_gpu_selection.py | 18 +-
.../tests/test_llama_cpp_no_context_shift.py | 12 +-
.../tests/test_llama_cpp_props_readback.py | 31 +-
.../backend/tests/test_llama_server_args.py | 45 +
studio/backend/tests/test_tensor_parallel.py | 5 +-
.../tests/test_tp_vision_regression.py | 17 +-
studio/backend/utils/models/gguf_metadata.py | 109 ++-
.../remembered-load-settings.ts | 24 +-
.../src/features/chat/api/chat-adapter.ts | 128 ++-
.../src/features/chat/api/chat-api.ts | 38 +-
.../frontend/src/features/chat/chat-page.tsx | 13 +-
.../src/features/chat/chat-settings-sheet.tsx | 450 ++++++++-
.../chat/hooks/use-chat-model-runtime.ts | 217 ++++-
.../hooks/use-staged-model-preparation.ts | 46 +-
.../lib/apply-inference-status-to-store.ts | 117 ++-
.../features/chat/presets/preset-policy.ts | 31 +
.../src/features/chat/shared-composer.tsx | 109 ++-
.../chat/stores/chat-runtime-store.ts | 369 +++++++-
.../frontend/src/features/chat/types/api.ts | 38 +
studio/frontend/src/hooks/use-gpu-info.ts | 186 ++--
studio/frontend/src/hooks/use-system.ts | 3 +
31 files changed, 4356 insertions(+), 401 deletions(-)
create mode 100644 studio/backend/tests/test_gpu_memory_mode.py
diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py
index c9ab7eb83b..d7c7eed518 100644
--- a/studio/backend/core/inference/llama_cpp.py
+++ b/studio/backend/core/inference/llama_cpp.py
@@ -11,6 +11,7 @@ import atexit
import contextlib
import functools
import json
+import math
import os
import re
import struct
@@ -23,11 +24,22 @@ import sys
import threading
import time
from pathlib import Path
-from typing import Callable, Collection, Generator, Iterable, List, Mapping, Optional, Union
+from typing import (
+ Callable,
+ Collection,
+ Generator,
+ Iterable,
+ List,
+ Literal,
+ Mapping,
+ Optional,
+ Union,
+)
import httpx
from core.inference.llama_server_args import (
+ _LAYER_OFFLOAD_FLAGS,
_effective_tensor_parallel,
_tensor_parallel_matches_loaded,
extra_args_disable_mmproj,
@@ -234,8 +246,7 @@ def _wsl_system_rocm_lib_dirs() -> "list[str]":
return out
-# Plan-without-action re-prompt state (intent signal, caps, message) now lives
-# in tool_call_parser, imported above under its old aliases.
+# Plan-without-action re-prompt state now lives in tool_call_parser (imported above).
# Default max_tokens to the effective context when known. The floor is high
# enough for reasoning-heavy GGUFs and max_tokens-omitting API clients.
@@ -1431,7 +1442,10 @@ def _extra_args_set_spec_type(extra_args: Optional[Iterable[str]]) -> bool:
return _extra_args_set_any_flag(extra_args, {"--spec-type", "--spec-default"})
-_GPU_OFFLOAD_OVERRIDE_FLAGS = frozenset({"-ngl", "--gpu-layers", "--n-gpu-layers", "-fit", "--fit"})
+# Layer-offload override detection. Single-sourced from llama_server_args, which
+# also strips these (plus the MoE flags) from inherited extras; sharing the layer
+# set keeps detection and stripping from drifting.
+_GPU_OFFLOAD_OVERRIDE_FLAGS = _LAYER_OFFLOAD_FLAGS
_THREAD_OVERRIDE_FLAGS = frozenset({"-t", "--threads"})
@@ -1895,6 +1909,17 @@ class LlamaCppBackend:
self._cache_type_kv: Optional[str] = None
# Whether --split-mode tensor was applied on the active load.
self._tensor_parallel: bool = False
+ # GPU memory strategy applied on the active load ("auto"/"manual").
+ self._gpu_memory_mode: str = "auto"
+ # Manual-mode load options (echoed back so the UI round-trips them).
+ self._gpu_layers: int = -1
+ # MoE expert layers to keep on CPU (--n-cpu-moe); 0 = none.
+ self._n_cpu_moe: int = 0
+ # Relative model share per GPU (--tensor-split), in GPU order; None =
+ # default (llama.cpp splits by free VRAM).
+ self._tensor_split: Optional[List[float]] = None
+ # User-picked physical GPU indices (None = automatic selection).
+ self._gpu_ids: Optional[List[int]] = None
# Layer load kept multi-GPU only to honor a downgraded tensor request, so a
# later explicit tensor-off reloads instead of deduping to it (#6659).
self._layer_preserves_tensor_intent: bool = False
@@ -1909,6 +1934,11 @@ class LlamaCppBackend:
self._spec_draft_n_max: Optional[int] = None
# KV-cache estimation fields (populated by _read_gguf_metadata)
self._n_layers: Optional[int] = None
+ # MoE metadata (populated by _read_gguf_metadata): expert count (>0 =
+ # MoE) and leading dense-layer count (offsets --n-cpu-moe, which counts
+ # from layer 0). See the n_moe_layers property.
+ self._n_experts: Optional[int] = None
+ self._leading_dense_block_count: Optional[int] = None
self._n_kv_heads: Optional[int] = None
self._n_kv_heads_by_layer: Optional[list[int]] = None
self._n_heads: Optional[int] = None
@@ -2329,6 +2359,79 @@ class LlamaCppBackend:
"""Whether --split-mode tensor is active on the loaded server."""
return self._tensor_parallel
+ @property
+ def gpu_memory_mode(self) -> str:
+ """Active GPU memory strategy: 'auto' or 'manual' (gpu_layers < 0 = Auto/--fit, >= 0 = pinned)."""
+ return self._gpu_memory_mode
+
+ @property
+ def gpu_layers(self) -> int:
+ """Requested --gpu-layers for manual mode (-1 when not manual)."""
+ return self._gpu_layers
+
+ @property
+ def n_cpu_moe(self) -> int:
+ """MoE expert layers manual mode kept on CPU (--n-cpu-moe); 0 = none."""
+ return self._n_cpu_moe
+
+ @property
+ def tensor_split(self) -> Optional[List[float]]:
+ """Manual-mode relative model share per GPU (--tensor-split); None =
+ default (split by free VRAM)."""
+ return self._tensor_split
+
+ @property
+ def gpu_ids(self) -> Optional[List[int]]:
+ """User-picked physical GPU indices, or None for automatic selection."""
+ return self._gpu_ids
+
+ @property
+ def n_layers(self) -> Optional[int]:
+ """Model layer count (GGUF block_count), or None if unknown."""
+ return self._n_layers
+
+ @property
+ def n_moe_layers(self) -> int:
+ """Number of MoE expert layers (the --n-cpu-moe ceiling), 0 if not MoE.
+
+ block_count minus the leading dense layers (which carry no experts):
+ --n-cpu-moe counts from layer 0, so those dense layers are no-ops.
+ """
+ if not self._n_experts or not self._n_layers:
+ return 0
+ return max(0, self._n_layers - (self._leading_dense_block_count or 0))
+
+ @staticmethod
+ def _resolve_cpu_moe_flag(
+ n_cpu_moe: int, n_moe_layers: int, leading_dense: int
+ ) -> Optional[int]:
+ """The --n-cpu-moe value (absolute first-N layers), or None to omit it.
+
+ Clamps the requested count to the model's MoE layers, then offsets past
+ the leading dense layers (--n-cpu-moe counts from layer 0). Returns None
+ for nothing-to-offload (0 requested) or a non-MoE model.
+ """
+ if n_cpu_moe <= 0 or n_moe_layers <= 0:
+ return None
+ return leading_dense + min(n_cpu_moe, n_moe_layers)
+
+ @staticmethod
+ def _sanitize_tensor_split(tensor_split: Optional[List[float]]) -> List[float]:
+ """Per-GPU shares with negative and non-finite entries clamped to 0.
+
+ A direct caller's negative entry would launch a placement different
+ from the ratio the UI showed, and inf would pass a plain ``> 0`` total
+ gate and emit ``--tensor-split inf,...``. Returns [] for input that
+ can't be read as floats (the length gate at the call site then drops
+ the split).
+ """
+ try:
+ return [
+ x if math.isfinite(x) and x > 0.0 else 0.0 for x in (float(v) for v in tensor_split)
+ ]
+ except (TypeError, ValueError, OverflowError):
+ return []
+
@property
def layer_preserves_tensor_intent(self) -> bool:
"""True when a downgraded tensor request kept this layer load multi-GPU."""
@@ -2530,6 +2633,7 @@ class LlamaCppBackend:
"spec_draft_n_max_flag": None,
"supports_kv_unified": False,
"supports_fit_ctx": False,
+ "supports_fit_target": False,
"supports_cache_ram": False,
"supports_ctx_checkpoints": False,
"supports_no_cache_prompt": False,
@@ -2549,6 +2653,7 @@ class LlamaCppBackend:
spec_draft_n_max_flag: Optional[str] = None
supports_kv_unified = False
supports_fit_ctx = False
+ supports_fit_target = False
supports_cache_ram = False
supports_ctx_checkpoints = False
supports_no_cache_prompt = False
@@ -2646,6 +2751,7 @@ class LlamaCppBackend:
supports_kv_unified = _is_real("--kv-unified")
supports_fit_ctx = _is_real("--fit-ctx")
+ supports_fit_target = _is_real("--fit-target")
supports_cache_ram = _is_real("--cache-ram")
supports_ctx_checkpoints = _is_real("--ctx-checkpoints")
supports_no_cache_prompt = _is_real("--no-cache-prompt")
@@ -2662,6 +2768,7 @@ class LlamaCppBackend:
"spec_draft_n_max_flag": spec_draft_n_max_flag,
"supports_kv_unified": supports_kv_unified,
"supports_fit_ctx": supports_fit_ctx,
+ "supports_fit_target": supports_fit_target,
"supports_cache_ram": supports_cache_ram,
"supports_ctx_checkpoints": supports_ctx_checkpoints,
"supports_no_cache_prompt": supports_no_cache_prompt,
@@ -2746,6 +2853,57 @@ class LlamaCppBackend:
except ValueError:
return None
+ @staticmethod
+ def _emit_child_gpu_visibility(env: dict, pinned: str) -> None:
+ """Write the child's GPU visibility mask (CUDA, plus the HIP mirror on
+ ROCm, where narrowing only CUDA_VISIBLE_DEVICES leaves an AMD child
+ seeing the full set). Do NOT also set ROCR_VISIBLE_DEVICES: ROCR and HIP
+ mask at different layers, so the same indices apply twice -- ROCR reduces
+ and re-indexes from 0, then a non-zero HIP pin points out of range, HIP
+ enumerates 0 devices, and llama.cpp falls back to CPU. The HIP mask alone
+ narrows correctly; clear any inherited ROCR mask so it can't double up."""
+ env["CUDA_VISIBLE_DEVICES"] = pinned
+ try:
+ import torch as _torch
+ if getattr(_torch.version, "hip", None) is not None:
+ env["HIP_VISIBLE_DEVICES"] = pinned
+ env.pop("ROCR_VISIBLE_DEVICES", None)
+ except Exception as e:
+ logger.debug("Failed to set ROCm visibility env vars for child: %s", e)
+
+ @staticmethod
+ def _pin_visible_gpu_order_for_split(env: dict) -> None:
+ """Pin the child's GPU enumeration to the picker's order for a manual
+ ``--tensor-split`` across the whole visible set. CUDA's default
+ FASTEST_FIRST enumeration applies the shares to the wrong cards on
+ heterogeneous hosts (#5025), and CUDA_DEVICE_ORDER only fixes the
+ numbering base: an inherited numeric visibility mask ALSO defines
+ enumeration order, so a reordered parent mask (CUDA_VISIBLE_DEVICES=3,1)
+ would still hand the shares to the wrong cards. The UI built the split
+ positionally over get_backend_visible_gpu_info's device list (ascending
+ physical via nvidia-smi, inherited mask order on the torch fallback), so
+ re-emit the same set in that report order -- not an assumed ascending
+ sort. The visible set itself never changes. No mask, an empty mask, or a
+ UUID/MIG mask (which resolves to None) is left alone -- the multi-GPU
+ controls are hidden for the latter."""
+ env["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"
+ inherited = LlamaCppBackend._resolve_visible_physical_ids()
+ if not inherited:
+ return
+ order = None
+ try:
+ from utils.hardware import get_backend_visible_gpu_info
+ info = get_backend_visible_gpu_info()
+ if info.get("available") and info.get("index_kind") == "physical":
+ reported = [d["index"] for d in info.get("devices", [])]
+ if sorted(reported) == sorted(inherited):
+ order = reported
+ except Exception as e:
+ logger.debug("Could not read reported GPU order for split pin: %s", e)
+ if order is None:
+ order = sorted(inherited)
+ LlamaCppBackend._emit_child_gpu_visibility(env, ",".join(str(i) for i in order))
+
@staticmethod
def _amd_apu_wants_unified_memory(gpu_indices = None) -> bool:
"""True only for AMD unified-memory APUs (gfx1150/gfx1151), where
@@ -3262,6 +3420,20 @@ class LlamaCppBackend:
# aborts a --split-mode tensor load, so it's dropped for the tensor attempt.
_TENSOR_PARALLEL_KV_TYPES = frozenset({"f16", "bf16", "f32"})
+ # Main-model placement settings that Manual mode owns. They must not leak
+ # from Studio's parent environment into llama-server and silently override
+ # the command assembled from the current request. Draft-model placement is
+ # intentionally separate and remains available to speculative decoding.
+ _MANUAL_PLACEMENT_ENV_VARS = (
+ "LLAMA_ARG_CPU_MOE",
+ "LLAMA_ARG_N_CPU_MOE",
+ "LLAMA_ARG_N_GPU_LAYERS",
+ "LLAMA_ARG_TENSOR_SPLIT",
+ "LLAMA_ARG_FIT",
+ "LLAMA_ARG_FIT_TARGET",
+ "LLAMA_ARG_FIT_CTX",
+ )
+
# (binary, mtime, model) that aborted on --split-mode tensor this process (#6415
# geometry limit, e.g. MQA n_head_kv=1). Model-keyed so one model's abort doesn't
# skip tensor for others; tensor is tried by default, recorded only on a real abort.
@@ -3426,6 +3598,12 @@ class LlamaCppBackend:
return env
+ @classmethod
+ def _clear_manual_placement_env(cls, env: dict[str, str]) -> None:
+ """Remove inherited main-model placement owned by Manual mode."""
+ for name in cls._MANUAL_PLACEMENT_ENV_VARS:
+ env.pop(name, None)
+
@staticmethod
def _select_gpus(
model_size_bytes: int,
@@ -4246,6 +4424,8 @@ class LlamaCppBackend:
self._supports_preserve_thinking = False
self._supports_tools = False
self._n_layers = None
+ self._n_experts = None
+ self._leading_dense_block_count = None
self._n_kv_heads = None
self._n_kv_heads_by_layer = None
self._n_heads = None
@@ -4335,6 +4515,8 @@ class LlamaCppBackend:
arch_keys = {
f"{arch}.context_length": "context_length",
f"{arch}.block_count": "n_layers",
+ f"{arch}.expert_count": "n_experts",
+ f"{arch}.leading_dense_block_count": "leading_dense_block_count",
f"{arch}.attention.head_count_kv": "n_kv_heads",
f"{arch}.attention.head_count": "n_heads",
f"{arch}.embedding_length": "embedding_length",
@@ -4523,6 +4705,28 @@ class LlamaCppBackend:
return None
+ @staticmethod
+ def _diffusion_gpu_arg(gpu_ids: Optional[List[int]], *, cpu_only: bool = False) -> str:
+ """Device token passed to the diffusion visual-server child.
+
+ The visual engine replaces its child's CUDA visibility mask with this
+ token, so an unpinned load must carry forward the first token from the
+ parent's mask rather than turning a parent-relative ordinal into a new
+ physical selection.
+ """
+ if gpu_ids:
+ return str(sorted(gpu_ids)[0])
+ if cpu_only:
+ return ""
+ if "DG_GPU" in os.environ:
+ return os.environ["DG_GPU"]
+ parent_mask = os.environ.get("CUDA_VISIBLE_DEVICES")
+ if parent_mask:
+ first = next((token.strip() for token in parent_mask.split(",") if token.strip()), "")
+ if first and first != "-1":
+ return first
+ return "0"
+
def _start_diffusion_server(
self,
*,
@@ -4533,6 +4737,7 @@ class LlamaCppBackend:
model_identifier: str,
n_ctx: int,
extra_args: Optional[List[str]],
+ gpu_ids: Optional[List[int]] = None,
) -> bool:
"""Launch the OpenAI-compat diffusion shim (which drives the on-device
visual decoder) and wait for health. Presents the same /v1 + /health
@@ -4558,7 +4763,11 @@ class LlamaCppBackend:
# CUDA_VISIBLE_DEVICES="" to force CPU serving. Keep the visual-server child
# CPU-masked (empty --gpu) so the shim does not re-expose GPU 0 via its default.
cpu_only = self._effective_gpu_count() == 0
- gpu = "" if cpu_only else os.environ.get("DG_GPU", "0")
+ # Honor the GPU picker first: the diffusion runner takes a single device,
+ # so use the lowest selected GPU (matches the sorted set recorded below, so
+ # the device used == the echoed gpu_ids[0]). With no pick, fall back to the
+ # CPU-only mask, else DG_GPU / 0.
+ gpu = self._diffusion_gpu_arg(gpu_ids, cpu_only = cpu_only)
cmd = list(shim_cmd) + [
"--gguf",
@@ -4586,6 +4795,11 @@ class LlamaCppBackend:
env.setdefault("UNSLOTH_ALLOW_CPU", "1")
env["DG_VISUAL_BIN"] = visual_bin
env["DG_GPU"] = gpu
+ if gpu_ids:
+ # The visual server remasks via CUDA_VISIBLE_DEVICES=; pin PCI
+ # order (as the llama-server path does) so the picked physical id maps
+ # to the GPU the picker showed, not CUDA's default fastest-first order.
+ env["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"
# The file-override shim imports its sibling visual_engine; put its dir on PYTHONPATH.
# (The zoo-package shim is an installed module and needs no PYTHONPATH change.)
if extra_pythonpath:
@@ -4631,6 +4845,23 @@ class LlamaCppBackend:
self._model_identifier = model_identifier
self._cache_type_kv = None
self._gpu_offload_active = True
+ # Diffusion doesn't use the llama.cpp GPU-memory knobs; reset them to
+ # defaults (the picked device is still recorded below) so /load, /status
+ # and reload dedup don't report a previous GGUF's manual settings.
+ self._gpu_memory_mode = "auto"
+ self._gpu_layers = -1
+ self._n_cpu_moe = 0
+ self._tensor_split = None
+ # Diffusion is never tensor-parallel; clear any state left by a prior TP
+ # chat load (load_model phase 1 only kills the process, it doesn't run
+ # the unload reset) so /status doesn't misreport TP and an identical
+ # re-Apply doesn't reload against stale tensor-parallel state.
+ self._tensor_parallel = False
+ # Record only the single device the runner actually uses (the lowest
+ # selected GPU, chosen above) -- not the whole pick. The diffusion runner
+ # is single-device, so echoing a multi-GPU list would misreport placement
+ # in /status and let a re-Apply dedup against GPUs the runner never used.
+ self._gpu_ids = [sorted(gpu_ids)[0]] if gpu_ids else None
if hf_variant:
self._hf_variant = hf_variant
elif gguf_path:
@@ -5721,6 +5952,11 @@ class LlamaCppBackend:
speculative_type: Optional[str] = None,
spec_draft_n_max: Optional[int] = None,
tensor_parallel: bool = False,
+ gpu_memory_mode: Literal["auto", "manual"] = "auto",
+ gpu_layers: int = -1,
+ n_cpu_moe: int = 0,
+ tensor_split: Optional[List[float]] = None,
+ gpu_ids: Optional[List[int]] = None,
n_threads: Optional[int] = None,
n_gpu_layers: Optional[int] = None, # caller compat, unused
n_parallel: int = 1,
@@ -5753,6 +5989,14 @@ class LlamaCppBackend:
"speculative_type": speculative_type,
"spec_draft_n_max": spec_draft_n_max,
"tensor_parallel": tensor_parallel,
+ # GPU-memory placement: replayed on respawn so a server SIGKILL'd by
+ # GPU/RAM pressure reloads onto the same devices with the same
+ # offload, not the auto defaults.
+ "gpu_memory_mode": gpu_memory_mode,
+ "gpu_layers": gpu_layers,
+ "n_cpu_moe": n_cpu_moe,
+ "tensor_split": list(tensor_split) if tensor_split is not None else None,
+ "gpu_ids": list(gpu_ids) if gpu_ids is not None else None,
"n_threads": n_threads,
"n_gpu_layers": n_gpu_layers,
"n_parallel": n_parallel,
@@ -5779,6 +6023,11 @@ class LlamaCppBackend:
speculative_type = speculative_type,
spec_draft_n_max = spec_draft_n_max,
tensor_parallel = tensor_parallel,
+ gpu_memory_mode = gpu_memory_mode,
+ gpu_layers = gpu_layers,
+ n_cpu_moe = n_cpu_moe,
+ tensor_split = tensor_split,
+ gpu_ids = gpu_ids,
chat_template_override = chat_template_override,
extra_args = extra_args,
is_vision = is_vision,
@@ -5899,6 +6148,7 @@ class LlamaCppBackend:
model_identifier = model_identifier,
n_ctx = n_ctx,
extra_args = extra_args,
+ gpu_ids = gpu_ids,
)
if not binary:
@@ -5960,6 +6210,59 @@ class LlamaCppBackend:
# use the same helper so a healthy env-driven tensor server matches.
split_mode_override = parse_split_mode_override(extra_args)
tensor_parallel = _effective_tensor_parallel(extra_args, tensor_parallel)
+ # gpu_layers=0 leaves nothing to split, yet --split-mode tensor or
+ # a per-GPU ratio still launches tensor mode -- and under the
+ # CPU-only mask below (no visible devices) that aborts the server
+ # instead of loading on CPU. Drop both here (nothing to split).
+ if gpu_memory_mode == "manual" and gpu_layers == 0:
+ if tensor_parallel or tensor_split:
+ logger.info(
+ "Manual gpu_layers=0: dropping tensor split/parallel "
+ "flags (nothing to split on the GPU)"
+ )
+ tensor_parallel = False
+ tensor_split = None
+ # Record the requested strategy for /status and the load
+ # response. 'manual' has no fallback, so the request value is the
+ # value actually applied.
+ self._gpu_memory_mode = gpu_memory_mode
+ # The layer/MoE/split knobs apply only with an explicit offload
+ # (manual + gpu_layers >= 0); else record defaults so /status and
+ # /load don't report knobs the server never applied.
+ if gpu_memory_mode == "manual" and gpu_layers >= 0:
+ self._gpu_layers = gpu_layers
+ self._n_cpu_moe = n_cpu_moe
+ self._tensor_split = tensor_split
+ else:
+ self._gpu_layers = -1
+ self._n_cpu_moe = 0
+ self._tensor_split = None
+ self._gpu_ids = sorted(gpu_ids) if gpu_ids else None
+ # Manual offload skips the TP planner but still emits --split-mode
+ # tensor at launch; drop it when fewer than 2 GPUs are in use --
+ # tensor split is a no-op there and aborts on some architectures.
+ # Done before the cache-drop below so a quantized KV survives.
+ if (
+ tensor_parallel
+ and gpu_memory_mode == "manual"
+ and gpu_layers >= 0
+ and self._effective_gpu_count(sorted(gpu_ids) if gpu_ids else None) < 2
+ ):
+ logger.info(
+ "Tensor parallelism requested in manual mode but fewer "
+ "than 2 GPUs are in use; ignoring (needs >= 2)."
+ )
+ tensor_parallel = False
+ # Drop TP for manual + Auto layers before the cache-drop below (like
+ # the <2-GPU guard above), so a requested quantized KV survives into
+ # the --fit load rather than being stripped for a tensor attempt.
+ if tensor_parallel and gpu_memory_mode == "manual" and gpu_layers < 0:
+ logger.info(
+ "Manual mode with Auto layers hands memory management to "
+ "llama.cpp --fit, which is incompatible with tensor "
+ "parallelism; ignoring the tensor split."
+ )
+ tensor_parallel = False
# Tensor mode aborts on a quantized KV cache, so drop it for the
# tensor attempt (and strip any inherited/explicit --cache-type
# that would re-impose it when appended last). Layer split does
@@ -6040,6 +6343,12 @@ class LlamaCppBackend:
"Vision-capable GGUF loaded without a usable mmproj; "
"image input will be disabled for this session"
)
+ # Seed before the try: the except (GPU-selection failure ->
+ # --fit on) falls through to the launch which reads this, and the
+ # probe that assigns it may throw first. Captured before manual
+ # empty `gpus` so the speculative defaults stay GPU-aware and the
+ # CPU-fallback check still knows GPUs were present.
+ _detected_gpus: list[tuple[int, int]] = []
model_size = None # set in the fit try; used by the APU RAM guard
# Layer-fallback min GPUs; raised below on a tensor downgrade. Bound
# before the try so the --fit-on except path still has it (no UnboundLocal).
@@ -6057,6 +6366,18 @@ class LlamaCppBackend:
_gpu_mem = self._get_gpu_memory(binary)
gpus = [(idx, free) for idx, free, _t in _gpu_mem]
total_by_idx = {idx: total for idx, _f, total in _gpu_mem}
+ # GPU picker: restrict every mode to the chosen devices, so
+ # auto selection only considers them and manual mask to
+ # them (the env block below pins CUDA/HIP_VISIBLE_DEVICES).
+ if gpu_ids:
+ _picked = set(gpu_ids)
+ gpus = [g for g in gpus if g[0] in _picked]
+
+ # GPUs the model will run on -- captured before manual
+ # empty `gpus` to bypass the planner. bool() drives the
+ # GPU-aware speculative defaults; the list feeds the
+ # CPU-fallback check.
+ _detected_gpus = list(gpus)
def _gpu_usable(g, frac = _CTX_FIT_VRAM_FRACTION):
# Per-GPU usable budget for ranking: free - (1-frac)*total.
@@ -6088,6 +6409,44 @@ class LlamaCppBackend:
# GPU/VRAM-fit logic below may shrink it on limited HW.
max_available_ctx = self._context_length or effective_ctx
+ # Manual + Auto layers (the Manual default): hand memory
+ # management to llama.cpp's --fit. Emptying the probed GPU set
+ # no-ops the selection/TP planning below, leaving gpu_indices
+ # None (an explicit gpu_ids pick still pins below) and use_fit
+ # True. An explicit context is honored (--fit optimizes around
+ # it); 0 lets --fit size it.
+ if gpu_memory_mode == "manual" and gpu_layers < 0:
+ # Tensor parallelism was already dropped above (before the
+ # cache-drop), so a quantized KV survives into this --fit load.
+ gpus = []
+ effective_ctx = requested_ctx if requested_ctx > 0 else 0
+ original_ctx = effective_ctx
+ # --fit aborts under --split-mode tensor; a raw extras
+ # --split-mode/--tensor-split (appended last) would
+ # otherwise reach llama-server. Strip it like the TP
+ # downgrade does.
+ extra_args = strip_split_mode_only(extra_args)
+ elif gpu_memory_mode == "manual":
+ # Manual offload (--gpu-layers + --fit off): no automatic
+ # device masking (a gpu_ids pick still pins below) or
+ # context cap -- the user owns both. tensor_parallel is
+ # honored but skips the memory-based planner (gpus = []);
+ # the toggle just emits --split-mode tensor (split by free
+ # VRAM, or by the Split ratio if set).
+ gpus = []
+ effective_ctx = (
+ requested_ctx if requested_ctx > 0 else (self._context_length or 0)
+ )
+ original_ctx = effective_ctx
+ # Strip the user --split-mode when the toggle owns the split
+ # (TP engaged -> Studio emits --split-mode tensor) or when the
+ # user asked for tensor (which aborts on a single GPU even if
+ # the manual <2-GPU guard downgraded TP). Otherwise keep their
+ # non-tensor mode (row/none/layer) -- the toggle can't express
+ # those.
+ if tensor_parallel or split_mode_override == "tensor":
+ extra_args = strip_split_mode_only(extra_args)
+
# Will MTP engage? If so, auto-fit reserves draft-model VRAM.
# Mirrors _build_speculative_flags: forced mtp/mtp+ngram always
# engage; auto only on an MTP model >= 3B; ngram/off never. A
@@ -6175,7 +6534,10 @@ class LlamaCppBackend:
_extra_n_max = _extra_args_spec_draft_n_max(extra_args)
_mtp_eff_n_max = _extra_n_max if _extra_n_max is not None else spec_draft_n_max
if _mtp_eff_n_max is None:
- _mtp_eff_n_max = 2 if gpus else 3
+ # _detected_gpus (not gpus) so manual -- which empty
+ # gpus to bypass the planner -- keep the GPU draft depth the
+ # launch flags also use, instead of the CPU default.
+ _mtp_eff_n_max = 2 if _detected_gpus else 3
# Separate-drafter weights live on GPU (an embedded head is
# already in model_size). Size the drafter the launch loads, by
# precedence: extras --model-draft (last-wins), else Unsloth's
@@ -6313,7 +6675,8 @@ class LlamaCppBackend:
# honor it, cap only if it fits no combination. Auto (native):
# prefer fewer GPUs with reduced context (multi-GPU is slower).
gpu_indices, use_fit = None, True
- # Per-GPU weight proportions for tensor mode (None = even).
+ # Per-GPU weight proportions for tensor mode (None lets
+ # llama.cpp split by free VRAM).
tp_tensor_split: Optional[list[int]] = None
explicit_ctx = requested_ctx > 0
# Flat MTP reserve fraction: used only as the fallback when the
@@ -6388,7 +6751,12 @@ class LlamaCppBackend:
# GPUs below that reserve from the set up front (gpu_indices
# becomes the CUDA_VISIBLE_DEVICES mask, fully excluding them).
tp_gpus = gpus
- if tensor_parallel:
+ # Manual mode owns the layer count and context, so it skips
+ # the memory-based planner; its toggle still emits
+ # --split-mode tensor below (split by free VRAM, or by the
+ # Split ratio if set). auto plans here.
+ plan_tp = tensor_parallel and gpu_memory_mode != "manual"
+ if plan_tp:
# Deterministic per-device compute buffer (replicated on
# every device in tensor mode); flat fallback when dims
# are unavailable. _plan_tensor_parallel uses the same.
@@ -6407,7 +6775,7 @@ class LlamaCppBackend:
# free yet have no budget left.
tp_gpus = [g for g in gpus if _gpu_usable(g) >= reserve_mib]
- if tensor_parallel and len(tp_gpus) < 2:
+ if plan_tp and len(tp_gpus) < 2:
# Tensor parallelism needs >= 2 usable GPUs. On a single
# GPU --split-mode tensor is a no-op; with 0 GPUs (CPU-only
# or probe failed) it must not reach llama-server; and a
@@ -6823,6 +7191,12 @@ class LlamaCppBackend:
tp_tensor_split = None
effective_ctx = requested_ctx # fall back to original
+ # GPU picker: when no narrower subset was chosen (manual, or
+ # a failed/file-size selection), pin the whole picked set so the
+ # model can't spill onto an unpicked GPU.
+ if gpu_ids and gpu_indices is None:
+ gpu_indices = sorted(gpu_ids)
+
# Unified-memory APUs load weights into system RAM (under WSL the VM
# cap, not the ROCm-reported VRAM, is the real ceiling); refuse an
# oversize load the OS would otherwise kill mid-flight. Base model
@@ -6859,8 +7233,6 @@ class LlamaCppBackend:
model_path,
"--port",
str(self._port),
- "-c",
- str(effective_ctx) if effective_ctx > 0 else "0",
"--parallel",
str(n_parallel),
"--flash-attn",
@@ -6868,6 +7240,17 @@ class LlamaCppBackend:
# Error out at n_ctx instead of silently rotating the KV cache; frontend catches it and points the user at "Context Length".
"--no-context-shift",
]
+ # A positive context is always passed (in auto-fit, --fit then
+ # optimizes the gpu-layer offload around it). When auto-fit has
+ # no explicit context, omit -c so --fit sizes it to fit VRAM:
+ # "-c 0" would instead pin the FULL native context (llama.cpp's
+ # -c handler sets fit_params_min_ctx = UINT32_MAX on value 0,
+ # disabling --fit's reduction). See gpu_memory_mode.
+ auto_fit = gpu_memory_mode == "manual" and gpu_layers < 0
+ if effective_ctx > 0:
+ cmd.extend(["-c", str(effective_ctx)])
+ elif not auto_fit:
+ cmd.extend(["-c", "0"])
# Report a clean public model id (matching GET /v1/models) rather
# than the raw -m path in llama-server's own /v1/models and the
@@ -6879,7 +7262,63 @@ class LlamaCppBackend:
cmd.extend(["--alias", _alias])
fully_gpu_offloaded = False
- if use_fit:
+ # Set when a positional --tensor-split is emitted, so the env block
+ # can pin CUDA to PCI order even without a GPU subset (see below).
+ manual_tensor_split_emitted = False
+ if gpu_memory_mode == "manual" and gpu_layers >= 0:
+ # Pin the user's layer count and disable auto-fit. --fit off
+ # also means _ctx_integrity_flags must not add --fit-ctx.
+ use_fit = False
+ cmd.extend(["--gpu-layers", str(gpu_layers), "--fit", "off"])
+ # Keep the first n_cpu_moe MoE layers' experts on CPU.
+ moe_flag = self._resolve_cpu_moe_flag(
+ n_cpu_moe,
+ self.n_moe_layers,
+ self._leading_dense_block_count or 0,
+ )
+ if moe_flag is not None:
+ cmd.extend(["--n-cpu-moe", str(moe_flag)])
+ elif n_cpu_moe:
+ # Requested on a dense model: nothing was emitted, so
+ # don't report a count llama-server never received.
+ self._n_cpu_moe = 0
+ # Distribute the model across GPUs by the user's per-GPU shares
+ # (--tensor-split). Works with layer split and tensor
+ # parallelism; --fit off means no fit/tensor abort. Only emit
+ # when >1 GPU is in use AND the list length matches that count:
+ # the field is hidden (not cleared) when the picker narrows to
+ # one, and a direct caller can send a stale ratio for a different
+ # GPU set. Studio drops any mismatch to the free-VRAM default
+ # (llama.cpp would silently zero-pad a short list, or abort past
+ # its 16-device cap).
+ _split_gpus = self._effective_gpu_count(gpu_indices)
+ if tensor_split and _split_gpus > 1:
+ # An all-zero/non-positive sanitized split assigns nothing
+ # anywhere, so fall through to the free-VRAM default in
+ # that case.
+ _sanitized_split = self._sanitize_tensor_split(tensor_split)
+ _split_total = sum(_sanitized_split)
+ if len(_sanitized_split) == _split_gpus and _split_total > 0:
+ cmd.extend(
+ ["--tensor-split", ",".join(f"{x:g}" for x in _sanitized_split)]
+ )
+ self._tensor_split = _sanitized_split
+ manual_tensor_split_emitted = True
+ else:
+ logger.warning(
+ "Dropping manual --tensor-split (%d entries for "
+ "%d GPUs, sanitized total %s); llama.cpp's "
+ "free-VRAM split applies instead",
+ len(tensor_split),
+ _split_gpus,
+ _split_total,
+ )
+ self._tensor_split = None
+ elif tensor_split:
+ # Single effective GPU: the split is never emitted, so
+ # don't report it as active via /status and /load.
+ self._tensor_split = None
+ elif use_fit:
cmd.extend(["--fit", "on"])
elif gpu_indices is not None:
# Fits on selected GPU(s) -- force all layers on GPU. --fit off is
@@ -6897,6 +7336,7 @@ class LlamaCppBackend:
self._ctx_integrity_flags(
n_parallel,
use_fit,
+ auto_fit,
requested_ctx,
effective_ctx,
server_caps,
@@ -6960,9 +7400,11 @@ class LlamaCppBackend:
self._cache_type_kv = None
# Tensor parallelism: split the model across GPUs by tensor
- # rather than by layer. Multi-GPU only -- a no-op on a single
- # GPU. Default (layer split) is left implicit by omitting the
- # flag. See llama.cpp --split-mode.
+ # rather than by layer. The UI only offers it on multi-GPU; a
+ # direct single-GPU caller is redundant (supported archs no-op,
+ # unsupported ones abort and the /load path retries layer split).
+ # Default (layer split) is left implicit by omitting the flag.
+ # See llama.cpp --split-mode.
if tensor_parallel:
cmd.extend(["--split-mode", "tensor"])
if tp_tensor_split and len(tp_tensor_split) > 1:
@@ -6994,7 +7436,7 @@ class LlamaCppBackend:
extra_args = extra_args,
model_identifier = model_identifier,
model_path = model_path,
- gpus = bool(gpus),
+ gpus = bool(_detected_gpus),
binary = binary,
mtp_draft_path = launch_mtp_draft_path,
)
@@ -7112,6 +7554,8 @@ class LlamaCppBackend:
# Library paths so llama-server finds its shared libs and CUDA DLLs.
env = self._llama_server_env_for_binary(binary)
+ if gpu_memory_mode == "manual":
+ self._clear_manual_placement_env(env)
# Omitting --threads relies on llama.cpp's physical-core default, so
# drop an inherited LLAMA_ARG_THREADS that would otherwise feed the
# arg handler and silently force hardware_concurrency(). #5692
@@ -7170,28 +7614,39 @@ class LlamaCppBackend:
# CUDA_VISIBLE_DEVICES leaves an AMD child seeing the full set, so
# set HIP_VISIBLE_DEVICES too. Vulkan is pinned via --device
# (above), not here.
- if gpu_indices is not None and not is_vulkan_backend:
- pinned = ",".join(str(i) for i in gpu_indices)
- env["CUDA_VISIBLE_DEVICES"] = pinned
- try:
- import torch as _torch
- if getattr(_torch.version, "hip", None) is not None:
- env["HIP_VISIBLE_DEVICES"] = pinned
- # Do NOT also set ROCR_VISIBLE_DEVICES to the same
- # value. ROCR_VISIBLE_DEVICES filters at the HSA/ROCr
- # layer and HIP_VISIBLE_DEVICES at the HIP layer, so
- # setting both with the same physical indices applies
- # the mask twice: ROCR reduces the visible set and
- # re-indexes it from 0, then HIP indexes into the
- # already-reduced set. A single non-zero pin (e.g.
- # "1") then points out of range at the HIP layer, HIP
- # enumerates 0 devices, and llama.cpp falls back to
- # CPU ("ggml_cuda_init: no ROCm-capable device is
- # detected"). The HIP mask alone narrows correctly;
- # clear any inherited ROCR mask so it can't double up.
- env.pop("ROCR_VISIBLE_DEVICES", None)
- except Exception as e:
- logger.debug("Failed to set ROCm visibility env vars for child: %s", e)
+ # A deliberate zero-offload load with no GPU companions runs
+ # entirely on CPU, yet a visible CUDA device still costs the child
+ # ~0.5 GB (context + compute scratch) that the CPU-only
+ # classification below reports as free. Hide the GPUs so the load
+ # is exactly what it claims: zero VRAM (verified: GPU stays at idle
+ # baseline and generation runs). Companion loads keep the normal
+ # masking, and a user device pin (in extras or an inherited
+ # LLAMA_ARG_DEVICE) keeps control of its own devices -- the child
+ # aborts on a pin it can't see. The draft-device forms count too:
+ # llama-server parses them even with no drafter loaded.
+ _cpu_only_zero_offload = (
+ gpu_memory_mode == "manual"
+ and gpu_layers == 0
+ and not is_vulkan_backend
+ and not self._zero_offload_keeps_gpu_visible(cmd, env)
+ )
+ if _cpu_only_zero_offload:
+ self._emit_child_gpu_visibility(env, "-1")
+ elif gpu_indices is not None and not is_vulkan_backend:
+ # When the user picked GPUs by index, align CUDA's ordering
+ # with the PCI-bus order the picker enumerated (nvidia-smi),
+ # so "GPU 1" in the UI is GPU 1 to llama.cpp -- not CUDA's
+ # default FASTEST_FIRST order (#5025).
+ if gpu_ids:
+ env["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"
+ self._emit_child_gpu_visibility(env, ",".join(str(i) for i in gpu_indices))
+ elif manual_tensor_split_emitted and not is_vulkan_backend:
+ # A manual per-GPU ratio across ALL GPUs (no explicit pick, so
+ # no CUDA_VISIBLE_DEVICES mask above): the UI built the
+ # --tensor-split list in ascending physical/PCI index order,
+ # so pin the child's enumeration to that order too. The whole
+ # visible set stays in use; only its ordering is fixed.
+ self._pin_visible_gpu_order_for_split(env)
# Captured before any text-only fallback strips it from cmd.
launched_with_mmproj = "--mmproj" in cmd
@@ -7350,7 +7805,6 @@ class LlamaCppBackend:
self._effective_context_length = (
effective_ctx if effective_ctx > 0 else self._context_length
)
- self._reconcile_effective_ctx_with_server()
self._max_context_length = (
max_available_ctx if max_available_ctx > 0 else self._effective_context_length
)
@@ -7535,6 +7989,10 @@ class LlamaCppBackend:
"session; run 'unsloth studio update' to enable vision."
)
cmd = self._strip_mmproj_args(_last_spawn_cmd)
+ # This retry bypasses _spawn_and_wait, so refresh the
+ # launched-argv snapshot itself -- the zero-offload
+ # classification below must not see the stripped --mmproj.
+ _last_spawn_cmd = list(cmd)
self._is_vision = False
self._mmproj_has_audio = False
self._start_llama_process(cmd, env)
@@ -7566,6 +8024,13 @@ class LlamaCppBackend:
self._healthy = True
self._commit_effective_parallel_slots(n_parallel)
+ # Server is up: adopt the real per-request context it allocated
+ # -- the length --fit chose, or a --parallel slot split -- so the
+ # reported context_length matches reality. (Querying /props
+ # before the spawn above always failed; the seeded value was the
+ # requested/native length.)
+ self._reconcile_effective_ctx_with_server()
+
# Commit caller intent only after _healthy=True so a failed start
# can't poison the next inheritance check. None keeps prior, []
# clears, list sets. Source records hf_variant for the route's
@@ -7580,11 +8045,24 @@ class LlamaCppBackend:
self._mtp_runtime_fallback_active = _mtp_active_for_launched_server
self._start_mtp_crash_watchdog()
- # Catch silent CPU fallback when GPU was intended (#5106).
- self._gpu_offload_active = self._classify_gpu_offload(
- gpu_indices is not None or use_fit, gpus or []
- )
- if self._gpu_offload_active is False:
+ # Catch silent CPU fallback when GPU was intended (#5106). Manual
+ # offload (no picker) leaves gpu_indices None and use_fit False, so
+ # include its GPU-layer intent; use the preserved probe since
+ # auto-layers/manual empty `gpus`. A deliberate zero-offload load
+ # classifies by its launched argv instead: the main model is
+ # CPU-only by construction and must read False (not None), or
+ # training needlessly unloads a server holding no VRAM.
+ _deliberate_cpu_only = gpu_memory_mode == "manual" and gpu_layers == 0
+ if _deliberate_cpu_only:
+ self._gpu_offload_active = self._zero_offload_gpu_flag(
+ _last_spawn_cmd, _detected_gpus, env
+ )
+ else:
+ self._gpu_offload_active = self._classify_gpu_offload(
+ gpu_indices is not None or use_fit or gpu_memory_mode == "manual",
+ _detected_gpus,
+ )
+ if self._gpu_offload_active is False and not _deliberate_cpu_only:
logger.warning(
"llama-server appears to have loaded the model entirely "
"on CPU even though Unsloth detected at least one GPU. "
@@ -7947,6 +8425,11 @@ class LlamaCppBackend:
gguf_path: Optional[str] = None,
spec_draft_n_max: Optional[int] = None,
tensor_parallel: bool = False,
+ gpu_memory_mode: Literal["auto", "manual"] = "auto",
+ gpu_layers: int = -1,
+ n_cpu_moe: int = 0,
+ tensor_split: Optional[List[float]] = None,
+ gpu_ids: Optional[List[int]] = None,
mtp_draft_path: Optional[str] = None,
preserve_multi_gpu_on_layer: bool = False,
) -> bool:
@@ -8003,6 +8486,38 @@ class LlamaCppBackend:
):
return False
+ # The diffusion runner is mode-agnostic (always "auto", ignores the
+ # layer/MoE/split knobs), so a standing manual preference in the
+ # request must not force a needless reload -- only the GPU pick matters.
+ if not self._is_diffusion:
+ # A GPU-memory-mode flip (Unsloth / manual) must always reload.
+ if self._gpu_memory_mode != gpu_memory_mode:
+ return False
+ # Manual: a layer-count change always reloads (covers Auto(-1) <-> a
+ # pinned count); MoE/split only matter with an explicit offload.
+ if gpu_memory_mode == "manual" and (
+ self._gpu_layers != gpu_layers
+ or (
+ gpu_layers >= 0
+ and (
+ self._n_cpu_moe != n_cpu_moe
+ or (self._tensor_split or None) != (tensor_split or None)
+ )
+ )
+ ):
+ return False
+ # A changed GPU pick must reload (compare order-insensitively; None/[]
+ # both mean automatic). The diffusion runner collapses a multi-GPU pick
+ # to its single lowest device, so self._gpu_ids holds just that device;
+ # normalize the request the same way, or a multi-GPU pick that resolves
+ # to the same device needlessly reloads.
+ if self._is_diffusion:
+ requested_gpu_pick = [sorted(gpu_ids)[0]] if gpu_ids else None
+ else:
+ requested_gpu_pick = sorted(gpu_ids) if gpu_ids else None
+ if (self._gpu_ids or None) != requested_gpu_pick:
+ return False
+
# Compare on the canonical requested mode. With --spec-type in
# extra_args the backend stores None; mirror that here.
if _extra_args_set_spec_type(extra_args):
@@ -8071,6 +8586,78 @@ class LlamaCppBackend:
return None
return classify_gpu_offload_lines(self._stdout_lines)
+ @staticmethod
+ def _cmd_has_gpu_companion(cmd: list, env: Optional[Mapping[str, str]] = None) -> bool:
+ """True when the argv/env carries a GPU companion: any --mmproj form, or
+ a drafter (Studio's --model-draft, the extras aliases, or the
+ LLAMA_ARG_SPEC_DRAFT_* env) -- these offload to the GPU regardless of
+ the main ``--gpu-layers``. A drafter explicitly forced to CPU
+ (--spec-draft-ngl 0 / --spec-draft-device cpu) doesn't count."""
+ if any(str(a).startswith("--mmproj") for a in cmd):
+ return True
+ if _extra_args_mtp_draft_path(cmd, env) is None:
+ return False
+ return not _extra_args_draft_offloaded_to_cpu(cmd, env)
+
+ @staticmethod
+ def _zero_offload_keeps_gpu_visible(cmd: list, env: Optional[Mapping[str, str]] = None) -> bool:
+ """Whether a zero-layer launch still has a reason to use visible GPUs.
+
+ Keep this shared by child masking and post-launch residency bookkeeping:
+ a device pin, surviving tensor mode, mmproj, or GPU drafter prevents the
+ launch from being a confirmed zero-VRAM server.
+ """
+ return (
+ LlamaCppBackend._cmd_has_gpu_device_pin(cmd, env)
+ or _effective_tensor_parallel(cmd, False, env)
+ or LlamaCppBackend._cmd_has_gpu_companion(cmd, env)
+ )
+
+ @staticmethod
+ def _cmd_has_gpu_device_pin(cmd: list, env: Optional[Mapping[str, str]] = None) -> bool:
+ """True when the effective main or draft ``--device`` pin names a GPU."""
+ main_flags = {"--device", "-dev"}
+ draft_flags = {"--spec-draft-device", "-devd", "--device-draft"}
+ last_main: Optional[str] = None
+ last_draft: Optional[str] = None
+ args = [str(arg) for arg in cmd]
+ for index, raw in enumerate(args):
+ flag, equals, inline = raw.partition("=")
+ if flag not in main_flags and flag not in draft_flags:
+ continue
+ value = inline if equals else (args[index + 1] if index + 1 < len(args) else "")
+ if flag in main_flags:
+ last_main = value
+ else:
+ last_draft = value
+ if last_main is None:
+ last_main = (env or {}).get("LLAMA_ARG_DEVICE")
+
+ def _names_gpu(value: Optional[str]) -> bool:
+ if value is None:
+ return False
+ devices = [item.strip().lower() for item in value.split(",") if item.strip()]
+ return not devices or any(item not in ("cpu", "none") for item in devices)
+
+ return _names_gpu(last_main) or _names_gpu(last_draft)
+
+ @staticmethod
+ def _zero_offload_gpu_flag(
+ spawn_cmd: list,
+ detected_gpus: list,
+ env: Optional[Mapping[str, str]] = None,
+ ) -> Optional[bool]:
+ """GPU-residency flag for a deliberate manual zero-offload load. The
+ main model is CPU-only by construction, but device pins, tensor mode,
+ mmproj, and GPU drafters can still make the server hold VRAM. The counted
+ offload classifier cannot see those allocations. This uses the same
+ predicate as the launch-time zero-VRAM mask; None means no GPU signal."""
+ if not detected_gpus:
+ return None
+ if LlamaCppBackend._is_vulkan_backend():
+ return True
+ return LlamaCppBackend._zero_offload_keeps_gpu_visible(spawn_cmd, env)
+
def load_cancelled(self) -> bool:
"""True if a load was cancelled (e.g. via unload/_cancel_event) and not
yet consumed by the next load_model. Lets the tensor->layer fallback
@@ -8114,11 +8701,18 @@ class LlamaCppBackend:
self._supports_tools = False
self._cache_type_kv = None
self._tensor_parallel = False
+ self._gpu_memory_mode = "auto"
+ self._gpu_layers = -1
+ self._n_cpu_moe = 0
+ self._tensor_split = None
+ self._gpu_ids = None
self._layer_preserves_tensor_intent = False
self._speculative_type = None
self._requested_spec_mode = None
self._spec_draft_n_max = None
self._n_layers = None
+ self._n_experts = None
+ self._leading_dense_block_count = None
self._n_kv_heads = None
self._n_kv_heads_by_layer = None
self._n_heads = None
@@ -8181,6 +8775,10 @@ class LlamaCppBackend:
# Clear healthy so a /load during the replacement's warm-up can't
# short-circuit against the previous server's health (#5401).
self._healthy = False
+ # Reset to unknown so the training guard treats the next (still
+ # loading) server as VRAM-resident rather than reading the killed
+ # server's stale zero-offload flag until the health probe reclassifies.
+ self._gpu_offload_active = None
# Drives _wait_for_vram_settle in the next load_model; set in finally
# so both in-process and frontend Apply paths record the kill.
self._last_kill_monotonic = time.monotonic()
@@ -8785,7 +9383,12 @@ class LlamaCppBackend:
@staticmethod
def _ctx_integrity_flags(
- n_parallel: int, use_fit: bool, requested_ctx: int, effective_ctx: int, caps: dict
+ n_parallel: int,
+ use_fit: bool,
+ auto_fit: bool,
+ requested_ctx: int,
+ effective_ctx: int,
+ caps: dict,
) -> list[str]:
"""Flags that keep the per-request window equal to the advertised ctx.
@@ -8793,14 +9396,28 @@ class LlamaCppBackend:
``--kv-unified`` default, silently splitting ``-c`` into per-slot
windows of ``-c / N``; restore the shared pool so one request can use
the full context. With ``--fit on``, ``--fit-ctx`` floors the fit step
- at an explicitly requested ctx (default floor is 4096) so it offloads
- or fails instead of silently shrinking the window.
+ at an explicitly requested ctx so it offloads or fails instead of
+ silently shrinking the window. The 8192 auto-floor and the tighter
+ ``--fit-target`` margin apply only under Manual + Auto (``auto_fit``),
+ which omits ``-c``: on the legacy auto path ``-c 0`` already pins the
+ native window and ``--fit-ctx 8192`` would override it down to 8192.
"""
flags: list[str] = []
if n_parallel > 1 and caps.get("supports_kv_unified"):
flags.append("--kv-unified")
- if use_fit and requested_ctx > 0 and effective_ctx > 0 and caps.get("supports_fit_ctx"):
- flags.extend(["--fit-ctx", str(effective_ctx)])
+ if use_fit and caps.get("supports_fit_ctx"):
+ if requested_ctx > 0 and effective_ctx > 0:
+ # Floor the fit step at the explicitly requested ctx.
+ flags.extend(["--fit-ctx", str(effective_ctx)])
+ elif auto_fit:
+ # Manual + Auto omits -c, so floor at 8192 so --fit doesn't
+ # shrink the window below a usable size.
+ flags.extend(["--fit-ctx", "8192"])
+ if use_fit and auto_fit and caps.get("supports_fit_target"):
+ # llama.cpp's --fit leaves 1 GiB free per device by default;
+ # tighten that to 512 MiB so it packs more of the model onto
+ # the GPU before spilling to system RAM.
+ flags.extend(["--fit-target", "512"])
return flags
def _query_server_n_ctx(self) -> Optional[int]:
diff --git a/studio/backend/core/inference/llama_server_args.py b/studio/backend/core/inference/llama_server_args.py
index 70d0dc774d..e72e10e071 100644
--- a/studio/backend/core/inference/llama_server_args.py
+++ b/studio/backend/core/inference/llama_server_args.py
@@ -186,12 +186,25 @@ _SPLIT_MODE_FLAGS: frozenset[str] = frozenset({"-sm", "--split-mode"})
_TENSOR_SPLIT_FLAGS: frozenset[str] = frozenset({"-ts", "--tensor-split"})
_SPLIT_SHADOWING_FLAGS: frozenset[str] = _SPLIT_MODE_FLAGS | _TENSOR_SPLIT_FLAGS
+# GPU-offload flags. Stripped only when the GPU Memory mode owns offload
+# (manual emits --fit / --gpu-layers / --n-cpu-moe); in auto, a user's
+# inherited -ngl is respected (the offload_overridden path), so this group is
+# opt-in, not default. Layer flags are shared with llama_cpp's override
+# detection; the MoE flags are strip-only (manual's --n-cpu-moe slider owns them).
+_LAYER_OFFLOAD_FLAGS: frozenset[str] = frozenset(
+ {"-ngl", "--gpu-layers", "--n-gpu-layers", "-fit", "--fit"}
+)
+_MOE_OFFLOAD_FLAGS: frozenset[str] = frozenset({"-ncmoe", "--n-cpu-moe", "-cmoe", "--cpu-moe"})
+_OFFLOAD_SHADOWING_FLAGS: frozenset[str] = _LAYER_OFFLOAD_FLAGS | _MOE_OFFLOAD_FLAGS
+
_SHADOWING_FLAGS: frozenset[str] = (
_CONTEXT_FLAGS | _CACHE_FLAGS | _SPEC_FLAGS | _TEMPLATE_FLAGS | _SPLIT_SHADOWING_FLAGS
)
# Shadowing flags that take no value -- strip the flag only, not the next token.
-_BOOLEAN_SHADOWING_FLAGS: frozenset[str] = frozenset({"--spec-default", "--jinja", "--no-jinja"})
+_BOOLEAN_SHADOWING_FLAGS: frozenset[str] = frozenset(
+ {"--spec-default", "--jinja", "--no-jinja", "-cmoe", "--cpu-moe"}
+)
def parse_ctx_override(args: Optional[Iterable[str]]) -> Optional[int]:
@@ -424,6 +437,8 @@ def strip_shadowing_flags(
strip_spec: bool = True,
strip_template: bool = True,
strip_split_mode: bool = True,
+ strip_tensor_split: bool = False,
+ strip_offload: bool = False,
) -> list[str]:
"""Strip flags that shadow first-class Unsloth settings.
@@ -432,6 +447,12 @@ def strip_shadowing_flags(
(same for cache / spec / template / split-mode). Each ``strip_*``
toggle controls one group; the route only strips groups whose
first-class field the caller actually supplied.
+
+ ``strip_split_mode`` removes both ``--split-mode`` and the coupled
+ ``--tensor-split`` (the Tensor Parallelism toggle owns the whole split).
+ ``strip_tensor_split`` removes ``--tensor-split`` *alone*, so manual mode can
+ replace an inherited per-GPU ratio while leaving the user's ``--split-mode``
+ row/none/layer choice intact.
"""
shadowing: set[str] = set()
if strip_context:
@@ -444,6 +465,10 @@ def strip_shadowing_flags(
shadowing |= _TEMPLATE_FLAGS
if strip_split_mode:
shadowing |= _SPLIT_SHADOWING_FLAGS
+ if strip_tensor_split:
+ shadowing |= _TENSOR_SPLIT_FLAGS
+ if strip_offload:
+ shadowing |= _OFFLOAD_SHADOWING_FLAGS
tokens = [str(a) for a in (args or [])]
out: list[str] = []
diff --git a/studio/backend/main.py b/studio/backend/main.py
index 4797764ce7..81d4c16e52 100644
--- a/studio/backend/main.py
+++ b/studio/backend/main.py
@@ -1156,9 +1156,23 @@ def _get_cached_system_gpu_info(logger) -> dict[str, Any]:
enriched_dev["vram_utilization_pct"] = util.get("vram_utilization_pct")
enriched_devices.append(enriched_dev)
+ # Whether GGUF loads accept an explicit gpu_ids pick: /load and
+ # /validate 400 picks on XPU hosts (no visibility mask speaks torch-xpu
+ # ordinals) and on Vulkan-only builds (--device pins ggml's own
+ # ordinals), so the picker must not offer them.
+ try:
+ from core.inference.llama_cpp import LlamaCppBackend
+ from utils.hardware import DeviceType, get_device
+ gpu_ids_supported = (
+ get_device() != DeviceType.XPU and not LlamaCppBackend._is_vulkan_backend()
+ )
+ except Exception as e:
+ logger.debug(f"Could not resolve gpu_ids support: {e}")
+ gpu_ids_supported = True
gpu_info = {
"available": visibility_info.get("available", False),
"devices": enriched_devices,
+ "gguf_gpu_ids_supported": gpu_ids_supported,
}
_system_gpu_cache = (time.monotonic(), gpu_info)
return gpu_info
diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py
index f3ae0f70df..d51d35189b 100644
--- a/studio/backend/models/inference.py
+++ b/studio/backend/models/inference.py
@@ -64,7 +64,7 @@ class LoadRequest(BaseModel):
)
gpu_ids: Optional[List[int]] = Field(
None,
- description = "Physical GPU indices to use, for example [0, 1]. Omit or pass [] to use automatic selection. Explicit gpu_ids are unsupported when the parent CUDA_VISIBLE_DEVICES uses UUID/MIG entries. Not supported for GGUF models.",
+ description = "Physical GPU indices to use, for example [0, 1]. Omit or pass [] to use automatic selection. Explicit gpu_ids are unsupported when the parent CUDA_VISIBLE_DEVICES uses UUID/MIG entries. For GGUF models the picked devices are pinned via CUDA/HIP_VISIBLE_DEVICES.",
)
speculative_type: Optional[str] = Field(
None,
@@ -100,6 +100,66 @@ class LoadRequest(BaseModel):
"No effect on a single GPU. Ignored for non-GGUF models."
),
)
+ gpu_memory_mode: Literal["auto", "manual"] = Field(
+ "auto",
+ description = (
+ "GPU memory strategy for GGUF models. 'auto' (default): Unsloth "
+ "selects GPUs and caps context to fit VRAM. 'manual': you own the "
+ "offload. Leave gpu_layers at -1 (Auto) to hand memory management to "
+ "llama.cpp's --fit (no device masking, no context auto-reduce, no "
+ "gpu-layer/tensor-split planning); set gpu_layers >= 0 to pin layers "
+ "and n_cpu_moe yourself (--fit off), with tensor_parallel still "
+ "applying (split by free VRAM unless tensor_split is set, no planner). "
+ "Ignored for non-GGUF."
+ ),
+ )
+ gpu_layers: int = Field(
+ -1,
+ ge = -1,
+ description = (
+ "Manual mode only: number of layers to offload to the GPU "
+ "(--gpu-layers, with --fit off). A value >= the model's layer count "
+ "offloads all of them. -1 = Auto: hand layer + context sizing to "
+ "llama.cpp's --fit. Ignored unless gpu_memory_mode is 'manual'."
+ ),
+ )
+ n_cpu_moe: int = Field(
+ 0,
+ ge = 0,
+ description = (
+ "Manual mode only: keep the first N MoE expert layers on the CPU "
+ "(--n-cpu-moe) to save VRAM on MoE models. 0 = none, N = number of "
+ "MoE layers offloaded (the backend offsets past any leading dense "
+ "layers). Ignored unless gpu_memory_mode is 'manual' with gpu_layers >= 0."
+ ),
+ )
+ tensor_split: Optional[List[float]] = Field(
+ None,
+ description = (
+ "Manual mode only: relative share of the model per GPU (--tensor-split), "
+ "in the order of the GPUs in use, e.g. [2, 1] for 2:1. Omit it to let "
+ "llama.cpp use its default, which splits by free VRAM. Any list given is "
+ "passed through as-is, so send [1, 1] to force an even split. Ignored "
+ "unless gpu_memory_mode is 'manual' with gpu_layers >= 0."
+ ),
+ )
+
+ @field_validator("tensor_split")
+ @classmethod
+ def _reject_degenerate_tensor_split(cls, value: Optional[List[float]]) -> Optional[List[float]]:
+ # A negative / non-finite / all-zero split is silently dropped at launch
+ # (stored as None) yet still compared raw in the reload dedupe, so an
+ # identical Apply reloads forever. Reject it up front; [] = no split.
+ if not value:
+ return value
+ import math
+
+ if any((not math.isfinite(v)) or v < 0 for v in value):
+ raise ValueError("tensor_split entries must be finite and non-negative")
+ if sum(value) <= 0:
+ raise ValueError("tensor_split must have a positive total")
+ return value
+
llama_extra_args: Optional[List[str]] = Field(
None,
description = (
@@ -133,6 +193,14 @@ class ValidateModelRequest(BaseModel):
max_seq_length: int = Field(0, ge = 0, le = 1048576)
load_in_4bit: bool = Field(True)
gpu_ids: Optional[List[int]] = Field(None)
+ gpu_memory_mode: Literal["auto", "manual"] = Field(
+ "auto",
+ description = (
+ "GGUF GPU-memory strategy intended for the follow-up load. Manual "
+ "placement bypasses the training coexistence estimate: Auto layers "
+ "delegate fitting to llama.cpp, while explicit layers are user-owned."
+ ),
+ )
include_context_length: bool = Field(
False,
description = "Also read the native context length from the local GGUF header. "
@@ -188,6 +256,16 @@ class ValidateModelResponse(BaseModel):
description = "Native training context length, read from the GGUF header when the file "
"is already downloaded locally; None for non-GGUF, gated, or not-yet-downloaded models.",
)
+ layer_count: Optional[int] = Field(
+ None,
+ description = "Total layer count (GGUF block_count), the manual gpu-layers ceiling, read "
+ "from the header alongside context_length; None when not read.",
+ )
+ moe_layer_count: Optional[int] = Field(
+ None,
+ description = "MoE expert-layer count (the manual --n-cpu-moe ceiling), read from the GGUF "
+ "header alongside context_length; 0 for dense models, None when not read.",
+ )
# Additive fields; the consuming consent dialog ships in a follow-up frontend PR.
requires_transformers_upgrade: bool = Field(
False,
@@ -333,6 +411,34 @@ class LoadResponse(BaseModel):
False,
description = "Whether tensor-parallel split (--split-mode tensor) is active.",
)
+ gpu_memory_mode: Literal["auto", "manual"] = Field(
+ "auto",
+ description = "Active GPU memory strategy ('auto' or 'manual').",
+ )
+ gpu_layers: int = Field(
+ -1,
+ description = "Manual mode: requested --gpu-layers value (-1 = Auto/--fit, or when not manual).",
+ )
+ n_cpu_moe: int = Field(
+ 0,
+ description = "Manual mode: MoE expert layers pinned to CPU (--n-cpu-moe); 0 = none.",
+ )
+ tensor_split: Optional[List[float]] = Field(
+ None,
+ description = "Manual mode: relative model share per GPU (--tensor-split); None = default (split by free VRAM).",
+ )
+ n_layers: Optional[int] = Field(
+ None,
+ description = "Model's layer count (GGUF block_count), for the manual gpu-layers ceiling.",
+ )
+ n_moe_layers: int = Field(
+ 0,
+ description = "Model's MoE expert-layer count (the n_cpu_moe ceiling); 0 if not an MoE model.",
+ )
+ gpu_ids: Optional[List[int]] = Field(
+ None,
+ description = "Physical GPU indices the model is pinned to, or None for automatic selection.",
+ )
class UnloadResponse(BaseModel):
@@ -461,6 +567,42 @@ class InferenceStatusResponse(BaseModel):
False,
description = "Whether tensor-parallel split (--split-mode tensor) is active.",
)
+ gpu_memory_mode: Literal["auto", "manual"] = Field(
+ "auto",
+ description = "Active GPU memory strategy ('auto' or 'manual').",
+ )
+ gpu_layers: int = Field(
+ -1,
+ description = "Manual mode: requested --gpu-layers value (-1 = Auto/--fit, or when not manual).",
+ )
+ n_cpu_moe: int = Field(
+ 0,
+ description = "Manual mode: MoE expert layers pinned to CPU (--n-cpu-moe); 0 = none.",
+ )
+ tensor_split: Optional[List[float]] = Field(
+ None,
+ description = "Manual mode: relative model share per GPU (--tensor-split); None = default (split by free VRAM).",
+ )
+ requested_context_length: Optional[int] = Field(
+ None,
+ description = (
+ "The n_ctx the active GGUF load was invoked with (0 = Auto). Lets the "
+ "UI re-seed a Manual + Auto-layers context pin on hydration, where "
+ "context_length only exposes the resolved value. None for non-GGUF."
+ ),
+ )
+ n_layers: Optional[int] = Field(
+ None,
+ description = "Model's layer count (GGUF block_count), for the manual gpu-layers ceiling.",
+ )
+ n_moe_layers: int = Field(
+ 0,
+ description = "Model's MoE expert-layer count (the n_cpu_moe ceiling); 0 if not an MoE model.",
+ )
+ gpu_ids: Optional[List[int]] = Field(
+ None,
+ description = "Physical GPU indices the model is pinned to, or None for automatic selection.",
+ )
llama_cpp_supports_mtp: bool = Field(
True,
description = (
diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py
index 3d527bf317..136e4f7645 100644
--- a/studio/backend/routes/inference.py
+++ b/studio/backend/routes/inference.py
@@ -13,7 +13,7 @@ from pathlib import Path
from fastapi import APIRouter, Depends, HTTPException, Request, status
from fastapi.responses import StreamingResponse, JSONResponse, Response
from starlette.requests import ClientDisconnect
-from typing import Any, Callable, List, Optional, Union
+from typing import Any, Callable, List, Literal, Optional, Union
import json
import httpx
from loggers import get_logger
@@ -3115,13 +3115,16 @@ def _normalise_settings_str(value: Optional[str]) -> Optional[str]:
def _should_strip_split_mode(request: LoadRequest, backend_extra: Optional[list[str]]) -> bool:
- """Whether an inherited --split-mode should be stripped on reload.
+ """Whether an inherited --split-mode (and its coupled --tensor-split) should
+ be stripped on reload.
The binary Tensor Parallelism toggle can't carry --split-mode's row/none/
layer modes, so only strip when the toggle overrides it: tensor being turned
on, or the inherited mode is tensor (toggle turning it off). Non-tensor modes
- survive. Shared by the inheritance strip and the already-loaded stale check
- so they agree on what reload would do.
+ survive. A manual per-GPU ratio is handled by _should_strip_tensor_split,
+ which strips only --tensor-split so the inherited mode is kept. Shared by the
+ inheritance strip and the already-loaded stale check so they agree on what
+ reload would do.
"""
fields_set = getattr(request, "model_fields_set", set())
return "tensor_parallel" in fields_set and (
@@ -3129,6 +3132,25 @@ def _should_strip_split_mode(request: LoadRequest, backend_extra: Optional[list[
)
+def _should_strip_tensor_split(request: LoadRequest) -> bool:
+ """Whether an inherited --tensor-split alone should be stripped on reload.
+
+ Manual explicit offload (gpu_layers >= 0) owns the per-GPU split: with a ratio
+ it emits its own --tensor-split (an inherited one, appended last, would
+ override it), and with the ratio cleared it wants llama.cpp's default
+ free-VRAM split. Either way an inherited --tensor-split must go, else the
+ cleared case silently keeps the stale ratio while status reports None.
+ Unlike _should_strip_split_mode this leaves --split-mode untouched, so a
+ user's row/none/layer mode survives a Studio split-ratio edit. When the
+ Tensor Parallelism toggle IS overriding the mode, _should_strip_split_mode
+ (called alongside this at every site) strips --split-mode anyway.
+ """
+ return (
+ getattr(request, "gpu_memory_mode", "auto") == "manual"
+ and getattr(request, "gpu_layers", -1) >= 0
+ )
+
+
def _carry_preserved_tensor_intent(
*, preserved: bool, same_model: bool, explicit_drop: bool
) -> bool:
@@ -3187,12 +3209,44 @@ def _request_matches_loaded_settings(
else strip_shadowing_flags(
backend_extra,
strip_split_mode = _should_strip_split_mode(request, backend_extra),
+ strip_tensor_split = _should_strip_tensor_split(request),
+ strip_offload = request.gpu_memory_mode == "manual",
)
)
if not _tensor_parallel_matches_loaded(
effective_extra, request.tensor_parallel, llama_backend.tensor_parallel
):
return False
+ # The diffusion runner is mode-agnostic (it always reports "auto" and ignores
+ # the layer/MoE/split knobs), so a standing manual preference in the request
+ # must not force a needless reload -- only the GPU pick matters.
+ if not llama_backend.is_diffusion:
+ if request.gpu_memory_mode != llama_backend.gpu_memory_mode:
+ return False
+ # Manual: a layer-count change always reloads; MoE/split only matter with
+ # an explicit offload (gpu_layers >= 0), so a leftover value under Auto
+ # must not force one. Mirrors LlamaCppBackend._already_in_target_state.
+ if request.gpu_memory_mode == "manual" and (
+ request.gpu_layers != llama_backend.gpu_layers
+ or (
+ request.gpu_layers >= 0
+ and (
+ request.n_cpu_moe != llama_backend.n_cpu_moe
+ or (request.tensor_split or None) != (llama_backend.tensor_split or None)
+ )
+ )
+ ):
+ return False
+ # A changed GPU pick must reload. The diffusion runner collapses a multi-GPU
+ # request to its single lowest device (it drives one device only), so the
+ # backend records just that device; compare the request the same way, or a
+ # multi-GPU pick that resolves to the same device needlessly reloads.
+ if llama_backend.is_diffusion:
+ _req_gpu_ids = [sorted(request.gpu_ids)[0]] if request.gpu_ids else None
+ else:
+ _req_gpu_ids = sorted(request.gpu_ids) if request.gpu_ids else None
+ if _req_gpu_ids != llama_backend.gpu_ids:
+ return False
# Preserved tensor->layer fallback (both report tensor=off, so the check above
# matches): if the user now explicitly drops tensor intent, reload so placement
# re-selects instead of keeping the all-GPU mask (#6659). The effective check
@@ -3235,14 +3289,17 @@ def _request_matches_loaded_settings(
# contain any shadow flag, so the reload path strips them rather than
# leaving a stale override in effect. (backend_extra computed above.)
if request.llama_extra_args is None:
- # Mirror the reload's conditional split-mode strip, so a preserved
- # non-tensor mode (row/none/layer) isn't seen as stale and doesn't
- # trigger a needless reload of a healthy server.
+ # Mirror the reload's conditional strips, so a preserved non-tensor mode
+ # (row/none/layer) isn't seen as stale and doesn't trigger a needless
+ # reload of a healthy server, while an inherited offload/ratio flag that
+ # the reload *would* strip is correctly seen as stale.
if (
backend_extra
and strip_shadowing_flags(
backend_extra,
strip_split_mode = _should_strip_split_mode(request, backend_extra),
+ strip_tensor_split = _should_strip_tensor_split(request),
+ strip_offload = request.gpu_memory_mode == "manual",
)
!= backend_extra
):
@@ -3861,6 +3918,46 @@ def _estimate_gguf_required_gb(
return None
+def _classify_diffusion_gguf(config: ModelConfig) -> Optional[bool]:
+ """Classify a GGUF as diffusion, normal, or unknown before it is loaded.
+
+ ``None`` is important here: a remote GGUF whose header is not cached can
+ still be routed to the single-GPU diffusion runner after download. Treating
+ that case as normal would let Manual mode skip the training guard even
+ though the runner ignores Manual's llama-server placement controls.
+ """
+ identity = " ".join(
+ str(getattr(config, attr, "") or "") for attr in ("identifier", "gguf_hf_repo", "gguf_file")
+ ).lower()
+ if "diffusion" in identity:
+ return True
+
+ try:
+ main = getattr(config, "gguf_file", None)
+ if not (main and Path(main).is_file()):
+ repo = getattr(config, "gguf_hf_repo", None)
+ variant = getattr(config, "gguf_variant", None)
+ if repo and variant:
+ from hub.utils.gguf import resolve_local_gguf_path
+ main = resolve_local_gguf_path(repo, variant)
+ if not main or not Path(main).is_file():
+ return None
+
+ probe = LlamaCppBackend()
+ probe._read_gguf_metadata(str(main))
+ if probe.is_diffusion:
+ return True
+ # A successfully decoded architecture proves that this is a normal
+ # llama-server GGUF. No architecture means the lightweight probe could
+ # not establish the routing decision, so preserve the unknown state.
+ if getattr(probe, "_architecture", None):
+ return False
+ return None
+ except Exception as e:
+ logger.debug("Could not identify diffusion GGUF for training guard: %s", e)
+ return None
+
+
def _guard_chat_load_against_training(
config: ModelConfig,
*,
@@ -3871,11 +3968,19 @@ def _guard_chat_load_against_training(
requested_gpu_ids: Optional[List[int]],
llama_extra_args: Optional[list[str]] = None,
n_parallel: int = 1,
+ gpu_memory_mode: Literal["auto", "manual"] = "auto",
) -> None:
- """Refuse loading a local chat model that would OOM an active training run.
+ """Protect active training from automatically placed chat-model loads.
+
No-op when training is inactive or unknown. `load_in_4bit` must be the
- effective quantization (see _effective_load_in_4bit). Raises HTTP 409 when the
- model would not fit alongside training."""
+ effective quantization (see _effective_load_in_4bit). Manual chat-GGUF
+ placement is an explicit override: Auto layers delegate fitting to
+ llama.cpp's ``--fit`` and pinned layers are owned by the user, so neither is
+ estimated here. Diffusion is still guarded because its mode-agnostic runner
+ ignores those controls and uses one GPU. An unclassified GGUF is guarded as
+ potentially diffusion until its local header proves otherwise. Other loads
+ raise HTTP 409 when they would not fit beside training.
+ """
from core.training import get_training_backend
from routes.training_vram import can_load_chat_during_training
@@ -3887,6 +3992,19 @@ def _guard_chat_load_against_training(
return
is_gguf = bool(getattr(config, "is_gguf", False))
+ diffusion_kind = _classify_diffusion_gguf(config) if is_gguf else False
+ if is_gguf and gpu_memory_mode == "manual" and diffusion_kind is False:
+ return
+
+ diffusion_gpu = None
+ if is_gguf and diffusion_kind is not False:
+ # Use the same token selection as the runner: an explicit pick wins,
+ # followed by DG_GPU, the first parent-visible token, then GPU 0.
+ diffusion_gpu = LlamaCppBackend._diffusion_gpu_arg(
+ requested_gpu_ids,
+ cpu_only = LlamaCppBackend._effective_gpu_count() == 0,
+ )
+
required_override_gb = (
_estimate_gguf_required_gb(
config,
@@ -3907,6 +4025,7 @@ def _guard_chat_load_against_training(
requested_gpu_ids = requested_gpu_ids,
is_gguf = is_gguf,
required_override_gb = required_override_gb,
+ single_device_gpu = diffusion_gpu,
)
if ok:
return
@@ -3934,6 +4053,98 @@ def _guard_chat_load_against_training(
raise HTTPException(status_code = 409, detail = detail)
+def _resolve_inherited_extra_args(
+ request,
+ config: ModelConfig,
+ model_identifier: str,
+ extra_llama_args: Optional[list[str]],
+ effective_chat_template_override: Optional[str] = None,
+) -> Optional[list[str]]:
+ """Effective pass-through extras for a GGUF request that omitted the field:
+ the previous same-model load's extras, shadow-stripped, so a settings-Apply
+ reload (which does not round-trip the extras field) keeps them (#5401)."""
+ if getattr(request, "llama_extra_args", None) is not None:
+ return extra_llama_args
+ if not getattr(config, "is_gguf", False):
+ return extra_llama_args
+ llama_backend = get_llama_cpp_backend()
+ if not llama_backend.extra_args:
+ return extra_llama_args
+ # Inherit the previous load's extras (the chat-settings Apply path doesn't
+ # round-trip them; an explicit [] still clears). Gated on (model_identifier,
+ # hf_variant) to refuse cross-model pickup, and shadowing flags are
+ # stripped so an inherited override can't win the last-wins CLI
+ # parse against a freshly-supplied first-class field.
+ source = llama_backend.extra_args_source
+ # Compare against the resolved variant, not the request field: callers
+ # commonly omit gguf_variant for local ``.gguf`` paths and HF auto-pick
+ # flows. ``config.gguf_variant`` is the variant load_model was actually
+ # invoked with, so both sides of the comparison key off the same string.
+ resolved_variant = (config.gguf_variant or "").lower()
+ request_variant = (request.gguf_variant or "").lower()
+ stored_variant = (source[1] or "").lower() if source else ""
+ same_model = bool(source and source[0] and source[0].lower() == model_identifier.lower())
+ if request.gguf_variant:
+ variant_mismatch = request_variant != stored_variant
+ else:
+ variant_mismatch = bool(stored_variant and resolved_variant != stored_variant)
+ same_source = same_model and not variant_mismatch
+ if not same_source:
+ logger.info(
+ "Not inheriting llama_extra_args: stored args came from %s, loading %s",
+ source,
+ (model_identifier, resolved_variant),
+ )
+ # Cross-model: clear explicitly so the backend doesn't
+ # inherit via "no opinion" semantics.
+ extra_llama_args = []
+ else:
+ # Strip only the groups whose first-class field was set by the caller, so
+ # an inherited --chat-template-file survives an Apply that omits
+ # chat_template_override. A bundled family template (e.g. gemma-4) counts as
+ # a first-class template even when the request omits chat_template_override,
+ # so strip the inherited --chat-template-file then too -- else the stale arg
+ # (appended last) shadows the bundled template while Studio reports its caps.
+ fields_set = getattr(request, "model_fields_set", set())
+ stripped = strip_shadowing_flags(
+ llama_backend.extra_args,
+ strip_context = "max_seq_length" in fields_set,
+ strip_cache = "cache_type_kv" in fields_set,
+ strip_spec = ("speculative_type" in fields_set or "spec_draft_n_max" in fields_set),
+ strip_template = (
+ "chat_template_override" in fields_set
+ or effective_chat_template_override is not None
+ ),
+ strip_split_mode = _should_strip_split_mode(request, llama_backend.extra_args),
+ # manual + per-GPU ratio emits its own --tensor-split; drop
+ # an inherited one (appended last would override it) while
+ # keeping the user's --split-mode row/none/layer choice.
+ strip_tensor_split = _should_strip_tensor_split(request),
+ # manual emits its own --fit/--gpu-layers, so an inherited offload flag
+ # must not last-wins-override it. auto leaves a user's inherited -ngl
+ # alone. getattr: a validate request reuses this resolver, no offload fields.
+ strip_offload = getattr(request, "gpu_memory_mode", "auto") == "manual",
+ )
+ try:
+ extra_llama_args = validate_extra_args(stripped)
+ except ValueError:
+ # Shouldn't happen on already-validated args; degrade to
+ # no-extras rather than 400 if managed flags changed.
+ logger.warning(
+ "Stored llama_extra_args failed revalidation; loading without them: %s",
+ stripped,
+ )
+ extra_llama_args = []
+ else:
+ if extra_llama_args:
+ logger.info(
+ "Inheriting llama_extra_args from previous "
+ "load (same model, shadow-stripped): %s",
+ extra_llama_args,
+ )
+ return extra_llama_args
+
+
def _model_json_response(model, status_code: int = 200) -> Response:
"""Serialize a pydantic response once via pydantic-core.
@@ -4040,6 +4251,35 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
None if request.llama_extra_args is None else extra_llama_args
)
+ # Manual mode owns the offload flags: strip them from EXPLICIT extras
+ # too (the inherited path already does), or a last-wins --gpu-layers /
+ # --fit in extras re-enables GPU offload on a load status reports as
+ # CPU-only. Manual + per-GPU ratio owns --tensor-split the same way.
+ if request.gpu_memory_mode == "manual" and extra_llama_args:
+ _stripped_explicit = strip_shadowing_flags(
+ extra_llama_args,
+ strip_context = False,
+ strip_cache = False,
+ strip_spec = False,
+ strip_template = False,
+ strip_split_mode = False,
+ strip_tensor_split = _should_strip_tensor_split(request),
+ strip_offload = True,
+ )
+ if _stripped_explicit != extra_llama_args:
+ logger.info(
+ "Manual GPU memory owns the offload flags; stripping them "
+ "from explicit llama_extra_args: %s -> %s",
+ extra_llama_args,
+ _stripped_explicit,
+ )
+ extra_llama_args = _stripped_explicit
+
+ # Keep every downstream consumer on the normalized explicit list. In
+ # particular, the already-loaded comparator must not compare the raw
+ # request's managed offload flags against the stripped launch state.
+ request = request.model_copy(update = {"llama_extra_args": extra_llama_args})
+
model_identifier, model_log_label, native_grant_backed = (
_resolve_model_identifier_for_request(request, operation = "load-model")
)
@@ -4121,6 +4361,13 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
speculative_type = llama_backend.requested_spec_mode,
spec_draft_n_max = llama_backend.spec_draft_n_max,
tensor_parallel = llama_backend.tensor_parallel,
+ gpu_memory_mode = llama_backend.gpu_memory_mode,
+ gpu_layers = llama_backend.gpu_layers,
+ n_cpu_moe = llama_backend.n_cpu_moe,
+ tensor_split = llama_backend.tensor_split,
+ n_layers = llama_backend.n_layers,
+ n_moe_layers = llama_backend.n_moe_layers,
+ gpu_ids = llama_backend.gpu_ids,
)
else:
if (
@@ -4187,12 +4434,41 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
# Normalize gpu_ids: empty list means auto-selection, same as None
effective_gpu_ids = request.gpu_ids if request.gpu_ids else None
- # Reject GGUF + gpu_ids first so the guard can't mask it with a VRAM 409.
+ # GGUF supports gpu_ids: validate the pick up front (before the training
+ # guard) so a bad pick is a clean 400, not masked by a VRAM 409. Rejects
+ # negative / out-of-range / duplicate ids and UUID/MIG parents. XPU hosts
+ # are rejected outright: the picker's indices are torch-xpu ordinals neither
+ # applicator speaks (CUDA/HIP masks don't apply, the Vulkan --device pin
+ # uses ggml's own Vulkan ordinals), so a pick could land on the wrong device.
if config.is_gguf and effective_gpu_ids is not None:
- raise HTTPException(
- status_code = 400,
- detail = "gpu_ids is not supported for GGUF models yet.",
- )
+ from utils.hardware import DeviceType, get_device
+ from utils.hardware.hardware import resolve_requested_gpu_ids
+
+ if get_device() == DeviceType.XPU:
+ raise HTTPException(
+ status_code = 400,
+ detail = (
+ "GPU selection (gpu_ids) is not supported on Intel XPU. "
+ "Omit gpu_ids to use all devices."
+ ),
+ )
+ # Same reasoning for a Vulkan-only build: --device pins ggml's own
+ # Vulkan ordinals, so a physical pick can land on the wrong card on
+ # masked or non-contiguous hosts.
+ if LlamaCppBackend._is_vulkan_backend():
+ raise HTTPException(
+ status_code = 400,
+ detail = (
+ "GPU selection (gpu_ids) is not supported with a Vulkan "
+ "llama.cpp build: physical GPU ids have no defined "
+ "mapping to Vulkan device ordinals. Omit gpu_ids to use "
+ "all devices."
+ ),
+ )
+ try:
+ resolve_requested_gpu_ids(effective_gpu_ids)
+ except ValueError as exc:
+ raise HTTPException(status_code = 400, detail = str(exc)) from exc
if not config.is_gguf and _mlx_distributed_launch_detected():
raise HTTPException(
status_code = 400,
@@ -4222,8 +4498,20 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
"architectures)"
)
- # Refuse a load that would OOM active training, before the unload step below
- # frees the resident model. Off-loop: guard does sync nvidia-smi / HF work.
+ # Inherit the previous same-model load's pass-through extras when this
+ # request omits the field (a settings-Apply reload doesn't round-trip
+ # them); shadow-stripped so an inherited flag can't override a
+ # first-class field the caller did set (#5401).
+ extra_llama_args = _resolve_inherited_extra_args(
+ request,
+ config,
+ model_identifier,
+ extra_llama_args,
+ effective_chat_template_override,
+ )
+
+ # Apply the training coexistence policy before the unload step below
+ # frees the resident model. Off-loop: the default-mode guard does sync work.
await asyncio.to_thread(
_guard_chat_load_against_training,
config,
@@ -4234,6 +4522,7 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
requested_gpu_ids = effective_gpu_ids,
llama_extra_args = extra_llama_args,
n_parallel = getattr(fastapi_request.app.state, "llama_parallel_slots", 1),
+ gpu_memory_mode = request.gpu_memory_mode,
)
# ── GGUF path: load via llama-server ──────────────────────
@@ -4245,84 +4534,6 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
from core.inference.llama_cpp import gguf_load_in_flight
gguf_load_stack.enter_context(gguf_load_in_flight(config.gguf_hf_repo))
- # Inherit llama_extra_args from the previous load when the request
- # omits the field (the chat-settings Apply path doesn't round-trip
- # them; explicit [] still clears). Gated on (model_identifier,
- # hf_variant) to refuse cross-model pickup, and shadowing flags are
- # stripped so an inherited override can't win the last-wins CLI
- # parse against a freshly-supplied first-class field.
- if request.llama_extra_args is None and llama_backend.extra_args:
- source = llama_backend.extra_args_source
- # Compare against the resolved variant, not the request
- # field: callers commonly omit gguf_variant for local
- # ``.gguf`` paths and HF auto-pick flows. ``config.gguf_
- # variant`` is the variant load_model was actually
- # invoked with (see the HF / local branches below), so
- # both sides of the comparison key off the same string.
- resolved_variant = (config.gguf_variant or "").lower()
- request_variant = (request.gguf_variant or "").lower()
- stored_variant = (source[1] or "").lower() if source else ""
- same_model = bool(
- source and source[0] and source[0].lower() == model_identifier.lower()
- )
- if request.gguf_variant:
- variant_mismatch = request_variant != stored_variant
- else:
- variant_mismatch = bool(stored_variant and resolved_variant != stored_variant)
- same_source = same_model and not variant_mismatch
- if not same_source:
- logger.info(
- "Not inheriting llama_extra_args: stored args came from %s, loading %s",
- source,
- (model_identifier, resolved_variant),
- )
- # Cross-model: clear explicitly so the backend doesn't
- # inherit via "no opinion" semantics.
- extra_llama_args = []
- else:
- # Strip only the groups whose first-class field was set by
- # the caller, so an inherited --chat-template-file survives
- # an Apply that omits chat_template_override. A bundled family
- # template (e.g. the gemma-4 override) is an effective
- # first-class template setting even when the raw request
- # omits chat_template_override, so strip the inherited
- # --chat-template-file in that case too -- otherwise the stale
- # extra arg (appended last) shadows the bundled template while
- # Unsloth reports the bundled template's capabilities.
- fields_set = getattr(request, "model_fields_set", set())
- stripped = strip_shadowing_flags(
- llama_backend.extra_args,
- strip_context = "max_seq_length" in fields_set,
- strip_cache = "cache_type_kv" in fields_set,
- strip_spec = (
- "speculative_type" in fields_set or "spec_draft_n_max" in fields_set
- ),
- strip_template = (
- "chat_template_override" in fields_set
- or effective_chat_template_override is not None
- ),
- strip_split_mode = _should_strip_split_mode(
- request, llama_backend.extra_args
- ),
- )
- try:
- extra_llama_args = validate_extra_args(stripped)
- except ValueError:
- # Shouldn't happen on already-validated args; degrade to
- # no-extras rather than 400 if managed flags changed.
- logger.warning(
- "Stored llama_extra_args failed revalidation; loading without them: %s",
- stripped,
- )
- extra_llama_args = []
- else:
- if extra_llama_args:
- logger.info(
- "Inheriting llama_extra_args from previous "
- "load (same model, shadow-stripped): %s",
- extra_llama_args,
- )
-
# Block cache writes that would race the download manager. This runs
# after pass-through argument inheritance so a carried --no-mmproj
# changes the companion requirement exactly as it does for the load.
@@ -4370,6 +4581,11 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
cache_type_kv = request.cache_type_kv,
speculative_type = request.speculative_type,
spec_draft_n_max = request.spec_draft_n_max,
+ gpu_memory_mode = request.gpu_memory_mode,
+ gpu_layers = request.gpu_layers,
+ n_cpu_moe = request.n_cpu_moe,
+ tensor_split = request.tensor_split,
+ gpu_ids = effective_gpu_ids,
n_parallel = _n_parallel,
)
if config.gguf_hf_repo:
@@ -4537,6 +4753,13 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
speculative_type = llama_backend.requested_spec_mode,
spec_draft_n_max = llama_backend.spec_draft_n_max,
tensor_parallel = llama_backend.tensor_parallel,
+ gpu_memory_mode = llama_backend.gpu_memory_mode,
+ gpu_layers = llama_backend.gpu_layers,
+ n_cpu_moe = llama_backend.n_cpu_moe,
+ tensor_split = llama_backend.tensor_split,
+ n_layers = llama_backend.n_layers,
+ n_moe_layers = llama_backend.n_moe_layers,
+ gpu_ids = llama_backend.gpu_ids,
)
# ── Standard path: load via Unsloth/transformers ──────────
@@ -4795,7 +5018,9 @@ def _requires_security_review_for_model(
@router.post("/validate", response_model = ValidateModelResponse)
async def validate_model(
- request: ValidateModelRequest, current_subject: str = Depends(get_current_subject)
+ request: ValidateModelRequest,
+ fastapi_request: Request = None,
+ current_subject: str = Depends(get_current_subject),
):
"""
Lightweight validation endpoint for model identifiers.
@@ -4823,15 +5048,39 @@ async def validate_model(
detail = f"Invalid model identifier: {model_log_label}",
)
- # Refuse early (before the frontend unloads to load this) if it can't fit
- # alongside training, using the same settings /load uses so they agree.
+ # Apply the same training coexistence policy as /load before the frontend
+ # unloads the current model.
effective_gpu_ids = request.gpu_ids if request.gpu_ids else None
- # Mirror /load: reject GGUF + gpu_ids before the guard so both return 400.
+ # Mirror /load: GGUF supports gpu_ids, so validate the pick (a bad one is
+ # a clean 400) before the guard sizes the model against training VRAM.
+ # XPU-host picks are rejected like /load (no defined mapping from the
+ # picker's torch-xpu ordinals to the launcher's device spaces).
if config.is_gguf and effective_gpu_ids is not None:
- raise HTTPException(
- status_code = 400,
- detail = "gpu_ids is not supported for GGUF models yet.",
- )
+ from utils.hardware import DeviceType, get_device
+ from utils.hardware.hardware import resolve_requested_gpu_ids
+
+ if get_device() == DeviceType.XPU:
+ raise HTTPException(
+ status_code = 400,
+ detail = (
+ "GPU selection (gpu_ids) is not supported on Intel XPU. "
+ "Omit gpu_ids to use all devices."
+ ),
+ )
+ if LlamaCppBackend._is_vulkan_backend():
+ raise HTTPException(
+ status_code = 400,
+ detail = (
+ "GPU selection (gpu_ids) is not supported with a Vulkan "
+ "llama.cpp build: physical GPU ids have no defined "
+ "mapping to Vulkan device ordinals. Omit gpu_ids to use "
+ "all devices."
+ ),
+ )
+ try:
+ resolve_requested_gpu_ids(effective_gpu_ids)
+ except ValueError as exc:
+ raise HTTPException(status_code = 400, detail = str(exc)) from exc
effective_load_in_4bit = _effective_load_in_4bit(config, request.load_in_4bit)
# Both checks cover the [adapter, base] set (matching the scan route and workers):
@@ -4895,16 +5144,32 @@ async def validate_model(
latest_tier_active_for, config.identifier, request.hf_token
):
effective_load_in_4bit = False
- # Off-loop: guard does sync nvidia-smi / HF work.
- await asyncio.to_thread(
- _guard_chat_load_against_training,
- config,
- model_identifier = model_identifier,
- hf_token = request.hf_token,
- load_in_4bit = effective_load_in_4bit,
- max_seq_length = request.max_seq_length,
- requested_gpu_ids = effective_gpu_ids,
- )
+ # A metadata-only probe just reads the GGUF header and allocates no VRAM,
+ # so it must not be refused by the training guard. Real loads validate
+ # without include_context_length and /load applies the guard again.
+ if not request.include_context_length:
+ # Match /load's inherited llama.cpp extras and parallel slot count so
+ # validation cannot pass a smaller estimate than the subsequent load.
+ effective_extra_args = _resolve_inherited_extra_args(
+ request, config, model_identifier, None
+ )
+ # Off-loop: guard does sync nvidia-smi / HF work.
+ await asyncio.to_thread(
+ _guard_chat_load_against_training,
+ config,
+ model_identifier = model_identifier,
+ hf_token = request.hf_token,
+ load_in_4bit = effective_load_in_4bit,
+ max_seq_length = request.max_seq_length,
+ requested_gpu_ids = effective_gpu_ids,
+ llama_extra_args = effective_extra_args,
+ n_parallel = (
+ getattr(fastapi_request.app.state, "llama_parallel_slots", 1)
+ if fastapi_request is not None
+ else 1
+ ),
+ gpu_memory_mode = request.gpu_memory_mode,
+ )
# A selected GGUF loads via llama.cpp: auto_map Python and root pickle weights in a
# mixed repo are inert for this load, so gating on them is a false positive. Only
@@ -4918,10 +5183,15 @@ async def validate_model(
# Native context length, read from the local GGUF header when present.
# Lets the staged ("Load on selection" off) flow populate the context
# slider before the GPU load; None until the file is downloaded.
+ # Staged header dims (one read): native context, total layer count, and
+ # MoE expert-layer count -- let the staged flow size the context, GPU-
+ # layers and manual --n-cpu-moe sliders before the load.
context_length: Optional[int] = None
+ layer_count: Optional[int] = None
+ moe_layer_count: Optional[int] = None
if request.include_context_length and is_gguf:
from hub.utils.gguf import resolve_local_gguf_path
- from utils.models.gguf_metadata import read_gguf_context_length
+ from utils.models.gguf_metadata import read_gguf_staged_dims
# Best-effort: a header-read failure must never fail validation of an
# otherwise-valid model (the outer except turns it into a 400).
@@ -4937,9 +5207,15 @@ async def validate_model(
model_identifier, request.gguf_variant
)
if local_gguf:
- context_length = read_gguf_context_length(local_gguf)
+ # Header walk reads tokenizer arrays for dense models (tens of
+ # ms); keep it off the event loop.
+ dims = await asyncio.to_thread(read_gguf_staged_dims, local_gguf)
+ if dims:
+ context_length = dims["context_length"]
+ layer_count = dims["layer_count"]
+ moe_layer_count = dims["moe_layer_count"]
except Exception as e:
- logger.debug("Context-length probe failed for %s: %s", model_log_label, e)
+ logger.debug("Header probe failed for %s: %s", model_log_label, e)
return ValidateModelResponse(
valid = True,
@@ -4954,6 +5230,8 @@ async def validate_model(
requires_trust_remote_code = requires_trust_remote_code,
requires_security_review = requires_security_review,
context_length = context_length,
+ layer_count = layer_count,
+ moe_layer_count = moe_layer_count,
requires_transformers_upgrade = transformers_upgrade is not None,
transformers_upgrade = transformers_upgrade,
)
@@ -5593,6 +5871,14 @@ async def get_status(current_subject: str = Depends(get_current_subject)):
speculative_type = llama_backend.requested_spec_mode,
spec_draft_n_max = llama_backend.spec_draft_n_max,
tensor_parallel = llama_backend.tensor_parallel,
+ gpu_memory_mode = llama_backend.gpu_memory_mode,
+ gpu_layers = llama_backend.gpu_layers,
+ n_cpu_moe = llama_backend.n_cpu_moe,
+ tensor_split = llama_backend.tensor_split,
+ requested_context_length = llama_backend.requested_n_ctx,
+ n_layers = llama_backend.n_layers,
+ n_moe_layers = llama_backend.n_moe_layers,
+ gpu_ids = llama_backend.gpu_ids,
llama_cpp_supports_mtp = _supports_mtp,
spec_fallback_reason = llama_backend.spec_fallback_reason,
llama_cpp_prebuilt_stale = _stale,
diff --git a/studio/backend/routes/models.py b/studio/backend/routes/models.py
index bb321695cd..0806c2f513 100644
--- a/studio/backend/routes/models.py
+++ b/studio/backend/routes/models.py
@@ -2731,7 +2731,11 @@ async def get_gguf_variants(
],
has_vision = response.has_vision,
default_variant = response.default_variant,
- context_length = _read_native_context_length(repo_id, is_local = local),
+ # The header walk reads tokenizer arrays on dense models (tens of
+ # ms per uncached file); keep it off the event loop.
+ context_length = await asyncio.to_thread(
+ _read_native_context_length, repo_id, is_local = local
+ ),
)
except HTTPException:
raise
diff --git a/studio/backend/routes/training_vram.py b/studio/backend/routes/training_vram.py
index fb361d3359..fd96fe2175 100644
--- a/studio/backend/routes/training_vram.py
+++ b/studio/backend/routes/training_vram.py
@@ -197,15 +197,18 @@ def can_load_chat_during_training(
requested_gpu_ids: Optional[List[int]],
is_gguf: bool = False,
required_override_gb: Optional[float] = None,
+ single_device_gpu: Optional[str] = None,
) -> Tuple[bool, Dict[str, Any]]:
"""Decide if a NEW chat model can load without OOMing active training (inverse
of can_keep_chat_during_training: training is already resident, so size the
chat model against the free VRAM that remains). Sizes/places it the same way
the loader will: HF auto reuses auto_select_gpu_ids; HF explicit requires an
even-share per-GPU floor for device_map="balanced"; GGUF sizes from
- required_override_gb over the visible pool. `load_in_4bit` must be effective
- (LoRA can flip 4-bit -> 16-bit). Non-CUDA allows the load; default-deny on any
- CUDA case it can't size, so a load never OOMs training."""
+ required_override_gb over the visible pool. ``single_device_gpu`` is the
+ exact physical device token selected by a single-device runner.
+ `load_in_4bit` must be effective (LoRA can flip 4-bit -> 16-bit). Non-CUDA
+ allows the load; default-deny on any CUDA case it can't size, so a load never
+ OOMs training."""
try:
from utils.hardware import (
DeviceType,
@@ -251,26 +254,49 @@ def can_load_chat_during_training(
}
# Explicit GPUs, or GGUF: size directly and check live free VRAM.
+ if single_device_gpu is not None:
+ mode = "single_device"
+ elif is_gguf:
+ mode = "gguf"
+ else:
+ mode = "explicit"
required_gb = required_override_gb
if required_gb is None:
required_gb, _meta = estimate_required_model_memory_gb(model_name, **est_kwargs)
if required_gb is None:
- mode = "explicit" if requested_gpu_ids else "gguf"
return False, {"mode": mode, "reason": "estimate_unavailable"}
free_by_index = _free_vram_by_index(get_visible_gpu_utilization().get("devices", []))
- if requested_gpu_ids:
+ if single_device_gpu is not None:
+ token = str(single_device_gpu).strip()
+ if not token:
+ # Empty token = a CPU-only single-device runner (e.g. a CPU
+ # diffusion GGUF): it uses no GPU VRAM, so it never threatens
+ # active training and can always load.
+ return True, {"mode": "single_device", "reason": "cpu_only"}
+ try:
+ selected_gpu = int(token)
+ if selected_gpu < 0:
+ raise ValueError
+ except (TypeError, ValueError):
+ # A non-numeric device token (e.g. a CUDA UUID / MIG handle)
+ # can't be mapped to a free-VRAM index, but the runner still
+ # drives ONE device. Size against the worst-case visible device
+ # (min free), never the aggregate pool, so a single-device load
+ # is never OK'd on capacity it can't use and OOMs training.
+ free_vals = [min(free_by_index.values())] if free_by_index else []
+ else:
+ free_vals = [free_by_index.get(selected_gpu, 0.0)]
+ elif requested_gpu_ids:
# Invalid ids -> load_model 400s first, so don't block; missing id = 0.
try:
resolved = resolve_requested_gpu_ids(requested_gpu_ids)
except ValueError:
- return True, {"mode": "explicit", "reason": "invalid_gpu_ids"}
+ return True, {"mode": mode, "reason": "invalid_gpu_ids"}
free_vals = [free_by_index.get(i, 0.0) for i in resolved]
- mode = "explicit"
else:
# GGUF: llama.cpp picks the GPU(s); any visible GPU is a candidate.
free_vals = list(free_by_index.values())
- mode = "gguf"
if not free_vals:
return False, {"mode": mode, "reason": "no_visible_gpus"}
diff --git a/studio/backend/tests/test_chat_load_during_training.py b/studio/backend/tests/test_chat_load_during_training.py
index 63dba8579c..7daa4224aa 100644
--- a/studio/backend/tests/test_chat_load_during_training.py
+++ b/studio/backend/tests/test_chat_load_during_training.py
@@ -168,11 +168,14 @@ class TestCanLoadGGUF(_GpuCacheResetMixin, unittest.TestCase):
devices,
required_override = None,
estimate = None,
+ single_device_gpu = None,
+ gpu_ids = None,
):
with (
patch("utils.hardware.get_device", return_value = DeviceType.CUDA),
patch("utils.hardware.estimate_required_model_memory_gb", return_value = (estimate, {})),
patch("utils.hardware.get_visible_gpu_utilization", return_value = {"devices": devices}),
+ patch("utils.hardware.resolve_requested_gpu_ids", return_value = gpu_ids),
patch("utils.hardware.auto_select_gpu_ids") as auto_mock,
):
ok, info = tv.can_load_chat_during_training(
@@ -180,9 +183,10 @@ class TestCanLoadGGUF(_GpuCacheResetMixin, unittest.TestCase):
hf_token = None,
load_in_4bit = True,
max_seq_length = 0,
- requested_gpu_ids = None,
+ requested_gpu_ids = gpu_ids,
is_gguf = True,
required_override_gb = required_override,
+ single_device_gpu = single_device_gpu,
)
return ok, info, auto_mock
@@ -198,6 +202,88 @@ class TestCanLoadGGUF(_GpuCacheResetMixin, unittest.TestCase):
ok, _, _ = self._run(devices = _devices((0, 80, 35), (1, 80, 70)), required_override = 20.0)
self.assertTrue(ok)
+ def test_no_per_gpu_floor_for_gguf_with_explicit_gpu_ids(self):
+ # gpu_ids narrows llama.cpp's candidate pool but does not turn its
+ # self-placement into HF device_map="balanced". The uneven selected
+ # pair therefore keeps the aggregate GGUF check without an even-share
+ # floor on the nearly-full card.
+ ok, info, _ = self._run(
+ devices = _devices((0, 80, 35), (1, 80, 70), (2, 80, 0)),
+ required_override = 20.0,
+ gpu_ids = [0, 1],
+ )
+ self.assertTrue(ok)
+ self.assertEqual(info["mode"], "gguf")
+
+ def test_single_device_uses_selected_gpu(self):
+ # The model needs 27 GB with headroom. GPU 0 has 45 GB free, while an
+ # unrelated training-heavy GPU 1 has only 10 GB free.
+ ok, info, _ = self._run(
+ devices = _devices((0, 80, 35), (1, 80, 70)),
+ required_override = 20.0,
+ single_device_gpu = "0",
+ )
+ self.assertTrue(ok)
+ self.assertEqual(info["usable_gb"], 45.0)
+
+ blocked, blocked_info, _ = self._run(
+ devices = _devices((0, 80, 35), (1, 80, 70)),
+ required_override = 20.0,
+ single_device_gpu = "1",
+ )
+ self.assertFalse(blocked)
+ self.assertEqual(blocked_info["usable_gb"], 10.0)
+
+ def test_single_device_unresolved_token_sizes_against_worst_device(self):
+ # A non-numeric device token (a CUDA UUID / MIG handle) can't map to a
+ # free-VRAM index. The runner still drives ONE device, so size against the
+ # worst-case visible device (min free), not the aggregate pool: one GPU
+ # with 80 GB free vs a 20 GB model -> allow.
+ ok, info, _ = self._run(
+ devices = _devices((0, 80, 0)),
+ required_override = 20.0,
+ single_device_gpu = "GPU-uuid",
+ )
+ self.assertTrue(ok)
+ self.assertEqual(info["mode"], "single_device")
+ self.assertNotIn("reason", info)
+
+ def test_single_device_unresolved_token_refuses_when_worst_device_full(self):
+ # Same UUID fallback, worst-case device nearly full (2 GB for a 20 GB
+ # model) -> refuse (default-deny), not on an unresolved-token technicality.
+ ok, info, _ = self._run(
+ devices = _devices((0, 80, 78)),
+ required_override = 20.0,
+ single_device_gpu = "GPU-uuid",
+ )
+ self.assertFalse(ok)
+ self.assertNotEqual(info.get("reason"), "unresolved_gpu_id")
+
+ def test_single_device_unresolved_token_uses_min_free_not_aggregate(self):
+ # The single-device runner uses ONE device but we can't tell which from a
+ # UUID token. Sizing against the aggregate pool would let a 20 GB model
+ # "fit" 160 GB of pooled free VRAM while landing on a 2 GB card and OOMing
+ # training. Min-free (2 GB) is the safe worst case -> refuse.
+ ok, info, _ = self._run(
+ devices = _devices((0, 80, 78), (1, 80, 0), (2, 80, 0)),
+ required_override = 20.0,
+ single_device_gpu = "GPU-uuid",
+ )
+ self.assertFalse(ok)
+ self.assertEqual(info["mode"], "single_device")
+
+ def test_single_device_cpu_token_allows(self):
+ # An empty device token = a CPU-only single-device runner (CPU diffusion
+ # GGUF): it uses no GPU VRAM, so it never threatens training -> allow
+ # regardless of how full the GPUs are.
+ ok, info, _ = self._run(
+ devices = _devices((0, 80, 78)),
+ required_override = 20.0,
+ single_device_gpu = "",
+ )
+ self.assertTrue(ok)
+ self.assertEqual(info["reason"], "cpu_only")
+
def test_estimate_unavailable_refuses(self):
# No override and the estimator can't size it -> default-deny.
ok, info, _ = self._run(devices = _devices((0, 80, 0)), required_override = None, estimate = None)
@@ -309,6 +395,8 @@ class TestChatLoadGuardRoute(unittest.TestCase):
captured = None,
training_active,
decision,
+ gpu_memory_mode = "auto",
+ requested_gpu_ids = None,
):
config = config or SimpleNamespace(is_gguf = False, is_lora = False, path = None)
with _stub_guard_deps(
@@ -320,7 +408,8 @@ class TestChatLoadGuardRoute(unittest.TestCase):
hf_token = None,
load_in_4bit = True,
max_seq_length = 0,
- requested_gpu_ids = None,
+ requested_gpu_ids = requested_gpu_ids,
+ gpu_memory_mode = gpu_memory_mode,
)
def test_noop_when_training_inactive(self):
@@ -332,6 +421,141 @@ class TestChatLoadGuardRoute(unittest.TestCase):
def test_allows_when_fits(self):
self._guard(training_active = True, decision = (True, {"mode": "auto"}))
+ def test_diffusion_detection_uses_name_before_download(self):
+ config = SimpleNamespace(
+ identifier = "unsloth/DiffusionGemma-GGUF",
+ gguf_hf_repo = "unsloth/DiffusionGemma-GGUF",
+ gguf_file = None,
+ )
+ self.assertTrue(self.route._classify_diffusion_gguf(config))
+
+ def test_uncached_gguf_classification_remains_unknown(self):
+ config = SimpleNamespace(
+ identifier = "owner/renamed-model",
+ gguf_hf_repo = "owner/renamed-model",
+ gguf_variant = "Q4_K_M",
+ gguf_file = None,
+ )
+ self.assertIsNone(self.route._classify_diffusion_gguf(config))
+
+ def test_diffusion_detection_reuses_loader_metadata_probe(self):
+ import tempfile
+
+ seen = []
+
+ class _Probe:
+ is_diffusion = False
+ _architecture = None
+
+ def _read_gguf_metadata(self, path):
+ seen.append(path)
+ self.is_diffusion = True
+
+ with tempfile.TemporaryDirectory() as d:
+ model = Path(d) / "renamed.gguf"
+ model.write_bytes(b"GGUF")
+ config = SimpleNamespace(identifier = "local", gguf_file = str(model))
+ with patch.object(self.route, "LlamaCppBackend", _Probe):
+ self.assertTrue(self.route._classify_diffusion_gguf(config))
+ self.assertEqual(seen, [str(model)])
+
+ def test_local_chat_gguf_classification_is_definitive(self):
+ import tempfile
+ class _Probe:
+ is_diffusion = False
+ _architecture = "llama"
+
+ def _read_gguf_metadata(self, _path):
+ pass
+
+ with tempfile.TemporaryDirectory() as d:
+ model = Path(d) / "renamed.gguf"
+ model.write_bytes(b"GGUF")
+ config = SimpleNamespace(identifier = "local", gguf_file = str(model))
+ with patch.object(self.route, "LlamaCppBackend", _Probe):
+ self.assertFalse(self.route._classify_diffusion_gguf(config))
+
+ def test_manual_known_normal_gguf_bypasses_training_estimate(self):
+ captured = []
+ config = SimpleNamespace(is_gguf = True)
+ with patch.object(self.route, "_classify_diffusion_gguf", return_value = False):
+ self._guard(
+ config = config,
+ captured = captured,
+ training_active = True,
+ decision = (False, {"reason": "must not run"}),
+ gpu_memory_mode = "manual",
+ )
+ self.assertEqual(captured, [])
+
+ def test_manual_unknown_gguf_keeps_single_device_training_guard(self):
+ captured = []
+ config = SimpleNamespace(is_gguf = True)
+ with (
+ patch.object(self.route, "_classify_diffusion_gguf", return_value = None),
+ patch.object(self.route, "_estimate_gguf_required_gb", return_value = 12.5),
+ patch.object(
+ self.route.LlamaCppBackend,
+ "_diffusion_gpu_arg",
+ return_value = "2",
+ ),
+ ):
+ self._guard(
+ config = config,
+ captured = captured,
+ training_active = True,
+ decision = (True, {"mode": "single_device"}),
+ gpu_memory_mode = "manual",
+ )
+ self.assertEqual(len(captured), 1)
+ self.assertEqual(captured[0]["single_device_gpu"], "2")
+
+ def test_manual_diffusion_uses_single_device_guard(self):
+ captured = []
+ config = SimpleNamespace(is_gguf = True)
+ with (
+ patch.object(self.route, "_classify_diffusion_gguf", return_value = True),
+ patch.object(self.route, "_estimate_gguf_required_gb", return_value = 12.5),
+ ):
+ self._guard(
+ config = config,
+ captured = captured,
+ training_active = True,
+ decision = (True, {"mode": "gguf"}),
+ gpu_memory_mode = "manual",
+ requested_gpu_ids = [3, 1],
+ )
+ self.assertEqual(len(captured), 1)
+ self.assertEqual(captured[0]["single_device_gpu"], "1")
+ self.assertEqual(captured[0]["requested_gpu_ids"], [3, 1])
+
+ def test_unpinned_diffusion_uses_runner_default_gpu(self):
+ captured = []
+ config = SimpleNamespace(is_gguf = True)
+ with (
+ patch.object(self.route, "_classify_diffusion_gguf", return_value = True),
+ patch.object(self.route, "_estimate_gguf_required_gb", return_value = 12.5),
+ patch.object(
+ self.route.LlamaCppBackend,
+ "_effective_gpu_count",
+ return_value = 2,
+ ),
+ patch.object(
+ self.route.LlamaCppBackend,
+ "_diffusion_gpu_arg",
+ return_value = "3",
+ ) as gpu_arg,
+ ):
+ self._guard(
+ config = config,
+ captured = captured,
+ training_active = True,
+ decision = (True, {"mode": "single_device"}),
+ gpu_memory_mode = "manual",
+ )
+ gpu_arg.assert_called_once_with(None, cpu_only = False)
+ self.assertEqual(captured[0]["single_device_gpu"], "3")
+
def test_refuses_with_headroom_number(self):
info = {"required_gb": 30.0, "usable_gb": 6.0, "needed_gb": 39.0, "mode": "auto"}
with self.assertRaises(HTTPException) as exc:
@@ -467,36 +691,115 @@ class TestValidateRefusesDuringTraining(unittest.TestCase):
self.assertEqual(captured[0]["load_in_4bit"], False)
self.assertEqual(captured[0]["max_seq_length"], 4096)
- def test_rejects_gguf_with_gpu_ids_before_guard(self):
- # /validate must mirror /load's GGUF + gpu_ids 400, before the VRAM guard.
+ def test_validate_forwards_manual_gpu_memory_mode_to_guard(self):
from models.inference import ValidateModelRequest
- request = ValidateModelRequest(model_path = "x.gguf", gpu_ids = [0])
+ request = ValidateModelRequest(
+ model_path = "unsloth/model-GGUF",
+ gguf_variant = "Q4_K_M",
+ gpu_memory_mode = "manual",
+ )
cfg = SimpleNamespace(
- identifier = "x.gguf",
- display_name = "x",
+ identifier = "unsloth/model-GGUF",
+ display_name = "model-GGUF",
is_gguf = True,
is_lora = False,
is_vision = False,
path = None,
base_model = None,
)
- captured = []
+ captured = {}
with (
patch.object(
self.route,
"_resolve_model_identifier_for_request",
- return_value = ("x.gguf", "x.gguf", False),
+ return_value = ("unsloth/model-GGUF", "unsloth/model-GGUF", False),
),
patch.object(self.route.ModelConfig, "from_identifier", return_value = cfg),
patch.object(self.route, "load_inference_config", return_value = {}),
- _stub_guard_deps(training_active = True, decision = (True, {}), captured = captured),
+ patch.object(
+ self.route,
+ "_guard_chat_load_against_training",
+ lambda config, **kw: captured.update(kw),
+ ),
):
- with self.assertRaises(HTTPException) as exc:
- asyncio.run(self.route.validate_model(request, current_subject = "u"))
- self.assertEqual(exc.exception.status_code, 400)
- self.assertIn("gpu_ids is not supported for GGUF", exc.exception.detail)
- self.assertEqual(captured, []) # guard never reached
+ asyncio.run(self.route.validate_model(request, current_subject = "u"))
+ self.assertEqual(captured.get("gpu_memory_mode"), "manual")
+
+ def test_validate_forwards_inherited_extras_and_parallel_to_guard(self):
+ # Regression: /load resolves inherited same-model extras and passes the
+ # real slot count to the guard; validate must do the same, else it sizes
+ # a smaller estimate (no inherited -c/--model-draft, n_parallel=1) and
+ # /load then 409s after the frontend has already unloaded.
+ from models.inference import ValidateModelRequest
+
+ request = ValidateModelRequest(model_path = "unsloth/Qwen3-1.7B", max_seq_length = 4096)
+ cfg = SimpleNamespace(
+ identifier = "unsloth/Qwen3-1.7B",
+ display_name = "Qwen3-1.7B",
+ is_gguf = False,
+ is_lora = False,
+ is_vision = False,
+ path = None,
+ base_model = None,
+ )
+ captured = {}
+ with (
+ patch.object(
+ self.route,
+ "_resolve_model_identifier_for_request",
+ return_value = ("unsloth/Qwen3-1.7B", "unsloth/Qwen3-1.7B", False),
+ ),
+ patch.object(self.route.ModelConfig, "from_identifier", return_value = cfg),
+ patch.object(self.route, "load_inference_config", return_value = {}),
+ patch.object(self.route, "_resolve_inherited_extra_args", return_value = ["-c", "32768"]),
+ patch.object(
+ self.route,
+ "_guard_chat_load_against_training",
+ lambda config, **kw: captured.update(kw),
+ ),
+ ):
+ asyncio.run(self.route.validate_model(request, current_subject = "u"))
+ self.assertEqual(captured.get("llama_extra_args"), ["-c", "32768"])
+ self.assertIn("n_parallel", captured)
+
+ def test_metadata_probe_skips_training_guard(self):
+ # A header-only probe (include_context_length) allocates no VRAM, so the
+ # training guard must not run -- else the staging GPU-layers / MoE sliders
+ # it feeds are hidden exactly when a during-training user needs them.
+ from models.inference import ValidateModelRequest
+
+ request = ValidateModelRequest(
+ model_path = "unsloth/Qwen3-1.7B",
+ max_seq_length = 4096,
+ include_context_length = True,
+ )
+ cfg = SimpleNamespace(
+ identifier = "unsloth/Qwen3-1.7B",
+ display_name = "Qwen3-1.7B",
+ is_gguf = False,
+ is_lora = False,
+ is_vision = False,
+ path = None,
+ base_model = None,
+ )
+ guard_called = []
+ with (
+ patch.object(
+ self.route,
+ "_resolve_model_identifier_for_request",
+ return_value = ("unsloth/Qwen3-1.7B", "unsloth/Qwen3-1.7B", False),
+ ),
+ patch.object(self.route.ModelConfig, "from_identifier", return_value = cfg),
+ patch.object(self.route, "load_inference_config", return_value = {}),
+ patch.object(
+ self.route,
+ "_guard_chat_load_against_training",
+ lambda *a, **kw: guard_called.append(True),
+ ),
+ ):
+ asyncio.run(self.route.validate_model(request, current_subject = "u"))
+ self.assertEqual(guard_called, [])
# ── _estimate_gguf_required_gb (sizes the same weights the loader loads) ──────
diff --git a/studio/backend/tests/test_gguf_metadata.py b/studio/backend/tests/test_gguf_metadata.py
index a5be07f8e3..ec0330ce05 100644
--- a/studio/backend/tests/test_gguf_metadata.py
+++ b/studio/backend/tests/test_gguf_metadata.py
@@ -15,6 +15,7 @@ from utils.models.gguf_metadata import (
pairing_score,
read_gguf_context_length,
read_gguf_general_metadata,
+ read_gguf_staged_dims,
read_mmproj_audio_capability,
)
@@ -153,6 +154,78 @@ def test_context_length_ignores_foreign_arch_key(tmp_path: Path):
assert read_gguf_context_length(str(p)) is None
+# --- read_gguf_staged_dims (one pass: context + layer + moe counts) ----
+
+
+def test_staged_dims_none_for_missing_or_non_gguf(tmp_path: Path):
+ assert read_gguf_staged_dims(str(tmp_path / "nope.gguf")) is None
+ p = tmp_path / "garbage.gguf"
+ p.write_bytes(b"not a gguf at all")
+ assert read_gguf_staged_dims(str(p)) is None
+
+
+def test_staged_dims_moe_with_leading_dense(tmp_path: Path):
+ # GLM-4.7-Flash shape: context + total layers + MoE layers in one read.
+ p = _write_synthetic_gguf(
+ tmp_path / "glm.gguf",
+ {"general.architecture": "deepseek2"},
+ extra_uint32 = {
+ "deepseek2.context_length": 202752,
+ "deepseek2.block_count": 47,
+ "deepseek2.expert_count": 64,
+ "deepseek2.leading_dense_block_count": 1,
+ },
+ )
+ assert read_gguf_staged_dims(str(p)) == {
+ "context_length": 202752,
+ "layer_count": 47,
+ "moe_layer_count": 46,
+ }
+
+
+def test_staged_dims_dense_model(tmp_path: Path):
+ # Dense: layer_count present, moe_layer_count 0 (slider hidden).
+ p = _write_synthetic_gguf(
+ tmp_path / "dense.gguf",
+ {"general.architecture": "qwen3"},
+ extra_uint32 = {"qwen3.context_length": 40960, "qwen3.block_count": 36},
+ )
+ assert read_gguf_staged_dims(str(p)) == {
+ "context_length": 40960,
+ "layer_count": 36,
+ "moe_layer_count": 0,
+ }
+
+
+def test_staged_dims_all_moe_no_leading_dense(tmp_path: Path):
+ # Experts present, no leading_dense key -> every block is a MoE layer.
+ p = _write_synthetic_gguf(
+ tmp_path / "moe.gguf",
+ {"general.architecture": "qwen35moe"},
+ extra_uint32 = {"qwen35moe.block_count": 40, "qwen35moe.expert_count": 256},
+ )
+ assert read_gguf_staged_dims(str(p)) == {
+ "context_length": None,
+ "layer_count": 40,
+ "moe_layer_count": 40,
+ }
+
+
+def test_staged_dims_uint64_block_count(tmp_path: Path):
+ # block_count stored as uint64 (vtype 10) still parses; moe == block_count.
+ p = _write_synthetic_gguf(
+ tmp_path / "moe64.gguf",
+ {"general.architecture": "gpt-oss"},
+ extra_uint32 = {"gpt-oss.expert_count": 32},
+ extra_uint64 = {"gpt-oss.block_count": 24},
+ )
+ assert read_gguf_staged_dims(str(p)) == {
+ "context_length": None,
+ "layer_count": 24,
+ "moe_layer_count": 24,
+ }
+
+
def test_context_length_read_from_uint64(tmp_path: Path):
# Some models store context_length as a uint64 (vtype 10).
p = _write_synthetic_gguf(
diff --git a/studio/backend/tests/test_gpu_memory_mode.py b/studio/backend/tests/test_gpu_memory_mode.py
new file mode 100644
index 0000000000..b17274197f
--- /dev/null
+++ b/studio/backend/tests/test_gpu_memory_mode.py
@@ -0,0 +1,879 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""Backend contract for the GPU Memory mode dropdown.
+
+The dropdown threads a single ``gpu_memory_mode`` ("auto" | "manual") from the
+chat UI through the load request. "manual" lets the user own the offload: with
+``gpu_layers < 0`` (Auto, the default) it hands all memory management to
+llama.cpp's ``--fit on`` (no CUDA/HIP device masking, no context auto-reduce, no
+gpu-layer or tensor-split planning); with ``gpu_layers >= 0`` it pins the layers
+and MoE offload itself (``--fit off``). These tests pin:
+
+ * the pydantic request/response/status contract (snake_case key, default
+ "auto", unknown values rejected),
+ * the backend ``gpu_memory_mode`` property and its reset on unload,
+ * the ``_already_in_target_state`` reload-detection branch, and
+ * that the manual + Auto-layers branch in ``load_model`` empties the probed
+ GPU set and drops tensor parallelism so the selection below no-ops, while
+ the explicit-offload branch emits ``--gpu-layers`` / ``--fit off``.
+"""
+
+from __future__ import annotations
+
+import inspect
+import sys
+import types as _types
+from pathlib import Path
+
+import pytest
+
+_BACKEND_DIR = str(Path(__file__).resolve().parent.parent)
+if _BACKEND_DIR not in sys.path:
+ sys.path.insert(0, _BACKEND_DIR)
+
+# Same external-dep stubs as the other llama_cpp unit tests so importing
+# the backend doesn't drag in structlog / httpx / loggers.
+_loggers_stub = _types.ModuleType("loggers")
+_loggers_stub.get_logger = lambda name: __import__("logging").getLogger(name)
+sys.modules.setdefault("loggers", _loggers_stub)
+
+_structlog_stub = _types.ModuleType("structlog")
+_structlog_stub.get_logger = lambda *a, **k: __import__("logging").getLogger("stub")
+sys.modules.setdefault("structlog", _structlog_stub)
+
+# httpx is a real, installed backend dependency: import it so the genuine module
+# is in sys.modules. A hand-rolled stub here is inevitably incomplete and, since
+# setdefault installs it before real httpx loads, would poison a combined pytest
+# run -- routes/inference references httpx.Response (and other attrs) at def time.
+import httpx # noqa: F401
+
+from core.inference import llama_cpp as llama_cpp_module
+from core.inference.llama_cpp import LlamaCppBackend
+from models.inference import (
+ InferenceStatusResponse,
+ LoadRequest,
+ LoadResponse,
+)
+
+
+# ── Pydantic contract (snake_case key, default "auto") ───────────────
+
+
+def test_load_request_defaults_gpu_memory_mode_auto():
+ assert LoadRequest(model_path = "owner/repo").gpu_memory_mode == "auto"
+
+
+def test_load_request_round_trips_json_key():
+ req = LoadRequest.model_validate({"model_path": "owner/repo", "gpu_memory_mode": "manual"})
+ assert req.gpu_memory_mode == "manual"
+ assert req.model_dump()["gpu_memory_mode"] == "manual"
+
+
+def test_load_request_rejects_unknown_mode():
+ with pytest.raises(ValueError):
+ LoadRequest(model_path = "owner/repo", gpu_memory_mode = "bogus")
+
+
+@pytest.mark.parametrize("model_cls", [LoadResponse, InferenceStatusResponse])
+def test_response_models_emit_gpu_memory_mode(model_cls):
+ if model_cls is LoadResponse:
+ default = model_cls(
+ status = "loaded",
+ model = "owner/repo",
+ display_name = "repo",
+ inference = {},
+ )
+ manual = model_cls(
+ status = "loaded",
+ model = "owner/repo",
+ display_name = "repo",
+ inference = {},
+ gpu_memory_mode = "manual",
+ )
+ else:
+ default = model_cls()
+ manual = model_cls(gpu_memory_mode = "manual")
+ assert default.model_dump()["gpu_memory_mode"] == "auto"
+ assert manual.model_dump()["gpu_memory_mode"] == "manual"
+
+
+# ── Backend property + reset ─────────────────────────────────────────
+
+
+class _FakeProcess:
+ """Stand-in for subprocess.Popen so _kill_process is a no-op."""
+
+ def terminate(self):
+ pass
+
+ def wait(self, timeout = None):
+ return 0
+
+ def kill(self):
+ pass
+
+ def poll(self):
+ return 0
+
+
+def test_gpu_memory_mode_property_defaults_auto():
+ assert LlamaCppBackend().gpu_memory_mode == "auto"
+
+
+def test_gpu_memory_mode_property_reflects_field():
+ backend = LlamaCppBackend()
+ backend._gpu_memory_mode = "manual"
+ assert backend.gpu_memory_mode == "manual"
+
+
+def test_unload_resets_gpu_memory_mode():
+ backend = LlamaCppBackend()
+ backend._process = _FakeProcess()
+ backend._gpu_memory_mode = "manual"
+ backend.unload_model()
+ assert backend.gpu_memory_mode == "auto"
+
+
+# ── _already_in_target_state reload-detection branch ─────────────────
+
+
+def _loaded_backend(gpu_memory_mode: str) -> LlamaCppBackend:
+ backend = LlamaCppBackend()
+ backend._process = _FakeProcess() # is_loaded only checks "is not None"
+ backend._healthy = True
+ backend._model_identifier = "owner/repo"
+ backend._hf_variant = "Q4_K_M"
+ backend._requested_n_ctx = 8192
+ backend._cache_type_kv = None
+ backend._requested_spec_mode = "auto"
+ backend._chat_template_override = None
+ backend._is_vision = False
+ backend._extra_args = None
+ backend._gguf_path = None
+ backend._gpu_memory_mode = gpu_memory_mode
+ return backend
+
+
+def _target_state(backend: LlamaCppBackend, gpu_memory_mode: str) -> bool:
+ return backend._already_in_target_state(
+ gguf_path = None,
+ model_identifier = "owner/repo",
+ hf_variant = "Q4_K_M",
+ n_ctx = 8192,
+ cache_type_kv = None,
+ speculative_type = "auto",
+ chat_template_override = None,
+ extra_args = None,
+ is_vision = False,
+ gpu_memory_mode = gpu_memory_mode,
+ )
+
+
+@pytest.mark.parametrize("mode", ["auto", "manual"])
+def test_already_in_target_state_matches_same_mode(mode):
+ assert _target_state(_loaded_backend(mode), mode) is True
+
+
+@pytest.mark.parametrize("loaded,requested", [("auto", "manual"), ("manual", "auto")])
+def test_already_in_target_state_reloads_on_mode_change(loaded, requested):
+ # Flipping the dropdown either direction must force a reload so the command
+ # is rebuilt with/without the Unsloth GPU masking.
+ assert _target_state(_loaded_backend(loaded), requested) is False
+
+
+def test_already_in_target_state_ignores_mode_for_diffusion():
+ # The diffusion runner is mode-agnostic (always "auto"), so a standing manual
+ # preference must not force a needless reload.
+ backend = _loaded_backend("auto")
+ backend._is_diffusion = True
+ assert _target_state(backend, "manual") is True
+
+
+# ── load_model: manual + Auto layers bypasses Unsloth GPU management ──
+
+
+def _load_model_source() -> str:
+ return inspect.getsource(llama_cpp_module.LlamaCppBackend.load_model)
+
+
+def test_auto_layers_branch_empties_gpus_and_drops_tensor_parallel():
+ # Emptying the probed set makes the selection / TP planning below no-op, so
+ # gpu_indices stays None and use_fit True (--fit on).
+ src = _load_model_source()
+ gate = src.find('if gpu_memory_mode == "manual" and gpu_layers < 0:')
+ assert gate != -1, "load_model must branch on manual + Auto layers (gpu_layers < 0)"
+ block = src[gate : gate + 1400]
+ assert "gpus = []" in block, "Auto-layers branch must empty the probed GPU set"
+ # --fit aborts under --split-mode tensor, so a raw-extras split-mode is stripped.
+ assert "strip_split_mode_only(extra_args)" in block
+ assert "requested_ctx if requested_ctx > 0 else 0" in block
+ # The branch sits before GPU selection assigns gpu_indices; --fit on is its emission.
+ assert gate < src.find("gpu_indices, use_fit = None, True")
+ assert 'cmd.extend(["--fit", "on"])' in src
+ # TP drops for this path, but at a guard BEFORE the quantized-KV cache-drop, so
+ # a requested quantized cache survives into the --fit load.
+ tp_drop = src.find('if tensor_parallel and gpu_memory_mode == "manual" and gpu_layers < 0:')
+ assert tp_drop != -1, "manual + Auto layers must drop tensor_parallel"
+ assert "tensor_parallel = False" in src[tp_drop : tp_drop + 400]
+ cache_drop = src.find("Tensor parallelism requires a non-quantized KV cache")
+ assert cache_drop != -1
+ assert (
+ tp_drop < cache_drop
+ ), "TP must drop before the cache-drop so a quantized KV survives --fit"
+
+
+def test_auto_layers_never_sends_ctx_size_zero():
+ # Sending "-c 0" sets fit_params_min_ctx = UINT32_MAX in llama.cpp, pinning
+ # the full native context and disabling --fit's reduction. So the base cmd
+ # must never carry -c, "-c 0" is emitted only outside the Auto-layers (--fit)
+ # case, and a positive context is passed through (which --fit optimizes
+ # layers around).
+ src = _load_model_source()
+ base_start = src.find("cmd = [")
+ base_end = src.find("\n ]", base_start)
+ base_block = src[base_start:base_end]
+ assert '"-c"' not in base_block, "-c must be conditional, not in the base cmd list"
+ assert 'cmd.extend(["-c", str(effective_ctx)])' in src, "positive ctx must pass -c"
+ assert 'auto_fit = gpu_memory_mode == "manual" and gpu_layers < 0' in src
+ zero = src.find('cmd.extend(["-c", "0"])')
+ assert zero != -1, '"-c 0" emission must exist outside the Auto-layers case'
+ guard = src.rfind("elif not auto_fit:", 0, zero)
+ assert guard != -1 and zero - guard < 120, '"-c 0" must sit under the not-auto_fit guard'
+
+
+def test_manual_mode_clears_inherited_main_model_placement_env():
+ env = {name: "inherited" for name in LlamaCppBackend._MANUAL_PLACEMENT_ENV_VARS}
+ env["LLAMA_ARG_N_GPU_LAYERS_DRAFT"] = "7"
+ env["UNRELATED"] = "kept"
+
+ LlamaCppBackend._clear_manual_placement_env(env)
+
+ assert not (set(env) & set(LlamaCppBackend._MANUAL_PLACEMENT_ENV_VARS))
+ assert env["LLAMA_ARG_N_GPU_LAYERS_DRAFT"] == "7"
+ assert env["UNRELATED"] == "kept"
+
+
+def test_load_model_sanitizes_manual_env_after_building_child_env():
+ src = _load_model_source()
+ env_build = src.find("env = self._llama_server_env_for_binary(binary)")
+ env_clear = src.find("self._clear_manual_placement_env(env)", env_build)
+ launch = src.find("subprocess.Popen", env_build)
+ assert env_build != -1
+ assert env_build < env_clear < launch
+
+
+# ── Manual offload (--gpu-layers + --fit off + --n-cpu-moe) ───────────
+
+
+def test_load_request_accepts_manual():
+ req = LoadRequest(
+ model_path = "owner/repo",
+ gpu_memory_mode = "manual",
+ gpu_layers = 20,
+ n_cpu_moe = 8,
+ tensor_split = [2, 1],
+ )
+ assert req.gpu_memory_mode == "manual"
+ assert req.gpu_layers == 20
+ assert req.n_cpu_moe == 8
+ assert req.tensor_split == [2, 1]
+
+
+def test_load_request_manual_defaults():
+ req = LoadRequest(model_path = "owner/repo")
+ assert req.gpu_layers == -1
+ assert req.n_cpu_moe == 0
+ assert req.tensor_split is None
+
+
+@pytest.mark.parametrize("bad", [[0, 0], [-1, 2], [float("inf"), 1], [float("nan"), 1]])
+def test_load_request_rejects_degenerate_tensor_split(bad):
+ # A negative/non-finite/all-zero split is dropped at launch but compared raw
+ # in the reload dedupe, so it would reload forever -- reject it up front.
+ with pytest.raises(ValueError):
+ LoadRequest(model_path = "owner/repo", tensor_split = bad)
+
+
+@pytest.mark.parametrize("good", [[2, 1], [1, 1], [], None])
+def test_load_request_accepts_valid_tensor_split(good):
+ assert LoadRequest(model_path = "owner/repo", tensor_split = good).tensor_split == good
+
+
+def test_route_normalizes_explicit_extras_before_reload_dedupe():
+ route_src = (Path(_BACKEND_DIR) / "routes" / "inference.py").read_text(encoding = "utf-8")
+ load_impl = route_src[route_src.index("async def _load_model_impl") :]
+ strip = load_impl.index("_stripped_explicit = strip_shadowing_flags")
+ normalize = load_impl.index(
+ 'request = request.model_copy(update = {"llama_extra_args": extra_llama_args})'
+ )
+ dedupe = load_impl.index("and _request_matches_loaded_settings(")
+ assert strip < normalize < dedupe
+
+
+@pytest.mark.parametrize("model_cls", [LoadResponse, InferenceStatusResponse])
+def test_response_models_emit_manual_fields(model_cls):
+ if model_cls is LoadResponse:
+ obj = model_cls(
+ status = "loaded",
+ model = "owner/repo",
+ display_name = "repo",
+ inference = {},
+ gpu_memory_mode = "manual",
+ gpu_layers = 20,
+ n_cpu_moe = 8,
+ tensor_split = [2, 1],
+ n_layers = 32,
+ n_moe_layers = 32,
+ )
+ else:
+ obj = model_cls(
+ gpu_memory_mode = "manual",
+ gpu_layers = 20,
+ n_cpu_moe = 8,
+ tensor_split = [2, 1],
+ n_layers = 32,
+ n_moe_layers = 32,
+ )
+ dumped = obj.model_dump()
+ assert dumped["gpu_memory_mode"] == "manual"
+ assert dumped["gpu_layers"] == 20
+ assert dumped["n_cpu_moe"] == 8
+ assert dumped["tensor_split"] == [2, 1]
+ assert dumped["n_layers"] == 32
+ assert dumped["n_moe_layers"] == 32
+
+
+def test_manual_properties_default_and_reflect_and_reset():
+ backend = LlamaCppBackend()
+ assert backend.gpu_layers == -1 and backend.n_cpu_moe == 0
+ assert backend.tensor_split is None
+ backend._gpu_layers = 20
+ backend._n_cpu_moe = 8
+ backend._tensor_split = [2, 1]
+ assert backend.gpu_layers == 20 and backend.n_cpu_moe == 8
+ assert backend.tensor_split == [2, 1]
+ backend._process = _FakeProcess()
+ backend.unload_model()
+ assert backend.gpu_layers == -1 and backend.n_cpu_moe == 0
+ assert backend.tensor_split is None
+
+
+def test_n_moe_layers_property():
+ # 0 for a dense model (hides the slider); block_count for all-MoE;
+ # block_count - leading_dense otherwise (GLM-4.7-Flash: 47 - 1 -> 46).
+ b = LlamaCppBackend()
+ b._n_layers = 36
+ b._n_experts = None
+ assert b.n_moe_layers == 0
+ b._n_experts = 128
+ b._leading_dense_block_count = None
+ assert b.n_moe_layers == 36
+ b._n_layers = 47
+ b._leading_dense_block_count = 1
+ assert b.n_moe_layers == 46
+
+
+def _target_state_manual(
+ backend,
+ *,
+ gpu_layers,
+ n_cpu_moe,
+ tensor_split = None,
+):
+ return backend._already_in_target_state(
+ gguf_path = None,
+ model_identifier = "owner/repo",
+ hf_variant = "Q4_K_M",
+ n_ctx = 8192,
+ cache_type_kv = None,
+ speculative_type = "auto",
+ chat_template_override = None,
+ extra_args = None,
+ is_vision = False,
+ gpu_memory_mode = "manual",
+ gpu_layers = gpu_layers,
+ n_cpu_moe = n_cpu_moe,
+ tensor_split = tensor_split,
+ )
+
+
+def test_manual_reloads_on_gpu_layers_or_n_cpu_moe_or_split_change():
+ backend = _loaded_backend("manual")
+ backend._gpu_layers = 20
+ backend._n_cpu_moe = 0
+ backend._tensor_split = None
+ # Same knobs -> no reload.
+ assert _target_state_manual(backend, gpu_layers = 20, n_cpu_moe = 0) is True
+ # Changed layer count -> reload.
+ assert _target_state_manual(backend, gpu_layers = 16, n_cpu_moe = 0) is False
+ # Changed MoE offload -> reload.
+ assert _target_state_manual(backend, gpu_layers = 20, n_cpu_moe = 8) is False
+ # Added a GPU split -> reload.
+ assert _target_state_manual(backend, gpu_layers = 20, n_cpu_moe = 0, tensor_split = [2, 1]) is False
+ # Same GPU split -> no reload.
+ backend._tensor_split = [2, 1]
+ assert _target_state_manual(backend, gpu_layers = 20, n_cpu_moe = 0, tensor_split = [2, 1]) is True
+
+
+def test_auto_layers_reload_tracks_only_gpu_layers():
+ # Under Auto (gpu_layers < 0) the MoE/split knobs don't apply, so a leftover
+ # request value must not reload -- only a gpu_layers change (Auto -> pinned) does.
+ backend = _loaded_backend("manual")
+ backend._gpu_layers = -1
+ backend._n_cpu_moe = 0
+ backend._tensor_split = None
+ # Same Auto, leftover MoE/split in the request -> still no reload.
+ assert _target_state_manual(backend, gpu_layers = -1, n_cpu_moe = 8, tensor_split = [2, 1]) is True
+ # Auto -> explicit offload reloads.
+ assert _target_state_manual(backend, gpu_layers = 20, n_cpu_moe = 0) is False
+
+
+def test_manual_offload_emits_gpu_layers_fit_off_and_n_cpu_moe():
+ src = _load_model_source()
+ gate = src.find('elif gpu_memory_mode == "manual":')
+ assert gate != -1, "load_model must have an explicit-offload manual branch"
+ block = src[gate : gate + 700]
+ # Empties the probed set (skips the planner) but keeps the user's TP choice
+ # (only the Auto-layers branch above drops TP).
+ assert "gpus = []" in block
+ assert "tensor_parallel = False" not in block
+ # The cmd emits the layer count with fit disabled, gated on gpu_layers >= 0.
+ assert 'if gpu_memory_mode == "manual" and gpu_layers >= 0:' in src
+ assert 'cmd.extend(["--gpu-layers", str(gpu_layers), "--fit", "off"])' in src
+ # MoE offload uses --n-cpu-moe via _resolve_cpu_moe_flag (tested behaviorally below).
+ assert "_resolve_cpu_moe_flag(" in src
+ assert 'cmd.extend(["--n-cpu-moe", str(moe_flag)])' in src
+ # A count requested on a dense model is never emitted, so it must also be
+ # dropped from the recorded state -- else /status and /load report a count
+ # llama-server never received (same rule as the tensor-split drop below).
+ moe_emit = src.find('cmd.extend(["--n-cpu-moe", str(moe_flag)])')
+ assert "elif n_cpu_moe:" in src[moe_emit : moe_emit + 300]
+ assert "self._n_cpu_moe = 0" in src[moe_emit : moe_emit + 300]
+ # The offload path forces use_fit False so --fit-ctx is never added under --fit off.
+ emit = src.find('cmd.extend(["--gpu-layers", str(gpu_layers), "--fit", "off"])')
+ assert "use_fit = False" in src[src.rfind("\n", 0, emit) - 200 : emit + 80]
+
+
+def test_status_reports_requested_context_length():
+ # The hydration path re-seeds a Manual+Auto context pin from the REQUESTED
+ # n_ctx (0 = Auto); context_length only exposes the resolved value.
+ assert "requested_context_length" in InferenceStatusResponse.model_fields
+ s = InferenceStatusResponse(requested_context_length = 8192)
+ assert s.model_dump()["requested_context_length"] == 8192
+ assert InferenceStatusResponse().model_dump()["requested_context_length"] is None
+ # The /status route must actually wire it from the backend (a declared-but-
+ # never-populated field would leave hydration silently reverting the pin).
+ from pathlib import Path as _P
+
+ route_src = (_P(_BACKEND_DIR) / "routes" / "inference.py").read_text(encoding = "utf-8")
+ assert "requested_context_length = llama_backend.requested_n_ctx" in route_src
+
+
+def test_manual_offload_emits_tensor_split():
+ # The offload path emits --tensor-split from the per-GPU shares, only when
+ # provided, with >1 GPU in use, AND matching that count (a stale ratio on a
+ # narrowed picker or a mismatched direct-API list must not emit -- llama-
+ # server aborts on a split/GPU-count mismatch).
+ src = _load_model_source()
+ assert "if tensor_split and _split_gpus > 1:" in src
+ # Emit only on a length match AND a positive sanitized total: a mismatched
+ # or all-zero split aborts llama-server / assigns nothing, so it's dropped.
+ # The emitted list is the sanitized one (clamping tested behaviorally below).
+ assert "_sanitized_split = self._sanitize_tensor_split(tensor_split)" in src
+ assert "if len(_sanitized_split) == _split_gpus and _split_total > 0:" in src
+ assert '"--tensor-split"' in src
+ # Joined as a comma list (e.g. "2,1") within the explicit-offload cmd branch.
+ gate = src.find('if gpu_memory_mode == "manual" and gpu_layers >= 0:')
+ nxt = src.find("elif use_fit:", gate)
+ assert '","' in src[gate:nxt] and "tensor_split" in src[gate:nxt]
+ # A split with a single effective GPU is never emitted, so it must also be
+ # dropped from the recorded state -- else /status and /load report a ratio
+ # llama-server never received and the dedupe baseline preserves it.
+ assert "elif tensor_split:" in src[gate:nxt]
+ drop = src.find("elif tensor_split:", gate, nxt)
+ assert "self._tensor_split = None" in src[drop : drop + 250]
+
+
+def test_sanitize_tensor_split_clamps_negative_and_non_finite():
+ # Negative entries would launch a placement different from the ratio the
+ # UI showed; inf passes a plain > 0 total gate and would emit
+ # "--tensor-split inf,..." (llama.cpp normalizes shares by the running
+ # total, so an inf poisons the shares from that entry on). Both clamp to 0.
+ sanitize = LlamaCppBackend._sanitize_tensor_split
+ assert sanitize([2, 1]) == [2.0, 1.0]
+ assert sanitize([-1, 2]) == [0.0, 2.0]
+ assert sanitize([float("inf"), 1]) == [0.0, 1.0]
+ assert sanitize([float("nan"), 1]) == [0.0, 1.0]
+ # All-zero survives sanitization; the call site's total gate drops it.
+ assert sanitize([0, 0]) == [0.0, 0.0]
+ # Unreadable input -> []; the call site's length gate drops it.
+ assert sanitize(["x", 1]) == []
+ assert sanitize([10**400, 1]) == []
+
+
+def test_zero_offload_mask_honors_device_pin_spellings():
+ # A user device pin must keep the GPUs visible: llama-server aborts on a
+ # pin it can't see ('error: invalid device'). The pin can arrive as
+ # --device or its -dev alias, as the draft forms (parsed even with no
+ # drafter loaded), or as an inherited LLAMA_ARG_DEVICE env var.
+ load_src = _load_model_source()
+ assert "self._zero_offload_keeps_gpu_visible(cmd, env)" in load_src
+ block = inspect.getsource(LlamaCppBackend._cmd_has_gpu_device_pin)
+ for flag in (
+ '"--device"',
+ '"-dev"',
+ '"--spec-draft-device"',
+ '"-devd"',
+ '"--device-draft"',
+ ):
+ assert flag in block
+ assert '"LLAMA_ARG_DEVICE"' in block
+
+
+def test_resolve_cpu_moe_flag():
+ # Clamp the requested MoE-layer count to the model's MoE layers, then offset
+ # past leading dense layers (--n-cpu-moe counts from layer 0).
+ R = LlamaCppBackend._resolve_cpu_moe_flag
+ assert R(0, 40, 0) is None # nothing requested
+ assert R(8, 0, 0) is None # dense model (no MoE layers)
+ assert R(8, 40, 0) == 8 # all-MoE: direct
+ assert R(100, 40, 0) == 40 # clamp to the MoE layer count
+ # GLM-4.7-Flash (deepseek2): block_count 47, leading_dense 1, n_moe 46.
+ assert R(5, 46, 1) == 6 # offset past the 1 dense layer
+ assert R(46, 46, 1) == 47 # all MoE on CPU == block_count
+
+
+def test_manual_allows_tensor_parallel_via_split_mode():
+ # Manual offload keeps the user's TP choice but skips the memory-based planner
+ # (plan_tp excludes manual, so its empty gpu set can't downgrade TP). The
+ # --split-mode tensor emission gates on tensor_parallel alone, so manual
+ # reaches it -- with tp_tensor_split None it's an even split (no
+ # --tensor-split). --fit off means no fit/tensor abort.
+ src = _load_model_source()
+ assert 'plan_tp = tensor_parallel and gpu_memory_mode != "manual"' in src
+ assert "if plan_tp:" in src
+ assert "if plan_tp and len(tp_gpus) < 2:" in src
+ sm = src.find('cmd.extend(["--split-mode", "tensor"])')
+ assert sm != -1, "TP must emit --split-mode tensor"
+ guard = src.rfind("if tensor_parallel:", 0, sm)
+ assert guard != -1 and sm - guard < 200, "split-mode gates on tensor_parallel"
+ # The tensor-split is only emitted for a planned (non-even) split, which
+ # manual never produces, so manual stays an even split.
+ assert "if tp_tensor_split and len(tp_tensor_split) > 1:" in src
+
+
+def test_fit_sets_target_margin():
+ # Manual + Auto (auto_fit) tightens the per-device VRAM margin to 512 MiB.
+ caps = {"supports_fit_target": True}
+ flags = LlamaCppBackend._ctx_integrity_flags(1, True, True, 0, 0, caps)
+ assert flags[flags.index("--fit-target") + 1] == "512"
+ # Not emitted on the legacy auto path (fit on but not auto_fit): -c 0 pins
+ # native there, so the tighter margin must not ride along.
+ assert "--fit-target" not in LlamaCppBackend._ctx_integrity_flags(1, True, False, 0, 0, caps)
+ # Not emitted when fit is off.
+ assert "--fit-target" not in LlamaCppBackend._ctx_integrity_flags(1, False, False, 0, 0, caps)
+ # Not emitted when the binary lacks support.
+ assert "--fit-target" not in LlamaCppBackend._ctx_integrity_flags(
+ 1, True, True, 0, 0, {"supports_fit_target": False}
+ )
+
+
+# ── GPU picker (gpu_ids -> CUDA_VISIBLE_DEVICES) ─────────────────────
+
+
+def test_load_request_accepts_gpu_ids():
+ req = LoadRequest(model_path = "owner/repo", gpu_ids = [1, 0])
+ assert req.gpu_ids == [1, 0]
+ assert LoadRequest(model_path = "owner/repo").gpu_ids is None
+
+
+@pytest.mark.parametrize("model_cls", [LoadResponse, InferenceStatusResponse])
+def test_response_models_emit_gpu_ids(model_cls):
+ if model_cls is LoadResponse:
+ obj = model_cls(status = "loaded", model = "m", display_name = "m", inference = {}, gpu_ids = [1])
+ else:
+ obj = model_cls(gpu_ids = [1])
+ assert obj.model_dump()["gpu_ids"] == [1]
+
+
+def test_gpu_ids_property_default_and_reset():
+ backend = LlamaCppBackend()
+ assert backend.gpu_ids is None
+ backend._gpu_ids = [0, 1]
+ assert backend.gpu_ids == [0, 1]
+ backend._process = _FakeProcess()
+ backend.unload_model()
+ assert backend.gpu_ids is None
+
+
+def _target_state_gpu_ids(backend, gpu_ids):
+ return backend._already_in_target_state(
+ gguf_path = None,
+ model_identifier = "owner/repo",
+ hf_variant = "Q4_K_M",
+ n_ctx = 8192,
+ cache_type_kv = None,
+ speculative_type = "auto",
+ chat_template_override = None,
+ extra_args = None,
+ is_vision = False,
+ gpu_ids = gpu_ids,
+ )
+
+
+def test_gpu_ids_reload_detection_is_order_insensitive():
+ backend = _loaded_backend("auto")
+ backend._gpu_ids = [0, 1]
+ # Same set, different order -> no reload.
+ assert _target_state_gpu_ids(backend, [1, 0]) is True
+ # Different set -> reload.
+ assert _target_state_gpu_ids(backend, [0]) is False
+ # Dropping the pick (auto) -> reload.
+ assert _target_state_gpu_ids(backend, None) is False
+
+
+def test_gpu_ids_reload_detection_collapses_diffusion_to_single_device():
+ # The diffusion runner drives only its single lowest device, so the backend
+ # records [lowest]. A later multi-GPU request that still resolves to that
+ # same lowest device must dedupe (no needless reload); a request whose lowest
+ # device moves, or that drops the pick, must reload.
+ backend = _loaded_backend("auto")
+ backend._is_diffusion = True
+ backend._gpu_ids = [1] # loaded on the lowest of an earlier [3, 1] pick
+ assert _target_state_gpu_ids(backend, [3, 1]) is True
+ assert _target_state_gpu_ids(backend, [1]) is True
+ # Lowest device changes (2, not 1) -> reload.
+ assert _target_state_gpu_ids(backend, [3, 2]) is False
+ # Dropping the pick (auto) -> reload.
+ assert _target_state_gpu_ids(backend, None) is False
+
+
+def test_start_diffusion_server_resets_tensor_parallel():
+ # A prior tensor-parallel chat load leaves self._tensor_parallel True (load_model
+ # phase 1 only kills the process, it skips the unload reset). Diffusion is never
+ # TP, so startup must clear it -- else /status misreports TP and an identical
+ # diffusion re-Apply reloads against stale tensor-parallel state.
+ src = inspect.getsource(llama_cpp_module.LlamaCppBackend._start_diffusion_server)
+ assert "self._tensor_parallel = False" in src
+
+
+def test_route_matches_loaded_settings_collapses_diffusion_gpu_ids():
+ # The route-level reload dedupe mirrors the backend: for a loaded diffusion
+ # model it compares the request against the single recorded device, not the
+ # full requested list, or a same-device multi-GPU pick reloads needlessly.
+ route_src = (Path(_BACKEND_DIR) / "routes" / "inference.py").read_text(encoding = "utf-8")
+ match_impl = route_src[route_src.index("def _request_matches_loaded_settings") :]
+ guard = match_impl.index("if llama_backend.is_diffusion:")
+ collapse = match_impl.index("[sorted(request.gpu_ids)[0]] if request.gpu_ids else None")
+ compare = match_impl.index("if _req_gpu_ids != llama_backend.gpu_ids:")
+ assert guard < collapse < compare
+
+
+# ── Manual tensor split: child enumeration pinned to the picker's order ──────
+
+
+def _patch_split_pin_env(monkeypatch, *, inherited, reported):
+ """Point the pin helper at a fake inherited mask and picker report.
+ ``reported`` None = enumeration unavailable (falls back to ascending)."""
+ import utils.hardware as hw
+
+ monkeypatch.setattr(
+ LlamaCppBackend, "_resolve_visible_physical_ids", staticmethod(lambda: inherited)
+ )
+ info = (
+ {"available": False}
+ if reported is None
+ else {
+ "available": True,
+ "index_kind": "physical",
+ "devices": [{"index": i} for i in reported],
+ }
+ )
+ monkeypatch.setattr(hw, "get_backend_visible_gpu_info", lambda: info)
+
+
+def test_split_pin_reorders_inherited_numeric_mask(monkeypatch):
+ # Parent CUDA_VISIBLE_DEVICES=3,1 makes the child enumerate dev0=phys3, but
+ # nvidia-smi reported the picker's list ascending -- the mask must be
+ # re-emitted in that order or the per-GPU shares land on the wrong cards.
+ _patch_split_pin_env(monkeypatch, inherited = [3, 1], reported = [1, 3])
+ env = {"CUDA_VISIBLE_DEVICES": "3,1"}
+ LlamaCppBackend._pin_visible_gpu_order_for_split(env)
+ assert env["CUDA_DEVICE_ORDER"] == "PCI_BUS_ID"
+ assert env["CUDA_VISIBLE_DEVICES"] == "1,3"
+
+
+def test_split_pin_keeps_mask_order_when_picker_reported_it(monkeypatch):
+ # Torch-fallback enumeration (no nvidia-smi) reports devices in inherited
+ # mask order, so the picker's split list follows the mask -- the pin must
+ # keep that order, not re-sort it into a mismatch.
+ _patch_split_pin_env(monkeypatch, inherited = [3, 1], reported = [3, 1])
+ env = {"CUDA_VISIBLE_DEVICES": "3,1"}
+ LlamaCppBackend._pin_visible_gpu_order_for_split(env)
+ assert env["CUDA_VISIBLE_DEVICES"] == "3,1"
+
+
+def test_split_pin_falls_back_to_ascending_without_report(monkeypatch):
+ # Enumeration unavailable: ascending physical is the best guess (it matches
+ # the dominant nvidia-smi report order).
+ _patch_split_pin_env(monkeypatch, inherited = [3, 1], reported = None)
+ env = {"CUDA_VISIBLE_DEVICES": "3,1"}
+ LlamaCppBackend._pin_visible_gpu_order_for_split(env)
+ assert env["CUDA_VISIBLE_DEVICES"] == "1,3"
+
+
+def test_split_pin_without_mask_only_sets_pci_order(monkeypatch):
+ # No inherited mask (or a UUID/MIG one resolving to None): enumeration order
+ # is fully fixed by CUDA_DEVICE_ORDER, so no mask is written.
+ _patch_split_pin_env(monkeypatch, inherited = None, reported = None)
+ env = {}
+ LlamaCppBackend._pin_visible_gpu_order_for_split(env)
+ assert env == {"CUDA_DEVICE_ORDER": "PCI_BUS_ID"}
+
+
+def test_split_pin_mirrors_hip_mask_on_rocm(monkeypatch):
+ # ROCm: the pin must land in HIP_VISIBLE_DEVICES too, and an inherited ROCR
+ # mask is cleared so the mask can't apply twice (ROCR re-indexes, then HIP
+ # would index into the already-reduced set).
+ _patch_split_pin_env(monkeypatch, inherited = [3, 1], reported = [1, 3])
+ torch_stub = _types.ModuleType("torch")
+ torch_stub.version = _types.SimpleNamespace(hip = "6.0")
+ monkeypatch.setitem(sys.modules, "torch", torch_stub)
+ env = {"CUDA_VISIBLE_DEVICES": "3,1", "ROCR_VISIBLE_DEVICES": "3,1"}
+ LlamaCppBackend._pin_visible_gpu_order_for_split(env)
+ assert env["CUDA_VISIBLE_DEVICES"] == "1,3"
+ assert env["HIP_VISIBLE_DEVICES"] == "1,3"
+ assert "ROCR_VISIBLE_DEVICES" not in env
+
+
+# ── Diffusion single-device selection ───────────────────────────────────────
+
+
+def test_diffusion_gpu_arg_uses_lowest_explicit_physical_id(monkeypatch):
+ monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "3,1")
+ monkeypatch.setenv("DG_GPU", "7")
+ assert LlamaCppBackend._diffusion_gpu_arg([3, 1]) == "1"
+
+
+def test_diffusion_gpu_arg_preserves_parent_mask_order(monkeypatch):
+ monkeypatch.delenv("DG_GPU", raising = False)
+ monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "3,1")
+ assert LlamaCppBackend._diffusion_gpu_arg(None) == "3"
+
+
+def test_diffusion_gpu_arg_honors_override_and_cpu_mask(monkeypatch):
+ monkeypatch.setenv("DG_GPU", "GPU-abc")
+ assert LlamaCppBackend._diffusion_gpu_arg(None) == "GPU-abc"
+ assert LlamaCppBackend._diffusion_gpu_arg(None, cpu_only = True) == ""
+
+
+# ── Deliberate zero-offload (manual gpu_layers=0): training-skip flag ─────────
+
+
+def test_zero_offload_flag_false_without_companions():
+ # CPU-only by construction: False lets training skip unloading a server that
+ # holds no VRAM.
+ cmd = ["llama-server", "-m", "model.gguf", "--gpu-layers", "0", "--fit", "off"]
+ assert LlamaCppBackend._zero_offload_gpu_flag(cmd, [(0, 8000, 24000)], {}) is False
+
+
+@pytest.mark.parametrize(
+ "companion",
+ ["--mmproj", "--model-draft", "-md", "--spec-draft-model", "-hfd"],
+)
+def test_zero_offload_flag_true_with_companion(companion):
+ # mmproj / a drafter offload to GPU regardless of --gpu-layers, so the
+ # server still holds VRAM and training must unload it. Drafter detection
+ # reuses the extras parser, so pass-through aliases count too.
+ cmd = ["llama-server", "-m", "model.gguf", "--gpu-layers", "0", companion, "x.gguf"]
+ assert LlamaCppBackend._zero_offload_gpu_flag(cmd, [(0, 8000, 24000)], {}) is True
+
+
+def test_zero_offload_flag_true_with_inline_companion_forms():
+ cmd = ["llama-server", "-m", "model.gguf", "--spec-draft-model=x.gguf"]
+ assert LlamaCppBackend._zero_offload_gpu_flag(cmd, [(0, 8000, 24000)], {}) is True
+ cmd = ["llama-server", "-m", "model.gguf", "--mmproj=proj.gguf"]
+ assert LlamaCppBackend._zero_offload_gpu_flag(cmd, [(0, 8000, 24000)], {}) is True
+
+
+def test_zero_offload_flag_true_with_env_drafter():
+ cmd = ["llama-server", "-m", "model.gguf", "--gpu-layers", "0"]
+ env = {"LLAMA_ARG_SPEC_DRAFT_MODEL": "x.gguf"}
+ assert LlamaCppBackend._zero_offload_gpu_flag(cmd, [(0, 8000, 24000)], env) is True
+
+
+@pytest.mark.parametrize(
+ "device_args",
+ [
+ ["--device", "CUDA0"],
+ ["--device=CUDA0"],
+ ["-dev", "CUDA0"],
+ ["--spec-draft-device", "CUDA0"],
+ ["--device-draft=CUDA0"],
+ ],
+)
+def test_zero_offload_flag_true_with_device_pin(device_args):
+ cmd = ["llama-server", "-m", "model.gguf", "--gpu-layers", "0", *device_args]
+ assert LlamaCppBackend._zero_offload_gpu_flag(cmd, [(0, 8000, 24000)], {}) is True
+
+
+def test_zero_offload_flag_true_with_env_device_pin():
+ cmd = ["llama-server", "-m", "model.gguf", "--gpu-layers", "0"]
+ env = {"LLAMA_ARG_DEVICE": "CUDA0"}
+ assert LlamaCppBackend._zero_offload_gpu_flag(cmd, [(0, 8000, 24000)], env) is True
+
+
+@pytest.mark.parametrize(
+ ("device_args", "env"),
+ [
+ (["--device", "cpu"], {}),
+ (["--device=none"], {}),
+ (["--spec-draft-device", "cpu"], {}),
+ ([], {"LLAMA_ARG_DEVICE": "none"}),
+ (["--device", "CUDA0", "--device", "cpu"], {}),
+ ],
+)
+def test_zero_offload_flag_false_with_cpu_device_pin(device_args, env):
+ cmd = ["llama-server", "-m", "model.gguf", "--gpu-layers", "0", *device_args]
+ assert LlamaCppBackend._zero_offload_gpu_flag(cmd, [(0, 8000, 24000)], env) is False
+
+
+def test_zero_offload_flag_true_with_surviving_tensor_mode():
+ cmd = ["llama-server", "-m", "model.gguf", "--gpu-layers", "0", "--split-mode", "tensor"]
+ assert LlamaCppBackend._zero_offload_gpu_flag(cmd, [(0, 8000, 24000)], {}) is True
+
+
+def test_zero_offload_flag_true_for_unmasked_vulkan(monkeypatch):
+ monkeypatch.setattr(LlamaCppBackend, "_is_vulkan_backend", staticmethod(lambda: True))
+ cmd = ["llama-server", "-m", "model.gguf", "--gpu-layers", "0"]
+ assert LlamaCppBackend._zero_offload_gpu_flag(cmd, [(0, 8000, 24000)], {}) is True
+
+
+def test_zero_offload_flag_none_without_gpus():
+ cmd = ["llama-server", "-m", "model.gguf", "--gpu-layers", "0"]
+ assert LlamaCppBackend._zero_offload_gpu_flag(cmd, [], {}) is None
+
+
+def test_cmd_has_gpu_companion_detection():
+ # The env mask for CPU-only zero-offload loads keys off this scan: any
+ # --mmproj form or a drafter (flag aliases / env) keeps the GPUs visible.
+ has = LlamaCppBackend._cmd_has_gpu_companion
+ assert has(["llama-server", "-m", "m.gguf"], {}) is False
+ assert has(["llama-server", "--mmproj", "p.gguf"], {}) is True
+ assert has(["llama-server", "--mmproj=p.gguf"], {}) is True
+ assert has(["llama-server", "-md", "d.gguf"], {}) is True
+ assert has(["llama-server"], {"LLAMA_ARG_SPEC_DRAFT_MODEL": "d.gguf"}) is True
+
+
+def test_cmd_companion_ignores_cpu_forced_drafter():
+ # A CPU-pinned drafter holds no VRAM: the zero-offload mask may hide the GPUs
+ # and training may leave the server alone.
+ has = LlamaCppBackend._cmd_has_gpu_companion
+ cmd = ["llama-server", "-md", "d.gguf", "--spec-draft-ngl", "0"]
+ assert has(cmd, {}) is False
+ cmd = ["llama-server", "-md", "d.gguf", "--spec-draft-device", "cpu"]
+ assert has(cmd, {}) is False
+ # mmproj still counts even alongside a CPU drafter.
+ cmd = ["llama-server", "-md", "d.gguf", "--spec-draft-ngl", "0", "--mmproj", "p.gguf"]
+ assert has(cmd, {}) is True
diff --git a/studio/backend/tests/test_gpu_selection.py b/studio/backend/tests/test_gpu_selection.py
index 69ad560788..d4f2fbe993 100644
--- a/studio/backend/tests/test_gpu_selection.py
+++ b/studio/backend/tests/test_gpu_selection.py
@@ -853,7 +853,13 @@ class TestRouteErrors(unittest.TestCase):
self.assertIn("only supported on CUDA devices", str(exc_info.exception))
- def test_inference_route_rejects_gpu_ids_for_gguf(self):
+ def test_inference_route_validates_gpu_ids_for_gguf(self):
+ # gpu_ids is now SUPPORTED for GGUF (the GPU picker), but still
+ # validated: a rejected pick surfaces as a clean 400, not the old
+ # "not supported for GGUF" rejection. Patch the validator so the test
+ # is deterministic regardless of the host's (or a prior test's) GPU env.
+ import utils.hardware.hardware as hardware_mod
+
inference_route = _load_route_module(
"inference_route_module_for_gguf_gpu_ids_test",
"routes/inference.py",
@@ -887,6 +893,11 @@ class TestRouteErrors(unittest.TestCase):
),
patch.object(inference_route.asyncio, "to_thread", new = _inline_to_thread),
patch.object(inference_route, "_hf_offline_if_dns_dead", nullcontext),
+ patch.object(
+ hardware_mod,
+ "resolve_requested_gpu_ids",
+ side_effect = ValueError("Invalid gpu_ids [0, 1]: rejected by test"),
+ ),
):
with self.assertRaises(HTTPException) as exc_info:
asyncio.run(
@@ -901,8 +912,11 @@ class TestRouteErrors(unittest.TestCase):
)
)
+ # The validator's ValueError becomes a clean 400 (not the removed
+ # "not supported for GGUF" rejection).
self.assertEqual(exc_info.exception.status_code, 400)
- self.assertIn("GGUF", exc_info.exception.detail)
+ self.assertIn("gpu_ids", exc_info.exception.detail.lower())
+ self.assertNotIn("not supported", exc_info.exception.detail.lower())
def test_training_route_returns_400_for_invalid_gpu_ids(self):
training_route = _load_route_module(
diff --git a/studio/backend/tests/test_llama_cpp_no_context_shift.py b/studio/backend/tests/test_llama_cpp_no_context_shift.py
index f320d29a02..662c918305 100644
--- a/studio/backend/tests/test_llama_cpp_no_context_shift.py
+++ b/studio/backend/tests/test_llama_cpp_no_context_shift.py
@@ -118,9 +118,17 @@ def test_flag_sits_inside_the_base_cmd_list():
"conditional branch -- otherwise some code paths would still "
"run with silent context shift enabled."
)
- # Pin that it sits next to -c / --ctx so the grouping makes sense.
- assert '"-c"' in block
assert '"--flash-attn"' in block
+ # -c is emitted in the conditional right after the base list, not inside
+ # it: auto-fit (--fit on with no pinned context) must omit -c entirely,
+ # because "-c 0" pins the full native context and disables --fit's
+ # VRAM-based sizing. Pin that it still sits next to the base block so the
+ # context grouping stays intact.
+ after = rest[end_rel : end_rel + 1000]
+ assert '"-c"' in after, (
+ "-c must still be emitted in the conditional immediately after the "
+ "base cmd list (omitted only in auto-fit, where --fit sizes context)."
+ )
def _iter_lines_with_offset(text: str):
diff --git a/studio/backend/tests/test_llama_cpp_props_readback.py b/studio/backend/tests/test_llama_cpp_props_readback.py
index 488645ee5a..fe1e67edad 100644
--- a/studio/backend/tests/test_llama_cpp_props_readback.py
+++ b/studio/backend/tests/test_llama_cpp_props_readback.py
@@ -225,31 +225,46 @@ def test_kv_unified_added_for_multi_slot():
"""Explicit --parallel N disables llama-server's auto-slots kv-unified
default, splitting -c into per-slot windows of -c/N; Unsloth must restore
the shared pool so one request can use the full advertised context."""
- flags = LlamaCppBackend._ctx_integrity_flags(4, False, 98304, 98304, _CAPS_ALL)
+ flags = LlamaCppBackend._ctx_integrity_flags(4, False, False, 98304, 98304, _CAPS_ALL)
assert "--kv-unified" in flags
def test_kv_unified_skipped_for_single_slot_or_old_build():
assert "--kv-unified" not in LlamaCppBackend._ctx_integrity_flags(
- 1, False, 98304, 98304, _CAPS_ALL
+ 1, False, False, 98304, 98304, _CAPS_ALL
)
assert "--kv-unified" not in LlamaCppBackend._ctx_integrity_flags(
- 4, False, 98304, 98304, _CAPS_NONE
+ 4, False, False, 98304, 98304, _CAPS_NONE
)
def test_fit_ctx_floors_explicit_request_under_fit():
- flags = LlamaCppBackend._ctx_integrity_flags(1, True, 98304, 98304, _CAPS_ALL)
+ # An explicit requested ctx floors --fit-ctx at that value on any --fit
+ # path, including legacy auto (auto_fit False).
+ flags = LlamaCppBackend._ctx_integrity_flags(1, True, False, 98304, 98304, _CAPS_ALL)
assert flags[flags.index("--fit-ctx") + 1] == "98304"
-def test_fit_ctx_skipped_without_fit_or_explicit_ctx_or_support():
+def test_fit_ctx_skipped_without_fit_or_support():
+ # No --fit on -> no --fit-ctx.
assert "--fit-ctx" not in LlamaCppBackend._ctx_integrity_flags(
- 1, False, 98304, 98304, _CAPS_ALL
+ 1, False, False, 98304, 98304, _CAPS_ALL
)
- assert "--fit-ctx" not in LlamaCppBackend._ctx_integrity_flags(1, True, 0, 262144, _CAPS_ALL)
+ # --fit on but the binary doesn't support --fit-ctx.
assert "--fit-ctx" not in LlamaCppBackend._ctx_integrity_flags(
- 1, True, 98304, 98304, _CAPS_NONE
+ 1, True, True, 98304, 98304, _CAPS_NONE
+ )
+
+
+def test_fit_ctx_floors_auto_request_at_8192_only_under_auto_fit():
+ # Manual + Auto (auto_fit) floors the auto window at 8192 so --fit can't
+ # shrink it to a tiny size.
+ flags = LlamaCppBackend._ctx_integrity_flags(1, True, True, 0, 262144, _CAPS_ALL)
+ assert flags[flags.index("--fit-ctx") + 1] == "8192"
+ # Legacy auto (fit on but not auto_fit) emits -c 0 to pin native, so the
+ # 8192 floor must NOT ride along and override that pin.
+ assert "--fit-ctx" not in LlamaCppBackend._ctx_integrity_flags(
+ 1, True, False, 0, 262144, _CAPS_ALL
)
diff --git a/studio/backend/tests/test_llama_server_args.py b/studio/backend/tests/test_llama_server_args.py
index ba52afad1c..c6d16363f8 100644
--- a/studio/backend/tests/test_llama_server_args.py
+++ b/studio/backend/tests/test_llama_server_args.py
@@ -747,6 +747,34 @@ def test_strip_shadowing_flags_defaults_strip_split_mode_too():
assert strip_shadowing_flags(["--split-mode", "tensor"]) == []
+def test_strip_offload_is_opt_in_and_covers_moe():
+ base = dict(
+ strip_context = False,
+ strip_cache = False,
+ strip_spec = False,
+ strip_template = False,
+ strip_split_mode = False,
+ )
+ # Default: offload (incl. MoE) flags are NOT stripped.
+ assert strip_shadowing_flags(["--n-cpu-moe", "8", "--top-k", "20"], **base) == [
+ "--n-cpu-moe",
+ "8",
+ "--top-k",
+ "20",
+ ]
+ # Opt-in strips layer AND MoE offload flags (value-aware), keeps the rest.
+ assert strip_shadowing_flags(
+ ["--n-cpu-moe", "8", "--gpu-layers", "33", "--fit", "off", "--top-k", "20"],
+ **base,
+ strip_offload = True,
+ ) == ["--top-k", "20"]
+ # Boolean --cpu-moe drops the flag only, not the following value.
+ assert strip_shadowing_flags(["--cpu-moe", "--seed", "-1"], **base, strip_offload = True) == [
+ "--seed",
+ "-1",
+ ]
+
+
@pytest.mark.parametrize(
"args",
[
@@ -796,6 +824,23 @@ def test_strip_split_mode_only_drops_tensor_split_too():
assert strip_split_mode_only(["-sm=tensor", "-ts=3,1"]) == []
+def test_strip_tensor_split_alone_preserves_split_mode():
+ # Manual mode emits its own --tensor-split, so an inherited ratio is dropped
+ # -- but the user's --split-mode row/none/layer choice (which the manual
+ # ratio toggle can't express) must survive. strip_tensor_split removes only
+ # the ratio, unlike strip_split_mode which removes the whole group.
+ out = strip_shadowing_flags(
+ ["--split-mode", "row", "--tensor-split", "1,1", "--top-k", "20"],
+ strip_context = False,
+ strip_cache = False,
+ strip_spec = False,
+ strip_template = False,
+ strip_split_mode = False,
+ strip_tensor_split = True,
+ )
+ assert out == ["--split-mode", "row", "--top-k", "20"]
+
+
def test_strip_shadowing_flags_keeps_model_draft_without_spec():
out = strip_shadowing_flags(
["--model-draft", "/custom/mtp.gguf"],
diff --git a/studio/backend/tests/test_tensor_parallel.py b/studio/backend/tests/test_tensor_parallel.py
index 06f72d3b9f..00c7aeac69 100644
--- a/studio/backend/tests/test_tensor_parallel.py
+++ b/studio/backend/tests/test_tensor_parallel.py
@@ -262,9 +262,12 @@ def test_proportional_tensor_split_is_emitted_in_tensor_mode():
src = _load_model_source()
assert '"--tensor-split"' in src
gate = src.find("if tensor_parallel:")
- ts = src.find('"--tensor-split"')
+ # Find the TP block's emission (after the gate); manual mode emits its own
+ # --tensor-split earlier in the source from the user's per-GPU shares.
+ ts = src.find('"--tensor-split"', gate)
nxt_else = src.find("self._tensor_parallel = False")
assert 0 <= gate < ts < nxt_else, "--tensor-split must be emitted under `if tensor_parallel:`"
+ assert "tp_tensor_split" in src[gate:nxt_else]
def test_mtp_decode_probe_wired_under_tensor_parallel():
diff --git a/studio/backend/tests/test_tp_vision_regression.py b/studio/backend/tests/test_tp_vision_regression.py
index fb0989b306..d1372ca415 100644
--- a/studio/backend/tests/test_tp_vision_regression.py
+++ b/studio/backend/tests/test_tp_vision_regression.py
@@ -126,10 +126,21 @@ _ALLOWED_TP_DROP_GUARDS = {
# Capability: --split-mode tensor aborted for this (binary, model) (#6415).
# Self-healing -- tried by default, skipped only after a real abort (vs #6416).
"tensor_parallel and self._tensor_split_aborts(binary, model_identifier)",
- # Capacity: tensor needs >= 2 GPUs clearing the compute-buffer reserve.
- "tensor_parallel and len(tp_gpus) < 2",
+ # Capacity: tensor needs >= 2 GPUs clearing the compute-buffer reserve. Gated
+ # on plan_tp (not raw tensor_parallel) so manual mode skips this planner (#6414).
+ "plan_tp and len(tp_gpus) < 2",
# Capacity: pooled usable VRAM can't hold weights + MTP reserve -> layer split.
"_tp_weight_budget_mib <= _tp_required_mib",
+ # Manual mode, Auto layers: --fit owns memory and is incompatible with a
+ # tensor split, so TP is dropped (surfaced via logger.info) before the
+ # cache-drop, so a quantized KV survives into the --fit load (#6414).
+ "tensor_parallel and gpu_memory_mode == 'manual' and (gpu_layers < 0)",
+ # Manual mode, explicit layers: a tensor split still needs >= 2 GPUs in use.
+ "tensor_parallel and gpu_memory_mode == 'manual' and (gpu_layers >= 0) and (self._effective_gpu_count(sorted(gpu_ids) if gpu_ids else None) < 2)",
+ # Manual mode, zero layers: nothing to split on the GPU, and a tensor-mode
+ # launch under the CPU-only GPU mask (no visible devices) aborts the server
+ # instead of the intended CPU-only load (#6414).
+ "gpu_memory_mode == 'manual' and gpu_layers == 0",
}
@@ -364,7 +375,7 @@ def test_compute_buffer_downgrade_preserves_multi_gpu_intent():
full GPU set too, so it is symmetric with the budget/geometry downgrades and
doesn't collapse a multi-GPU layer load to one card (reviewer.py P1 on #6659)."""
src = inspect.getsource(LlamaCppBackend.load_model)
- gate = src.find("tensor_parallel and len(tp_gpus) < 2")
+ gate = src.find("plan_tp and len(tp_gpus) < 2")
assert gate != -1
# Bound to exactly this block: from its gate to the next (budget) downgrade.
nxt = src.find("_tp_weight_budget_mib <= _tp_required_mib", gate)
diff --git a/studio/backend/utils/models/gguf_metadata.py b/studio/backend/utils/models/gguf_metadata.py
index c24ec28e1d..50b3cd3513 100644
--- a/studio/backend/utils/models/gguf_metadata.py
+++ b/studio/backend/utils/models/gguf_metadata.py
@@ -50,9 +50,11 @@ _CACHE_MAX_ENTRIES = 4096
# keyed by (file cache key, wanted key). None = key absent / file unreadable.
_BOOL_CACHE: Dict[Tuple[_CacheKey, str], Optional[bool]] = {}
-# Native training context length (``{arch}.context_length``). None = absent /
-# unreadable. Lets the UI show the real context ceiling before a model loads.
-_CONTEXT_CACHE: Dict[_CacheKey, Optional[int]] = {}
+# GGUF header dims for the staged/deferred-load UI: context_length, layer_count
+# (block_count), and moe_layer_count (block_count minus leading dense layers; 0
+# if not MoE). One cached pass fills all three so the staged sheet can size every
+# slider before the model loads. None = unreadable / not a GGUF.
+_DIMS_CACHE: Dict[_CacheKey, Optional[Dict[str, Optional[int]]]] = {}
def _cache_key(path: str) -> Optional[_CacheKey]:
@@ -142,32 +144,45 @@ def _parse_gguf_header(path: str) -> Optional[Dict[str, str]]:
return out
-def read_gguf_context_length(path: str) -> Optional[int]:
- """Return the GGUF's native training context length (``{arch}.context_length``),
- or ``None`` if missing/unreadable/not a GGUF. Cached by (path, mtime, size).
- Lets the UI populate the context slider before the model is loaded."""
+def read_gguf_staged_dims(path: str) -> Optional[Dict[str, Optional[int]]]:
+ """GGUF header dims for the staged-load UI in one cached pass:
+ ``{"context_length", "layer_count", "moe_layer_count"}``. Each may be None
+ when absent (moe_layer_count is 0 for a dense model). Returns ``None`` if not
+ a GGUF / unreadable. Cached by (path, mtime, size). Lets the staged sheet size
+ the context, GPU-layers and MoE sliders before the model loads."""
key = _cache_key(path)
if key is None:
return None
with _CACHE_LOCK:
- if key in _CONTEXT_CACHE:
- return _CONTEXT_CACHE[key]
- result = _parse_gguf_context_length(path)
+ if key in _DIMS_CACHE:
+ return _DIMS_CACHE[key]
+ result = _parse_gguf_staged_dims(path)
with _CACHE_LOCK:
- while len(_CONTEXT_CACHE) >= _CACHE_MAX_ENTRIES:
+ while len(_DIMS_CACHE) >= _CACHE_MAX_ENTRIES:
try:
- _CONTEXT_CACHE.pop(next(iter(_CONTEXT_CACHE)))
+ _DIMS_CACHE.pop(next(iter(_DIMS_CACHE)))
except StopIteration:
break
- _CONTEXT_CACHE[key] = result
+ _DIMS_CACHE[key] = result
return result
-def _parse_gguf_context_length(path: str) -> Optional[int]:
- # The context key is architecture-namespaced (``llama.context_length`` etc.),
- # so we learn the key only after reading ``general.architecture``. GGUF writes
- # general.* before arch.* keys, matching the loader's own parser.
- ctx_key: Optional[str] = None
+def read_gguf_context_length(path: str) -> Optional[int]:
+ """Native training context length (``{arch}.context_length``), or ``None``.
+ Thin accessor over read_gguf_staged_dims."""
+ dims = read_gguf_staged_dims(path)
+ return dims["context_length"] if dims else None
+
+
+def _parse_gguf_arch_uints(path: str, wanted_suffixes: frozenset[str]) -> Optional[Dict[str, int]]:
+ """Walk a GGUF header once and return the requested architecture-namespaced
+ uint (vtype 4/10) keys, e.g. ``{"block_count": 32}``. Keys are
+ ``{arch}.``; the arch is learned from ``general.architecture`` (GGUF
+ writes general.* before arch.* keys, matching the loader's own parser).
+ Returns ``None`` if not a GGUF / unreadable, else a dict (possibly empty or
+ partial when some keys are absent)."""
+ arch: Optional[str] = None
+ found: Dict[str, int] = {}
try:
with open(path, "rb") as f:
head = f.read(24)
@@ -204,28 +219,68 @@ def _parse_gguf_context_length(path: str) -> Optional[int]:
sbytes = f.read(slen)
if len(sbytes) < slen:
break
- ctx_key = f"{sbytes.decode('utf-8', 'replace')}.context_length"
- elif ctx_key is not None and key == ctx_key and vtype in (4, 10):
+ arch = sbytes.decode("utf-8", "replace")
+ elif (
+ arch is not None
+ and vtype in (4, 10)
+ and key.startswith(f"{arch}.")
+ and key[len(arch) + 1 :] in wanted_suffixes
+ ):
width = 4 if vtype == 4 else 8
n_bytes = f.read(width)
if len(n_bytes) < width:
break
- value = struct.unpack(" 0 else None
+ found[key[len(arch) + 1 :]] = struct.unpack(
+ " Optional[Dict[str, Optional[int]]]:
+ vals = _parse_gguf_arch_uints(
+ path,
+ frozenset(
+ {
+ "context_length",
+ "block_count",
+ "expert_count",
+ "leading_dense_block_count",
+ }
+ ),
+ )
+ if vals is None:
+ return None
+ ctx = vals.get("context_length")
+ block = vals.get("block_count")
+ # A real context/layer count is positive; treat 0/garbage as absent so the
+ # UI never builds a slider with max < min.
+ context_length = ctx if ctx and ctx > 0 else None
+ layer_count = block if block and block > 0 else None
+ # MoE layer count = block_count - leading dense layers, only when experts
+ # exist; else 0 (dense -> slider hidden). Mirrors n_moe_layers in
+ # core/inference/llama_cpp.py.
+ if not vals.get("expert_count") or not block:
+ moe_layer_count: Optional[int] = 0
+ else:
+ moe_layer_count = max(0, block - (vals.get("leading_dense_block_count") or 0))
+ return {
+ "context_length": context_length,
+ "layer_count": layer_count,
+ "moe_layer_count": moe_layer_count,
+ }
# Strings (8) and arrays (9) are handled inline.
diff --git a/studio/frontend/src/components/assistant-ui/model-selector/remembered-load-settings.ts b/studio/frontend/src/components/assistant-ui/model-selector/remembered-load-settings.ts
index ec75b17f20..08492ab480 100644
--- a/studio/frontend/src/components/assistant-ui/model-selector/remembered-load-settings.ts
+++ b/studio/frontend/src/components/assistant-ui/model-selector/remembered-load-settings.ts
@@ -2,7 +2,9 @@
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
// Per-model pre-load inference settings, persisted in localStorage so the load
-// dialog can offer "Remember settings for ".
+// dialog can offer "Remember settings for ". GGUF picks only: every
+// field is a llama.cpp load knob, so all save/restore call sites gate on
+// GGUF-ness (a non-GGUF blob would only snapshot leftover standing values).
const KEY = "unsloth_load_settings";
@@ -12,14 +14,22 @@ export interface RememberedLoadSettings {
speculativeType: string | null;
specDraftNMax: number | null;
tensorParallel: boolean;
+ // GPU Memory controls. Optional so an older blob (which lacked them) still
+ // parses, leaving the live knobs untouched on apply. The mode is kept with the
+ // manual knobs (gpuLayers/nCpuMoe are ignored outside Manual mode). A null
+ // selectedGpuIds is meaningful (all GPUs), so it's distinguished from absent.
+ // The per-GPU split ratio is deliberately NOT remembered: it's positionally
+ // bound to the exact GPU set/order and unvalidated, so it would mismatch.
+ gpuMemoryMode?: "auto" | "manual";
+ gpuLayers?: number;
+ nCpuMoe?: number;
+ selectedGpuIds?: number[] | null;
}
-// Storage key for a pick's remembered settings. The remembered knobs are
-// VRAM-budget driven (context override, KV-cache dtype, tensor-parallel), so the
-// right values differ per quant. An HF repo collapses all its GGUF variants into
-// one `id`, so fold the variant in to scope settings per quant. Local .gguf
-// paths key by their file path (already file-specific); native drag-drop files
-// key by display label, so same-named files in different folders share an entry.
+// Storage key for a pick's remembered settings, scoped per quant (the VRAM-budget
+// knobs differ per quant). An HF repo collapses its GGUF variants into one `id`,
+// so fold the variant in. Local .gguf paths are already file-specific; native
+// drag-drop files key by display label, so same-named files share an entry.
export function rememberedLoadSettingsKey(selection: {
id: string;
ggufVariant?: string | null;
diff --git a/studio/frontend/src/features/chat/api/chat-adapter.ts b/studio/frontend/src/features/chat/api/chat-adapter.ts
index 0bf46e7343..7083f02288 100644
--- a/studio/frontend/src/features/chat/api/chat-adapter.ts
+++ b/studio/frontend/src/features/chat/api/chat-adapter.ts
@@ -45,12 +45,18 @@ import {
import {
type PendingImageEditReference,
type RagAutoInject,
+ GPU_LAYERS_AUTO,
+ loadedGpuMemoryFieldsUnlessStaged,
+ reconcilePersistedGpuIds,
resolveLoadedSpeculativeSettings,
resolveSpeculativeSettingsForLoad,
+ persistGpuMemoryModeOnLoad,
resolveToolsEnabledOnLoad,
saveSpeculativeType,
useChatRuntimeStore,
} from "../stores/chat-runtime-store";
+import { resolveFitMaxSeqLength, resolveManualAutoCtxPin } from "../presets/preset-policy";
+import { ensureGpuDeviceCache } from "@/hooks/use-gpu-info";
import { useExternalProvidersStore } from "../stores/external-providers-store";
import {
shouldPreserveFullOutput,
@@ -1489,6 +1495,13 @@ async function autoLoadSmallestModel(): Promise<{
max_seq_length: number;
is_lora: boolean;
gguf_variant?: string | null;
+ // GGUF-only: scopes the training guard to the same placement policy /load
+ // will use. Manual mode must match because it makes placement user-owned.
+ // The layer/MoE/split/KV/spec knobs are deliberately not sent: Auto mode's
+ // guard sizes conservatively, while Manual mode bypasses that estimate.
+ // The safetensors fallback omits both fields and uses HF auto-placement.
+ gpu_ids?: number[];
+ gpu_memory_mode?: "auto" | "manual";
}): Promise {
const validation = await validateModel({
...payload,
@@ -1520,12 +1533,18 @@ async function autoLoadSmallestModel(): Promise<{
return false;
}
const currentStore = useChatRuntimeStore.getState();
- const remembered = loadRememberedLoadSettings(
- rememberedLoadSettingsKey({
- id: candidate.id,
- ggufVariant: candidate.ggufVariant,
- }),
- );
+ // Blobs are saved for GGUF picks only (the sheet gates on it), so don't
+ // let a legacy non-GGUF blob feed a stale context/spec choice into a
+ // safetensors auto-load.
+ const remembered =
+ candidate.kind === "gguf"
+ ? loadRememberedLoadSettings(
+ rememberedLoadSettingsKey({
+ id: candidate.id,
+ ggufVariant: candidate.ggufVariant,
+ }),
+ )
+ : null;
const effectiveMaxSeqLength = resolveLoadMaxSeqLength({
modelId: candidate.id,
ggufVariant: candidate.ggufVariant,
@@ -1537,6 +1556,38 @@ async function autoLoadSmallestModel(): Promise<{
maxSeqLength: candidate.maxSeqLength,
presetSource: currentStore.activePresetSource,
});
+ // The GPU knobs are per-model, so read them from the same remembered
+ // settings that fed effectiveMaxSeqLength -- on a background auto-load the
+ // live store holds session defaults, not the saved Manual mode / layer pin /
+ // GPU pick. Absent fields fall back like applyRememberedLoadSettings: the
+ // mode to the store (a persisted standing preference), the per-model knobs to
+ // their defaults. The saved GPU pick is reconciled against the GPUs present
+ // now, like the interactive restore.
+ const effectiveGpuMemoryMode =
+ remembered?.gpuMemoryMode ?? currentStore.gpuMemoryMode;
+ const effectiveGpuLayers = remembered?.gpuLayers ?? GPU_LAYERS_AUTO;
+ const effectiveNCpuMoe = remembered?.nCpuMoe ?? 0;
+ if (remembered?.selectedGpuIds != null) {
+ // Warm the device cache first: on a cold cache the reconcile passes the
+ // saved pick through unvalidated, and a stale cross-host pick then fails
+ // the load with the picker hidden.
+ await ensureGpuDeviceCache();
+ }
+ const effectiveGpuIds =
+ remembered?.selectedGpuIds !== undefined
+ ? reconcilePersistedGpuIds(remembered.selectedGpuIds)
+ : null;
+ // Under Manual GPU memory + Auto layers, llama.cpp's --fit owns context
+ // sizing, so send 0 (or the pinned length). GGUF-only; a no-op otherwise.
+ // The context pin is per-model too, so it comes from remembered settings,
+ // not the live store.
+ const fitMaxSeqLength = resolveFitMaxSeqLength(
+ candidate.kind === "gguf",
+ effectiveGpuMemoryMode,
+ effectiveGpuLayers,
+ remembered?.contextLength ?? null,
+ effectiveMaxSeqLength,
+ );
const effectiveSpeculativeType =
remembered?.speculativeType ?? specSettings.speculativeType;
const effectiveSpecDraftNMax =
@@ -1544,9 +1595,16 @@ async function autoLoadSmallestModel(): Promise<{
if (
!(await canAutoLoad({
model_path: candidate.id,
- max_seq_length: effectiveMaxSeqLength,
+ max_seq_length: fitMaxSeqLength,
is_lora: false,
gguf_variant: candidate.ggufVariant,
+ // The same remembered-derived GPU pick the load below sends.
+ ...(candidate.kind === "gguf"
+ ? {
+ gpu_ids: effectiveGpuIds ?? undefined,
+ gpu_memory_mode: effectiveGpuMemoryMode,
+ }
+ : {}),
}))
) {
skippedAutoLoadCandidates.add(
@@ -1558,7 +1616,7 @@ async function autoLoadSmallestModel(): Promise<{
const loadResp = await loadModel({
model_path: candidate.id,
hf_token: hfToken,
- max_seq_length: effectiveMaxSeqLength,
+ max_seq_length: fitMaxSeqLength,
load_in_4bit: true,
is_lora: false,
gguf_variant: candidate.ggufVariant,
@@ -1567,8 +1625,22 @@ async function autoLoadSmallestModel(): Promise<{
speculative_type: effectiveSpeculativeType,
spec_draft_n_max: effectiveSpecDraftNMax,
tensor_parallel: remembered?.tensorParallel ?? false,
+ // GGUF-only: the safetensors fallback loads via HF auto-placement (no
+ // explicit pins). The split ratio is deliberately never remembered
+ // (positionally bound to an exact GPU set), so auto-load leaves llama.cpp's
+ // free-VRAM default in charge rather than sending a stale store value.
+ ...(candidate.kind === "gguf"
+ ? {
+ gpu_memory_mode: effectiveGpuMemoryMode,
+ gpu_layers: effectiveGpuLayers,
+ n_cpu_moe: effectiveNCpuMoe,
+ gpu_ids: effectiveGpuIds ?? undefined,
+ }
+ : {}),
});
saveSpeculativeType(effectiveSpeculativeType);
+ // Self-gates on is_gguf (skips diffusion), so persists only for a real GGUF load.
+ persistGpuMemoryModeOnLoad(loadResp, effectiveGpuMemoryMode);
useChatRuntimeStore
.getState()
.setCheckpoint(candidate.id, candidate.ggufVariant ?? undefined);
@@ -1597,6 +1669,15 @@ async function autoLoadSmallestModel(): Promise<{
store.setModels([...store.models, autoModel]);
}
if (candidate.kind === "gguf") {
+ // Keep an explicit Manual+Auto context pin the load just applied (so a
+ // later Apply doesn't silently revert it to auto-fit sizing), mirroring
+ // the interactive path's keepCustomCtx; other cases baseline on
+ // ggufContextLength.
+ const keepCustomCtx = resolveManualAutoCtxPin(
+ effectiveGpuMemoryMode,
+ effectiveGpuLayers,
+ remembered?.contextLength ?? null,
+ );
useChatRuntimeStore.setState({
ggufContextLength: loadResp.context_length ?? 131072,
ggufMaxContextLength:
@@ -1613,6 +1694,10 @@ async function autoLoadSmallestModel(): Promise<{
loadedKvCacheDtype: loadResp.cache_type_kv ?? null,
tensorParallel: loadResp.tensor_parallel ?? false,
loadedTensorParallel: loadResp.tensor_parallel ?? false,
+ ...loadedGpuMemoryFieldsUnlessStaged(loadResp, {
+ customContextLength: keepCustomCtx,
+ }),
+ loadedCustomContextLength: keepCustomCtx,
defaultChatTemplate: loadResp.chat_template ?? null,
chatTemplateOverride: null,
loadedChatTemplateOverride: null,
@@ -1633,6 +1718,9 @@ async function autoLoadSmallestModel(): Promise<{
loadedKvCacheDtype: loadResp.cache_type_kv ?? null,
tensorParallel: loadResp.tensor_parallel ?? false,
loadedTensorParallel: loadResp.tensor_parallel ?? false,
+ // Non-GGUF response: clears any stale GPU baseline a prior manual-GPU
+ // GGUF load left, matching the interactive/status sibling load paths.
+ ...loadedGpuMemoryFieldsUnlessStaged(loadResp),
defaultChatTemplate: loadResp.chat_template ?? null,
chatTemplateOverride: null,
loadedChatTemplateOverride: null,
@@ -1820,12 +1908,17 @@ async function autoLoadSmallestModel(): Promise<{
duration: 30000,
});
try {
+ const rt = useChatRuntimeStore.getState();
if (
!(await canAutoLoad({
model_path: "unsloth/Qwen3.5-4B-MTP-GGUF",
max_seq_length: 0,
is_lora: false,
gguf_variant: "UD-Q4_K_XL",
+ // The same live-store GPU pick the load below sends (a fresh default
+ // model has no remembered settings to prefer).
+ gpu_ids: rt.selectedGpuIds ?? undefined,
+ gpu_memory_mode: rt.gpuMemoryMode,
}))
) {
toast.dismiss(toastId);
@@ -1835,6 +1928,9 @@ async function autoLoadSmallestModel(): Promise<{
const loadResp = await loadModel({
model_path: "unsloth/Qwen3.5-4B-MTP-GGUF",
hf_token: hfToken,
+ // Model default under both modes: Auto layers + no pin means
+ // resolveFitMaxSeqLength returns 0 for every mode (the canAutoLoad
+ // preflight above sends the same).
max_seq_length: 0,
load_in_4bit: true,
is_lora: false,
@@ -1842,8 +1938,20 @@ async function autoLoadSmallestModel(): Promise<{
trust_remote_code: trustRemoteCode,
speculative_type: specSettings.speculativeType,
spec_draft_n_max: specSettings.specDraftNMax,
+ // GPU Memory mode is a standing preference, so honor it on auto-load.
+ // The layer/MoE/split knobs and the context pin are per-model: the live
+ // store may hold edits drafted for a staged pick, and a fresh default
+ // model has no remembered settings, so those stay at their defaults like
+ // the cached-candidate path. The GPU pick deliberately differs (it's the
+ // picker's current on-screen selection, which the canAutoLoad preflight
+ // above already committed to).
+ gpu_memory_mode: rt.gpuMemoryMode,
+ gpu_layers: GPU_LAYERS_AUTO,
+ n_cpu_moe: 0,
+ gpu_ids: rt.selectedGpuIds ?? undefined,
});
saveSpeculativeType(specSettings.speculativeType);
+ persistGpuMemoryModeOnLoad(loadResp, rt.gpuMemoryMode);
useChatRuntimeStore
.getState()
.setCheckpoint("unsloth/Qwen3.5-4B-MTP-GGUF", "UD-Q4_K_XL");
@@ -1880,6 +1988,10 @@ async function autoLoadSmallestModel(): Promise<{
loadedKvCacheDtype: loadResp.cache_type_kv ?? null,
tensorParallel: loadResp.tensor_parallel ?? false,
loadedTensorParallel: loadResp.tensor_parallel ?? false,
+ ...loadedGpuMemoryFieldsUnlessStaged(loadResp),
+ // Drives the GPU Memory controls' diffusion gate; set alongside the
+ // GPU fields on every load path so the gate can't read stale.
+ loadedIsDiffusion: loadResp.is_diffusion ?? false,
defaultChatTemplate: loadResp.chat_template ?? null,
chatTemplateOverride: null,
loadedIsMultimodal: isMultimodalResponse(loadResp),
diff --git a/studio/frontend/src/features/chat/api/chat-api.ts b/studio/frontend/src/features/chat/api/chat-api.ts
index ebf9461172..0f6af38033 100644
--- a/studio/frontend/src/features/chat/api/chat-api.ts
+++ b/studio/frontend/src/features/chat/api/chat-api.ts
@@ -127,28 +127,38 @@ export async function validateModel(
native_path_lease: payload.nativePathLease ?? null,
hf_token: payload.hf_token,
gguf_variant: payload.gguf_variant ?? null,
- // Send the intended load settings so validate's VRAM check matches the
- // follow-up /load and doesn't unload for a load /load would then reject.
+ // Intended load settings so validate's preflight matches the follow-up
+ // /load. Default placement is sized against the selected GPUs.
max_seq_length: payload.max_seq_length,
load_in_4bit: payload.load_in_4bit,
+ gpu_ids: payload.gpu_ids,
+ // Manual placement is an explicit override: Auto layers use llama.cpp
+ // --fit, while a pinned layer count is owned by the user. Tell validate
+ // so it applies the same training-guard policy as /load.
+ gpu_memory_mode: payload.gpu_memory_mode,
}),
});
return parseJsonOrThrow(response);
}
/**
- * Read a GGUF's native context length from its local header (no GPU load, no
- * download). Returns null when the file isn't downloaded yet, the model isn't a
- * GGUF, or it's gated. For a native (drag-drop / picked) file, pass
- * `nativePathToken` so the backend reads the granted local path. Used by the
- * deferred-load staging flow to fill the context slider before the single load.
+ * Read a GGUF's header dims (native context length, total layer count, MoE
+ * expert-layer count) from its local file (no GPU load, no download). All are
+ * null when the file isn't downloaded yet, the model isn't a GGUF, or it's
+ * gated. For a native (drag-drop / picked) file, pass `nativePathToken` so the
+ * backend reads the granted local path. Used by the deferred-load staging flow
+ * to size the context, GPU-layers and MoE sliders before the single load.
*/
-export async function fetchGgufContextLength(payload: {
+export async function fetchGgufStagedMetadata(payload: {
model_path: string;
gguf_variant?: string | null;
hf_token?: string | null;
nativePathToken?: string | null;
-}): Promise {
+}): Promise<{
+ contextLength: number | null;
+ layerCount: number | null;
+ moeLayerCount: number | null;
+}> {
let nativePathLease: string | null = null;
if (payload.nativePathToken) {
try {
@@ -156,8 +166,8 @@ export async function fetchGgufContextLength(payload: {
await consumeNativePathToken(payload.nativePathToken, "validate-model")
).nativePathLease;
} catch {
- // Lease expired / revoked: degrade to no context (the load can re-mint).
- return null;
+ // Lease expired / revoked: degrade to no metadata (the load can re-mint).
+ return { contextLength: null, layerCount: null, moeLayerCount: null };
}
}
const response = await authFetch("/api/inference/validate", {
@@ -172,7 +182,11 @@ export async function fetchGgufContextLength(payload: {
}),
});
const res = await parseJsonOrThrow(response);
- return res.context_length ?? null;
+ return {
+ contextLength: res.context_length ?? null,
+ layerCount: res.layer_count ?? null,
+ moeLayerCount: res.moe_layer_count ?? null,
+ };
}
export async function unloadModel(payload: UnloadModelRequest): Promise {
diff --git a/studio/frontend/src/features/chat/chat-page.tsx b/studio/frontend/src/features/chat/chat-page.tsx
index ec0ad977bf..217eaf8b6d 100644
--- a/studio/frontend/src/features/chat/chat-page.tsx
+++ b/studio/frontend/src/features/chat/chat-page.tsx
@@ -1445,9 +1445,11 @@ export function ChatPage({
// were already seeded on stage, so keepSpeculative only when a config was
// saved -- otherwise the standing speculative preference should win.
autoLoadStagedRef.current = (pending) => {
- const remembered = loadRememberedLoadSettings(
- rememberedLoadSettingsKey(pending),
- );
+ // Blobs are saved for GGUF picks only (the sheet gates on it), so don't
+ // let a legacy non-GGUF blob claim a seeded config here.
+ const remembered = hasGgufSource(pending)
+ ? loadRememberedLoadSettings(rememberedLoadSettingsKey(pending))
+ : null;
void selectModel({
...pending,
isDownloaded: true,
@@ -2813,6 +2815,11 @@ export function ChatPage({
selectModel({
id: state.params.checkpoint,
ggufVariant: state.activeGgufVariant ?? undefined,
+ // A native (drag-drop / picked) GGUF's checkpoint is only a display
+ // label, so the reload needs its path token to re-mint a lease --
+ // else applying the now-exposed GPU/context controls can't resolve
+ // the file. Null for non-native loads, which reload by id as before.
+ nativePathToken: state.activeNativePathToken ?? undefined,
forceReload: true,
isDownloaded: true,
loadingDescription: "Reloading with updated chat template.",
diff --git a/studio/frontend/src/features/chat/chat-settings-sheet.tsx b/studio/frontend/src/features/chat/chat-settings-sheet.tsx
index cedd298ecf..b368a811fa 100644
--- a/studio/frontend/src/features/chat/chat-settings-sheet.tsx
+++ b/studio/frontend/src/features/chat/chat-settings-sheet.tsx
@@ -55,6 +55,7 @@ import { Switch } from "@/components/ui/switch";
import { Textarea } from "@/components/ui/textarea";
import { InfoHint } from "@/components/ui/info-hint";
import { Tooltip, TooltipContent } from "@/components/ui/tooltip";
+import { useGpuDevices } from "@/hooks/use-gpu-info";
import { useIsMobile } from "@/hooks/use-mobile";
import { useLlamaUpdateCheck } from "@/hooks/use-llama-update-check";
import { cn } from "@/lib/utils";
@@ -99,8 +100,11 @@ import {
providerSupportsFastMode,
} from "./provider-capabilities";
import {
+ GPU_LAYERS_AUTO,
+ distributeByWeight,
isPendingGguf,
pendingSelectionMatches,
+ rebalanceSplit,
useChatRuntimeStore,
} from "./stores/chat-runtime-store";
import { RetrievalSettingsSection } from "@/features/rag/components/retrieval-settings-section";
@@ -250,6 +254,7 @@ function ParamSlider({
displayValue,
info,
valueSize,
+ disabled,
}: {
label: string;
value: number;
@@ -260,6 +265,7 @@ function ParamSlider({
displayValue?: string;
info?: ReactNode;
valueSize?: number;
+ disabled?: boolean;
}) {
return (
onChange(snapToStep(v, step, min, max))}
className="panel-slider"
+ disabled={disabled}
/>
);
@@ -540,8 +548,17 @@ export function ChatSettingsPanel({
const base = slash >= 0 ? id.slice(slash + 1) : id;
return base || id;
})();
+ const activeNativePathToken = useChatRuntimeStore(
+ (s) => s.activeNativePathToken,
+ );
+ const loadedGgufContextLength = useChatRuntimeStore((s) => s.ggufContextLength);
+ // A GGUF loaded from a native path / direct .gguf has no HF variant, so key
+ // off the same signal the status hydration uses -- variant OR native token OR
+ // a GGUF context -- else the GPU Memory controls hide for a loaded local GGUF.
const isLoadedGguf =
- useChatRuntimeStore((s) => s.activeGgufVariant) != null;
+ useChatRuntimeStore((s) => s.activeGgufVariant) != null ||
+ activeNativePathToken != null ||
+ loadedGgufContextLength != null;
// While a pick is staged the sheet configures *that* model, so its GGUF-ness
// (not the currently loaded model's) decides whether the GGUF-only controls
// show. Otherwise a staged non-GGUF Hub repo would inherit the loaded GGUF's
@@ -607,6 +624,25 @@ export function ChatSettingsPanel({
const loadedTensorParallel = useChatRuntimeStore(
(s) => s.loadedTensorParallel,
);
+ const gpuMemoryMode = useChatRuntimeStore((s) => s.gpuMemoryMode);
+ const setGpuMemoryMode = useChatRuntimeStore((s) => s.setGpuMemoryMode);
+ const loadedGpuMemoryMode = useChatRuntimeStore((s) => s.loadedGpuMemoryMode);
+ const loadedIsDiffusion = useChatRuntimeStore((s) => s.loadedIsDiffusion);
+ const gpuLayers = useChatRuntimeStore((s) => s.gpuLayers);
+ const setGpuLayers = useChatRuntimeStore((s) => s.setGpuLayers);
+ const loadedGpuLayers = useChatRuntimeStore((s) => s.loadedGpuLayers);
+ const nCpuMoe = useChatRuntimeStore((s) => s.nCpuMoe);
+ const setNCpuMoe = useChatRuntimeStore((s) => s.setNCpuMoe);
+ const loadedNCpuMoe = useChatRuntimeStore((s) => s.loadedNCpuMoe);
+ const splitRatio = useChatRuntimeStore((s) => s.splitRatio);
+ const setSplitRatio = useChatRuntimeStore((s) => s.setSplitRatio);
+ const loadedSplitRatio = useChatRuntimeStore((s) => s.loadedSplitRatio);
+ const ggufLayerCount = useChatRuntimeStore((s) => s.ggufLayerCount);
+ const moeLayerCount = useChatRuntimeStore((s) => s.moeLayerCount);
+ const selectedGpuIds = useChatRuntimeStore((s) => s.selectedGpuIds);
+ const setSelectedGpuIds = useChatRuntimeStore((s) => s.setSelectedGpuIds);
+ const loadedGpuIds = useChatRuntimeStore((s) => s.loadedGpuIds);
+ const gpuDevices = useGpuDevices();
const chatTemplateOverride = useChatRuntimeStore(
(s) => s.chatTemplateOverride,
);
@@ -614,6 +650,9 @@ export function ChatSettingsPanel({
(s) => s.loadedChatTemplateOverride,
);
const customContextLength = useChatRuntimeStore((s) => s.customContextLength);
+ const loadedCustomContextLength = useChatRuntimeStore(
+ (s) => s.loadedCustomContextLength,
+ );
const setCustomContextLength = useChatRuntimeStore(
(s) => s.setCustomContextLength,
);
@@ -641,10 +680,14 @@ export function ChatSettingsPanel({
: null;
useEffect(() => {
if (!pendingKey) return;
- const saved = loadRememberedLoadSettings(pendingKey);
+ // GGUF-only, like the stageOrLoad / Hub restore paths: every remembered
+ // field is a llama.cpp knob, so a non-GGUF pick has nothing to restore --
+ // and applying its blob would clobber the standing gpuMemoryMode with a
+ // stale snapshot (the save on Load below is gated the same way).
+ const saved = pendingIsGguf ? loadRememberedLoadSettings(pendingKey) : null;
setRemember(saved != null);
if (saved) applyRememberedLoadSettings(saved);
- }, [pendingKey, applyRememberedLoadSettings]);
+ }, [pendingKey, pendingIsGguf, applyRememberedLoadSettings]);
// While staging, the sheet reflects the STAGED model, so its header context
// takes precedence over the loaded model's (which may differ or be larger).
const baseContext = pendingIsGguf ? stagedContextLength : ggufContextLength;
@@ -661,15 +704,132 @@ export function ChatSettingsPanel({
const ctxDisplayValue = customContextLength ?? baseContext ?? "";
const ctxMaxValue = baseNativeContext ?? baseContext ?? null;
const kvDirty = kvCacheDtype !== loadedKvCacheDtype;
- const ctxDirty = customContextLength !== null;
+ const ctxDirty = customContextLength !== loadedCustomContextLength;
const specDirty = speculativeType !== loadedSpeculativeType;
const specDraftDirty = specDraftNMax !== loadedSpecDraftNMax;
const tpDirty = tensorParallel !== (loadedTensorParallel ?? false);
+ // A loaded diffusion GGUF runs mode-agnostic (pins all layers on one GPU,
+ // ignores --fit/--gpu-layers), so the GPU Memory mode + manual controls don't
+ // apply -- hide them and don't let the preserved standing mode read as dirty.
+ // The GPU picker still applies (diffusion pins the chosen device). A staged pick
+ // keeps the controls (a pending pick's diffusion-ness isn't known until load).
+ const gpuModeApplies =
+ isGguf && (pendingSelection != null || !loadedIsDiffusion);
+ const gpuDirty =
+ gpuModeApplies && gpuMemoryMode !== (loadedGpuMemoryMode ?? "auto");
+ const isManual = gpuModeApplies && gpuMemoryMode === "manual";
+ // Manual with the GPU Layers slider at "Auto" (leftmost): --fit owns the whole
+ // layout, so the offload knobs (MoE, split, TP) don't apply.
+ const autoLayers = isManual && gpuLayers < 0;
+ // GPUs actually in use: the picked subset, or all visible when none picked.
+ const gpusInUse = selectedGpuIds ?? gpuDevices.map((d) => d.index);
+ // TP is off with fewer than 2 GPUs in use (single GPU, or the picker narrowed
+ // to one): tensor split is a no-op there and aborts on some archs. Mirrors the
+ // multi-GPU gate on the GPU picker / Split ratio. (Under Auto layers the whole
+ // TP control is hidden -- llama.cpp's --fit aborts under --split-mode tensor.)
+ const tpDisabled = gpusInUse.length <= 1;
+ // Manual gpu-layers ceiling = model layer count + 1 (else a safe fallback):
+ // llama.cpp counts the output layer as one more offloadable layer past the
+ // repeating blocks ("offloaded 33/33" needs -ngl 33 on a 32-block model), so
+ // the slider max must reach it or full offload is unreachable. While staging,
+ // use the staged model's layer count (read from its header).
+ const stagedLayerCount = pendingSelection?.layerCount ?? null;
+ const modelLayerCount = pendingIsGguf ? stagedLayerCount : ggufLayerCount;
+ const gpuLayersMax = modelLayerCount != null ? modelLayerCount + 1 : 256;
+ // MoE-offload slider: shown only for MoE models, capped at their MoE-layer
+ // count. While staging, use the staged model's count (read from its header);
+ // otherwise the loaded model's.
+ const stagedMoeLayerCount = pendingSelection?.moeLayerCount ?? null;
+ const moeLayersMax = pendingIsGguf
+ ? (stagedMoeLayerCount ?? 0)
+ : (moeLayerCount ?? 0);
+ const showMoeSlider = isManual && !autoLayers && moeLayersMax > 0;
+ // gpuLayers always counts; MoE only with an explicit layer count (see above).
+ const manualDirty =
+ isManual &&
+ (gpuLayers !== loadedGpuLayers ||
+ (!autoLayers && nCpuMoe !== (loadedNCpuMoe ?? 0)));
+ // GPU picker: only meaningful on multi-GPU, and only when the reported
+ // indices are physical (relative ordinals from a parent CUDA_VISIBLE_DEVICES
+ // mask can't be mapped back to pin a device). null = use all (auto).
+ const showGpuPicker =
+ isGguf &&
+ gpuDevices.length > 1 &&
+ gpuDevices.every((d) => d.physicalIndex);
+ const isGpuChecked = (index: number) =>
+ selectedGpuIds === null || selectedGpuIds.includes(index);
+ const toggleGpu = (index: number) => {
+ const all = gpuDevices.map((d) => d.index);
+ const current = selectedGpuIds ?? all;
+ const next = current.includes(index)
+ ? current.filter((i) => i !== index)
+ : [...current, index].sort((a, b) => a - b);
+ if (next.length === 0) return; // keep at least one GPU selected
+ setSelectedGpuIds(next.length === all.length ? null : next);
+ // The per-GPU split is positional, so any change to the set of GPUs in use
+ // invalidates it: drop it (the sliders fall back to the VRAM-weighted
+ // default). TP needs 2+ GPUs, so disable it when only one remains.
+ setSplitRatio(null);
+ if (next.length <= 1) {
+ setTensorParallel(false);
+ }
+ };
+ const gpuIdsKey = (ids: number[] | null) => (ids === null ? "auto" : ids.join(","));
+ const gpuIdsDirty = gpuIdsKey(selectedGpuIds) !== gpuIdsKey(loadedGpuIds);
+ // Per-GPU layer split (--tensor-split): manual + 2+ GPUs in use. One slider
+ // per GPU, each a layer count; together they sum to the GPU Layers total.
+ const showSplitRatio =
+ isManual && !autoLayers && showGpuPicker && gpusInUse.length > 1;
+ // The total the per-GPU counts sum to (the GPU Layers slider value); 0 under
+ // Auto, where the split is hidden. The devices behind the GPUs in use, for
+ // labels + the VRAM-weighted default.
+ const splitTotal = Math.max(0, Math.min(gpuLayers, gpuLayersMax));
+ const gpusInUseDevices = gpusInUse.map(
+ (i) => gpuDevices.find((d) => d.index === i) ?? null,
+ );
+ // Displayed per-GPU counts. splitRatio is a stable reference balance (only a
+ // slider edit changes it), rescaled to the current total; deriving rather than
+ // mutating it on GPU Layers changes keeps the balance intact when the total
+ // passes through low values or Auto. No saved split: free-VRAM-weighted default
+ // (llama.cpp's unset default splits by free VRAM, so the first edit starts from
+ // the default's placement, not a total-VRAM ratio that can land layers on a
+ // busy GPU). A genuine 0 (a full GPU) is a real weight, not missing data: the
+ // probe's no-data case degrades to the total server-side, and an all-zero list
+ // falls back to an even split in distributeByWeight. Not yet sent.
+ const splitCounts =
+ splitRatio && splitRatio.length === gpusInUse.length
+ ? distributeByWeight(splitTotal, splitRatio)
+ : distributeByWeight(
+ splitTotal,
+ gpusInUseDevices.map((d) => d?.memoryFreeGb ?? d?.memoryTotalGb ?? 1),
+ );
+ const setSplitCount = (k: number, v: number) =>
+ setSplitRatio(rebalanceSplit(splitTotal, splitCounts, k, v));
+ const splitRatioDirty =
+ isManual &&
+ !autoLayers &&
+ JSON.stringify(splitRatio ?? null) !== JSON.stringify(loadedSplitRatio ?? null);
+ // Auto-fit context (Manual + Auto layers): <= 0 means "Auto" (--fit sizes it);
+ // a positive value pins it. Surface the length --fit chose once it's loaded.
+ const fitCtxAuto = autoLayers && (customContextLength ?? 0) <= 0;
+ const loadedAutoLayers =
+ loadedGpuMemoryMode === "manual" && (loadedGpuLayers ?? GPU_LAYERS_AUTO) < 0;
+ const fitResolvedCtx =
+ fitCtxAuto && loadedAutoLayers ? ggufContextLength : null;
// A saved chat-template override is a reload-time setting too, so surface
// Apply for a template-only edit (otherwise it could never be applied).
const templateDirty = chatTemplateOverride !== loadedChatTemplateOverride;
const modelSettingsDirty =
- kvDirty || ctxDirty || specDirty || specDraftDirty || tpDirty || templateDirty;
+ kvDirty ||
+ ctxDirty ||
+ specDirty ||
+ specDraftDirty ||
+ tpDirty ||
+ gpuDirty ||
+ manualDirty ||
+ gpuIdsDirty ||
+ splitRatioDirty ||
+ templateDirty;
const [presetNameInput, setPresetNameInput] = useState(activePreset);
const [systemPromptEditorOpen, setSystemPromptEditorOpen] = useState(false);
const [systemPromptDraft, setSystemPromptDraft] = useState("");
@@ -980,7 +1140,64 @@ export function ChatSettingsPanel({
)}
{isGguf && (
<>
- {showContextControl && (
+ {showContextControl && (autoLayers ? (
+
+
+
+
+ Context Length
+
+
+ Auto: llama.cpp's --fit sizes the context to fit VRAM.
+ Set a length to pin it instead -- --fit then optimizes
+ GPU layer offload around it. The length --fit chose
+ shows here after loading.
+
+
+ Default: Unsloth
+ fits the model and context to your GPUs.
+
+
+ Manual: set GPU
+ Layers yourself. Leave it on Auto to let llama.cpp size
+ the context and offload overflow (including MoE experts)
+ to RAM.
+
+
+
+
+
+
+
+
+ )}
+ {isManual && (
+ <>
+
+ Layers to keep on the GPU (--gpu-layers); the rest run
+ on CPU. Auto lets llama.cpp size the split (and the
+ context) to fit VRAM. At the maximum, the whole model
+ is on the GPU.
+ >
+ }
+ />
+ {showMoeSlider && (
+
+ Keep the experts of this many MoE layers on the CPU
+ (--n-cpu-moe) to save VRAM. 0 = all experts on the
+ GPU; at the maximum, all are on the CPU.
+ >
+ }
+ />
+ )}
+ {showSplitRatio && (
+
+
+
+ Layers per GPU
+
+
+ Splits GPU Layers across GPUs (--tensor-split).
+ Without Tensor Parallelism each value is the layer
+ count on that GPU; with it, every GPU holds a slice
+ of each layer, so the values are only a ratio.
+
+
+
+ GPUs
+
+
+ Which GPUs this model may use. Unchecked GPUs are hidden
+ from llama.cpp (CUDA_VISIBLE_DEVICES, or
+ HIP_VISIBLE_DEVICES on ROCm). Leave all checked to use
+ every GPU.
+
+
+ )}
>
)}
{/* No persistent "enable custom code" toggle: it is consented per model
@@ -1228,14 +1603,21 @@ export function ChatSettingsPanel({
{Math.round((stagedDownloadFraction ?? 0) * 100)}%
)}
-
+ {/* GGUF picks only: a non-GGUF pick shows none of the load
+ knobs the blob captures, so there is nothing to remember. */}
+ {pendingIsGguf && (
+
+ )}
{stagedLoading ? (
// Mid-load: nothing to load or abandon until it settles, so disable.
) : null}
-
+ {/* The template override is a load-time knob too (applied on the next
+ reload) and the in-flight load already snapshotted it, so lock its
+ editors like the sibling controls -- a mid-load save would be
+ silently clobbered by the load response despite its toast. */}
+
diff --git a/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts b/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts
index 10e0904e4f..3003b52230 100644
--- a/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts
+++ b/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts
@@ -29,9 +29,14 @@ import {
} from "../api/chat-api";
import { formatEta, formatRate } from "../utils/format-transfer";
import {
+ GPU_LAYERS_AUTO,
isLocalModelPath,
+ loadedGpuMemoryFields,
+ loadedGpuMemoryFieldsUnlessStaged,
pendingSelectionMatches,
+ persistGpuMemoryModeOnLoad,
readPersistedSpeculativeType,
+ reconcilePersistedGpuIds,
resolveToolsEnabledOnLoad,
saveSpeculativeType,
useChatRuntimeStore,
@@ -46,9 +51,12 @@ import {
} from "../lib/apply-inference-status-to-store";
import {
mergeBackendRecommendedInference,
+ resolveFitMaxSeqLength,
resolveLoadMaxSeqLength,
+ resolveManualAutoCtxPin,
} from "../presets/preset-policy";
import { recordLastLocalModelLoad } from "../utils/last-local-model-load";
+import { ensureGpuDeviceCache } from "@/hooks/use-gpu-info";
import {
isMultimodalResponse,
} from "../types/api";
@@ -291,9 +299,12 @@ async function syncInferenceStatusToStore(options?: {
if (statusRes.active_model && !isExternalSelectionActive) {
const checkpointId = resolveInferenceCheckpointId(statusRes);
if (checkpointId) {
+ const previousGgufVariant =
+ useChatRuntimeStore.getState().activeGgufVariant;
setCheckpoint(checkpointId, statusRes.gguf_variant);
applyActiveModelStatusToStore(statusRes, {
previousCheckpoint: selectedCheckpoint,
+ previousGgufVariant,
});
// setModels(listRes...) above used catalog data, which omits audio
// capability. Re-apply live status so attach gates survive a refresh.
@@ -511,7 +522,11 @@ export function useChatModelRuntime() {
typeof selection === "string" ? false : selection.isDownloaded ?? false;
const model = models.find((entry) => entry.id === modelId);
const lora = loras.find((entry) => entry.id === modelId);
- const isGguf = explicitIsGguf ?? model?.isGguf ?? false;
+ // A native path-token selection is a local GGUF by construction (the
+ // native model intents only grant .gguf files), but its id is a display
+ // label that need not end in ".gguf" -- without this, Manual + Auto
+ // layers would pin the UI context instead of letting --fit size it.
+ const isGguf = explicitIsGguf ?? model?.isGguf ?? nativePathToken != null;
const loraIsAdapter = lora?.exportType === "lora";
const isLora =
explicitIsLora ?? model?.isLora ?? loraIsAdapter ?? false;
@@ -578,18 +593,27 @@ export function useChatModelRuntime() {
let trustRemoteCode = stateBeforeUnload.params.trustRemoteCode ?? false;
let approvedRemoteCodeFingerprint: string | null = null;
const maxSeqLength = stateBeforeUnload.params.maxSeqLength;
+ const previousActiveNativePathToken =
+ stateBeforeUnload.activeNativePathToken;
const previousIsGguf =
previousModel?.isGguf === true
|| previousVariant != null
+ || previousActiveNativePathToken != null
|| (previousCheckpoint?.toLowerCase().endsWith(".gguf") ?? false);
- const rollbackMaxSeqLength = previousIsGguf
- ? (stateBeforeUnload.ggufContextLength ?? 0)
- : maxSeqLength;
+ // Respect the rolled-back model's auto-layers mode: a Manual+Auto model
+ // with an unpinned (auto) context must reload with 0 (so --fit
+ // re-auto-sizes), not the positive context it happened to pick (which
+ // the backend would treat as a pin).
+ const rollbackMaxSeqLength = resolveFitMaxSeqLength(
+ previousIsGguf,
+ stateBeforeUnload.loadedGpuMemoryMode ?? "auto",
+ stateBeforeUnload.loadedGpuLayers ?? GPU_LAYERS_AUTO,
+ stateBeforeUnload.loadedCustomContextLength,
+ previousIsGguf ? (stateBeforeUnload.ggufContextLength ?? 0) : maxSeqLength,
+ );
const hfToken = stateBeforeUnload.hfToken || null;
const previousModelRequiresTrustRemoteCode =
stateBeforeUnload.modelRequiresTrustRemoteCode;
- const previousActiveNativePathToken =
- stateBeforeUnload.activeNativePathToken;
// Snapshot the load settings at click time, before the awaits below
// (validation, the trust dialog, unload). For a staged Load these knobs
// stay editable and a sheet-close revert (abandonStagedModel) can fire
@@ -598,11 +622,29 @@ export function useChatModelRuntime() {
// updates this snapshot in lock-step so non-staged loads are unchanged.
const loadChatTemplateOverride = stateBeforeUnload.chatTemplateOverride;
const loadKvCacheDtype = stateBeforeUnload.kvCacheDtype;
- const loadCustomContextLength = stateBeforeUnload.customContextLength;
+ // gpuMemoryMode is a standing preference (kept across a model switch);
+ // the rest are per-model knobs the reset below clears, so they are
+ // re-baselined there in lock-step with the store.
+ let loadCustomContextLength = stateBeforeUnload.customContextLength;
const loadGgufContextLength = stateBeforeUnload.ggufContextLength;
const loadTensorParallel = stateBeforeUnload.tensorParallel;
const loadActivePresetSource = stateBeforeUnload.activePresetSource;
const loadActiveGgufVariant = stateBeforeUnload.activeGgufVariant;
+ const loadGpuMemoryMode = stateBeforeUnload.gpuMemoryMode;
+ let loadGpuLayers = stateBeforeUnload.gpuLayers;
+ let loadNCpuMoe = stateBeforeUnload.nCpuMoe;
+ let loadSplitRatio = stateBeforeUnload.splitRatio;
+ // Reconcile the persisted pick against the GPUs present now, so a stale
+ // cross-host / now-hidden pick is dropped before /load rather than
+ // rejected there. Warm the device cache first: load-on-selection can
+ // run before any GPU hook mounted, and a cold cache would pass the
+ // pick through unvalidated. validateGpuIds derives from this too.
+ if (stateBeforeUnload.selectedGpuIds != null) {
+ await ensureGpuDeviceCache();
+ }
+ let loadSelectedGpuIds = reconcilePersistedGpuIds(
+ stateBeforeUnload.selectedGpuIds,
+ );
let loadSpeculativeType = stateBeforeUnload.speculativeType;
let loadSpecDraftNMax = stateBeforeUnload.specDraftNMax;
try {
@@ -615,16 +657,47 @@ export function useChatModelRuntime() {
// context can exceed maxSeqLength, so sizing on raw maxSeqLength could
// pass, unload, then have /load refuse it. Uses the click-time
// snapshot (same values loadModel uses below), so the two agree.
- const validateMaxSeqLength = resolveLoadMaxSeqLength({
- modelId,
- ggufVariant,
- customContextLength: loadCustomContextLength,
- ggufContextLength: loadGgufContextLength,
- currentCheckpoint,
- activeGgufVariant: loadActiveGgufVariant,
- maxSeqLength,
- presetSource: loadActivePresetSource,
- });
+ // Mirror what /load does on a cross-model switch: the reset below
+ // clears the per-model Auto-layers context pin + GPU pick, and
+ // Manual+Auto sizes context through resolveFitMaxSeqLength.
+ // gpuMemoryMode is a standing preference, kept across the switch.
+ // A same-repo quant switch (same checkpoint, different gguf_variant)
+ // is a different model for per-model knobs: the pinned context,
+ // gpuLayers, GPU pick, and MoE offload are scoped per variant, so
+ // treat a variant change like a model switch and re-baseline them.
+ const switchingModelOrVariant =
+ currentCheckpoint !== modelId ||
+ (loadActiveGgufVariant ?? null) !== (ggufVariant ?? null);
+ const resetsPerModelSettings = Boolean(
+ currentCheckpoint && switchingModelOrVariant && !keepSpeculative,
+ );
+ const validateCustomContextLength = resetsPerModelSettings
+ ? null
+ : loadCustomContextLength;
+ const validateGpuIds = resetsPerModelSettings
+ ? null
+ : loadSelectedGpuIds;
+ // The reset below re-baselines gpuLayers to Auto; mirror it here.
+ const validateGpuLayers = resetsPerModelSettings
+ ? GPU_LAYERS_AUTO
+ : loadGpuLayers;
+ const validateMaxSeqLength = resolveFitMaxSeqLength(
+ isGguf,
+ loadGpuMemoryMode,
+ validateGpuLayers,
+ validateCustomContextLength,
+ resolveLoadMaxSeqLength({
+ modelId,
+ ggufVariant,
+ isGguf,
+ customContextLength: validateCustomContextLength,
+ ggufContextLength: loadGgufContextLength,
+ currentCheckpoint,
+ activeGgufVariant: loadActiveGgufVariant,
+ maxSeqLength,
+ presetSource: loadActivePresetSource,
+ }),
+ );
const validation = await validateModel({
model_path: modelId,
nativePathLease: validateNativePathLease,
@@ -633,6 +706,8 @@ export function useChatModelRuntime() {
load_in_4bit: true,
is_lora: isLora,
gguf_variant: ggufVariant ?? null,
+ gpu_ids: validateGpuIds ?? undefined,
+ ...(isGguf ? { gpu_memory_mode: loadGpuMemoryMode } : {}),
});
// Upgrade consent runs before the security dialogs; Accept installs and the load continues.
if (validation.requires_transformers_upgrade) {
@@ -697,18 +772,52 @@ export function useChatModelRuntime() {
// keepSpeculative skips this for a staged Load: the user picked the
// mode for this model on the sidebar, so honor it (the backend still
// falls back at runtime if the model has no MTP head).
- if (currentCheckpoint && currentCheckpoint !== modelId && !keepSpeculative) {
+ if (resetsPerModelSettings) {
const persistedSpeculativeType = readPersistedSpeculativeType();
useChatRuntimeStore.setState({
speculativeType: persistedSpeculativeType,
loadedSpeculativeType: persistedSpeculativeType,
specDraftNMax: null,
loadedSpecDraftNMax: null,
+ // Per-model GPU knobs must not follow onto a different model
+ // (gpuMemoryMode is a standing preference and is kept).
+ selectedGpuIds: null,
+ gpuLayers: GPU_LAYERS_AUTO,
+ nCpuMoe: 0,
+ splitRatio: null,
+ // A Manual+Auto context pin is per-model; clear it so a different
+ // model loads at Auto/native, not the previous model's pin.
+ customContextLength: null,
});
loadSpeculativeType = persistedSpeculativeType;
loadSpecDraftNMax = null;
+ // Keep the click-time snapshot in lock-step with the store reset so
+ // the load below sizes against the cleared per-model knobs, not the
+ // previous model's (gpuMemoryMode is standing, so left as captured).
+ loadCustomContextLength = null;
+ loadSelectedGpuIds = null;
+ loadGpuLayers = GPU_LAYERS_AUTO;
+ loadNCpuMoe = 0;
+ loadSplitRatio = null;
}
+ // Pinning layers on the SAME model keeps the currently resolved
+ // context: with no explicit pin, a manual+pinned reload would send 0,
+ // which the backend's --fit off branch treats as the NATIVE context --
+ // far larger than the sheet shows when the load was fit-sized (Default
+ // or Manual + Auto layers may auto-reduce context to fit VRAM), a
+ // likely OOM. ggufContextLength is that resolved value; a model already
+ // at native reloads unchanged, so this is safe for any prior mode.
+ if (
+ isGguf &&
+ !switchingModelOrVariant &&
+ loadGpuMemoryMode === "manual" &&
+ loadGpuLayers >= 0 &&
+ loadCustomContextLength == null &&
+ (loadGgufContextLength ?? 0) > 0
+ ) {
+ loadCustomContextLength = loadGgufContextLength;
+ }
const effectiveMaxSeqLength = resolveLoadMaxSeqLength({
modelId,
ggufVariant,
@@ -720,13 +829,20 @@ export function useChatModelRuntime() {
maxSeqLength,
presetSource: loadActivePresetSource,
});
+ const loadMaxSeqLength = resolveFitMaxSeqLength(
+ isGguf,
+ loadGpuMemoryMode,
+ loadGpuLayers,
+ loadCustomContextLength,
+ effectiveMaxSeqLength,
+ );
const effectiveChatTemplateOverride =
loadChatTemplateOverride?.trim() ? loadChatTemplateOverride : null;
const loadResponse = await loadModel({
model_path: modelId,
nativePathLease: loadNativePathLease,
hf_token: hfToken,
- max_seq_length: effectiveMaxSeqLength,
+ max_seq_length: loadMaxSeqLength,
load_in_4bit: true,
is_lora: isLora,
gguf_variant: ggufVariant ?? null,
@@ -737,6 +853,11 @@ export function useChatModelRuntime() {
speculative_type: loadSpeculativeType,
spec_draft_n_max: loadSpecDraftNMax,
tensor_parallel: loadTensorParallel,
+ gpu_memory_mode: loadGpuMemoryMode,
+ gpu_layers: loadGpuLayers,
+ n_cpu_moe: loadNCpuMoe,
+ tensor_split: loadSplitRatio ?? undefined,
+ gpu_ids: loadSelectedGpuIds ?? undefined,
});
// If cancelled while loading, don't update UI to show
@@ -747,6 +868,9 @@ export function useChatModelRuntime() {
// preference now (the requested intent, not the resolved echo;
// saveSpeculativeType keeps only the universal auto/ngram/off).
saveSpeculativeType(loadSpeculativeType);
+ // Persist the GPU Memory mode only on a successful load (not on
+ // dropdown change), so an abandoned selection doesn't stick.
+ persistGpuMemoryModeOnLoad(loadResponse, loadGpuMemoryMode);
const currentParams = useChatRuntimeStore.getState().params;
setParams(
@@ -782,9 +906,13 @@ export function useChatModelRuntime() {
const reportedNativeCtx = loadResponse.is_gguf
? (loadResponse.native_context_length ?? null)
: null;
- // A successful reload has applied settings, so clear pending custom
- // context state and display the backend-reported effective context.
- const keepCustomCtx = null;
+ // Keep an explicit Manual+Auto context pin (so a later Apply doesn't
+ // revert it to Auto); other cases baseline on ggufContextLength.
+ const keepCustomCtx = resolveManualAutoCtxPin(
+ loadGpuMemoryMode,
+ loadGpuLayers,
+ loadCustomContextLength,
+ );
const reasoningAlwaysOn = loadResponse.reasoning_always_on ?? false;
const reasoningStyle = loadResponse.reasoning_style ?? "enable_thinking";
const supportsReasoning = loadResponse.supports_reasoning ?? false;
@@ -837,11 +965,13 @@ export function useChatModelRuntime() {
loadedKvCacheDtype: loadedKv,
tensorParallel: loadedTp,
loadedTensorParallel: loadedTp,
+ ...loadedGpuMemoryFields(loadResponse),
speculativeType: loadedSpec,
loadedSpeculativeType: loadedSpec,
specDraftNMax: loadResponse.spec_draft_n_max ?? null,
loadedSpecDraftNMax: loadResponse.spec_draft_n_max ?? null,
customContextLength: keepCustomCtx,
+ loadedCustomContextLength: keepCustomCtx,
defaultChatTemplate: loadResponse.chat_template ?? null,
chatTemplateOverride: effectiveChatTemplateOverride,
loadedChatTemplateOverride: effectiveChatTemplateOverride,
@@ -938,7 +1068,7 @@ export function useChatModelRuntime() {
}
}
try {
- await loadModel({
+ const rollbackResponse = await loadModel({
model_path: previousCheckpoint,
nativePathLease: rollbackNativePathLease,
hf_token: hfToken,
@@ -951,14 +1081,51 @@ export function useChatModelRuntime() {
// Resend the previous model's pinned approval so restoring it is not re-blocked.
approved_remote_code_fingerprint:
approvedRemoteCodeFingerprints.get(previousCheckpoint) ?? null,
+ chat_template_override:
+ stateBeforeUnload.loadedChatTemplateOverride,
+ cache_type_kv: stateBeforeUnload.loadedKvCacheDtype,
+ speculative_type:
+ stateBeforeUnload.loadedSpeculativeType,
+ spec_draft_n_max:
+ stateBeforeUnload.loadedSpecDraftNMax,
// Restore the previous model in the split mode it was running,
// not the default layer split.
tensor_parallel: stateBeforeUnload.loadedTensorParallel ?? false,
+ gpu_memory_mode: stateBeforeUnload.loadedGpuMemoryMode ?? "auto",
+ gpu_layers: stateBeforeUnload.loadedGpuLayers ?? -1,
+ n_cpu_moe: stateBeforeUnload.loadedNCpuMoe ?? 0,
+ tensor_split: stateBeforeUnload.loadedSplitRatio ?? undefined,
+ gpu_ids: stateBeforeUnload.loadedGpuIds ?? undefined,
});
+ const rollbackSpeculativeType = normalizeSpeculativeType(
+ rollbackResponse.speculative_type,
+ );
useChatRuntimeStore.setState({
activeNativePathToken: previousActiveNativePathToken ?? null,
- loadedSpeculativeType: null,
- loadedSpecDraftNMax: null,
+ loadedSpeculativeType: rollbackSpeculativeType,
+ loadedSpecDraftNMax:
+ rollbackResponse.spec_draft_n_max ?? null,
+ loadedKvCacheDtype: rollbackResponse.cache_type_kv ?? null,
+ loadedChatTemplateOverride:
+ stateBeforeUnload.loadedChatTemplateOverride,
+ // Re-baseline the GPU knobs from the rolled-back load's own
+ // response (the shared seeding every load path uses): the
+ // refresh() below can't do it, since the status reseed is
+ // gated off while modelLoading is still true. A failed staged
+ // Load stays staged for retry, so the staged hold applies.
+ ...loadedGpuMemoryFieldsUnlessStaged(rollbackResponse, {
+ tensorParallel: rollbackResponse.tensor_parallel ?? false,
+ loadedTensorParallel:
+ rollbackResponse.tensor_parallel ?? false,
+ // refresh() is held while modelLoading remains true, so
+ // restore the rolled-back model's context pin directly.
+ customContextLength:
+ stateBeforeUnload.loadedCustomContextLength,
+ }),
+ loadedTensorParallel:
+ rollbackResponse.tensor_parallel ?? false,
+ loadedCustomContextLength:
+ stateBeforeUnload.loadedCustomContextLength,
});
await refresh();
} catch {
diff --git a/studio/frontend/src/features/chat/hooks/use-staged-model-preparation.ts b/studio/frontend/src/features/chat/hooks/use-staged-model-preparation.ts
index d8076c720b..a3e7a2d264 100644
--- a/studio/frontend/src/features/chat/hooks/use-staged-model-preparation.ts
+++ b/studio/frontend/src/features/chat/hooks/use-staged-model-preparation.ts
@@ -7,7 +7,7 @@ import { useRepoDownload } from "@/features/hub/download-manager/use-repo-downlo
import type { DownloadJob } from "@/features/hub/download-manager/use-repo-download";
import { useLatestRef } from "@/features/hub/hooks/use-latest-ref";
-import { fetchGgufContextLength } from "../api/chat-api";
+import { fetchGgufStagedMetadata } from "../api/chat-api";
import {
isPendingGguf,
pendingSelectionMatches,
@@ -46,8 +46,16 @@ export function useStagedModelPreparation(opts?: {
const pendingDownloaded = useChatRuntimeStore(
(s) => s.pendingSelection?.isDownloaded ?? false,
);
- const pendingHasContext = useChatRuntimeStore(
- (s) => s.pendingSelection?.contextLength != null,
+ // "Already probed" must key off layerCount / moeLayerCount, which only the
+ // full header probe fills (it sets all three together, so either is a
+ // reliable marker). contextLength alone can be list-seeded from
+ // /gguf-variants, which returns no layer/MoE counts -- treating it as
+ // complete would skip the probe and leave the GPU Layers slider at its 256
+ // fallback and the MoE slider hidden until the model loads.
+ const pendingHasMetadata = useChatRuntimeStore(
+ (s) =>
+ s.pendingSelection?.layerCount != null ||
+ s.pendingSelection?.moeLayerCount != null,
);
const setPendingSelection = useChatRuntimeStore((s) => s.setPendingSelection);
const onAutoLoadRef = useLatestRef(opts?.onAutoLoad);
@@ -69,25 +77,31 @@ export function useStagedModelPreparation(opts?: {
if (!current?.id || !isPendingGguf(current)) return;
const { id, ggufVariant, nativePathToken } = current;
try {
- const contextLength = await fetchGgufContextLength({
- model_path: id,
- gguf_variant: ggufVariant,
- hf_token: useChatRuntimeStore.getState().hfToken || null,
- nativePathToken,
- });
+ const { contextLength, layerCount, moeLayerCount } =
+ await fetchGgufStagedMetadata({
+ model_path: id,
+ gguf_variant: ggufVariant,
+ hf_token: useChatRuntimeStore.getState().hfToken || null,
+ nativePathToken,
+ });
// Apply only if the same model is still staged (the user may have switched
// picks or loaded/cancelled while the request was in flight).
const latest = useChatRuntimeStore.getState().pendingSelection;
if (
latest &&
- contextLength != null &&
- pendingSelectionMatches(latest, { id, ggufVariant, nativePathToken })
+ pendingSelectionMatches(latest, { id, ggufVariant, nativePathToken }) &&
+ (contextLength != null || layerCount != null || moeLayerCount != null)
) {
- setPendingSelection({ ...latest, contextLength });
+ setPendingSelection({
+ ...latest,
+ contextLength,
+ layerCount,
+ moeLayerCount,
+ });
}
} catch {
- // Leave contextLength null: the context slider stays hidden and the user
- // can still load (context fills in from the load response afterwards).
+ // Leave metadata null: the context/MoE sliders stay hidden and the user
+ // can still load (they fill in from the load response afterwards).
}
}, [setPendingSelection]);
@@ -125,7 +139,7 @@ export function useStagedModelPreparation(opts?: {
if (
!pendingId ||
(!pendingIsGguf && !pendingIsHubRepo) ||
- pendingHasContext
+ pendingHasMetadata
) {
return;
}
@@ -146,7 +160,7 @@ export function useStagedModelPreparation(opts?: {
pendingIsGguf,
pendingIsHubRepo,
pendingDownloaded,
- pendingHasContext,
+ pendingHasMetadata,
startDownloadRef,
fetchMetadataRef,
]);
diff --git a/studio/frontend/src/features/chat/lib/apply-inference-status-to-store.ts b/studio/frontend/src/features/chat/lib/apply-inference-status-to-store.ts
index a4b5f848e2..69bb38bbbe 100644
--- a/studio/frontend/src/features/chat/lib/apply-inference-status-to-store.ts
+++ b/studio/frontend/src/features/chat/lib/apply-inference-status-to-store.ts
@@ -2,13 +2,17 @@
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
import { getInferenceStatus } from "../api/chat-api";
-import { mergeBackendRecommendedInference } from "../presets/preset-policy";
+import {
+ mergeBackendRecommendedInference,
+ resolveManualAutoCtxPin,
+} from "../presets/preset-policy";
import { clampReasoningEffortToLevels } from "../provider-capabilities";
import {
CHAT_REASONING_ENABLED_KEY,
type ReasoningEffort,
type ReasoningStyle,
loadOptionalBool,
+ loadedGpuMemoryFields,
resolveToolsEnabledOnLoad,
useChatRuntimeStore,
} from "../stores/chat-runtime-store";
@@ -20,6 +24,10 @@ import type { ChatModelSummary } from "../types/runtime";
type LocalReasoningEffort = Extract;
+function sameArray(a: T[] | null, b: T[] | null): boolean {
+ return JSON.stringify(a) === JSON.stringify(b);
+}
+
// Canonicalises backend / persisted speculative mode values onto the UI modes.
export function normalizeSpeculativeType(
v: string | null | undefined,
@@ -119,6 +127,10 @@ function ensureActiveModelInStoreList(
export type ApplyInferenceStatusOptions = {
previousCheckpoint?: string;
+ /** activeGgufVariant BEFORE the caller's setCheckpoint synced it to the
+ * status -- without it a variant-only switch underneath the tab reads as
+ * steady state and the hydration reseed keeps the old quant's baselines. */
+ previousGgufVariant?: string | null;
};
/** Mirror refresh() hydration so adopted CLI models get reasoning/tools flags. */
@@ -144,9 +156,13 @@ export function applyActiveModelStatusToStore(
);
}
+ const previousGgufVariant =
+ options.previousGgufVariant !== undefined
+ ? options.previousGgufVariant
+ : store.activeGgufVariant;
const hydratingExistingModel =
previousCheckpoint !== checkpointId ||
- store.activeGgufVariant !== (status.gguf_variant ?? null);
+ previousGgufVariant !== (status.gguf_variant ?? null);
const supportsReasoning = status.supports_reasoning ?? false;
const reasoningAlwaysOn = status.reasoning_always_on ?? false;
const reasoningStyle = status.reasoning_style ?? "enable_thinking";
@@ -185,6 +201,66 @@ export function applyActiveModelStatusToStore(
// While a load is in flight, performLoad owns the load params. Seeding them
// from a stale poll here would clobber the values the load dialog just set.
const seedLoadParams = !prevState.modelLoading;
+ // A Manual + Auto-layers load sent its positive context pin as max_seq_length,
+ // and status only exposes the RESOLVED context; re-seed the pin from the
+ // requested value (parity with the load paths' keepCustomCtx). Baselines
+ // unconditionally: anything but an applicable pin is null, so a previous
+ // model's pin can't survive a model change underneath and reload at the old length.
+ const gpuPin = status.is_gguf
+ ? resolveManualAutoCtxPin(
+ status.gpu_memory_mode ?? "auto",
+ status.gpu_layers ?? -1,
+ status.requested_context_length ?? null,
+ )
+ : null;
+ const incomingGpuMode = status.is_gguf
+ ? (status.gpu_memory_mode ?? "auto")
+ : null;
+ const incomingGpuLayers =
+ incomingGpuMode === "manual" ? (status.gpu_layers ?? null) : null;
+ const incomingNCpuMoe =
+ incomingGpuMode === "manual" ? (status.n_cpu_moe ?? null) : null;
+ const incomingSplit =
+ incomingGpuMode === "manual" ? (status.tensor_split ?? null) : null;
+ const incomingGpuIds = status.is_gguf ? (status.gpu_ids ?? null) : null;
+ const gpuStatusChanged =
+ prevState.loadedGpuMemoryMode !== incomingGpuMode ||
+ prevState.loadedGpuLayers !== incomingGpuLayers ||
+ prevState.loadedNCpuMoe !== incomingNCpuMoe ||
+ !sameArray(prevState.loadedSplitRatio, incomingSplit) ||
+ !sameArray(prevState.loadedGpuIds, incomingGpuIds) ||
+ prevState.loadedCustomContextLength !== gpuPin;
+ const gpuMemoryEditsPending =
+ (prevState.loadedGpuMemoryMode !== null &&
+ prevState.gpuMemoryMode !== prevState.loadedGpuMemoryMode) ||
+ (prevState.loadedGpuMemoryMode === "manual" &&
+ (prevState.gpuLayers !== prevState.loadedGpuLayers ||
+ prevState.nCpuMoe !== prevState.loadedNCpuMoe ||
+ !sameArray(prevState.splitRatio, prevState.loadedSplitRatio))) ||
+ prevState.customContextLength !== prevState.loadedCustomContextLength;
+ const gpuIdsEditPending = !sameArray(
+ prevState.selectedGpuIds,
+ prevState.loadedGpuIds,
+ );
+ const incomingGpuFields = loadedGpuMemoryFields(status);
+ // A same-model reload from another client advances every loaded baseline.
+ // Preserve each editable group only when this tab has an unapplied change.
+ const preserveSameModelEdits = gpuStatusChanged && !hydratingExistingModel;
+ const gpuStatusFields = {
+ ...incomingGpuFields,
+ customContextLength: gpuPin,
+ loadedCustomContextLength: gpuPin,
+ ...(preserveSameModelEdits &&
+ gpuMemoryEditsPending && {
+ gpuMemoryMode: prevState.gpuMemoryMode,
+ gpuLayers: prevState.gpuLayers,
+ nCpuMoe: prevState.nCpuMoe,
+ splitRatio: prevState.splitRatio,
+ customContextLength: prevState.customContextLength,
+ }),
+ ...(preserveSameModelEdits &&
+ gpuIdsEditPending && { selectedGpuIds: prevState.selectedGpuIds }),
+ };
useChatRuntimeStore.setState({
supportsReasoning,
@@ -215,30 +291,51 @@ export function applyActiveModelStatusToStore(
loadedIsMultimodal: isMultimodalResponse(status),
loadedIsDiffusion: status.is_diffusion ?? false,
specFallbackReason: status.spec_fallback_reason ?? null,
+ // The spec / KV seeds share the GPU-fields reseed mechanism below: a
+ // non-GGUF status leaves their loaded baselines null, so the "unseeded"
+ // guard re-fires every refresh -- hold them too while a staged pick's
+ // settings are being edited, or the refresh resets the staged edit.
+ // hydratingExistingModel reopens every load-param seed: when the active
+ // model changed underneath this tab (auto-switch, another client), the
+ // old model's baselines are stale and must adopt the new status.
...(seedLoadParams &&
- prevState.loadedSpeculativeType === null && {
+ prevState.pendingSelection == null &&
+ (prevState.loadedSpeculativeType === null || hydratingExistingModel) && {
speculativeType: currentSpecType,
loadedSpeculativeType: currentSpecType,
}),
...(seedLoadParams &&
+ prevState.pendingSelection == null &&
status.spec_draft_n_max !== undefined &&
- prevState.loadedSpecDraftNMax === null &&
- prevState.specDraftNMax === null && {
+ (hydratingExistingModel ||
+ (prevState.loadedSpecDraftNMax === null &&
+ prevState.specDraftNMax === null)) && {
specDraftNMax: status.spec_draft_n_max ?? null,
loadedSpecDraftNMax: status.spec_draft_n_max ?? null,
}),
...(seedLoadParams &&
+ prevState.pendingSelection == null &&
status.cache_type_kv !== undefined &&
- prevState.loadedKvCacheDtype === null && {
+ (prevState.loadedKvCacheDtype === null || hydratingExistingModel) && {
kvCacheDtype: status.cache_type_kv,
loadedKvCacheDtype: status.cache_type_kv,
}),
...(seedLoadParams &&
+ prevState.pendingSelection == null &&
status.tensor_parallel !== undefined &&
- prevState.loadedTensorParallel === null && {
+ (prevState.loadedTensorParallel === null || hydratingExistingModel) && {
tensorParallel: status.tensor_parallel,
loadedTensorParallel: status.tensor_parallel,
}),
+ // Re-seed on first hydration, model/variant changes, or a same-model backend
+ // placement change. gpuStatusFields preserves dirty local edits in the last
+ // case while advancing their loaded baselines.
+ ...(seedLoadParams &&
+ prevState.pendingSelection == null &&
+ (prevState.loadedGpuMemoryMode === null ||
+ hydratingExistingModel ||
+ gpuStatusChanged) &&
+ gpuStatusFields),
...(status.chat_template_override !== undefined &&
prevState.loadedChatTemplateOverride === null &&
prevState.chatTemplateOverride === null && {
@@ -298,7 +395,11 @@ export async function tryAdoptServerActiveModel(): Promise {
if (previousCheckpoint) {
return true;
}
+ const previousGgufVariant = useChatRuntimeStore.getState().activeGgufVariant;
store.setCheckpoint(checkpointId, status.gguf_variant);
- applyActiveModelStatusToStore(status, { previousCheckpoint });
+ applyActiveModelStatusToStore(status, {
+ previousCheckpoint,
+ previousGgufVariant,
+ });
return true;
}
diff --git a/studio/frontend/src/features/chat/presets/preset-policy.ts b/studio/frontend/src/features/chat/presets/preset-policy.ts
index f96ee91f1b..23d79a35e1 100644
--- a/studio/frontend/src/features/chat/presets/preset-policy.ts
+++ b/studio/frontend/src/features/chat/presets/preset-policy.ts
@@ -339,3 +339,34 @@ export function resolveLoadMaxSeqLength({
}
return maxSeqLength;
}
+
+/**
+ * Adjust a resolved max-seq-length for the GPU Memory mode. Under Manual + Auto
+ * layers (GGUF, gpuLayers < 0) llama.cpp's --fit owns context sizing, so send 0
+ * (the backend omits -c) unless the user pinned a length; every other case keeps
+ * the resolved fallback. Shared by every GGUF load path so they can't drift.
+ */
+export function resolveFitMaxSeqLength(
+ isGguf: boolean | null | undefined,
+ gpuMemoryMode: "auto" | "manual",
+ gpuLayers: number,
+ customContextLength: number | null,
+ fallback: number,
+): number {
+ if (!isGguf || gpuMemoryMode !== "manual" || gpuLayers >= 0) return fallback;
+ return customContextLength && customContextLength > 0 ? customContextLength : 0;
+}
+
+// A Manual + Auto-layers load sends its positive context pin as max_seq_length;
+// keep it across a status reseed/Apply so the model isn't reverted to auto-fit
+// sizing. Anything else (Auto mode, pinned layers, no pin) baselines to null.
+// The caller keeps its own isGguf/targetIsGguf guard inline.
+export function resolveManualAutoCtxPin(
+ gpuMemoryMode: "auto" | "manual",
+ gpuLayers: number,
+ customContextLength: number | null,
+): number | null {
+ return gpuMemoryMode === "manual" && gpuLayers < 0 && (customContextLength ?? 0) > 0
+ ? customContextLength
+ : null;
+}
diff --git a/studio/frontend/src/features/chat/shared-composer.tsx b/studio/frontend/src/features/chat/shared-composer.tsx
index a0813fe27b..31e50ee60c 100644
--- a/studio/frontend/src/features/chat/shared-composer.tsx
+++ b/studio/frontend/src/features/chat/shared-composer.tsx
@@ -84,6 +84,8 @@ import {
useTransformersUpgradeDialogStore,
} from "@/features/transformers-upgrade";
import { loadModel, validateModel } from "./api/chat-api";
+import { resolveFitMaxSeqLength, resolveManualAutoCtxPin } from "./presets/preset-policy";
+import { ensureGpuDeviceCache } from "@/hooks/use-gpu-info";
import {
parseExternalModelId,
providerTypeSupportsVision,
@@ -95,8 +97,11 @@ import {
usePlusMenuPrefsStore,
} from "./stores/plus-menu-prefs-store";
import {
+ loadedGpuMemoryFieldsUnlessStaged,
type ReasoningEffort,
+ reconcilePersistedGpuIds,
resolveLoadedSpeculativeSettings,
+ persistGpuMemoryModeOnLoad,
resolveSpeculativeSettingsForLoad,
saveSpeculativeType,
useChatRuntimeStore,
@@ -1037,10 +1042,32 @@ export function SharedComposer({
return parts[parts.length - 1] || id;
}
+ // Warm the device cache before the snapshot below reconciles the GPU
+ // pick: on a cold cache the reconcile passes a stale pick through.
+ if (store.selectedGpuIds != null) {
+ await ensureGpuDeviceCache();
+ }
+ // The GPU/offload knobs both compare loads must use, snapshotted at Send.
+ // ensureModelLoaded runs sequentially and the first load's response echo
+ // (loadedGpuMemoryFields) rewrites the live store -- a non-GGUF or Auto
+ // first model resets gpuLayers/nCpuMoe/split/pick to defaults -- so
+ // reading the store per load would hand model 2 the first model's echoed
+ // defaults instead of the settings the user pressed Send with.
+ const compareLoadKnobs = {
+ gpuMemoryMode: store.gpuMemoryMode,
+ gpuLayers: store.gpuLayers,
+ nCpuMoe: store.nCpuMoe,
+ splitRatio: store.splitRatio,
+ // Reconcile the pick against the GPUs present now, like the model-switch
+ // path: an early remember-restore can hold a stale cross-host pick that
+ // /load would reject (the device cache is populated by send time).
+ selectedGpuIds: reconcilePersistedGpuIds(store.selectedGpuIds),
+ tensorParallel: store.tensorParallel,
+ customContextLength: store.customContextLength,
+ };
// Set when an accepted transformers install unloaded the active model
// server-side; a later failure must then clear the stale checkpoint.
let upgradeUnloadedActive = false;
-
// Helper: load a model and update store checkpoint
async function ensureModelLoaded(
sel: CompareModelSelection,
@@ -1057,15 +1084,35 @@ export function SharedComposer({
if (isAlreadyActive) {
return "ready";
}
+ const targetIsGguf =
+ sel.id.toLowerCase().endsWith(".gguf") || sel.ggufVariant != null;
+ // Size validation exactly as the load below, so the training-guard
+ // preflight checks the footprint that actually loads (under Manual + Auto
+ // layers the load sends 0 / the pinned context, not raw maxSeqLength).
+ const compareMaxSeqLength = resolveFitMaxSeqLength(
+ targetIsGguf,
+ compareLoadKnobs.gpuMemoryMode,
+ compareLoadKnobs.gpuLayers,
+ compareLoadKnobs.customContextLength,
+ maxSeqLength,
+ );
const validation = await validateModel({
model_path: sel.id,
hf_token: currentStore.hfToken || null,
- max_seq_length: maxSeqLength,
+ max_seq_length: compareMaxSeqLength,
load_in_4bit: true,
is_lora: sel.isLora,
gguf_variant: sel.ggufVariant ?? null,
trust_remote_code: loadTrustRemoteCode,
chat_template_override: effectiveChatTemplateOverride,
+ // Scope the validate to the picked GPUs. GGUF-only, like the load
+ // below: a non-GGUF target must not inherit a hidden GGUF GPU pick.
+ ...(targetIsGguf
+ ? {
+ gpu_ids: compareLoadKnobs.selectedGpuIds ?? undefined,
+ gpu_memory_mode: compareLoadKnobs.gpuMemoryMode,
+ }
+ : {}),
});
// Upgrade dialog first (mirrors the primary load path).
if (validation.requires_transformers_upgrade) {
@@ -1114,7 +1161,7 @@ export function SharedComposer({
const resp = await loadModel({
model_path: sel.id,
hf_token: useChatRuntimeStore.getState().hfToken || null,
- max_seq_length: maxSeqLength,
+ max_seq_length: compareMaxSeqLength,
load_in_4bit: true,
is_lora: sel.isLora,
gguf_variant: sel.ggufVariant ?? null,
@@ -1123,10 +1170,25 @@ export function SharedComposer({
chat_template_override: effectiveChatTemplateOverride,
speculative_type: specSettings.speculativeType,
spec_draft_n_max: specSettings.specDraftNMax,
- // Honor the Tensor Parallelism toggle on compare loads too.
- tensor_parallel: currentStore.tensorParallel,
+ // Honor the Tensor Parallelism + GPU Memory choices on compare loads.
+ // GGUF-only, like the auto-load path: the picker is a GGUF control,
+ // so a non-GGUF target loads via HF auto-placement instead of being
+ // pinned to a leftover GGUF pick it can't even show.
+ tensor_parallel: compareLoadKnobs.tensorParallel,
+ ...(targetIsGguf
+ ? {
+ gpu_memory_mode: compareLoadKnobs.gpuMemoryMode,
+ gpu_layers: compareLoadKnobs.gpuLayers,
+ n_cpu_moe: compareLoadKnobs.nCpuMoe,
+ tensor_split: compareLoadKnobs.splitRatio ?? undefined,
+ gpu_ids: compareLoadKnobs.selectedGpuIds ?? undefined,
+ }
+ : {}),
});
saveSpeculativeType(specSettings.speculativeType);
+ // Persist the GPU Memory mode on a non-diffusion GGUF compare-load too,
+ // so an applied manual choice survives a restart.
+ persistGpuMemoryModeOnLoad(resp, compareLoadKnobs.gpuMemoryMode);
upgradeUnloadedActive = false;
const store = useChatRuntimeStore.getState();
store.setCheckpoint(
@@ -1136,6 +1198,17 @@ export function SharedComposer({
store.setModelRequiresTrustRemoteCode(
resp.requires_trust_remote_code ?? false,
);
+ // Keep an explicit Manual+Auto context pin the load just applied (so a
+ // later Apply/Reset doesn't silently revert the model to auto-fit
+ // sizing), mirroring the interactive path's keepCustomCtx. Non-GGUF
+ // compare loads don't send the pin, so their baseline clears.
+ const keepCustomCtx = targetIsGguf
+ ? resolveManualAutoCtxPin(
+ compareLoadKnobs.gpuMemoryMode,
+ compareLoadKnobs.gpuLayers,
+ compareLoadKnobs.customContextLength,
+ )
+ : null;
useChatRuntimeStore.setState({
supportsReasoning: resp.supports_reasoning ?? false,
reasoningAlwaysOn: resp.reasoning_always_on ?? false,
@@ -1144,6 +1217,32 @@ export function SharedComposer({
supportsTools: resp.supports_tools ?? false,
tensorParallel: resp.tensor_parallel ?? false,
loadedTensorParallel: resp.tensor_parallel ?? false,
+ customContextLength: keepCustomCtx,
+ loadedCustomContextLength: keepCustomCtx,
+ // Seed the loaded GGUF context (interactive/auto-load parity): the
+ // settings sheet keys the GGUF GPU controls off it for a direct .gguf
+ // with no variant, and a later Apply reads it as the resolved context.
+ ...(targetIsGguf
+ ? {
+ ggufContextLength: resp.context_length ?? 131072,
+ ggufMaxContextLength:
+ resp.max_context_length ?? resp.context_length ?? 131072,
+ ggufNativeContextLength: resp.native_context_length ?? null,
+ }
+ : { ggufContextLength: null }),
+ // Compare loads resolve by id (HF repo / local path), never through a
+ // native-path lease, so a token left by a previously loaded native
+ // GGUF is stale here -- isLoadedGguf keys off it, and a stale token
+ // would dress a non-GGUF compare load in GGUF controls. Mirror the
+ // interactive path, which writes it on every load success.
+ activeNativePathToken: null,
+ // Held under an open staged pick: setCheckpoint preserves a stage on
+ // the empty->active transition, so a compare load can complete with
+ // staged GPU edits still on screen.
+ ...loadedGpuMemoryFieldsUnlessStaged(resp),
+ // Drives the GPU Memory controls' diffusion gate; set alongside the
+ // GPU fields on every load path so the gate can't read stale.
+ loadedIsDiffusion: resp.is_diffusion ?? false,
loadedIsMultimodal: isMultimodalResponse(resp),
...resolveLoadedSpeculativeSettings(resp),
});
diff --git a/studio/frontend/src/features/chat/stores/chat-runtime-store.ts b/studio/frontend/src/features/chat/stores/chat-runtime-store.ts
index 192ce1ec69..5786947118 100644
--- a/studio/frontend/src/features/chat/stores/chat-runtime-store.ts
+++ b/studio/frontend/src/features/chat/stores/chat-runtime-store.ts
@@ -7,6 +7,10 @@ import {
mirrorHfTokenInto,
useHfTokenStore,
} from "@/features/hub";
+import {
+ cachedPinnableGpuIndices,
+ ensureGpuDeviceCache,
+} from "@/hooks/use-gpu-info";
import { toast } from "@/lib/toast";
import { create } from "zustand";
import { isExternalModelId, parseExternalModelId } from "../external-providers";
@@ -74,6 +78,7 @@ export const CHAT_RAG_AUTOINJECT_MIN_SCORE_KEY =
export const CHAT_RAG_OCR_KEY = "unsloth_chat_rag_ocr_scanned";
export const CHAT_RAG_CAPTION_KEY = "unsloth_chat_rag_caption_figures";
export const CHAT_SPECULATIVE_TYPE_KEY = "unsloth_chat_speculative_type";
+export const CHAT_GPU_MEMORY_MODE_KEY = "unsloth_chat_gpu_memory_mode";
// Persist only the model-agnostic intents (auto/ngram/off). MTP modes
// (mtp/mtp+ngram) and spec_draft_n_max stay session-only: a persisted MTP
@@ -497,6 +502,213 @@ export function saveSpeculativeType(value: string | null): void {
}
}
+// GPU Memory strategy is a standing preference (like speculative type), not a
+// per-model setting: a "manual" choice persists across model switches and reloads.
+export function readPersistedGpuMemoryMode(): "auto" | "manual" {
+ return loadString(CHAT_GPU_MEMORY_MODE_KEY, "auto") === "manual" ? "manual" : "auto";
+}
+
+export function saveGpuMemoryMode(value: "auto" | "manual"): void {
+ saveString(CHAT_GPU_MEMORY_MODE_KEY, value);
+}
+
+/** Persist the GPU Memory mode after a load, but only for a non-diffusion GGUF:
+ * non-GGUF has no such mode, and diffusion runs mode-agnostic (reports "auto"),
+ * so neither must clobber the standing manual preference. */
+export function persistGpuMemoryModeOnLoad(
+ resp: { is_gguf?: boolean; is_diffusion?: boolean },
+ mode: "auto" | "manual",
+): void {
+ if (resp.is_gguf && !resp.is_diffusion) saveGpuMemoryMode(mode);
+}
+
+// Manual-mode gpu_layers sentinel: -1 = Auto (hand layer + context sizing to
+// llama.cpp's --fit). The Manual default; "all on GPU" is the slider's max.
+export const GPU_LAYERS_AUTO = -1;
+
+// Round real-valued shares to integers summing exactly to `total`, giving the
+// leftover units to the largest fractional parts (largest-remainder method).
+function largestRemainder(shares: number[], total: number): number[] {
+ const out = shares.map((x) => Math.floor(x));
+ let rem = total - out.reduce((a, b) => a + b, 0);
+ const byFrac = shares
+ .map((x, i) => ({ i, frac: x - Math.floor(x) }))
+ .sort((a, b) => b.frac - a.frac);
+ for (let k = 0; rem > 0 && k < byFrac.length; k++, rem--) out[byFrac[k].i] += 1;
+ return out;
+}
+
+// Spread `total` layers across GPUs in proportion to `weights` (e.g. per-GPU
+// VRAM), as integers summing exactly to `total`; even split for all-zero/empty
+// weights. Default per-GPU layer split before the user edits it (mirrors
+// llama.cpp's free-VRAM default).
+export function distributeByWeight(total: number, weights: number[]): number[] {
+ if (weights.length === 0) return [];
+ const t = Math.max(0, Math.floor(total));
+ const sum = weights.reduce((a, b) => a + b, 0);
+ const w = sum > 0 ? weights : weights.map(() => 1);
+ const wSum = w.reduce((a, b) => a + b, 0);
+ return largestRemainder(
+ w.map((x) => (t * x) / wSum),
+ t,
+ );
+}
+
+// Set GPU `index` to `value` and rebalance the rest so per-GPU counts still sum
+// to `total`; others absorb the remainder in proportion to their counts (evenly
+// if all zero). The --tensor-split editor: counts are sent verbatim, and
+// llama.cpp gives each GPU exactly its count when gpu_layers == sum(counts).
+export function rebalanceSplit(
+ total: number,
+ counts: number[],
+ index: number,
+ value: number,
+): number[] {
+ const v = Math.max(0, Math.min(value, total));
+ const out = counts.slice();
+ const otherIdx = counts.map((_, i) => i).filter((i) => i !== index);
+ // No other GPU to absorb the remainder: this one holds everything.
+ if (otherIdx.length === 0) {
+ out[index] = total;
+ return out;
+ }
+ out[index] = v;
+ const dist = distributeByWeight(
+ total - v,
+ otherIdx.map((i) => counts[i]),
+ );
+ otherIdx.forEach((i, k) => (out[i] = dist[k]));
+ return out;
+}
+
+// Validate a persisted gpu_ids pick against the GPUs present right now, before
+// restoring it from remembered settings. Returns null (= automatic) when the
+// pick is stale (none of the saved ids exist, or the host can't pin a multi-GPU
+// set), so a saved [1] on a now-1-GPU host doesn't get sent and rejected with no
+// way to clear it. A null pick (= automatic) passes through unchanged, and an
+// unpopulated device cache leaves the pick alone (the backend still guards).
+export function reconcilePersistedGpuIds(
+ ids: number[] | null,
+): number[] | null {
+ if (ids == null) return ids;
+ const pinnable = cachedPinnableGpuIndices();
+ if (pinnable === null) return ids; // cache not ready: can't validate, keep it
+ const kept = ids.filter((i) => pinnable.includes(i));
+ return kept.length > 0 ? kept : null;
+}
+
+// Store fields derived from a load/status response's GPU-memory settings.
+// Shared by every load path so the manual-knob round-trip can't drift.
+export function loadedGpuMemoryFields(resp: {
+ is_gguf?: boolean;
+ is_diffusion?: boolean;
+ gpu_memory_mode?: "auto" | "manual";
+ gpu_layers?: number;
+ n_cpu_moe?: number;
+ tensor_split?: number[] | null;
+ n_layers?: number | null;
+ n_moe_layers?: number;
+ gpu_ids?: number[] | null;
+}) {
+ // GPU-memory state is meaningful only for a GGUF chat load. A non-GGUF response
+ // still carries gpu_memory_mode (its default "auto" is serialized), so gate on
+ // the authoritative is_gguf flag, not the field's presence -- otherwise loading
+ // a transformers model would reset the standing manual preference.
+ if (!resp.is_gguf) {
+ // Clear the GPU pick / offload baseline a prior GGUF load may have left, so it
+ // reflects the non-GGUF model (no pin) -- else a stale loadedGpuIds reads as
+ // dirty (gpuIdsDirty is ungated) and Reset restores it while the picker is
+ // hidden. gpuMemoryMode (the standing preference) is kept, but its loaded
+ // baseline clears to null so Reset preserves the preference, not a stale mode.
+ return {
+ selectedGpuIds: null,
+ loadedGpuIds: null,
+ loadedGpuMemoryMode: null,
+ gpuLayers: GPU_LAYERS_AUTO,
+ loadedGpuLayers: null,
+ nCpuMoe: 0,
+ loadedNCpuMoe: null,
+ splitRatio: null,
+ loadedSplitRatio: null,
+ ggufLayerCount: null,
+ moeLayerCount: null,
+ };
+ }
+ const mode = resp.gpu_memory_mode ?? "auto";
+ const gpuIds = resp.gpu_ids ?? null;
+ // Layer/MoE/split knobs apply (and are reported) only in manual mode; in auto
+ // the server ignores them, so don't seed the loaded baseline or the editable
+ // knobs with values it never applied. In manual, the server reports gpu_layers
+ // = -1 under Auto, which round-trips the slider back to its Auto position.
+ const manualKnobs =
+ mode === "manual"
+ ? {
+ loadedGpuLayers: resp.gpu_layers ?? null,
+ loadedNCpuMoe: resp.n_cpu_moe ?? null,
+ loadedSplitRatio: resp.tensor_split ?? null,
+ gpuLayers: resp.gpu_layers ?? GPU_LAYERS_AUTO,
+ nCpuMoe: resp.n_cpu_moe ?? 0,
+ splitRatio: resp.tensor_split ?? null,
+ }
+ : {
+ loadedGpuLayers: null,
+ loadedNCpuMoe: null,
+ loadedSplitRatio: null,
+ // Auto ignores these, so reset the editable knobs too (not just the
+ // loaded baseline) -- else a later switch back to Manual would snapshot
+ // and send a previous model's stale gpuLayers/nCpuMoe/split that this
+ // load never applied. Mirrors the non-GGUF branch above.
+ gpuLayers: GPU_LAYERS_AUTO,
+ nCpuMoe: 0,
+ splitRatio: null,
+ };
+ return {
+ // A diffusion GGUF runs mode-agnostic (pins all layers on one GPU, reports
+ // "auto"), so adopt everything a chat GGUF does EXCEPT the live standing
+ // preference -- the next chat load must still honor the user's manual choice.
+ // The loaded baseline is still "auto", but the UI hides mode controls for a
+ // loaded diffusion model so it can't read as dirty against the preference.
+ ...(resp.is_diffusion ? {} : { gpuMemoryMode: mode }),
+ loadedGpuMemoryMode: mode,
+ ggufLayerCount: resp.n_layers ?? null,
+ // MoE expert-layer count: the n_cpu_moe slider max, and 0 hides the slider.
+ moeLayerCount: resp.n_moe_layers ?? null,
+ // The picker reflects what loaded (the request sent the user's pick).
+ selectedGpuIds: gpuIds,
+ loadedGpuIds: gpuIds,
+ ...manualKnobs,
+ };
+}
+
+/** loadedGpuMemoryFields (plus any seedExtras), unless a staged pick is open.
+ *
+ * With a staged pick open (the load fired mid-staging), preserve its editable
+ * GPU knobs and seedExtras, but still advance every loaded baseline. Otherwise
+ * cancelling the stage restores its edits onto the newly loaded model. The
+ * status reseed cannot repair that while pendingSelection holds it off.
+ */
+export function loadedGpuMemoryFieldsUnlessStaged(
+ resp: Parameters[0],
+ seedExtras?: T,
+) {
+ const fields = loadedGpuMemoryFields(resp);
+ if (useChatRuntimeStore.getState().pendingSelection != null) {
+ return {
+ loadedGpuMemoryMode: fields.loadedGpuMemoryMode,
+ loadedGpuLayers: fields.loadedGpuLayers,
+ loadedNCpuMoe: fields.loadedNCpuMoe,
+ loadedSplitRatio: fields.loadedSplitRatio,
+ loadedGpuIds: fields.loadedGpuIds,
+ // These are metadata ceilings for the model that actually loaded, not
+ // editable values from the open stage. Advance them with the baselines
+ // so abandoning the stage cannot expose the previous model's limits.
+ ggufLayerCount: fields.ggufLayerCount,
+ moeLayerCount: fields.moeLayerCount,
+ };
+ }
+ return { ...fields, ...seedExtras };
+}
+
/** A local model staged for a deferred load (see `pendingSelection`). Shape is
* a subset of the load hook's `SelectedModelInput`, structurally assignable. */
export type PendingModelSelection = {
@@ -515,6 +727,13 @@ export type PendingModelSelection = {
* Scoped here (not the shared `ggufContextLength`) so a staged model's
* metadata never pollutes the currently-loaded model's context display. */
contextLength?: number | null;
+ /** Total layer count (GGUF block_count); the manual gpu-layers ceiling is
+ * this + 1 (llama.cpp counts the output layer as offloadable too);
+ * scoped here like contextLength. */
+ layerCount?: number | null;
+ /** MoE expert-layer count from the GGUF header (manual --n-cpu-moe ceiling);
+ * 0 for dense models, scoped here like contextLength. */
+ moeLayerCount?: number | null;
/** "Load on selection" on + un-cached GGUF: download via the manager (global
* indicator) without opening the sheet, then load once the download finishes. */
autoLoad?: boolean;
@@ -743,6 +962,32 @@ type ChatRuntimeStore = {
tensorParallel: boolean;
/** Backend-reported tensor-parallel state; null until first hydrated. */
loadedTensorParallel: boolean | null;
+ /** GPU memory strategy for GGUF loads. "auto" = Unsloth picks GPUs and context
+ * to fit; "manual" = you own the offload (gpuLayers < 0 = Auto/--fit, >= 0
+ * pins layers + nCpuMoe). */
+ gpuMemoryMode: "auto" | "manual";
+ /** Backend-reported gpu memory mode; null until first hydrated. */
+ loadedGpuMemoryMode: "auto" | "manual" | null;
+ /** Manual mode: layers to offload to GPU. -1 = Auto (--fit); >= model layer
+ * count = all. */
+ gpuLayers: number;
+ loadedGpuLayers: number | null;
+ /** Manual mode: MoE expert layers to keep on CPU (--n-cpu-moe); 0 = none. */
+ nCpuMoe: number;
+ loadedNCpuMoe: number | null;
+ /** Manual mode: per-GPU layer counts (--tensor-split), in GPU-in-use order;
+ * null = unset (llama.cpp splits by free VRAM). */
+ splitRatio: number[] | null;
+ /** Backend-reported per-GPU split ratio (--tensor-split); null = unset. */
+ loadedSplitRatio: number[] | null;
+ /** Model layer count (GGUF block_count); the manual gpu-layers ceiling is
+ * this + 1 (the output layer is offloadable too). */
+ ggufLayerCount: number | null;
+ /** MoE expert-layer count: the nCpuMoe slider max; 0/null hides the slider. */
+ moeLayerCount: number | null;
+ /** Picked physical GPU indices (null = use all / automatic). */
+ selectedGpuIds: number[] | null;
+ loadedGpuIds: number[] | null;
/** Persisted: when false, picking a local model stages it as
* `pendingSelection` (and opens settings) instead of loading immediately,
* so load settings can be set before the single load. */
@@ -766,6 +1011,9 @@ type ChatRuntimeStore = {
* per step, cleared when the run ends, never persisted into the transcript. */
activeDiffusionCanvas: DiffusionCanvasFrame | null;
customContextLength: number | null;
+ /** The pinned context the loaded model used (null = Auto), so dirty-tracking
+ * and a later fit Apply can tell an explicit pin apart from Auto. */
+ loadedCustomContextLength: number | null;
defaultChatTemplate: string | null;
chatTemplateOverride: string | null;
loadedChatTemplateOverride: string | null;
@@ -884,6 +1132,11 @@ type ChatRuntimeStore = {
* which skip the sheet but must still honor a saved config. */
applyRememberedLoadSettings: (settings: RememberedLoadSettings) => void;
setTensorParallel: (value: boolean) => void;
+ setGpuMemoryMode: (mode: "auto" | "manual") => void;
+ setGpuLayers: (value: number) => void;
+ setNCpuMoe: (value: number) => void;
+ setSplitRatio: (value: number[] | null) => void;
+ setSelectedGpuIds: (ids: number[] | null) => void;
setLoadOnSelection: (value: boolean) => void;
setExpandQuantizations: (value: boolean) => void;
setShowAllQuantizations: (value: boolean) => void;
@@ -1101,11 +1354,12 @@ function setScalarSettingVersion(
/** The "revert to the loaded model" baseline for the editable load knobs.
* Shared by resetModelSettingsToLoaded (full revert) and stageModel (which
- * overrides speculative to start a fresh pick from the standing default). */
+ * overrides speculative and the per-model GPU knobs to start a fresh pick). */
function loadedBaselineSettings(s: ChatRuntimeStore) {
const hasLoadedModel = Boolean(s.params.checkpoint);
return {
- customContextLength: null,
+ // Revert to the loaded model's pin (null = Auto), not a blanket Auto.
+ customContextLength: s.loadedCustomContextLength,
kvCacheDtype: s.loadedKvCacheDtype,
tensorParallel: s.loadedTensorParallel ?? false,
speculativeType: hasLoadedModel
@@ -1113,6 +1367,20 @@ function loadedBaselineSettings(s: ChatRuntimeStore) {
: readPersistedSpeculativeType(),
specDraftNMax: hasLoadedModel ? s.loadedSpecDraftNMax : null,
chatTemplateOverride: s.loadedChatTemplateOverride,
+ // GPU memory mode is a standing preference; revert to the loaded model's
+ // mode (or the persisted default when nothing is loaded). Manual knobs and
+ // the GPU pick are per-model and revert to their loaded baseline. A loaded
+ // model with no applicable mode -- diffusion ("auto" baseline) or non-GGUF
+ // (null baseline) -- keeps the live preference so Reset can't drop it.
+ gpuMemoryMode: !hasLoadedModel
+ ? readPersistedGpuMemoryMode()
+ : s.loadedIsDiffusion
+ ? s.gpuMemoryMode
+ : (s.loadedGpuMemoryMode ?? s.gpuMemoryMode),
+ gpuLayers: s.loadedGpuLayers ?? GPU_LAYERS_AUTO,
+ nCpuMoe: s.loadedNCpuMoe ?? 0,
+ splitRatio: s.loadedSplitRatio ?? null,
+ selectedGpuIds: s.loadedGpuIds,
};
}
@@ -1213,6 +1481,18 @@ export const useChatRuntimeStore = create((set, get) => ({
loadedSpecDraftNMax: null,
tensorParallel: false,
loadedTensorParallel: null,
+ gpuMemoryMode: readPersistedGpuMemoryMode(),
+ loadedGpuMemoryMode: null,
+ gpuLayers: GPU_LAYERS_AUTO,
+ loadedGpuLayers: null,
+ nCpuMoe: 0,
+ loadedNCpuMoe: null,
+ splitRatio: null,
+ loadedSplitRatio: null,
+ ggufLayerCount: null,
+ moeLayerCount: null,
+ selectedGpuIds: null,
+ loadedGpuIds: null,
loadOnSelection: loadBool(CHAT_LOAD_ON_SELECTION_KEY, true),
expandQuantizations: loadBool(CHAT_EXPAND_QUANTIZATIONS_KEY, false),
showAllQuantizations: loadBool(CHAT_SHOW_ALL_QUANTIZATIONS_KEY, true),
@@ -1221,6 +1501,7 @@ export const useChatRuntimeStore = create((set, get) => ({
loadedIsMultimodal: false,
loadedIsDiffusion: false,
customContextLength: null,
+ loadedCustomContextLength: null,
defaultChatTemplate: null,
chatTemplateOverride: null,
loadedChatTemplateOverride: null,
@@ -1455,9 +1736,23 @@ export const useChatRuntimeStore = create((set, get) => ({
loadedSpecDraftNMax: null,
tensorParallel: false,
loadedTensorParallel: null,
+ // Standing preference: survives unload, unlike the per-model knobs above.
+ gpuMemoryMode: readPersistedGpuMemoryMode(),
+ loadedGpuMemoryMode: null,
+ gpuLayers: GPU_LAYERS_AUTO,
+ loadedGpuLayers: null,
+ nCpuMoe: 0,
+ loadedNCpuMoe: null,
+ splitRatio: null,
+ loadedSplitRatio: null,
+ ggufLayerCount: null,
+ moeLayerCount: null,
+ selectedGpuIds: null,
+ loadedGpuIds: null,
loadedIsMultimodal: false,
loadedIsDiffusion: false,
customContextLength: null,
+ loadedCustomContextLength: null,
defaultChatTemplate: null,
chatTemplateOverride: null,
loadedChatTemplateOverride: null,
@@ -1753,17 +2048,67 @@ export const useChatRuntimeStore = create((set, get) => ({
setSpeculativeType: (speculativeType) => set({ speculativeType }),
setSpecDraftNMax: (specDraftNMax) => set({ specDraftNMax }),
setTensorParallel: (tensorParallel) => set({ tensorParallel }),
+ // Standing preference, but persisted only on a successful load (see
+ // use-chat-model-runtime), not on selection -- so an unapplied pick the user
+ // resets/abandons doesn't stick to the next session.
+ setGpuMemoryMode: (gpuMemoryMode) => set({ gpuMemoryMode }),
+ setGpuLayers: (gpuLayers) => set({ gpuLayers }),
+ setNCpuMoe: (nCpuMoe) => set({ nCpuMoe }),
+ setSplitRatio: (splitRatio) => set({ splitRatio }),
+ setSelectedGpuIds: (selectedGpuIds) => set({ selectedGpuIds }),
resetModelSettingsToLoaded: () => set((s) => loadedBaselineSettings(s)),
- applyRememberedLoadSettings: (settings) =>
+ applyRememberedLoadSettings: (settings) => {
+ const gpuCacheWasCold = cachedPinnableGpuIndices() === null;
+ const restoredGpuIds =
+ settings.selectedGpuIds !== undefined
+ ? reconcilePersistedGpuIds(settings.selectedGpuIds)
+ : undefined;
// Coalesce every field: a blob persisted by an older/newer build can omit
// keys, and a raw spread would push `undefined` into fields typed non-null.
+ // The GPU knobs are spread only when present, but first reset the per-model
+ // ones to defaults: this path (load-on-selection) starts from the loaded
+ // model's baseline and skips the model-switch reset, so a blob omitting
+ // gpuLayers/nCpuMoe/selectedGpuIds (older build) or splitRatio (never
+ // remembered) must not inherit the previous model's placement. gpuMemoryMode
+ // (standing preference) is NOT reset, only applied when the blob carries it;
+ // selectedGpuIds keeps a meaningful null (all GPUs), so it keys off undefined.
set({
+ gpuLayers: GPU_LAYERS_AUTO,
+ nCpuMoe: 0,
+ splitRatio: null,
+ selectedGpuIds: null,
customContextLength: settings.contextLength ?? null,
kvCacheDtype: settings.kvCacheDtype ?? null,
speculativeType: settings.speculativeType ?? "auto",
specDraftNMax: settings.specDraftNMax ?? null,
tensorParallel: settings.tensorParallel ?? false,
- }),
+ ...(settings.gpuMemoryMode != null && {
+ gpuMemoryMode: settings.gpuMemoryMode,
+ }),
+ ...(settings.gpuLayers != null && { gpuLayers: settings.gpuLayers }),
+ ...(settings.nCpuMoe != null && { nCpuMoe: settings.nCpuMoe }),
+ ...(restoredGpuIds !== undefined && {
+ // Reconcile against the GPUs present now (see reconcilePersistedGpuIds):
+ // a saved [1] on a 1-GPU host (or under relative/UUID visibility) would
+ // hide the picker yet still send gpu_ids, which the backend rejects.
+ selectedGpuIds: restoredGpuIds,
+ }),
+ });
+ // A cold cache makes the synchronous restore provisional. Reconcile again
+ // when the shared fetch completes, but only if this exact restored array is
+ // still current so a user edit, stage change, or load cannot be overwritten.
+ if (gpuCacheWasCold && restoredGpuIds != null) {
+ void ensureGpuDeviceCache().then(() => {
+ set((state) => {
+ if (state.selectedGpuIds !== restoredGpuIds) return state;
+ const reconciled = reconcilePersistedGpuIds(restoredGpuIds);
+ return reconciled === restoredGpuIds
+ ? state
+ : { selectedGpuIds: reconciled };
+ });
+ });
+ }
+ },
setLoadOnSelection: (loadOnSelection) => {
saveBool(CHAT_LOAD_ON_SELECTION_KEY, loadOnSelection);
set({ loadOnSelection });
@@ -1798,6 +2143,22 @@ export const useChatRuntimeStore = create((set, get) => ({
// Load's keepSpeculative) a forced MTP mode onto a model that may lack it.
speculativeType: readPersistedSpeculativeType(),
specDraftNMax: null,
+ // Keep the on-screen GPU Memory selection (loadedBaselineSettings would
+ // otherwise revert it to the loaded model's mode, dropping a Manual choice
+ // just made). Use the live store value, not the persisted one, which can
+ // lag a mode hydrated from an out-of-band load.
+ gpuMemoryMode: s.gpuMemoryMode,
+ // Per-model GPU knobs start from defaults too so a fresh pick doesn't
+ // inherit the loaded model's layer/MoE/split/GPU choices, matching the
+ // immediate-switch reset.
+ gpuLayers: GPU_LAYERS_AUTO,
+ nCpuMoe: 0,
+ splitRatio: null,
+ selectedGpuIds: null,
+ // Fresh pick starts at Auto context (loadedBaselineSettings would
+ // otherwise restore the current model's pin). Leaves the baseline
+ // intact, like the GPU knobs, so abandoning restores the loaded pin.
+ customContextLength: null,
};
});
},
diff --git a/studio/frontend/src/features/chat/types/api.ts b/studio/frontend/src/features/chat/types/api.ts
index d72c406fdd..c24ddde5f5 100644
--- a/studio/frontend/src/features/chat/types/api.ts
+++ b/studio/frontend/src/features/chat/types/api.ts
@@ -65,6 +65,18 @@ export interface LoadModelRequest {
* of by layer for GGUF models. Multi-GPU only; no effect on a single GPU.
*/
tensor_parallel?: boolean | null;
+ /** GPU memory strategy for GGUF models. "auto" (default): Unsloth selects GPUs
+ * and caps context to fit VRAM. "manual": you own the offload -- gpu_layers
+ * -1 (Auto) hands sizing to llama.cpp's --fit, >= 0 pins layers/n_cpu_moe. */
+ gpu_memory_mode?: "auto" | "manual";
+ /** Manual mode: layers to offload to GPU (--gpu-layers, --fit off); -1 = Auto (--fit). */
+ gpu_layers?: number;
+ /** Manual mode: MoE expert layers to keep on CPU (--n-cpu-moe); 0 = none. */
+ n_cpu_moe?: number;
+ /** Manual mode: relative model share per GPU (--tensor-split), in GPU order. */
+ tensor_split?: number[] | null;
+ /** Picked physical GPU indices (omit/empty = automatic). */
+ gpu_ids?: number[];
}
export interface ValidateModelResponse {
@@ -80,6 +92,13 @@ export interface ValidateModelResponse {
requires_security_review?: boolean;
/** Native context length from the local GGUF header; null until downloaded. */
context_length?: number | null;
+ /** Total layer count (GGUF block_count); the manual gpu-layers ceiling is
+ * this + 1 (llama.cpp counts the output layer as offloadable too); null
+ * until downloaded. */
+ layer_count?: number | null;
+ /** MoE expert-layer count from the GGUF header (manual --n-cpu-moe ceiling);
+ * 0 for dense models, null until downloaded. */
+ moe_layer_count?: number | null;
/** Architecture only shipped by a newer transformers; UI pauses on the upgrade dialog. */
requires_transformers_upgrade?: boolean;
/** Set only when requires_transformers_upgrade. */
@@ -159,6 +178,14 @@ export interface LoadModelResponse {
spec_draft_n_max?: number | null;
/** Whether tensor-parallel split (--split-mode tensor) is active. */
tensor_parallel?: boolean;
+ gpu_memory_mode?: "auto" | "manual";
+ gpu_layers?: number;
+ n_cpu_moe?: number;
+ tensor_split?: number[] | null;
+ n_layers?: number | null;
+ /** Model's MoE expert-layer count (the n_cpu_moe ceiling); 0 if not MoE. */
+ n_moe_layers?: number;
+ gpu_ids?: number[] | null;
}
export interface UnloadModelRequest {
@@ -203,6 +230,17 @@ export interface InferenceStatusResponse {
spec_draft_n_max?: number | null;
/** Whether tensor-parallel split (--split-mode tensor) is active. */
tensor_parallel?: boolean;
+ gpu_memory_mode?: "auto" | "manual";
+ gpu_layers?: number;
+ n_cpu_moe?: number;
+ tensor_split?: number[] | null;
+ /** n_ctx the active GGUF load was invoked with (0 = Auto); re-seeds a
+ * Manual + Auto-layers context pin on hydration. Null for non-GGUF. */
+ requested_context_length?: number | null;
+ gpu_ids?: number[] | null;
+ n_layers?: number | null;
+ /** Model's MoE expert-layer count (the n_cpu_moe ceiling); 0 if not MoE. */
+ n_moe_layers?: number;
/**
* Why MTP was disabled on the loaded model despite being requested.
* "binary_no_mtp" / "binary_outdated" -> updating llama.cpp would re-enable
diff --git a/studio/frontend/src/hooks/use-gpu-info.ts b/studio/frontend/src/hooks/use-gpu-info.ts
index 1e313acdf3..db2cc021be 100644
--- a/studio/frontend/src/hooks/use-gpu-info.ts
+++ b/studio/frontend/src/hooks/use-gpu-info.ts
@@ -15,6 +15,19 @@ export interface GpuInfo {
systemRamTotalGb: number
}
+export interface SystemGpuDevice {
+ index: number;
+ name: string;
+ memoryTotalGb: number;
+ /** Free VRAM at fetch time. Degrades to the total when the utilization
+ * probe had no usage data; 0 only when the total is unknown too. */
+ memoryFreeGb: number;
+ /** "physical" = `index` is a stable physical/PCI id safe to pin via gpu_ids;
+ * "relative" = an ordinal into a parent CUDA_VISIBLE_DEVICES mask, which the
+ * backend can't map back, so the picker must not offer it. */
+ physicalIndex: boolean;
+}
+
const DEFAULT_GPU: GpuInfo = {
available: false,
name: "Unknown",
@@ -25,70 +38,135 @@ const DEFAULT_GPU: GpuInfo = {
systemRamTotalGb: 0
};
-// Module-level cache so multiple components share one fetch.
-let cachedGpu: GpuInfo | null = null;
-let fetchPromise: Promise | null = null;
+// One module-level cache so every GPU hook shares a single /api/system fetch.
+let cachedSystem: SystemInfoResponse | null = null;
+let systemPromise: Promise | null = null;
-async function fetchGpuOnce(): Promise {
- if (cachedGpu) return cachedGpu;
- if (fetchPromise) return fetchPromise;
-
- fetchPromise = (async () => {
+async function fetchSystemOnce(): Promise {
+ if (cachedSystem) return cachedSystem;
+ if (systemPromise) return systemPromise;
+ systemPromise = (async () => {
try {
const res = await authFetch("/api/system");
if (!res.ok) throw new Error(`HTTP ${res.status}`);
-
- const data = await res.json() as SystemInfoResponse;
- const gpuData = data?.gpu;
-
- // CPU/RAM exist even on hosts without a GPU, so populate them on every path.
- // No discrete GPU (e.g. Mac): still surface system RAM so memory math
- // (unified memory) has a budget to work with.
- const base = {
- cpuCore: data?.cpu?.physical_count ?? 0,
- cpuThread: data?.cpu?.logical_count ?? 0,
- systemRamAvailableGb: data?.memory?.available_gb ?? 0,
- systemRamTotalGb: data?.memory?.total_gb ?? 0,
- };
-
- const devices = gpuData?.devices ?? [];
- const info: GpuInfo =
- gpuData?.available && devices.length
- ? {
- ...base,
- available: true,
- name: devices[0]?.name ?? "Unknown",
- memoryTotalGb: devices.reduce((sum, d) => sum + (d.memory_total_gb ?? 0), 0),
- }
- : { ...DEFAULT_GPU, ...base };
- cachedGpu = info;
- return info;
+ cachedSystem = (await res.json()) as SystemInfoResponse;
+ return cachedSystem;
} catch {
- // Reset promise so subsequent calls retry (e.g. backend wasn't ready)
- fetchPromise = null;
- return DEFAULT_GPU;
+ systemPromise = null; // reset so a later call retries (backend not ready)
+ return null;
}
})();
+ return systemPromise;
+}
- return fetchPromise;
+function toGpuInfo(data: SystemInfoResponse | null): GpuInfo {
+ // CPU/RAM exist even on GPU-less hosts (e.g. Mac), so populate them on every
+ // path: unified-memory math still needs a RAM budget to work with.
+ const base = {
+ cpuCore: data?.cpu?.physical_count ?? 0,
+ cpuThread: data?.cpu?.logical_count ?? 0,
+ systemRamAvailableGb: data?.memory?.available_gb ?? 0,
+ systemRamTotalGb: data?.memory?.total_gb ?? 0,
+ };
+ const gpuData = data?.gpu;
+ const devices = gpuData?.devices ?? [];
+ if (!gpuData?.available || !devices.length) {
+ return { ...DEFAULT_GPU, ...base };
+ }
+ return {
+ ...base,
+ available: true,
+ name: devices[0]?.name ?? "Unknown",
+ memoryTotalGb: devices.reduce((sum, d) => sum + (d.memory_total_gb ?? 0), 0),
+ };
+}
+
+function toGpuDevices(data: SystemInfoResponse | null): SystemGpuDevice[] {
+ // Unpinnable configurations must hide every pick surface: XPU indices are
+ // torch-xpu ordinals no applicator speaks, and Vulkan-only builds pin ggml's
+ // own ordinals -- /load and /validate 400 picks on both, so the backend
+ // reports gpu.gguf_gpu_ids_supported and every gate keyed on physicalIndex
+ // (picker, persisted-pick reconcile) follows it. The device flavor lives on
+ // the TOP-LEVEL device_backend field; absent support info defaults to
+ // pinnable (older backend).
+ const pinnableBackend =
+ data?.device_backend !== "xpu" &&
+ data?.gpu?.gguf_gpu_ids_supported !== false;
+ return (data?.gpu?.devices ?? [])
+ .filter((d) => typeof d.index === "number")
+ .map((d) => ({
+ index: d.index as number,
+ name: d.name ?? `GPU ${d.index}`,
+ memoryTotalGb: d.memory_total_gb ?? 0,
+ memoryFreeGb: d.vram_free_gb ?? 0,
+ physicalIndex: pinnableBackend && d.index_kind === "physical",
+ }));
+}
+
+/** Aggregate GPU info from /api/system; shares one module-level fetch across all GPU hooks. */
+export function useGpuInfo(): GpuInfo {
+ const [gpu, setGpu] = useState(
+ cachedSystem ? toGpuInfo(cachedSystem) : DEFAULT_GPU,
+ );
+ useEffect(() => {
+ // No early return on cachedSystem: a consumer mounting as the cache fills
+ // (between render and effect) would otherwise stay stuck at the default.
+ let cancelled = false;
+ fetchSystemOnce().then((d) => {
+ if (!cancelled) setGpu(toGpuInfo(d));
+ });
+ return () => {
+ cancelled = true;
+ };
+ }, []);
+ return gpu;
+}
+
+/** All backend-visible GPUs (index, name, total VRAM); shares the same fetch. */
+export function useGpuDevices(): SystemGpuDevice[] {
+ const [devices, setDevices] = useState(
+ cachedSystem ? toGpuDevices(cachedSystem) : [],
+ );
+ useEffect(() => {
+ // No early return on cachedSystem: a consumer mounting as the cache fills
+ // (between render and effect) would otherwise stay stuck at the default.
+ let cancelled = false;
+ fetchSystemOnce().then((d) => {
+ if (!cancelled) setDevices(toGpuDevices(d));
+ });
+ return () => {
+ cancelled = true;
+ };
+ }, []);
+ return devices;
}
/**
- * Fetch GPU info from /api/system. Cached at module level, so only one request
- * is made no matter how many components call this hook.
+ * Await the shared /api/system fetch so cachedPinnableGpuIndices (and the
+ * store's reconcilePersistedGpuIds) can validate a persisted pick before a
+ * load path sends it -- on a cold cache the reconcile passes ids through
+ * unvalidated, and a stale cross-host pick then fails /load with the picker
+ * hidden. Resolves immediately once the module cache is warm; a failed fetch
+ * keeps the cache cold, preserving the "can't validate, backend guards"
+ * degradation.
*/
-export function useGpuInfo(): GpuInfo {
- const [gpu, setGpu] = useState(cachedGpu ?? DEFAULT_GPU);
+export async function ensureGpuDeviceCache(): Promise {
+ await fetchSystemOnce();
+}
- useEffect(() => {
- if (cachedGpu) return;
-
- let cancelled = false;
- fetchGpuOnce().then((info) => {
- if (!cancelled) setGpu(info);
- });
- return () => { cancelled = true; };
- }, []);
-
- return gpu;
-}
\ No newline at end of file
+/**
+ * Pinnable physical GPU indices from the already-fetched /api/system cache, for
+ * non-React code (the store) that needs to validate a persisted `gpu_ids` pick
+ * without triggering a fetch. Returns:
+ * - `null` when the cache isn't populated yet (caller can't validate, so keep
+ * the pick and let the backend guard reject a truly bad one);
+ * - `[]` when the host has no pinnable multi-GPU set (single GPU, or relative/
+ * UUID-masked indices) -- the picker is hidden, so any saved pick is stale;
+ * - the physical indices otherwise.
+ */
+export function cachedPinnableGpuIndices(): number[] | null {
+ if (!cachedSystem) return null;
+ const physical = toGpuDevices(cachedSystem).filter((d) => d.physicalIndex);
+ // Mirrors the sheet's showGpuPicker gate: only a 2+ physical-GPU host can pin.
+ return physical.length > 1 ? physical.map((d) => d.index) : [];
+}
diff --git a/studio/frontend/src/hooks/use-system.ts b/studio/frontend/src/hooks/use-system.ts
index a135cce86e..8cfe2bace4 100644
--- a/studio/frontend/src/hooks/use-system.ts
+++ b/studio/frontend/src/hooks/use-system.ts
@@ -40,6 +40,9 @@ export interface SystemInfoResponse {
gpu: {
available: boolean;
backend?: string;
+ /** Whether GGUF loads accept an explicit gpu_ids pick (false on XPU hosts
+ * and Vulkan-only builds, where /load and /validate 400 picks). */
+ gguf_gpu_ids_supported?: boolean;
backend_cuda_visible_devices?: string | null;
parent_visible_gpu_ids?: number[];
index_kind?: string;
From 03590f696e97401361d59d61e1b9b367238ea229 Mon Sep 17 00:00:00 2001
From: Daniel Han
Date: Sun, 19 Jul 2026 06:08:54 -0700
Subject: [PATCH 02/41] Give opencode real timeout headroom in Local Agent
Guides CI (#7235)
* Raise the opencode invoke timeout in Local Agent Guides CI
The connection (opencode) cell flakes with a 600s timeout reported as guide drift, but it is not a hang: in a passing run the same opencode run finishes in ~482s (08:12:31 to 08:20:33), right against the shared AGENT_INVOKE_TIMEOUT of 600s, so about one run in six drifts past the cap.
opencode is the slow outlier. The print-mode agents (claude -p, codex exec) run one turn against a minimal injected system prompt, while opencode run runs its own full turn with opencode's large system prompt plus a separate small_model call to name the session (start.py pins small_model to the same 4B the server hosts). On a CPU-served gemma-4-E4B that is about 8 minutes, leaving no margin under 600s.
Double opencode's per-invoke timeout in agent-guides-drive.sh and keep the tight 600s cap for the fast agents, so a genuine headless-TTY hang still fails quickly. 1200s stays well under the 40-minute job budget.
* Normalize the agent invoke timeout before doubling it for opencode
Strip an optional trailing 's' from AGENT_INVOKE_TIMEOUT so the opencode
arithmetic, and the "${TIMEOUT}s" timeout message, stay valid if a
timeout(1)-style suffix is ever configured.
* Only double the opencode timeout for a bare-integer seconds value
Guard the arithmetic so a GNU timeout(1) duration suffix (s/m/h/d, including
floats like 0.5s) is passed through unchanged instead of breaking the
expansion; timeout(1) parses those directly. Bare seconds still double.
---------
Co-authored-by: danielhanchen
---
.github/scripts/agent-guides-drive.sh | 17 +++++++++++++++++
1 file changed, 17 insertions(+)
diff --git a/.github/scripts/agent-guides-drive.sh b/.github/scripts/agent-guides-drive.sh
index 2457f08407..b63ac94b93 100755
--- a/.github/scripts/agent-guides-drive.sh
+++ b/.github/scripts/agent-guides-drive.sh
@@ -36,6 +36,23 @@ AGENT="${2:?usage: agent-guides-drive.sh }"
# Determinism (seed/temp) is applied at the server level by
# serve-unsloth-run.sh --extra; agents inherit it through the API.
TIMEOUT="${AGENT_INVOKE_TIMEOUT:-180}"
+# opencode is the slow outlier. Unlike the print-mode agents (claude -p, codex
+# exec) it runs a full turn AND a separate small_model call to name the session,
+# so one connection reply takes ~8 min on a CPU-served 4B -- right at the shared
+# 600s cap, so the cell flaked when a run drifted past a ~480s success. Give it
+# headroom (still well under the 40-min job budget); the fast agents keep the
+# tight cap that still catches a real headless-TTY hang.
+case "$AGENT" in
+ opencode)
+ # Double it, but only for a bare-integer seconds value. A GNU timeout(1)
+ # duration suffix (s/m/h/d, including floats like 0.5s) is left unchanged so
+ # the arithmetic never sees a non-number; timeout(1) parses it directly.
+ case "$TIMEOUT" in
+ *[!0-9]*) ;;
+ *) TIMEOUT=$(( TIMEOUT * 2 )) ;;
+ esac
+ ;;
+esac
# Claude refuses --dangerously-skip-permissions outside a sandbox; the CI runner
# IS the sandbox, so declare it (mirrors unslothai/scripts launcher.sh). Harmless
From a9be36830eb4ec731dcd10008d2f6dc3bb102d40 Mon Sep 17 00:00:00 2001
From: Daniel Han
Date: Sun, 19 Jul 2026 06:19:29 -0700
Subject: [PATCH 03/41] Installer: allow torch 2.11.x on the CUDA install path
(fresh install + studio) (#6959)
* Studio: allow torch 2.11.x on the CUDA install path
The CUDA torch repair path (_ensure_cuda_torch) installs torch/torchvision/
torchaudio from an exclusive --index-url, so _CUDA_TORCH_PKG_SPEC decides
exactly which torch the Studio venv gets. It was capped at torch<2.11.0, so on
a cu128/cu130 host the venv resolved torch 2.10.x even though the CUDA indexes
now publish torch 2.11.0. That left the Studio venv a torch minor behind the
torch 2.11.0 Docker base image, so the CUDA dedup step would relink base libs
under a mismatched torch.
Raise the upper bound to <2.12.0 (torchvision <0.27.0, torchaudio <2.12.0) so
the CUDA install path lands on torch 2.11.x, matching the rocm7.2 spec and the
base image. The torchao selector already maps torch 2.11 -> torchao 0.17.0, and
_ensure_flash_attn degrades gracefully when no prebuilt wheel matches (Blackwell
skips it outright; non-Blackwell prints a warning and continues), so no other
pin needs to move.
Add test_cuda_torch_spec.py to lock the bound (torch 2.11.x in, 2.12.x out) and
assert the CUDA and rocm7.2 upper bounds stay in lockstep.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* test: use zip(strict=True) so a spec length mismatch fails loudly
* install.sh: widen the CUDA torch ceiling to <2.12.0 so a fresh install matches the base
Raising _CUDA_TORCH_PKG_SPEC alone was not enough: that spec only feeds
_ensure_cuda_torch(), the ROCm-poisoning repair path that early-returns on a
normal NVIDIA host. A fresh CUDA install (including the studio Docker build,
which runs `bash install.sh --local`) takes its torch from install.sh's
TORCH_CONSTRAINT, which was still capped at torch>=2.4,<2.11.0, so cu12x/cu13x
resolved torch 2.10.x and the venv landed a minor behind the torch 2.11.0 base
image.
Extend the existing `case "$TORCH_INDEX_URL"` block (which already relaxes
rocm7.2) with a `*/cu[0-9]*` branch that widens the ceiling to <2.12.0, keeping
the >=2.4 floor so an older CUDA index (e.g. cu118) that tops out below 2.11
still resolves. The CPU wheel and older ROCm tags stay on <2.11.0 (the glob
does not match /cpu). torchvision/torchaudio are bare on this install line and
resolve their compatible companions via wheel metadata, matching the rocm7.2
pattern.
Add behavioral tests (Python + shell) exercising the case block: cu118/124/126/
128/130 widen to <2.12.0, rocm7.2 stays 2.11.x, and /cpu plus older ROCm keep
the default <2.11.0.
* install.sh: key the CUDA torch widening off the index leaf, not the full URL
The `*/cu[0-9]*` glob matched a `cu` segment anywhere in TORCH_INDEX_URL,
so a custom UNSLOTH_PYTORCH_MIRROR whose base path contains e.g. cu128 but whose
final leaf is cpu or an older ROCm tag would still widen TORCH_CONSTRAINT to
<2.12.0, contradicting the block's own comment and letting a CPU / older-ROCm
mirror resolve torch 2.11.x. Match on _torch_index_leaf (the final path segment
the backend classification just above already computes) so only a real cu*/
rocm7.2 leaf is affected; cpu and older ROCm keep the default <2.11.0. Update
the Python + shell tests to mirror the leaf-anchored case and add regression
cases for a mirror base that contains cu128 but resolves to a cpu / rocm7.1 leaf.
* install: freeze the torch trio during the with-deps unsloth installs
Released unsloth wheels can pin an older torch than Step 1 installed
(unsloth 2026.7.2 declares torch<2.11.0), so the with-deps resolve from
PyPI silently downgrades the pinned +cuXXX torch trio to PyPI's default
wheel. The flavor guard cannot catch every such swap: PyPI's torch 2.10
default is itself cu128-flavored, so the cuXXX tag comparison still
matches while the version silently drops. Freeze the just-installed trio
with uv --overrides (overrides replace dependency requirements during
resolution), keeping torch 2.11.0+cuXXX in place while unsloth's other
dependencies resolve normally. Verified on the cu128 path: without the
override torch drops 2.11.0+cu128 -> 2.10.0; with it the trio survives
and unsloth 2026.7.2 + unsloth-zoo install cleanly.
* install: fold UV_OVERRIDE env files into the torch-trio overrides file
The CLI --overrides flag is the command-line form of UV_OVERRIDE, so
passing it replaced any overrides file already exported for the process;
macOS arm64 exports UV_OVERRIDE=overrides-darwin-arm64.txt for the same
generic install path and would have lost those pins. Concatenate any
UV_OVERRIDE files into the temp trio file so both keep applying.
* install: extend the torch-trio overrides guard to migrated installs
Four follow-ups to the Step-2 --overrides guard, all empirically verified:
1. The migrated-environment with-deps unsloth install resolved
unsloth>=2026.7.2 (which pins torch<2.11.0) without the overrides file,
so a migrated CUDA venv on torch 2.11 was silently downgraded -- the
exact bug this branch fixes on the fresh path. The overrides build is
now a function (_build_unsloth_torch_overrides, reading the trio
installed at call time) invoked by both with-deps paths; the migrated
no-torch path installs --no-deps and stays unguarded.
2. The overrides temp file is now cleaned by the EXIT trap (same pattern
as _UV_OVERRIDE_TMPDIR, pre-initialized empty so an inherited value can
never reach the trap's rm); previously any Step-2 failure leaked it.
3. Folding UV_OVERRIDE files used cat, which joins the last requirement of
a file lacking a trailing newline onto the next file's first requirement
(reproduced: idna==3.10certifi==2025.1.31 makes uv fail parsing).
4. Inherited torch/torchvision/torchaudio override lines are now filtered
out when folding: uv intersects duplicate overrides rather than
last-wins (verified on uv 0.10.12: direct conflict is unsatisfiable,
transitive conflict silently backtracks), so a conflicting inherited
trio pin would break the resolve the generated exact pins protect.
Both 3 and 4 are handled by a single newline-terminating awk filter
that preserves non-trio overrides (torchmetrics, torchao, ...).
test_unsloth_torch_override.sh extended: migrated-path coverage, trap
assertion, and a functional fold test (14 checks).
* installer: tighten comments
* install: keep the existing torch release when re-running the installer
Re-running `curl -fsSL https://unsloth.ai/install.sh | sh` over an existing
install rebuilds the venv for clean state, which silently moved users to the
newest torch in range (2.10 -> 2.11 once the constraint widened). A torch the
user already validated must survive an unsloth update.
Before the old venv is moved aside for rollback, its torch version is probed
(last stdout line only, so sitecustomize noise cannot corrupt it). After the
index leaf is chosen, _previous_torch_pin turns that version into a
torch==X.Y.Z pin, but only when it cannot do harm:
- cu*/cpu leaves only; rocm leaves keep their floors (rocm7.2 must land 2.11
for the Strix _grouped_mm fix) and the Radeon wheel-matching path is
untouched.
- The wheel's flavor tag must match the freshly chosen leaf, so a flavor
change (cpu -> cuda, cu126 -> cu130) still installs the correct new build.
- The base must look like a release, so probe noise never becomes a pin.
- UNSLOTH_TORCH_UPGRADE=1 opts out and restores the old always-newest
behavior; the substep line advertises it.
The supported range is kept in _PREV_FALLBACK_CONSTRAINT: if the exact
release is not resolvable from the chosen index (custom mirrors prune old
wheels), the install warns and falls back to the newest supported release
instead of failing the whole run. The later flavor-mismatch repair reuses
TORCH_CONSTRAINT, so a mid-install clobber is repaired back to the kept
release rather than the newest one.
Verified end to end: a venv seeded with torch 2.10.0+cu130 re-run through the
full installer finishes with torch 2.10.0+cu130 (previously 2.11.0+cu130).
Tests: tests/sh/test_previous_torch_pin.sh covers keep/flavor-change/rocm/
noise/opt-out plus wiring (probe ordering before venv replacement, fallback
present, SKIP_TORCH gate).
* install: constrain kept torch pins to the supported window
Review caught that _previous_torch_pin pinned the previous venv's torch on
flavor match alone, so a release outside the installer's active range (a
2.3.x manual install below the >=2.4 floor, or a 2.12.x manual upgrade above
the ceiling) replaced the bounds computed just above it and a rerun kept a
torch the installer otherwise deliberately excludes.
New _torch_release_in_window checks the probed base against the active
TORCH_CONSTRAINT ("torch>=A.B[,
---
install.sh | 168 +++++++++++++++++-
studio/backend/tests/test_cuda_torch_spec.py | 73 ++++++++
.../test_tokenizers_and_torch_constraint.py | 80 +++++++++
tests/sh/test_previous_torch_pin.sh | 101 +++++++++++
tests/sh/test_torch_constraint.sh | 14 ++
tests/sh/test_unsloth_torch_override.sh | 131 ++++++++++++++
6 files changed, 558 insertions(+), 9 deletions(-)
create mode 100644 studio/backend/tests/test_cuda_torch_spec.py
create mode 100644 tests/sh/test_previous_torch_pin.sh
create mode 100644 tests/sh/test_unsloth_torch_override.sh
diff --git a/install.sh b/install.sh
index 5972379d26..6076721540 100755
--- a/install.sh
+++ b/install.sh
@@ -472,11 +472,13 @@ _on_install_exit() {
_restore_studio_venv_replacement
fi
[ -n "${_UV_OVERRIDE_TMPDIR:-}" ] && rm -rf "$_UV_OVERRIDE_TMPDIR" 2>/dev/null || true
+ [ -n "${_UNSLOTH_TORCH_OVERRIDES:-}" ] && rm -f "$_UNSLOTH_TORCH_OVERRIDES" 2>/dev/null || true
exit "$_status"
}
-# Empty so an inherited value can never reach the trap's rm; only a temp dir
-# this script creates below (Apple Silicon, spaced path) is ever removed.
+# Empty so an inherited value never reaches the trap's rm; only temp paths this
+# script creates below (spaced-path dir, torch-trio overrides) are removed.
_UV_OVERRIDE_TMPDIR=""
+_UNSLOTH_TORCH_OVERRIDES=""
trap _on_install_exit EXIT
# ── Helper: download a URL to a file (supports curl and wget) ──
@@ -1821,6 +1823,8 @@ tauri_log "STEP" "Creating virtual environment"
mkdir -p "$STUDIO_HOME"
_MIGRATED=false
+# Empty so an inherited value can never masquerade as a probed torch version.
+_PREV_TORCH_VER=""
if [ -x "$VENV_DIR/bin/python" ]; then
# why: matching guard to the .venv branch below -- in env-mode
@@ -1838,6 +1842,12 @@ if [ -x "$VENV_DIR/bin/python" ]; then
echo " Move it aside or choose an empty UNSLOTH_STUDIO_HOME." >&2
exit 1
fi
+ # Record the existing venv's torch BEFORE the replacement moves it aside: a re-run
+ # rebuilds the venv for clean state, but must keep the torch release the user
+ # already has (see _previous_torch_pin below). Last line only: sitecustomize or
+ # import-hook noise on stdout must not corrupt the version.
+ _PREV_TORCH_VER=$("$VENV_DIR/bin/python" -c \
+ "import torch; print(torch.__version__)" 2>/dev/null | tail -n 1 || true)
# New layout already exists — replace only after preserving rollback copy.
substep "preserving existing environment for rollback..."
_start_studio_venv_replacement "$VENV_DIR"
@@ -2187,6 +2197,68 @@ _torch_flavor_tag() {
esac
}
+# Whether release base $1 (X.Y[.Z...]) falls inside constraint window $2
+# ("torch>=A.B[.C],="*",<"*) ;;
+ *) echo "no"; return ;;
+ esac
+ _trw_floor="${_trw_con#torch>=}"; _trw_floor="${_trw_floor%%,*}"
+ _trw_ceil="${_trw_con##*,<}"
+ _v_maj="${1%%.*}"; _v_rest="${1#*.}"; _v_min="${_v_rest%%.*}"
+ _f_maj="${_trw_floor%%.*}"; _f_rest="${_trw_floor#*.}"; _f_min="${_f_rest%%.*}"
+ _c_maj="${_trw_ceil%%.*}"; _c_rest="${_trw_ceil#*.}"; _c_min="${_c_rest%%.*}"
+ for _trw_n in "$_v_maj" "$_v_min" "$_f_maj" "$_f_min" "$_c_maj" "$_c_min"; do
+ case "$_trw_n" in ''|*[!0-9]*) echo "no"; return ;; esac
+ done
+ if [ "$_v_maj" -gt "$_f_maj" ] || { [ "$_v_maj" -eq "$_f_maj" ] && [ "$_v_min" -ge "$_f_min" ]; }; then
+ if [ "$_v_maj" -lt "$_c_maj" ] || { [ "$_v_maj" -eq "$_c_maj" ] && [ "$_v_min" -lt "$_c_min" ]; }; then
+ echo "yes"
+ return
+ fi
+ fi
+ echo "no"
+}
+
+# Whether a re-run should keep the previous venv's torch: echo "torch==X.Y.Z" when the
+# probed previous version ($1) has a flavor tag matching the freshly chosen cu*/cpu index
+# leaf ($2) AND sits inside the active constraint window ($3), else "". Re-running
+# `curl | sh` rebuilds the venv for clean state, but a healthy torch the user already
+# validated must not be silently moved to a newer release (2.10 -> 2.11); a flavor
+# change (cpu <-> cuda, cu126 -> cu130) still installs the correct new build, rocm
+# leaves keep their floors (rocm7.2 must land 2.11 for the Strix _grouped_mm fix), and
+# a release outside the window (2.3.x manual install, 2.12.x manual upgrade) is never
+# kept: the installer's own bounds win. Opt out with UNSLOTH_TORCH_UPGRADE=1 to get
+# the newest release.
+_previous_torch_pin() {
+ _ptp_ver="$1"
+ _ptp_leaf="$2"
+ _ptp_con="$3"
+ [ -n "$_ptp_ver" ] || { echo ""; return; }
+ [ "${UNSLOTH_TORCH_UPGRADE:-0}" = "1" ] && { echo ""; return; }
+ case "$_ptp_leaf" in
+ cu[0-9]*|cpu) ;;
+ *) echo ""; return ;;
+ esac
+ _ptp_base="${_ptp_ver%%+*}"
+ # The base must look like a release (probe noise / garbage must never become a pin).
+ case "$_ptp_base" in
+ [0-9]*.[0-9]*) ;;
+ *) echo ""; return ;;
+ esac
+ [ "$(_torch_release_in_window "$_ptp_base" "$_ptp_con")" = "yes" ] || { echo ""; return; }
+ if [ "$(_torch_flavor_tag "$_ptp_ver")" = "$_ptp_leaf" ]; then
+ echo "torch==$_ptp_base"
+ else
+ echo ""
+ fi
+}
+
# Expected tag from the index leaf ($1): cuXXX / cpu / rocm (rocmX.Y and gfx* ->
# rocm). Empty on an unknown leaf (odd mirror) so the repair safely no-ops.
_expected_torch_flavor_tag() {
@@ -2478,12 +2550,32 @@ case "$_torch_index_leaf" in
*) export UNSLOTH_TORCH_BACKEND="cuda" ;;
esac
-# 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" ;;
+# rocm7.2 and the CUDA cu12x/cu13x indexes now ship torch 2.11.x, so widen the
+# ceiling to <2.12.0 (matches the base image and _CUDA_TORCH_PKG_SPEC in
+# studio/install_python_stack.py). Keep the >=2.4 floor so an older CUDA index
+# (e.g. cu118) still resolves. Match on _torch_index_leaf, not the full URL, so
+# a mirror whose base path contains cu*/rocm7.2 but resolves to a cpu/older-rocm
+# leaf keeps the default <2.11.0.
+case "$_torch_index_leaf" in
+ rocm7.2) TORCH_CONSTRAINT="torch>=2.11.0,<2.12.0" ;;
+ cu[0-9]*) TORCH_CONSTRAINT="torch>=2.4,<2.12.0" ;;
esac
+# Re-run over an existing install: keep the previous venv's torch release instead of
+# resolving the newest in range. The range stays in _PREV_FALLBACK_CONSTRAINT so the
+# install can fall back when the exact release is not on the chosen index (custom
+# mirrors may prune old wheels). Skipped for --no-torch (no previous probe runs).
+_PREV_TORCH_PIN=""
+_PREV_FALLBACK_CONSTRAINT="$TORCH_CONSTRAINT"
+if [ "$SKIP_TORCH" = false ]; then
+ _prev_pin=$(_previous_torch_pin "$_PREV_TORCH_VER" "$_torch_index_leaf" "$TORCH_CONSTRAINT")
+ if [ -n "$_prev_pin" ]; then
+ _PREV_TORCH_PIN="$_prev_pin"
+ TORCH_CONSTRAINT="$_prev_pin"
+ substep "existing install has torch $_PREV_TORCH_VER -- keeping it (set UNSLOTH_TORCH_UPGRADE=1 to get the newest release)"
+ fi
+fi
+
# 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".
@@ -2705,6 +2797,43 @@ esac
# ── Install unsloth directly into the venv (no activation needed) ──
tauri_log "STEP" "Installing PyTorch"
_VENV_PY="$VENV_DIR/bin/python"
+
+# A released unsloth wheel can pin an older torch (unsloth 2026.7.2 declares
+# torch<2.11.0); a with-deps PyPI resolve then downgrades the whole trio,
+# swapping the pinned +cuXXX/+rocm build for PyPI's default. The flavor guard
+# below misses this (PyPI's torch 2.10 default is itself cu128-flavored), so
+# freeze the trio via uv --overrides (overrides replace dependency requirements
+# during resolution) while unsloth's other deps resolve normally. Sets
+# _UNSLOTH_TORCH_OVERRIDES from the trio in the venv; every with-deps unsloth
+# install (migrated and fresh) must call this before resolving and rm it after.
+_build_unsloth_torch_overrides() {
+ _UNSLOTH_TORCH_OVERRIDES=""
+ [ "$SKIP_TORCH" = false ] || return 0
+ _torch_trio_pins=$("$_VENV_PY" -c "
+from importlib.metadata import version, PackageNotFoundError
+for _p in ('torch', 'torchvision', 'torchaudio'):
+ try:
+ print(_p + '==' + version(_p))
+ except PackageNotFoundError:
+ pass
+" 2>/dev/null) || _torch_trio_pins=""
+ case "$_torch_trio_pins" in
+ torch==*)
+ _UNSLOTH_TORCH_OVERRIDES=$(mktemp)
+ printf '%s\n' "$_torch_trio_pins" > "$_UNSLOTH_TORCH_OVERRIDES"
+ # The CLI --overrides flag replaces any UV_OVERRIDE env file (same
+ # uv setting; macOS arm64 exports one here), so fold its pins in.
+ # awk, not cat: it drops inherited torch-trio lines (uv intersects
+ # duplicate overrides, so a conflicting pin would make resolution
+ # unsatisfiable) and newline-terminates the last line so an
+ # unterminated file cannot join two requirements into one.
+ for _ov_file in ${UV_OVERRIDE:-}; do
+ [ -f "$_ov_file" ] && awk '!/^[[:space:]]*torch(vision|audio)?([[:space:]<>=!~;@[]|$)/' "$_ov_file" >> "$_UNSLOTH_TORCH_OVERRIDES"
+ done
+ ;;
+ esac
+}
+
if [ "$_MIGRATED" = true ]; then
# Migrated env: force-reinstall unsloth+unsloth-zoo to ensure clean state
# in the new venv location, while preserving existing torch/CUDA
@@ -2729,9 +2858,13 @@ if [ "$_MIGRATED" = true ]; then
else
# Pin mlx-lm away from 0.31.3 here too: a curl-piped migration has no
# overrides file, so UV_OVERRIDE is unset and this positional is the only cover.
+ _build_unsloth_torch_overrides
run_install_cmd_retry "install unsloth (migrated)" uv pip install --python "$_VENV_PY" \
+ ${_UNSLOTH_TORCH_OVERRIDES:+--overrides "$_UNSLOTH_TORCH_OVERRIDES"} \
--reinstall-package unsloth --reinstall-package unsloth-zoo \
"unsloth>=2026.7.3" "unsloth-zoo>=2026.7.3" ${_MLX_LM_EXCLUDE_ARG:-}
+ [ -n "$_UNSLOTH_TORCH_OVERRIDES" ] && rm -f "$_UNSLOTH_TORCH_OVERRIDES"
+ _UNSLOTH_TORCH_OVERRIDES=""
fi
if [ "$STUDIO_LOCAL_INSTALL" = true ]; then
substep "overlaying local repo (editable)..."
@@ -2913,8 +3046,20 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
fi
else
substep "installing PyTorch ($TORCH_INDEX_URL)..."
- run_install_cmd_retry "install PyTorch" uv pip install --python "$_VENV_PY" "$TORCH_CONSTRAINT" torchvision torchaudio \
- --default-index "$TORCH_INDEX_URL"
+ if [ -n "$_PREV_TORCH_PIN" ]; then
+ # Kept previous release: fall back to the supported range if the exact
+ # release is not resolvable from the chosen index (pruned mirror).
+ if ! run_install_cmd_retry "install PyTorch (kept release)" uv pip install --python "$_VENV_PY" "$TORCH_CONSTRAINT" torchvision torchaudio \
+ --default-index "$TORCH_INDEX_URL"; then
+ substep "[WARN] $_PREV_TORCH_PIN is not installable from $TORCH_INDEX_URL -- installing the newest supported release instead" "$C_WARN"
+ TORCH_CONSTRAINT="$_PREV_FALLBACK_CONSTRAINT"
+ run_install_cmd_retry "install PyTorch" uv pip install --python "$_VENV_PY" "$TORCH_CONSTRAINT" torchvision torchaudio \
+ --default-index "$TORCH_INDEX_URL"
+ fi
+ else
+ run_install_cmd_retry "install PyTorch" uv pip install --python "$_VENV_PY" "$TORCH_CONSTRAINT" torchvision torchaudio \
+ --default-index "$TORCH_INDEX_URL"
+ fi
fi
# AMD ROCm: install bitsandbytes (once, after torch, for all ROCm paths).
# Gate on SKIP_TORCH=false so a user running with --no-torch on a ROCm
@@ -2927,9 +3072,10 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
;;
esac
fi
- # Fresh: Step 2 - install unsloth, preserving pre-installed torch
+ # Fresh: Step 2 - install unsloth, preserving the torch Step 1 installed
tauri_log "STEP" "Installing Unsloth"
substep "installing unsloth (this may take a few minutes)..."
+ _build_unsloth_torch_overrides
if [ "$SKIP_TORCH" = true ]; then
# No-torch: install unsloth + unsloth-zoo with --no-deps, then
# runtime deps (typer, safetensors, transformers, etc.) with --no-deps.
@@ -2953,6 +3099,7 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
fi
elif [ "$STUDIO_LOCAL_INSTALL" = true ]; then
run_install_cmd_retry "install unsloth (local)" uv pip install --python "$_VENV_PY" \
+ ${_UNSLOTH_TORCH_OVERRIDES:+--overrides "$_UNSLOTH_TORCH_OVERRIDES"} \
--upgrade-package unsloth "unsloth>=2026.7.3" "unsloth-zoo>=2026.7.3"
substep "overlaying local repo (editable)..."
run_install_cmd "overlay local repo" uv pip install --python "$_VENV_PY" -e "$_REPO_ROOT" --no-deps
@@ -2962,8 +3109,11 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
"unsloth-zoo @ git+https://github.com/unslothai/unsloth-zoo"
else
run_install_cmd_retry "install unsloth" uv pip install --python "$_VENV_PY" \
+ ${_UNSLOTH_TORCH_OVERRIDES:+--overrides "$_UNSLOTH_TORCH_OVERRIDES"} \
--upgrade-package unsloth -- "$PACKAGE_NAME" ${_MLX_LM_EXCLUDE_ARG:-}
fi
+ [ -n "$_UNSLOTH_TORCH_OVERRIDES" ] && rm -f "$_UNSLOTH_TORCH_OVERRIDES"
+ _UNSLOTH_TORCH_OVERRIDES=""
# AMD ROCm: repair torch if the unsloth/unsloth-zoo install pulled in
# CUDA torch from PyPI, overwriting the ROCm wheels installed in Step 1.
if [ "$SKIP_TORCH" = false ]; then
diff --git a/studio/backend/tests/test_cuda_torch_spec.py b/studio/backend/tests/test_cuda_torch_spec.py
new file mode 100644
index 0000000000..928cef787e
--- /dev/null
+++ b/studio/backend/tests/test_cuda_torch_spec.py
@@ -0,0 +1,73 @@
+# 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 _CUDA_TORCH_PKG_SPEC in install_python_stack.py.
+
+The CUDA repair path installs the torch trio from an exclusive --index-url (no
+PyPI fallback), so these pinned ranges decide which torch the venv gets. The
+upper bound is locked to the 2.11.x family to match the base image and rocm7.2
+spec and to keep the companions off a torch-2.12 wheel that would ABI-mismatch.
+"""
+
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+
+import pytest
+from packaging.requirements import Requirement
+
+# install_python_stack.py lives at repo_root/studio/install_python_stack.py
+_INSTALL_SCRIPT = Path(__file__).resolve().parents[2] / "install_python_stack.py"
+
+
+def _load_module(monkeypatch):
+ """(Re-)import and return install_python_stack (mirrors test_torchao_select)."""
+ sys.modules.pop("install_python_stack", None)
+ monkeypatch.syspath_prepend(str(_INSTALL_SCRIPT.parent))
+ import install_python_stack
+
+ return install_python_stack
+
+
+def _spec_of(pkg_spec: str):
+ """Parse 'torch>=2.4,<2.12.0' into a packaging SpecifierSet."""
+ return Requirement(pkg_spec).specifier
+
+
+@pytest.mark.parametrize(
+ "index, allowed, rejected",
+ [
+ # torch: 2.11.x allowed (matches base image); 2.12.x excluded.
+ (0, ["2.11.0", "2.11.2", "2.10.0", "2.4.0"], ["2.12.0", "2.3.0", "1.13.1"]),
+ # torchvision: 0.26.x (torch 2.11 companion) allowed; 0.27.x (torch 2.12) out.
+ (1, ["0.26.0", "0.26.1", "0.19.0"], ["0.27.0", "0.18.0"]),
+ # torchaudio: same 2.11.x window as torch.
+ (2, ["2.11.0", "2.10.0", "2.4.0"], ["2.12.0", "2.3.0"]),
+ ],
+)
+def test_cuda_spec_bounds(monkeypatch, index, allowed, rejected):
+ mod = _load_module(monkeypatch)
+ spec = _spec_of(mod._CUDA_TORCH_PKG_SPEC[index])
+ for v in allowed:
+ assert spec.contains(v, prereleases = True), f"{v} should satisfy {spec}"
+ for v in rejected:
+ assert not spec.contains(v, prereleases = True), f"{v} should not satisfy {spec}"
+
+
+def test_cuda_spec_matches_rocm72_upper_bound(monkeypatch):
+ """CUDA and rocm7.2 target the same torch 2.11.x family, so their upper
+ bounds must stay in lockstep (bump both together at 2.12.x)."""
+ mod = _load_module(monkeypatch)
+ rocm72 = mod._ROCM_TORCH_PKG_SPECS["rocm7.2"]
+
+ def _upper(pkg_spec: str) -> str:
+ for clause in _spec_of(pkg_spec):
+ if clause.operator == "<":
+ return clause.version
+ raise AssertionError(f"no upper bound in {pkg_spec!r}")
+
+ for cuda_pkg, rocm_pkg in zip(mod._CUDA_TORCH_PKG_SPEC, rocm72, strict = True):
+ assert _upper(cuda_pkg) == _upper(
+ rocm_pkg
+ ), f"CUDA {cuda_pkg!r} upper bound must match rocm7.2 {rocm_pkg!r}"
diff --git a/tests/python/test_tokenizers_and_torch_constraint.py b/tests/python/test_tokenizers_and_torch_constraint.py
index 4322f0c7d6..c58808689b 100644
--- a/tests/python/test_tokenizers_and_torch_constraint.py
+++ b/tests/python/test_tokenizers_and_torch_constraint.py
@@ -69,6 +69,21 @@ class TestStructuralTorchConstraint:
def test_tightened_assignment_exists(self):
assert 'TORCH_CONSTRAINT="torch>=2.6,<2.11.0"' in self._sh
+ def test_cuda_constraint_widened_to_2_12(self):
+ """A fresh CUDA install widens the ceiling to <2.12.0 so cu12x/cu13x
+ land torch 2.11.x (matches the base image and _CUDA_TORCH_PKG_SPEC);
+ without it cu128/cu130 resolves torch 2.10.x."""
+ assert 'TORCH_CONSTRAINT="torch>=2.4,<2.12.0"' in self._sh
+
+ def test_cuda_case_widens_via_index_leaf(self):
+ """The cu* branch of the _torch_index_leaf case sets the widened
+ constraint (parallel to rocm7.2), anchored on the leaf."""
+ m = re.search(
+ r'cu\[0-9\]\*\)\s*TORCH_CONSTRAINT="torch>=2\.4,<2\.12\.0"',
+ self._sh,
+ )
+ assert m is not None, "CUDA (cu*) TORCH_CONSTRAINT widening case not found"
+
def test_variable_used_in_pip_install(self):
"""$TORCH_CONSTRAINT must appear in a uv pip install line."""
assert '"$TORCH_CONSTRAINT"' in self._sh
@@ -384,6 +399,71 @@ class TestTorchConstraintShell:
logged = log_file.read_text()
assert "torch>=2.4,<2.11.0" in logged, f"uv log: {logged}"
+ # Mirrors the _torch_index_leaf case in install.sh: rocm7.2 -> 2.11.x floor,
+ # CUDA -> widened <2.12.0 ceiling, else (CPU/older ROCm) -> default. Anchored
+ # on the final path segment, so a mirror base path containing cu*/rocm7.2 but
+ # ending in a cpu/older-rocm leaf keeps the default.
+ _INDEX_SNIPPET = textwrap.dedent(r"""
+ #!/bin/bash
+ set -e
+ TORCH_INDEX_URL="{index_url}"
+ TORCH_CONSTRAINT="torch>=2.4,<2.11.0"
+ _torch_index_leaf="${TORCH_INDEX_URL%/}"
+ _torch_index_leaf="${_torch_index_leaf##*/}"
+ case "$_torch_index_leaf" in
+ rocm7.2) TORCH_CONSTRAINT="torch>=2.11.0,<2.12.0" ;;
+ cu[0-9]*) TORCH_CONSTRAINT="torch>=2.4,<2.12.0" ;;
+ esac
+ echo "$TORCH_CONSTRAINT"
+ """).strip()
+
+ def _resolve_index(self, tmp_path: pathlib.Path, index_url: str) -> str:
+ script_file = tmp_path / "index_snippet.sh"
+ script_file.write_text(self._INDEX_SNIPPET.replace("{index_url}", index_url))
+ script_file.chmod(0o755)
+ result = subprocess.run(
+ ["bash", str(script_file)],
+ capture_output = True,
+ text = True,
+ timeout = 10,
+ )
+ assert result.returncode == 0, f"Script failed: {result.stderr}"
+ return result.stdout.strip()
+
+ @pytest.mark.parametrize("leaf", ["cu118", "cu124", "cu126", "cu128", "cu130"])
+ def test_cuda_index_widens_to_2_12(self, tmp_path, leaf):
+ url = f"https://download.pytorch.org/whl/{leaf}"
+ assert self._resolve_index(tmp_path, url) == "torch>=2.4,<2.12.0"
+
+ def test_rocm72_index_uses_211_floor(self, tmp_path):
+ url = "https://download.pytorch.org/whl/rocm7.2"
+ assert self._resolve_index(tmp_path, url) == "torch>=2.11.0,<2.12.0"
+
+ def test_cpu_index_keeps_default(self, tmp_path):
+ # /cpu must NOT match the */cu[0-9]* branch.
+ url = "https://download.pytorch.org/whl/cpu"
+ assert self._resolve_index(tmp_path, url) == "torch>=2.4,<2.11.0"
+
+ def test_older_rocm_index_keeps_default(self, tmp_path):
+ url = "https://download.pytorch.org/whl/rocm7.1"
+ assert self._resolve_index(tmp_path, url) == "torch>=2.4,<2.11.0"
+
+ def test_cuda_index_custom_mirror_widens(self, tmp_path):
+ url = "https://internal.example.com/pytorch/cu128"
+ assert self._resolve_index(tmp_path, url) == "torch>=2.4,<2.12.0"
+
+ @pytest.mark.parametrize(
+ "url",
+ [
+ "https://internal.example.com/pytorch/cu128/cpu",
+ "https://internal.example.com/cu128/whl/rocm7.1",
+ ],
+ )
+ def test_cuda_in_mirror_path_but_noncuda_leaf_keeps_default(self, tmp_path, url):
+ # A cu128 in the mirror base path must not widen when the leaf is cpu /
+ # older ROCm: the case anchors on _torch_index_leaf, not the whole URL.
+ assert self._resolve_index(tmp_path, url) == "torch>=2.4,<2.11.0"
+
# Group 3 -- E2E tokenizers fix (requires network, ~2-5 min)
@pytest.mark.e2e
diff --git a/tests/sh/test_previous_torch_pin.sh b/tests/sh/test_previous_torch_pin.sh
new file mode 100644
index 0000000000..253ede8a27
--- /dev/null
+++ b/tests/sh/test_previous_torch_pin.sh
@@ -0,0 +1,101 @@
+#!/bin/bash
+# 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 install.sh's _previous_torch_pin, which keeps the previous
+# venv's torch release on a re-run (curl | sh over an existing install) instead
+# of silently moving the user to a newer release. Helpers are extracted from
+# install.sh and sourced.
+set -e
+
+SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
+INSTALL_SH="$SCRIPT_DIR/../../install.sh"
+PASS=0
+FAIL=0
+
+# Extract _previous_torch_pin and its dependencies _torch_flavor_tag and
+# _torch_release_in_window.
+_FUNC_FILE=$(mktemp)
+{
+ sed -n '/^_torch_flavor_tag()/,/^}/p' "$INSTALL_SH"
+ echo ""
+ sed -n '/^_torch_release_in_window()/,/^}/p' "$INSTALL_SH"
+ echo ""
+ sed -n '/^_previous_torch_pin()/,/^}/p' "$INSTALL_SH"
+} > "$_FUNC_FILE"
+# shellcheck disable=SC1090
+. "$_FUNC_FILE"
+rm -f "$_FUNC_FILE"
+
+assert_eq() {
+ _label="$1"; _expected="$2"; _actual="$3"
+ if [ "$_actual" = "$_expected" ]; then
+ echo " PASS: $_label"; PASS=$((PASS + 1))
+ else
+ echo " FAIL: $_label (expected '$_expected', got '$_actual')"; FAIL=$((FAIL + 1))
+ fi
+}
+
+unset UNSLOTH_TORCH_UPGRADE
+
+echo "=== _previous_torch_pin: matching flavor keeps the release ==="
+assert_eq "cu126 wheel on cu126 leaf" "torch==2.10.0" "$(_previous_torch_pin '2.10.0+cu126' 'cu126' 'torch>=2.4,<2.12.0')"
+assert_eq "cu130 wheel on cu130 leaf" "torch==2.10.0" "$(_previous_torch_pin '2.10.0+cu130' 'cu130' 'torch>=2.4,<2.12.0')"
+assert_eq "cpu wheel on cpu leaf" "torch==2.10.0" "$(_previous_torch_pin '2.10.0+cpu' 'cpu' 'torch>=2.4,<2.12.0')"
+assert_eq "untagged wheel on cpu leaf" "torch==2.10.0" "$(_previous_torch_pin '2.10.0' 'cpu' 'torch>=2.4,<2.12.0')"
+assert_eq "local suffix stripped" "torch==2.9.1" "$(_previous_torch_pin '2.9.1+cu128' 'cu128' 'torch>=2.4,<2.12.0')"
+
+echo "=== _previous_torch_pin: flavor change installs the new build ==="
+assert_eq "cu126 wheel on cu130 leaf" "" "$(_previous_torch_pin '2.10.0+cu126' 'cu130' 'torch>=2.4,<2.12.0')"
+assert_eq "cpu wheel on cu126 leaf" "" "$(_previous_torch_pin '2.10.0+cpu' 'cu126' 'torch>=2.4,<2.12.0')"
+assert_eq "cu126 wheel on cpu leaf" "" "$(_previous_torch_pin '2.10.0+cu126' 'cpu' 'torch>=2.4,<2.12.0')"
+
+echo "=== _previous_torch_pin: rocm and unknown leaves never pin ==="
+assert_eq "rocm7.2 leaf keeps its floor" "" "$(_previous_torch_pin '2.11.0+rocm7.2' 'rocm7.2' 'torch>=2.4,<2.12.0')"
+assert_eq "gfx leaf keeps its floor" "" "$(_previous_torch_pin '2.11.0+rocm7.2' 'gfx120X-all' 'torch>=2.4,<2.12.0')"
+assert_eq "unknown mirror leaf" "" "$(_previous_torch_pin '2.10.0+cu126' 'simple' 'torch>=2.4,<2.12.0')"
+
+echo "=== _previous_torch_pin: probe noise never becomes a pin ==="
+assert_eq "empty version" "" "$(_previous_torch_pin '' 'cu126' 'torch>=2.4,<2.12.0')"
+assert_eq "garbage version" "" "$(_previous_torch_pin 'not-a-version' 'cpu' 'torch>=2.4,<2.12.0')"
+assert_eq "traceback fragment" "" "$(_previous_torch_pin "ModuleNotFoundError: No module named 'torch'" 'cpu' 'torch>=2.4,<2.12.0')"
+
+echo "=== _previous_torch_pin: out-of-window releases never pin ==="
+assert_eq "2.3.x below the cu floor" "" "$(_previous_torch_pin '2.3.1+cu118' 'cu118' 'torch>=2.4,<2.12.0')"
+assert_eq "2.12.x above the cu ceiling" "" "$(_previous_torch_pin '2.12.0+cu130' 'cu130' 'torch>=2.4,<2.12.0')"
+assert_eq "floor boundary 2.4.0 kept" "torch==2.4.0" "$(_previous_torch_pin '2.4.0+cu126' 'cu126' 'torch>=2.4,<2.12.0')"
+assert_eq "ceiling-adjacent 2.11.x kept" "torch==2.11.1" "$(_previous_torch_pin '2.11.1+cu130' 'cu130' 'torch>=2.4,<2.12.0')"
+assert_eq "cpu window excludes 2.11.x" "" "$(_previous_torch_pin '2.11.0+cpu' 'cpu' 'torch>=2.4,<2.11.0')"
+assert_eq "mac floor excludes 2.5.x" "" "$(_previous_torch_pin '2.5.1' 'cpu' 'torch>=2.6,<2.11.0')"
+assert_eq "malformed window never pins" "" "$(_previous_torch_pin '2.10.0+cu126' 'cu126' 'torch')"
+assert_eq "empty window never pins" "" "$(_previous_torch_pin '2.10.0+cu126' 'cu126' '')"
+
+echo "=== _torch_release_in_window ==="
+assert_eq "in window" "yes" "$(_torch_release_in_window '2.10.0' 'torch>=2.4,<2.12.0')"
+assert_eq "at floor" "yes" "$(_torch_release_in_window '2.4.0' 'torch>=2.4,<2.12.0')"
+assert_eq "below floor" "no" "$(_torch_release_in_window '2.3.1' 'torch>=2.4,<2.12.0')"
+assert_eq "at ceiling" "no" "$(_torch_release_in_window '2.12.0' 'torch>=2.4,<2.12.0')"
+assert_eq "next major" "no" "$(_torch_release_in_window '3.0.0' 'torch>=2.4,<2.12.0')"
+assert_eq "patch-level floor" "yes" "$(_torch_release_in_window '2.11.5' 'torch>=2.11.0,<2.12.0')"
+assert_eq "no ceiling -> no" "no" "$(_torch_release_in_window '2.10.0' 'torch>=2.4')"
+assert_eq "garbage minor -> no" "no" "$(_torch_release_in_window '2.x' 'torch>=2.4,<2.12.0')"
+
+echo "=== _previous_torch_pin: UNSLOTH_TORCH_UPGRADE=1 opts out ==="
+assert_eq "upgrade env set" "" "$(UNSLOTH_TORCH_UPGRADE=1 _previous_torch_pin '2.10.0+cu126' 'cu126' 'torch>=2.4,<2.12.0')"
+assert_eq "upgrade env 0" "torch==2.10.0" "$(UNSLOTH_TORCH_UPGRADE=0 _previous_torch_pin '2.10.0+cu126' 'cu126' 'torch>=2.4,<2.12.0')"
+
+echo "=== install.sh wiring ==="
+# The probe must run against the OLD venv, before it is moved aside for rollback.
+_probe_line=$(grep -n '_PREV_TORCH_VER=\$(' "$INSTALL_SH" | head -1 | cut -d: -f1)
+_move_line=$(grep -n '_start_studio_venv_replacement "\$VENV_DIR"' "$INSTALL_SH" | head -1 | cut -d: -f1)
+assert_eq "probe exists" "yes" "$([ -n "$_probe_line" ] && echo yes)"
+assert_eq "probe before venv replacement" "yes" "$([ -n "$_probe_line" ] && [ -n "$_move_line" ] && [ "$_probe_line" -lt "$_move_line" ] && echo yes)"
+# A kept release that vanished from the index must fall back to the supported range.
+assert_eq "resolve-failure fallback wired" "yes" "$(grep -q 'TORCH_CONSTRAINT="\$_PREV_FALLBACK_CONSTRAINT"' "$INSTALL_SH" && echo yes)"
+assert_eq "pin gated on SKIP_TORCH" "yes" "$(grep -q 'if \[ "\$SKIP_TORCH" = false \]; then' "$INSTALL_SH" && echo yes)"
+
+echo ""
+if [ "$FAIL" -gt 0 ]; then
+ echo "$FAIL check(s) FAILED"
+ exit 1
+fi
+echo "All $PASS checks passed"
diff --git a/tests/sh/test_torch_constraint.sh b/tests/sh/test_torch_constraint.sh
index 293a709360..d60dfc9f90 100644
--- a/tests/sh/test_torch_constraint.sh
+++ b/tests/sh/test_torch_constraint.sh
@@ -108,6 +108,20 @@ assert_eq "\$TORCH_CONSTRAINT used in pip install" "yes" "$_has_var"
_hardcoded=$(grep -c '"torch>=2.4,<2.11.0"' "$INSTALL_SH" || true)
assert_eq "hardcoded torch>=2.4 appears exactly once" "1" "$_hardcoded"
+# A fresh CUDA install widens the ceiling to <2.12.0 so cu12x/cu13x land torch
+# 2.11.x (matches the base image and _CUDA_TORCH_PKG_SPEC).
+_cuda_widen=$(grep -c 'TORCH_CONSTRAINT="torch>=2.4,<2.12.0"' "$INSTALL_SH" || true)
+assert_eq "CUDA TORCH_CONSTRAINT widened to <2.12.0" "1" "$_cuda_widen"
+
+# Widening keys off the final leaf (_torch_index_leaf), not the full URL, so a
+# mirror base path with cu*/rocm7.2 but a cpu/older-rocm leaf is not mis-widened.
+_cuda_case=$(grep -c 'cu\[0-9\]\*)' "$INSTALL_SH" || true)
+_has_cuda_case=$([ "$_cuda_case" -ge 1 ] && echo "yes" || echo "no")
+assert_eq "cu* index case adjusts TORCH_CONSTRAINT" "yes" "$_has_cuda_case"
+_leaf_case=$(grep -c 'case "\$_torch_index_leaf" in' "$INSTALL_SH" || true)
+_has_leaf_constraint=$([ "$_leaf_case" -ge 2 ] && echo "yes" || echo "no")
+assert_eq "constraint case anchors on _torch_index_leaf" "yes" "$_has_leaf_constraint"
+
echo ""
echo "=== Structural: tokenizers in no-torch-runtime.txt ==="
diff --git a/tests/sh/test_unsloth_torch_override.sh b/tests/sh/test_unsloth_torch_override.sh
new file mode 100644
index 0000000000..7e8e3f5b5b
--- /dev/null
+++ b/tests/sh/test_unsloth_torch_override.sh
@@ -0,0 +1,131 @@
+#!/bin/bash
+# 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 the torch-trio --overrides guard on the Step-2 unsloth installs in
+# install.sh. A released unsloth wheel can pin an older torch (2026.7.2 declares
+# torch<2.11.0); without the overrides file a with-deps PyPI resolve downgrades
+# the trio Step 1 installed, and the flavor guard misses it (PyPI's torch 2.10
+# default is itself cu128-flavored). Same assertion pattern as test_torch_constraint.sh.
+set -e
+
+SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
+INSTALL_SH="$SCRIPT_DIR/../../install.sh"
+PASS=0
+FAIL=0
+
+assert_true() {
+ _label="$1"; _ok="$2"
+ if [ "$_ok" = "0" ]; then
+ echo " PASS: $_label"
+ PASS=$((PASS + 1))
+ else
+ echo " FAIL: $_label"
+ FAIL=$((FAIL + 1))
+ fi
+}
+
+echo "=== test_unsloth_torch_override ==="
+
+# 1. Every with-deps unsloth install carries the overrides expansion (local,
+# generic, migrated); the --no-deps no-torch paths need no guard.
+_local_block=$(grep -A2 '"install unsloth (local)"' "$INSTALL_SH")
+printf '%s' "$_local_block" | grep -q -- '--overrides "\$_UNSLOTH_TORCH_OVERRIDES"'
+assert_true "local (with-deps) unsloth install passes --overrides" "$?"
+
+_generic_block=$(grep -A2 '"install unsloth" uv pip install' "$INSTALL_SH")
+printf '%s' "$_generic_block" | grep -q -- '--overrides "\$_UNSLOTH_TORCH_OVERRIDES"'
+assert_true "generic (with-deps) unsloth install passes --overrides" "$?"
+
+_migrated_block=$(grep -A3 '"install unsloth (migrated)"' "$INSTALL_SH")
+printf '%s' "$_migrated_block" | grep -q -- '--overrides "\$_UNSLOTH_TORCH_OVERRIDES"'
+assert_true "migrated (with-deps) unsloth install passes --overrides" "$?"
+
+_no_torch_block=$(grep -A2 '"install unsloth (no-torch)"' "$INSTALL_SH")
+if printf '%s' "$_no_torch_block" | grep -q -- '--overrides'; then _rc=1; else _rc=0; fi
+assert_true "no-torch (--no-deps) unsloth install has no overrides" "$_rc"
+
+_migrated_nt_block=$(grep -A2 '"install unsloth (migrated no-torch)"' "$INSTALL_SH")
+if printf '%s' "$_migrated_nt_block" | grep -q -- '--overrides'; then _rc=1; else _rc=0; fi
+assert_true "migrated no-torch (--no-deps) unsloth install has no overrides" "$_rc"
+
+# 2. The overrides file is only built when SKIP_TORCH=false.
+grep -B2 '_torch_trio_pins=\$(' "$INSTALL_SH" | grep -q 'SKIP_TORCH" = false'
+assert_true "overrides file build is gated on SKIP_TORCH=false" "$?"
+
+# 3. The pin-collection snippet emits exact ==pins for the installed trio (run
+# the embedded python against this test's interpreter).
+_snippet=$(sed -n '/_torch_trio_pins=\$("\$_VENV_PY" -c "/,/^" 2>\/dev\/null)/p' "$INSTALL_SH" \
+ | sed '1s/.*-c "//' | sed '$d')
+_out=$(python3 -c "$_snippet" 2>&1) || true
+# torch may or may not be importable on the test host; the snippet must not
+# crash and every line it does emit must be an exact pkg==version pin.
+if [ -n "$_out" ]; then
+ printf '%s\n' "$_out" | grep -vqE '^(torch|torchvision|torchaudio)==.+$' && _rc=1 || _rc=0
+else
+ _rc=0
+fi
+assert_true "pin snippet emits only exact trio ==pins (or nothing)" "$_rc"
+
+# 4. The temp overrides file is cleaned up after Step 2.
+grep -q 'rm -f "\$_UNSLOTH_TORCH_OVERRIDES"' "$INSTALL_SH"
+assert_true "overrides temp file is removed after the unsloth installs" "$?"
+
+# 5. Any UV_OVERRIDE env file is folded in (the CLI --overrides flag would
+# otherwise replace it, dropping e.g. the macOS arm64 darwin overrides).
+grep -q 'for _ov_file in \${UV_OVERRIDE:-}' "$INSTALL_SH"
+assert_true "UV_OVERRIDE env files are merged into the overrides file" "$?"
+
+# 6. The EXIT trap also removes the overrides file, so a failed Step 2 (set -e
+# fires before the normal-path rm) cannot leak it.
+sed -n '/_on_install_exit() {/,/^}/p' "$INSTALL_SH" \
+ | grep -q 'rm -f "\$_UNSLOTH_TORCH_OVERRIDES"'
+assert_true "EXIT trap removes the overrides temp file on failure" "$?"
+
+# 7. The UV_OVERRIDE fold filters inherited files instead of cat-ing them (run
+# the extracted awk program on sample files): (a) inherited torch-trio lines
+# are dropped so the generated exact pins win (uv intersects duplicates);
+# (b) every line is newline-terminated so an unterminated file cannot join
+# two requirements into one.
+_awk_prog=$(sed -n "s/.*awk '\(.*\)' \"\$_ov_file\".*/\1/p" "$INSTALL_SH")
+[ -n "$_awk_prog" ]
+assert_true "UV_OVERRIDE fold uses the trio-filtering awk program" "$?"
+
+_ov_dir=$(mktemp -d)
+printf '%s' 'transformers>=4.57.6' > "$_ov_dir/ov1.txt" # no trailing newline
+cat > "$_ov_dir/ov2.txt" <<'EOF'
+# comment survives
+torch<2.11.0
+torchvision==0.25.0
+torchaudio!=2.11.0
+torchmetrics==1.0
+anyio<4.14.0
+EOF
+_merged="$_ov_dir/merged.txt"
+printf '%s\n' 'torch==2.11.0+cu128' > "$_merged"
+for _f in "$_ov_dir/ov1.txt" "$_ov_dir/ov2.txt"; do
+ awk "$_awk_prog" "$_f" >> "$_merged"
+done
+
+grep -qx 'transformers>=4.57.6' "$_merged"
+assert_true "no-trailing-newline override stays a separate requirement line" "$?"
+
+if grep -qx 'torchmetrics==1.0' "$_merged" && grep -qx 'anyio<4.14.0' "$_merged"; then
+ _rc=0
+else
+ _rc=1
+fi
+assert_true "unrelated inherited overrides are preserved" "$_rc"
+
+if grep -qE '^(torch|torchvision|torchaudio)([[:space:]<>=!~;@[]|$)' "$_merged" \
+ && [ "$(grep -cE '^(torch|torchvision|torchaudio)([[:space:]<>=!~;@[]|$)' "$_merged")" != "1" ]; then
+ _rc=1
+else
+ _rc=0
+fi
+grep -qx 'torch==2.11.0+cu128' "$_merged" || _rc=1
+assert_true "inherited torch-trio lines are dropped; generated pin wins" "$_rc"
+rm -rf "$_ov_dir"
+
+echo ""
+echo "Results: $PASS passed, $FAIL failed"
+[ "$FAIL" -eq 0 ] || exit 1
From b307823b1daf7013632340bceac5b2f70dbc04a8 Mon Sep 17 00:00:00 2001
From: Andrew Chen <48723787+chuenchen309@users.noreply.github.com>
Date: Sun, 19 Jul 2026 21:33:48 +0800
Subject: [PATCH 04/41] fix(chat_templates): bind loop_messages when
default_system_message is None (#7199)
* fix(chat_templates): bind loop_messages when default_system_message is None
construct_chat_template(default_system_message=None) built a system part that
binds loop_messages only inside the `{% if messages[0]['role'] == 'system' %}`
arm. The `Fix missing loop_messages` step right below then found no
unconditional `{% set loop_messages = messages %}`, concluded loop_messages was
missing, and rewrote `{% for message in loop_messages %}` back to
`{% for message in messages %}` -- undoing the `messages[1:]` skip.
A caller-supplied system message therefore reached the loop and tripped
raise_exception:
Only user and assistant roles are supported!
Add the `{% else %}` arm so loop_messages is always bound, mirroring the
default_system_message is not None branch minus the default text. That also
stops the rewrite from firing, since the unconditional binding is now present.
Renders before / after, same template, same inputs:
default_system_message input before after
None system msg raise_exception 'Be terse.\n### User: Hi\n'
None no system '### User: Hi\n' unchanged
'You are helpful.' system msg 'Be terse.\n### User: Hi\n' unchanged
'You are helpful.' no system 'You are helpful.\n...' unchanged
The rewrite still fires for templates with no {SYSTEM} part, which is what it
was there for -- verified unchanged.
Co-Authored-By: Claude Opus 4.8
* Scope loop_messages binding to {SYSTEM} templates for PR #7199
The None branch now only adds the else arm when system_part contains
{SYSTEM}, so a static prefix with no {SYSTEM} placeholder keeps raising on a
caller system message instead of silently dropping it. Strengthen the tests:
assert the default does not leak when a caller system message is present, and
add a regression test for the static prefix case.
---------
Co-authored-by: Claude Opus 4.8
Co-authored-by: danielhanchen
---
...test_construct_chat_template_validation.py | 90 +++++++++++++++++++
unsloth/chat_templates.py | 6 ++
2 files changed, 96 insertions(+)
diff --git a/tests/python/test_construct_chat_template_validation.py b/tests/python/test_construct_chat_template_validation.py
index 66d3d80920..53d281d435 100644
--- a/tests/python/test_construct_chat_template_validation.py
+++ b/tests/python/test_construct_chat_template_validation.py
@@ -104,3 +104,93 @@ def test_chat_template_does_not_leak_sentinel_when_section_starts_with_it(chat_t
)
assert "{INPUT}" not in jinja_template
assert "{OUTPUT}" not in jinja_template
+
+
+_SYSTEM_CHAT_TEMPLATE = (
+ "{SYSTEM}\n"
+ "### User: {INPUT}\n### Assistant: {OUTPUT}"
+ "### User: {INPUT}\n### Assistant: {OUTPUT}"
+)
+
+
+def _render(jinja_template, messages):
+ from jinja2.sandbox import ImmutableSandboxedEnvironment
+
+ env = ImmutableSandboxedEnvironment()
+ env.globals["raise_exception"] = lambda message: (_ for _ in ()).throw(RuntimeError(message))
+ return env.from_string(jinja_template).render(
+ messages = messages,
+ bos_token = "",
+ eos_token = "",
+ add_generation_prompt = False,
+ )
+
+
+@pytest.mark.parametrize("default_system_message", [None, "You are helpful."])
+def test_system_message_is_consumed_by_the_system_part(default_system_message):
+ """A caller-supplied system message must be rendered by the system part and
+ skipped by the message loop, whatever `default_system_message` is.
+
+ With `default_system_message = None` the generated template used to bind
+ `loop_messages` only inside the `{% if %}` arm. The `Fix missing
+ loop_messages` step then saw no unconditional binding, rewrote the loop back
+ to `messages`, and the system message reached the loop and tripped
+ `raise_exception`.
+ """
+ _, jinja_template, _, _ = construct_chat_template(
+ tokenizer = _SuccessFakeTokenizer(),
+ chat_template = _SYSTEM_CHAT_TEMPLATE,
+ default_system_message = default_system_message,
+ extra_eos_tokens = [""],
+ )
+ rendered = _render(
+ jinja_template,
+ [
+ {"role": "system", "content": "Be terse."},
+ {"role": "user", "content": "Hi"},
+ ],
+ )
+ assert rendered.count("Be terse.") == 1
+ assert rendered.count("Hi") == 1
+ # A caller system message overrides the default; the default must not leak in.
+ if default_system_message is not None:
+ assert default_system_message not in rendered
+
+
+def test_absent_system_message_still_renders_without_default():
+ """`default_system_message = None` with no system message in the input must
+ keep working -- the `{% else %}` arm has to bind `loop_messages = messages`."""
+ _, jinja_template, _, _ = construct_chat_template(
+ tokenizer = _SuccessFakeTokenizer(),
+ chat_template = _SYSTEM_CHAT_TEMPLATE,
+ default_system_message = None,
+ extra_eos_tokens = [""],
+ )
+ rendered = _render(jinja_template, [{"role": "user", "content": "Hi"}])
+ assert "Hi" in rendered
+
+
+_NO_SYSTEM_CHAT_TEMPLATE = (
+ "PREAMBLE\n"
+ "### User: {INPUT}\n### Assistant: {OUTPUT}"
+ "### User: {INPUT}\n### Assistant: {OUTPUT}"
+)
+
+
+def test_static_prefix_without_system_still_rejects_system_message():
+ """A template with a static prefix but no {SYSTEM} placeholder cannot render a
+ caller system message, so it must still raise rather than silently drop it."""
+ _, jinja_template, _, _ = construct_chat_template(
+ tokenizer = _SuccessFakeTokenizer(),
+ chat_template = _NO_SYSTEM_CHAT_TEMPLATE,
+ default_system_message = None,
+ extra_eos_tokens = [""],
+ )
+ with pytest.raises(RuntimeError, match = "Only user and assistant roles are supported!"):
+ _render(
+ jinja_template,
+ [
+ {"role": "system", "content": "Be terse."},
+ {"role": "user", "content": "Hi"},
+ ],
+ )
diff --git a/unsloth/chat_templates.py b/unsloth/chat_templates.py
index f47c78ba80..b857c34bcb 100644
--- a/unsloth/chat_templates.py
+++ b/unsloth/chat_templates.py
@@ -2652,6 +2652,12 @@ extra_eos_tokens = None,
"{{ '" + full_system + "' }}"\
"{% set loop_messages = messages %}"\
"{% endif %}"
+ elif "{SYSTEM}" in system_part:
+ # Only bind loop_messages when the template can render a caller system
+ # message. A static prefix with no {SYSTEM} must still raise, not drop it.
+ partial_system += "{% else %}"\
+ "{% set loop_messages = messages %}"\
+ "{% endif %}"
else:
partial_system += "{% endif %}"
From b3c0259cffdccb91362e7a16dc856632319f7304 Mon Sep 17 00:00:00 2001
From: Daniel Han
Date: Sun, 19 Jul 2026 07:55:06 -0700
Subject: [PATCH 05/41] Installer: preserve the previous torch release across
every flavor and vendor on re-runs (#7250)
* install: preserve the previous torch release across every flavor and vendor
A re-run of curl | sh over an existing install was supposed to keep the
user's validated torch release, but the pin required the old build's
local flavor tag to match the freshly chosen index leaf. That gate was
wrong in practice: a PyPI-sourced torch reports a BARE version (on Linux
the PyPI wheel IS a CUDA build), which classified as cpu and never
matched a cu leaf, so a healthy 2.10 on a cu130 host was silently moved
to 2.11 (reproduced end to end); the same happened for any flavor drift
such as cu128 to cu130 after a driver upgrade, and AMD ROCm leaves were
excluded from preservation entirely.
The rule is now release-based and flavor-agnostic: the probed previous
release is pinned whenever it sits inside the final constraint window,
and the pin installs from the freshly chosen index, so the flavor always
follows the machine (NVIDIA cu*, AMD rocm/gfx, Intel/CPU, mac) while the
release follows the user. The pin is evaluated AFTER every index and
constraint decision including the Strix reroute, so raised floors
(rocm7.2 / Strix gfx need torch 2.11 for the _grouped_mm fix) correctly
reject an older release and win. UNSLOTH_TORCH_UPGRADE=1 still opts out,
out-of-window releases are never kept, and probe noise never becomes a
pin.
The kept-release install with its range fallback (for indexes that do
not carry the exact release) is factored into
_install_torch_default_index and used by every --default-index torch
path: the default NVIDIA/CPU/mac path and all three ROCm-index
fallbacks, which previously bypassed the fallback. The Radeon-repo
direct-wheel path keeps its curated per-arch wheel set (those wheels are
already exact-pinned per rocm release).
Platform coverage: install.sh serves Linux, WSL (including the WoA
fallback), and macOS for all vendors; native Windows install.ps1 still
caps at <2.11.0 everywhere, so the silent 2.10-to-2.11 move cannot occur
there (2.11 alignment is a separate follow-up).
Verified: 35-check unit suite rewritten to the new spec (any-flavor
keep, floor rejection, noise, window edges, opt-out, wiring including
pin-after-reroute and helper coverage); end-to-end matrix against
sandboxed UNSLOTH_STUDIO_HOME installs on a cu130 host covering PyPI
bare, cu128 drift, cu130 same-flavor, out-of-window 2.3, the upgrade
opt-out, the hidden-GPU cpu leaf, and a fresh-install control.
* install: honor the kept torch release on the Radeon direct-wheel path
The Radeon repo path installs an explicit wheel trio selected by
_pick_radeon_wheel, bypassing --default-index, so the kept-release pin
only took effect when the listing failed and the install fell back to
the ROCm index. On a re-run over an in-window Radeon install the trio
search started at the newest common minor and silently moved the user
forward (2.9 to 2.10 whenever the repo offered both).
The trio search now starts at the kept release's minor when
_PREV_TORCH_PIN is set and the listing still offers a torch wheel for
that minor. Radeon wheels are patch-curated per rocm release, so the
minor is the unit of preservation there; the raised rocm7.2 / Strix
floors still win because the pin is window-checked against the final
constraint before this point, and gaps keep the existing downward
search / ROCm-index fallback.
Verified with a simulated listing carrying both a 2.9 and a 2.10 trio:
no pin selects the 2.10 trio, a kept 2.9 release selects the matched
2.9 / 0.24 / 2.9 trio, and an unavailable minor degrades to the newest
trio. Added a structural wiring check to test_previous_torch_pin.sh
(now 36 checks).
* install: tighten comments in the torch preservation paths
* install: exact kept release on the Radeon path, pin fallback in ROCm repairs
The minor-level clamp on the Radeon direct-wheel path still allowed
patch drift (a kept 2.10.0 could become 2.10.1 when the listing carried
both) and the downward gap search could settle below the kept minor,
both breaking the exact preservation guarantee the other vendor paths
honor. The kept release now gets an exact-first trio attempt before the
newest-trio search: pick the kept patch (else the newest patch of the
kept minor, for listings that pruned the exact patch) together with the
paired torchvision/torchaudio wheels for that minor. Any gap warns and
falls back to the unchanged newest-trio search, mirroring
_install_torch_default_index, so a rerun installs either the kept
release or the same set a fresh install would choose, never something
in between.
The two ROCm torch repair sites (torch overwritten by dependency
resolution, on the migrated and fresh paths) installed TORCH_CONSTRAINT
directly, so a pinned release missing from the generic ROCm index would
abort the rerun instead of falling back. Both now route through
_install_torch_default_index, which passes extra uv args through
(--force-reinstall) and clears the pin once the fallback fires so later
paths stay consistent.
Verified against synthetic listings: both patches listed keeps exactly
2.10.0; a kept minor missing vision/audio warns and yields the newest
complete trio rather than a silent undercut; a pruned patch stays on
the kept minor; no pin keeps the existing newest-trio behavior. Unit
suite now 39 checks, all passing.
* install: never pin nightly/dev/source torch builds on a rerun
A survey of published torch version strings (PyPI bare, +cpu, +cu116
through +cu132, +rocmX.Y and +rocmX.Y.Z, +xpu, nightly .devYYYYMMDD,
source a0+git, rc tags) showed one gap: nightly, dev, rc, and source
builds passed the loose release-shape check, producing a pin such as
torch==2.11.0.dev20250704 that no stable index carries. The range
fallback rescued the install, but it printed "keeping it" and then
burned a doomed resolve first. The base must now be a plain numeric
X.Y[.Z] release, so those builds skip the pin and go straight to the
newest supported release.
Added unit checks for +xpu and three-component +rocm7.2.1 tags (both
already preserved correctly) and for nightly, a0 source, and rc builds
(never pinned). Suite now 44 checks, all passing.
* install: pair kept-release companions, protect the flavor repair, note substitutions
Three fixes from a 12-way review pass over the preservation work:
The kept-release install left torchvision and torchaudio unconstrained
next to the exact torch pin. torchvision exact-pins its torch in wheel
metadata so it always paired correctly, but torchaudio no longer does:
a kept torch 2.9.0 on cu130 resolved torchaudio 2.11.0 (verified with
uv dry-runs). The helper now pairs both companions to the kept minor
(torchvision 0.minor+15, torchaudio 2.minor); if the index lacks the
paired set the existing range fallback fires. Verified resolving
correctly on cu130, cu126, and rocm6.4.
The wrong-flavor repair at the end of the install was the one remaining
default-index torch install outside the helper. It runs under set -e,
so a retained pin absent from the repair index (reachable when the
Radeon direct-wheel path installed the kept release and dependency
resolution later overwrote it) aborted the installer at the last step
instead of falling back. It now routes through the helper with its
reinstall flags passed through.
The Radeon kept-release path installed a same-series build silently
when the listing had pruned the exact patch; it now prints what it is
substituting.
Unit suite extended with wiring checks for all three (46 checks, all
passing).
---
install.sh | 180 +++++++++++++++++-----------
tests/sh/test_previous_torch_pin.sh | 106 ++++++++++------
2 files changed, 181 insertions(+), 105 deletions(-)
diff --git a/install.sh b/install.sh
index 6076721540..7918a2bd23 100755
--- a/install.sh
+++ b/install.sh
@@ -2225,37 +2225,67 @@ _torch_release_in_window() {
echo "no"
}
-# Whether a re-run should keep the previous venv's torch: echo "torch==X.Y.Z" when the
-# probed previous version ($1) has a flavor tag matching the freshly chosen cu*/cpu index
-# leaf ($2) AND sits inside the active constraint window ($3), else "". Re-running
-# `curl | sh` rebuilds the venv for clean state, but a healthy torch the user already
-# validated must not be silently moved to a newer release (2.10 -> 2.11); a flavor
-# change (cpu <-> cuda, cu126 -> cu130) still installs the correct new build, rocm
-# leaves keep their floors (rocm7.2 must land 2.11 for the Strix _grouped_mm fix), and
-# a release outside the window (2.3.x manual install, 2.12.x manual upgrade) is never
-# kept: the installer's own bounds win. Opt out with UNSLOTH_TORCH_UPGRADE=1 to get
-# the newest release.
+# Keep the previous venv's torch on a re-run: echo "torch==X.Y.Z" when the probed
+# version ($1) is inside the active constraint window ($2), else "". The RELEASE is kept
+# regardless of flavor tag; the pin installs from the freshly chosen index, so flavor
+# follows the machine (cpu <-> cuda, cu126 -> cu130, PyPI bare -> +cu130) while the
+# release follows the user. Gating on flavor was wrong: a PyPI torch reports a BARE
+# version (on Linux the PyPI wheel IS CUDA), misclassified "cpu", so a healthy 2.10 on a
+# cu130 host was moved to 2.11. Per-leaf floors still win (rocm7.2 / gfx >=2.11 for the
+# Strix _grouped_mm fix, out-of-window manual installs) and are never pinned; the caller's
+# _PREV_FALLBACK_CONSTRAINT installs the newest supported release when the index lacks the
+# exact one. Opt out with UNSLOTH_TORCH_UPGRADE=1.
_previous_torch_pin() {
_ptp_ver="$1"
- _ptp_leaf="$2"
- _ptp_con="$3"
+ _ptp_con="$2"
[ -n "$_ptp_ver" ] || { echo ""; return; }
[ "${UNSLOTH_TORCH_UPGRADE:-0}" = "1" ] && { echo ""; return; }
- case "$_ptp_leaf" in
- cu[0-9]*|cpu) ;;
- *) echo ""; return ;;
- esac
_ptp_base="${_ptp_ver%%+*}"
- # The base must look like a release (probe noise / garbage must never become a pin).
+ # Base must be a plain numeric release (X.Y[.Z]); probe noise and
+ # nightly/dev/source builds (2.11.0.dev20250704, 2.9.0a0) must never
+ # become a pin -- no stable index carries them, so pinning would only
+ # print "keeping it" and then burn a doomed resolve before falling back.
case "$_ptp_base" in
+ *[!0-9.]* | *..* | .* | *.) echo ""; return ;;
[0-9]*.[0-9]*) ;;
*) echo ""; return ;;
esac
[ "$(_torch_release_in_window "$_ptp_base" "$_ptp_con")" = "yes" ] || { echo ""; return; }
- if [ "$(_torch_flavor_tag "$_ptp_ver")" = "$_ptp_leaf" ]; then
- echo "torch==$_ptp_base"
+ echo "torch==$_ptp_base"
+}
+
+# Install torch from TORCH_INDEX_URL honoring a kept-release pin: with _PREV_TORCH_PIN
+# set, TORCH_CONSTRAINT is the exact previous release; fall back to the supported range
+# if the index lacks it (pruned mirror) rather than failing. Used by every --default-index
+# path (NVIDIA cu*, AMD rocm/gfx fallbacks, cpu/mac, ROCm repairs) so preservation is
+# uniform. Extra args (e.g. --force-reinstall) are passed through to uv.
+_install_torch_default_index() {
+ if [ -n "$_PREV_TORCH_PIN" ]; then
+ # Pair the companions with the kept torch minor: torchaudio no longer
+ # exact-pins torch in its metadata, so leaving it unconstrained resolves
+ # a newer mismatched build (a kept torch 2.9.0 pulled torchaudio 2.11.0).
+ _itdi_base="${_PREV_TORCH_PIN#torch==}"
+ _itdi_minor="${_itdi_base#*.}"
+ _itdi_minor="${_itdi_minor%%.*}"
+ _itdi_tv="torchvision"
+ _itdi_ta="torchaudio"
+ case "$_itdi_base" in
+ 2.*)
+ _itdi_tv="torchvision==0.$((_itdi_minor + 15)).*"
+ _itdi_ta="torchaudio==2.${_itdi_minor}.*"
+ ;;
+ esac
+ if ! run_install_cmd_retry "install PyTorch (kept release)" uv pip install --python "$_VENV_PY" "$TORCH_CONSTRAINT" "$_itdi_tv" "$_itdi_ta" \
+ --default-index "$TORCH_INDEX_URL" "$@"; then
+ substep "[WARN] $_PREV_TORCH_PIN is not installable from $TORCH_INDEX_URL -- installing the newest supported release instead" "$C_WARN"
+ TORCH_CONSTRAINT="$_PREV_FALLBACK_CONSTRAINT"
+ _PREV_TORCH_PIN=""
+ run_install_cmd_retry "install PyTorch" uv pip install --python "$_VENV_PY" "$TORCH_CONSTRAINT" torchvision torchaudio \
+ --default-index "$TORCH_INDEX_URL" "$@"
+ fi
else
- echo ""
+ run_install_cmd_retry "install PyTorch" uv pip install --python "$_VENV_PY" "$TORCH_CONSTRAINT" torchvision torchaudio \
+ --default-index "$TORCH_INDEX_URL" "$@"
fi
}
@@ -2561,21 +2591,6 @@ case "$_torch_index_leaf" in
cu[0-9]*) TORCH_CONSTRAINT="torch>=2.4,<2.12.0" ;;
esac
-# Re-run over an existing install: keep the previous venv's torch release instead of
-# resolving the newest in range. The range stays in _PREV_FALLBACK_CONSTRAINT so the
-# install can fall back when the exact release is not on the chosen index (custom
-# mirrors may prune old wheels). Skipped for --no-torch (no previous probe runs).
-_PREV_TORCH_PIN=""
-_PREV_FALLBACK_CONSTRAINT="$TORCH_CONSTRAINT"
-if [ "$SKIP_TORCH" = false ]; then
- _prev_pin=$(_previous_torch_pin "$_PREV_TORCH_VER" "$_torch_index_leaf" "$TORCH_CONSTRAINT")
- if [ -n "$_prev_pin" ]; then
- _PREV_TORCH_PIN="$_prev_pin"
- TORCH_CONSTRAINT="$_prev_pin"
- substep "existing install has torch $_PREV_TORCH_VER -- keeping it (set UNSLOTH_TORCH_UPGRADE=1 to get the newest release)"
- fi
-fi
-
# 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".
@@ -2660,6 +2675,23 @@ case "$TORCH_INDEX_URL" in
fi
;;
esac
+# Re-run over an existing install: keep the previous venv's torch RELEASE; the fresh
+# index above supplies the right flavor for this machine. Evaluated HERE, after every
+# index/constraint decision including the Strix reroute, so the window checked is the
+# final one and a raised floor (rocm7.2 / Strix gfx) rejects an older release.
+# _PREV_FALLBACK_CONSTRAINT keeps the range so the install can fall back when the exact
+# release is not on the chosen index (mirrors may prune old wheels). Skipped for --no-torch.
+_PREV_TORCH_PIN=""
+_PREV_FALLBACK_CONSTRAINT="$TORCH_CONSTRAINT"
+if [ "$SKIP_TORCH" = false ]; then
+ _prev_pin=$(_previous_torch_pin "$_PREV_TORCH_VER" "$TORCH_CONSTRAINT")
+ if [ -n "$_prev_pin" ]; then
+ _PREV_TORCH_PIN="$_prev_pin"
+ TORCH_CONSTRAINT="$_prev_pin"
+ substep "existing install has torch $_PREV_TORCH_VER -- keeping it (set UNSLOTH_TORCH_UPGRADE=1 to get the newest release)"
+ fi
+fi
+
_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"
@@ -2885,10 +2917,7 @@ if [ "$_MIGRATED" = true ]; then
_has_hip=$("$_VENV_PY" -c "import torch; print(getattr(torch.version,'hip','') or '')" 2>/dev/null || true)
if [ -z "$_has_hip" ]; then
substep "repairing ROCm torch (overwritten by dependency resolution)..."
- run_install_cmd_retry "repair ROCm torch" uv pip install --python "$_VENV_PY" \
- "$TORCH_CONSTRAINT" torchvision torchaudio \
- --default-index "$TORCH_INDEX_URL" \
- --force-reinstall
+ _install_torch_default_index --force-reinstall
fi
;;
esac
@@ -2953,7 +2982,42 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
_ta_ver=$(_extract_version "$_ta_whl" "torchaudio")
_radeon_versions_match=false
- if [ -n "$_torch_ver" ] && [ -n "$_tv_ver" ] && [ -n "$_ta_ver" ]; then
+ # Kept release (_PREV_TORCH_PIN) wins here too: pick its exact
+ # patch (else the newest patch of its minor) plus the paired
+ # vision/audio wheels. Any gap falls back to the newest-trio
+ # search below, mirroring _install_torch_default_index, so a
+ # rerun never drifts to another release nor below the kept one.
+ if [ -n "$_PREV_TORCH_PIN" ]; then
+ _prev_kept_base="${_PREV_TORCH_PIN#torch==}"
+ _prev_kept_minor="${_prev_kept_base#*.}"
+ _prev_kept_minor="${_prev_kept_minor%%.*}"
+ case "$_prev_kept_minor" in
+ ''|*[!0-9]*) ;;
+ *)
+ _kept_torch=$(_pick_radeon_wheel "torch" "${_prev_kept_base}" 2>/dev/null) || _kept_torch=""
+ [ -z "$_kept_torch" ] && { _kept_torch=$(_pick_radeon_wheel "torch" "2.${_prev_kept_minor}." 2>/dev/null) || _kept_torch=""; }
+ _kept_tv=$(_pick_radeon_wheel "torchvision" "0.$((_prev_kept_minor + 15))." 2>/dev/null) || _kept_tv=""
+ _kept_ta=$(_pick_radeon_wheel "torchaudio" "2.${_prev_kept_minor}." 2>/dev/null) || _kept_ta=""
+ if [ -n "$_kept_torch" ] && [ -n "$_kept_tv" ] && [ -n "$_kept_ta" ]; then
+ _torch_whl=$_kept_torch
+ _tv_whl=$_kept_tv
+ _ta_whl=$_kept_ta
+ _tri_whl=""
+ _radeon_versions_match=true
+ # Say so when the listing pruned the exact patch
+ # and a same-series build is installed instead.
+ case "$(printf '%s' "${_kept_torch##*/}" | sed 's/%2[Bb]/+/g')" in
+ "torch-${_prev_kept_base}"[+-]*) ;;
+ *) substep "kept release ${_prev_kept_base} is not in the Radeon listing -- installing the closest 2.${_prev_kept_minor} series build instead" ;;
+ esac
+ else
+ substep "[WARN] Radeon repo lacks a complete wheel set for kept $_PREV_TORCH_PIN -- installing the newest compatible set instead" "$C_WARN"
+ fi
+ ;;
+ esac
+ fi
+ if [ "$_radeon_versions_match" != true ] && \
+ [ -n "$_torch_ver" ] && [ -n "$_tv_ver" ] && [ -n "$_ta_ver" ]; then
_torch_minor=${_torch_ver#*.}
_ta_minor=${_ta_ver#*.}
_tv_minor=${_tv_ver#*.}
@@ -3011,9 +3075,7 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
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"
- run_install_cmd_retry "install PyTorch" uv pip install --python "$_VENV_PY" \
- "$TORCH_CONSTRAINT" torchvision torchaudio \
- --default-index "$TORCH_INDEX_URL"
+ _install_torch_default_index
else
substep "installing PyTorch from Radeon repo (${_RADEON_BASE_URL})..."
# Pass explicit wheel URLs so the matched trio is
@@ -3034,32 +3096,15 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
fi
else
substep "[WARN] Radeon repo unavailable; falling back to ROCm index ($TORCH_INDEX_URL)" "$C_WARN"
- run_install_cmd_retry "install PyTorch" uv pip install --python "$_VENV_PY" \
- "$TORCH_CONSTRAINT" torchvision torchaudio \
- --default-index "$TORCH_INDEX_URL"
+ _install_torch_default_index
fi
else
substep "[WARN] Radeon GPU detected but could not detect full ROCm version; falling back to ROCm index" "$C_WARN"
- run_install_cmd_retry "install PyTorch" uv pip install --python "$_VENV_PY" \
- "$TORCH_CONSTRAINT" torchvision torchaudio \
- --default-index "$TORCH_INDEX_URL"
+ _install_torch_default_index
fi
else
substep "installing PyTorch ($TORCH_INDEX_URL)..."
- if [ -n "$_PREV_TORCH_PIN" ]; then
- # Kept previous release: fall back to the supported range if the exact
- # release is not resolvable from the chosen index (pruned mirror).
- if ! run_install_cmd_retry "install PyTorch (kept release)" uv pip install --python "$_VENV_PY" "$TORCH_CONSTRAINT" torchvision torchaudio \
- --default-index "$TORCH_INDEX_URL"; then
- substep "[WARN] $_PREV_TORCH_PIN is not installable from $TORCH_INDEX_URL -- installing the newest supported release instead" "$C_WARN"
- TORCH_CONSTRAINT="$_PREV_FALLBACK_CONSTRAINT"
- run_install_cmd_retry "install PyTorch" uv pip install --python "$_VENV_PY" "$TORCH_CONSTRAINT" torchvision torchaudio \
- --default-index "$TORCH_INDEX_URL"
- fi
- else
- run_install_cmd_retry "install PyTorch" uv pip install --python "$_VENV_PY" "$TORCH_CONSTRAINT" torchvision torchaudio \
- --default-index "$TORCH_INDEX_URL"
- fi
+ _install_torch_default_index
fi
# AMD ROCm: install bitsandbytes (once, after torch, for all ROCm paths).
# Gate on SKIP_TORCH=false so a user running with --no-torch on a ROCm
@@ -3122,10 +3167,7 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
_has_hip=$("$_VENV_PY" -c "import torch; print(getattr(torch.version,'hip','') or '')" 2>/dev/null || true)
if [ -z "$_has_hip" ]; then
substep "repairing ROCm torch (overwritten by dependency resolution)..."
- run_install_cmd_retry "repair ROCm torch" uv pip install --python "$_VENV_PY" \
- "$TORCH_CONSTRAINT" torchvision torchaudio \
- --default-index "$TORCH_INDEX_URL" \
- --force-reinstall
+ _install_torch_default_index --force-reinstall
fi
;;
esac
@@ -3164,9 +3206,7 @@ if [ "$SKIP_TORCH" = false ] && [ -n "${TORCH_INDEX_URL:-}" ]; then
if [ -n "$_installed_torch_tag" ] && [ "$_installed_torch_tag" != "$_expected_torch_tag" ] \
&& [ "$(_torch_index_repairable "$TORCH_INDEX_URL")" = "yes" ]; then
substep "PyTorch flavor mismatch (installed $_installed_torch_tag, need $_expected_torch_tag) -- reinstalling correct build..."
- run_install_cmd "reinstall PyTorch ($_expected_torch_tag)" uv pip install --python "$_VENV_PY" \
- "$TORCH_CONSTRAINT" torchvision torchaudio \
- --default-index "$TORCH_INDEX_URL" \
+ _install_torch_default_index \
--reinstall-package torch --reinstall-package torchvision --reinstall-package torchaudio
_installed_torch_ver=$("$_VENV_PY" -c "import torch; print(torch.__version__)" 2>/dev/null || true)
_installed_torch_tag=""
diff --git a/tests/sh/test_previous_torch_pin.sh b/tests/sh/test_previous_torch_pin.sh
index 253ede8a27..1bc0d1f27f 100644
--- a/tests/sh/test_previous_torch_pin.sh
+++ b/tests/sh/test_previous_torch_pin.sh
@@ -2,9 +2,12 @@
# 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 install.sh's _previous_torch_pin, which keeps the previous
-# venv's torch release on a re-run (curl | sh over an existing install) instead
-# of silently moving the user to a newer release. Helpers are extracted from
-# install.sh and sourced.
+# venv's torch RELEASE on a re-run instead of moving the user to a newer one.
+# The release is kept regardless of the old build's flavor tag (PyPI bare,
+# +cuXXX, +rocm, +cpu): the pin installs from the freshly chosen index, so the
+# flavor follows the machine while the release follows the user. Per-leaf
+# windows still win (rocm7.2 / Strix floors, out-of-window manual installs).
+# Helpers are extracted from install.sh and sourced.
set -e
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
@@ -12,12 +15,9 @@ INSTALL_SH="$SCRIPT_DIR/../../install.sh"
PASS=0
FAIL=0
-# Extract _previous_torch_pin and its dependencies _torch_flavor_tag and
-# _torch_release_in_window.
+# Extract _previous_torch_pin and its dependency _torch_release_in_window.
_FUNC_FILE=$(mktemp)
{
- sed -n '/^_torch_flavor_tag()/,/^}/p' "$INSTALL_SH"
- echo ""
sed -n '/^_torch_release_in_window()/,/^}/p' "$INSTALL_SH"
echo ""
sed -n '/^_previous_torch_pin()/,/^}/p' "$INSTALL_SH"
@@ -37,37 +37,43 @@ assert_eq() {
unset UNSLOTH_TORCH_UPGRADE
-echo "=== _previous_torch_pin: matching flavor keeps the release ==="
-assert_eq "cu126 wheel on cu126 leaf" "torch==2.10.0" "$(_previous_torch_pin '2.10.0+cu126' 'cu126' 'torch>=2.4,<2.12.0')"
-assert_eq "cu130 wheel on cu130 leaf" "torch==2.10.0" "$(_previous_torch_pin '2.10.0+cu130' 'cu130' 'torch>=2.4,<2.12.0')"
-assert_eq "cpu wheel on cpu leaf" "torch==2.10.0" "$(_previous_torch_pin '2.10.0+cpu' 'cpu' 'torch>=2.4,<2.12.0')"
-assert_eq "untagged wheel on cpu leaf" "torch==2.10.0" "$(_previous_torch_pin '2.10.0' 'cpu' 'torch>=2.4,<2.12.0')"
-assert_eq "local suffix stripped" "torch==2.9.1" "$(_previous_torch_pin '2.9.1+cu128' 'cu128' 'torch>=2.4,<2.12.0')"
+echo "=== _previous_torch_pin: in-window releases are kept, any flavor ==="
+assert_eq "cu126 wheel" "torch==2.10.0" "$(_previous_torch_pin '2.10.0+cu126' 'torch>=2.4,<2.12.0')"
+assert_eq "cu130 wheel" "torch==2.10.0" "$(_previous_torch_pin '2.10.0+cu130' 'torch>=2.4,<2.12.0')"
+assert_eq "cpu wheel" "torch==2.10.0" "$(_previous_torch_pin '2.10.0+cpu' 'torch>=2.4,<2.12.0')"
+assert_eq "PyPI bare version (CUDA build on Linux)" "torch==2.10.0" "$(_previous_torch_pin '2.10.0' 'torch>=2.4,<2.12.0')"
+assert_eq "rocm wheel" "torch==2.10.0" "$(_previous_torch_pin '2.10.0+rocm6.4' 'torch>=2.4,<2.11.0')"
+assert_eq "rocm three-component tag" "torch==2.9.1" "$(_previous_torch_pin '2.9.1+rocm7.2.1' 'torch>=2.4,<2.12.0')"
+assert_eq "Intel xpu wheel" "torch==2.9.0" "$(_previous_torch_pin '2.9.0+xpu' 'torch>=2.4,<2.12.0')"
+assert_eq "local suffix stripped" "torch==2.9.1" "$(_previous_torch_pin '2.9.1+cu128' 'torch>=2.4,<2.12.0')"
-echo "=== _previous_torch_pin: flavor change installs the new build ==="
-assert_eq "cu126 wheel on cu130 leaf" "" "$(_previous_torch_pin '2.10.0+cu126' 'cu130' 'torch>=2.4,<2.12.0')"
-assert_eq "cpu wheel on cu126 leaf" "" "$(_previous_torch_pin '2.10.0+cpu' 'cu126' 'torch>=2.4,<2.12.0')"
-assert_eq "cu126 wheel on cpu leaf" "" "$(_previous_torch_pin '2.10.0+cu126' 'cpu' 'torch>=2.4,<2.12.0')"
-
-echo "=== _previous_torch_pin: rocm and unknown leaves never pin ==="
-assert_eq "rocm7.2 leaf keeps its floor" "" "$(_previous_torch_pin '2.11.0+rocm7.2' 'rocm7.2' 'torch>=2.4,<2.12.0')"
-assert_eq "gfx leaf keeps its floor" "" "$(_previous_torch_pin '2.11.0+rocm7.2' 'gfx120X-all' 'torch>=2.4,<2.12.0')"
-assert_eq "unknown mirror leaf" "" "$(_previous_torch_pin '2.10.0+cu126' 'simple' 'torch>=2.4,<2.12.0')"
+echo "=== _previous_torch_pin: raised floors reject older releases ==="
+# rocm7.2 / Strix gfx leaves raise TORCH_CONSTRAINT to >=2.11.0 BEFORE the pin
+# is evaluated, so an old 2.10 is out of window there and the floor wins.
+assert_eq "old 2.10 vs rocm7.2 floor" "" "$(_previous_torch_pin '2.10.0+rocm7.1' 'torch>=2.11.0,<2.12.0')"
+assert_eq "2.11 passes the rocm7.2 floor" "torch==2.11.0" "$(_previous_torch_pin '2.11.0+rocm7.2' 'torch>=2.11.0,<2.12.0')"
echo "=== _previous_torch_pin: probe noise never becomes a pin ==="
-assert_eq "empty version" "" "$(_previous_torch_pin '' 'cu126' 'torch>=2.4,<2.12.0')"
-assert_eq "garbage version" "" "$(_previous_torch_pin 'not-a-version' 'cpu' 'torch>=2.4,<2.12.0')"
-assert_eq "traceback fragment" "" "$(_previous_torch_pin "ModuleNotFoundError: No module named 'torch'" 'cpu' 'torch>=2.4,<2.12.0')"
+assert_eq "empty version" "" "$(_previous_torch_pin '' 'torch>=2.4,<2.12.0')"
+assert_eq "garbage version" "" "$(_previous_torch_pin 'not-a-version' 'torch>=2.4,<2.12.0')"
+assert_eq "traceback fragment" "" "$(_previous_torch_pin "ModuleNotFoundError: No module named 'torch'" 'torch>=2.4,<2.12.0')"
+
+echo "=== _previous_torch_pin: nightly / dev / source builds never pin ==="
+# No stable index carries these, so pinning would print "keeping it" and then
+# burn a doomed resolve before the range fallback rescues the install.
+assert_eq "nightly dev build" "" "$(_previous_torch_pin '2.11.0.dev20250704+cu128' 'torch>=2.4,<2.12.0')"
+assert_eq "source build a0 tag" "" "$(_previous_torch_pin '2.9.0a0+gitabc1234' 'torch>=2.4,<2.12.0')"
+assert_eq "release candidate" "" "$(_previous_torch_pin '2.11.0rc1+cu130' 'torch>=2.4,<2.12.0')"
echo "=== _previous_torch_pin: out-of-window releases never pin ==="
-assert_eq "2.3.x below the cu floor" "" "$(_previous_torch_pin '2.3.1+cu118' 'cu118' 'torch>=2.4,<2.12.0')"
-assert_eq "2.12.x above the cu ceiling" "" "$(_previous_torch_pin '2.12.0+cu130' 'cu130' 'torch>=2.4,<2.12.0')"
-assert_eq "floor boundary 2.4.0 kept" "torch==2.4.0" "$(_previous_torch_pin '2.4.0+cu126' 'cu126' 'torch>=2.4,<2.12.0')"
-assert_eq "ceiling-adjacent 2.11.x kept" "torch==2.11.1" "$(_previous_torch_pin '2.11.1+cu130' 'cu130' 'torch>=2.4,<2.12.0')"
-assert_eq "cpu window excludes 2.11.x" "" "$(_previous_torch_pin '2.11.0+cpu' 'cpu' 'torch>=2.4,<2.11.0')"
-assert_eq "mac floor excludes 2.5.x" "" "$(_previous_torch_pin '2.5.1' 'cpu' 'torch>=2.6,<2.11.0')"
-assert_eq "malformed window never pins" "" "$(_previous_torch_pin '2.10.0+cu126' 'cu126' 'torch')"
-assert_eq "empty window never pins" "" "$(_previous_torch_pin '2.10.0+cu126' 'cu126' '')"
+assert_eq "2.3.x below the cu floor" "" "$(_previous_torch_pin '2.3.1+cu118' 'torch>=2.4,<2.12.0')"
+assert_eq "2.12.x above the cu ceiling" "" "$(_previous_torch_pin '2.12.0+cu130' 'torch>=2.4,<2.12.0')"
+assert_eq "floor boundary 2.4.0 kept" "torch==2.4.0" "$(_previous_torch_pin '2.4.0+cu126' 'torch>=2.4,<2.12.0')"
+assert_eq "ceiling-adjacent 2.11.x kept" "torch==2.11.1" "$(_previous_torch_pin '2.11.1+cu130' 'torch>=2.4,<2.12.0')"
+assert_eq "cpu window excludes 2.11.x" "" "$(_previous_torch_pin '2.11.0+cpu' 'torch>=2.4,<2.11.0')"
+assert_eq "mac floor excludes 2.5.x" "" "$(_previous_torch_pin '2.5.1' 'torch>=2.6,<2.11.0')"
+assert_eq "malformed window never pins" "" "$(_previous_torch_pin '2.10.0+cu126' 'torch')"
+assert_eq "empty window never pins" "" "$(_previous_torch_pin '2.10.0+cu126' '')"
echo "=== _torch_release_in_window ==="
assert_eq "in window" "yes" "$(_torch_release_in_window '2.10.0' 'torch>=2.4,<2.12.0')"
@@ -80,8 +86,8 @@ assert_eq "no ceiling -> no" "no" "$(_torch_release_in_window '2.10.0' 'tor
assert_eq "garbage minor -> no" "no" "$(_torch_release_in_window '2.x' 'torch>=2.4,<2.12.0')"
echo "=== _previous_torch_pin: UNSLOTH_TORCH_UPGRADE=1 opts out ==="
-assert_eq "upgrade env set" "" "$(UNSLOTH_TORCH_UPGRADE=1 _previous_torch_pin '2.10.0+cu126' 'cu126' 'torch>=2.4,<2.12.0')"
-assert_eq "upgrade env 0" "torch==2.10.0" "$(UNSLOTH_TORCH_UPGRADE=0 _previous_torch_pin '2.10.0+cu126' 'cu126' 'torch>=2.4,<2.12.0')"
+assert_eq "upgrade env set" "" "$(UNSLOTH_TORCH_UPGRADE=1 _previous_torch_pin '2.10.0+cu126' 'torch>=2.4,<2.12.0')"
+assert_eq "upgrade env 0" "torch==2.10.0" "$(UNSLOTH_TORCH_UPGRADE=0 _previous_torch_pin '2.10.0+cu126' 'torch>=2.4,<2.12.0')"
echo "=== install.sh wiring ==="
# The probe must run against the OLD venv, before it is moved aside for rollback.
@@ -89,9 +95,39 @@ _probe_line=$(grep -n '_PREV_TORCH_VER=\$(' "$INSTALL_SH" | head -1 | cut -d: -f
_move_line=$(grep -n '_start_studio_venv_replacement "\$VENV_DIR"' "$INSTALL_SH" | head -1 | cut -d: -f1)
assert_eq "probe exists" "yes" "$([ -n "$_probe_line" ] && echo yes)"
assert_eq "probe before venv replacement" "yes" "$([ -n "$_probe_line" ] && [ -n "$_move_line" ] && [ "$_probe_line" -lt "$_move_line" ] && echo yes)"
+# The pin must be evaluated AFTER the last index/constraint decision (the Strix
+# reroute raises the floor), so a raised floor rejects an older kept release.
+_pin_line=$(grep -n '_prev_pin=\$(_previous_torch_pin' "$INSTALL_SH" | head -1 | cut -d: -f1)
+_strix_line=$(grep -n 'Strix Halo / Strix Point: force rocm7.2 wheels' "$INSTALL_SH" | head -1 | cut -d: -f1)
+assert_eq "pin evaluated after the Strix reroute" "yes" "$([ -n "$_pin_line" ] && [ -n "$_strix_line" ] && [ "$_pin_line" -gt "$_strix_line" ] && echo yes)"
# A kept release that vanished from the index must fall back to the supported range.
assert_eq "resolve-failure fallback wired" "yes" "$(grep -q 'TORCH_CONSTRAINT="\$_PREV_FALLBACK_CONSTRAINT"' "$INSTALL_SH" && echo yes)"
assert_eq "pin gated on SKIP_TORCH" "yes" "$(grep -q 'if \[ "\$SKIP_TORCH" = false \]; then' "$INSTALL_SH" && echo yes)"
+# Every --default-index torch install path must go through the kept-release
+# helper (definition + default path + three ROCm-index fallbacks + two ROCm
+# repairs + the flavor repair), so a pinned release missing from the index
+# never aborts a rerun.
+_helper_uses=$(grep -c '_install_torch_default_index' "$INSTALL_SH")
+assert_eq "kept-release helper used by all default-index paths" "yes" "$([ "$_helper_uses" -ge 8 ] && echo yes)"
+_repair_uses=$(grep -c '_install_torch_default_index --force-reinstall' "$INSTALL_SH")
+assert_eq "ROCm repairs routed through the kept-release helper" "yes" "$([ "$_repair_uses" -ge 2 ] && echo yes)"
+# The wrong-flavor repair must use the helper too (it runs under set -e, so a
+# direct uv call with an unresolvable pin would abort the whole installer).
+assert_eq "flavor repair routed through the kept-release helper" "yes" "$(grep -q '_install_torch_default_index \\' "$INSTALL_SH" && grep -q -- '--reinstall-package torch --reinstall-package torchvision --reinstall-package torchaudio' "$INSTALL_SH" && echo yes)"
+# The kept-release install must pair the companions with the kept minor:
+# torchaudio no longer exact-pins torch, so unconstrained it resolves a newer
+# mismatched build (verified: torch==2.9.0 pulled torchaudio 2.11.0 on cu130).
+assert_eq "kept-release install pairs torchvision/torchaudio to the kept minor" "yes" "$(grep -q 'torchaudio==2.\${_itdi_minor}.\*' "$INSTALL_SH" && grep -q 'torchvision==0.\$((_itdi_minor + 15)).\*' "$INSTALL_SH" && echo yes)"
+# The Radeon direct-wheel path must also honor the pin: an exact-first kept-trio
+# attempt (exact patch, else the kept minor's newest patch, with paired
+# vision/audio) runs BEFORE the newest-trio search, and the newest-trio search
+# only runs when that attempt did not produce a match, so a kept release can
+# neither drift to another patch/minor nor be undercut by the gap search.
+_radeon_kept_line=$(grep -n '_kept_torch=\$(_pick_radeon_wheel "torch" *"\${_prev_kept_base}"' "$INSTALL_SH" | head -1 | cut -d: -f1)
+_radeon_loop_line=$(grep -n 'Loop downwards to find the first complete matching trio' "$INSTALL_SH" | head -1 | cut -d: -f1)
+assert_eq "Radeon kept-trio attempt before the newest-trio search" "yes" "$([ -n "$_radeon_kept_line" ] && [ -n "$_radeon_loop_line" ] && [ "$_radeon_kept_line" -lt "$_radeon_loop_line" ] && echo yes)"
+assert_eq "Radeon newest-trio search gated on no kept match" "yes" "$(grep -q 'if \[ "\$_radeon_versions_match" != true \] &&' "$INSTALL_SH" && echo yes)"
+assert_eq "Radeon kept-trio gap falls back with a warning" "yes" "$(grep -q 'lacks a complete wheel set for kept' "$INSTALL_SH" && echo yes)"
echo ""
if [ "$FAIL" -gt 0 ]; then
From 17fd6c8ec6c3788c1da2f9e86452fc340e451444 Mon Sep 17 00:00:00 2001
From: Daniel Han
Date: Sun, 19 Jul 2026 17:19:15 -0700
Subject: [PATCH 06/41] studio: fix stale GGUF load-marker ordering test after
inheritance relocation (#7252)
#6414 moved the llama_extra_args inheritance out of the GGUF branch in
_load_model_impl into _guard_chat_load_against_training, which runs before the
branch, so 'if request.llama_extra_args is None' is no longer inside the
gguf_branch slice that test_load_marker_precedes_hub_guard_and_unload checks.
The assertion failed on that now-missing landmark even though the guarantee it
protects (the gguf_load_in_flight marker is entered before the hub-download
guard and the unload) is intact. Drop the relocated landmark from the ordering
so the test matches the current structure.
Co-authored-by: danielhanchen
---
studio/backend/tests/test_gguf_load_cache_reuse.py | 6 +++++-
1 file changed, 5 insertions(+), 1 deletion(-)
diff --git a/studio/backend/tests/test_gguf_load_cache_reuse.py b/studio/backend/tests/test_gguf_load_cache_reuse.py
index 15d91cd324..62596fcc8a 100644
--- a/studio/backend/tests/test_gguf_load_cache_reuse.py
+++ b/studio/backend/tests/test_gguf_load_cache_reuse.py
@@ -728,9 +728,13 @@ class TestLoadHubDownloadExclusion:
source = (Path(__file__).resolve().parent.parent / "routes" / "inference.py").read_text()
gguf_branch = source[source.index("if config.is_gguf:") :]
+ # The gguf_load_in_flight marker must be entered before the hub-download
+ # guard and the unload so a concurrent load can't race the download
+ # manager. The llama_extra_args inheritance that used to sit between the
+ # marker and the guard now runs in _guard_chat_load_against_training, ahead
+ # of the GGUF branch, so it is no longer a landmark inside this slice.
assert (
gguf_branch.index("enter_context(gguf_load_in_flight")
- < gguf_branch.index("if request.llama_extra_args is None")
< gguf_branch.index("_hub_download_blocks_gguf_load")
< gguf_branch.index("unsloth_backend.unload_model")
)
From 8fab1c5310e6d4117a939f30c8d1546ffca023bd Mon Sep 17 00:00:00 2001
From: Lee Jackson <130007945+Imagineer99@users.noreply.github.com>
Date: Mon, 20 Jul 2026 05:03:50 +0100
Subject: [PATCH 07/41] Route OpenCode yolo aliases to native auto mode (#7187)
Route --yolo to OpenCode native --auto for the default TUI and run; keep the config permission fallback for no-auto subcommands (including hidden console/generate) and for --mini, which ignores --auto.
---
unsloth_cli/commands/start.py | 115 ++++++++++++++++++--
unsloth_cli/tests/test_start.py | 185 +++++++++++++++++++++++++++++++-
2 files changed, 285 insertions(+), 15 deletions(-)
diff --git a/unsloth_cli/commands/start.py b/unsloth_cli/commands/start.py
index 9257a7fbcb..8128447a02 100644
--- a/unsloth_cli/commands/start.py
+++ b/unsloth_cli/commands/start.py
@@ -149,8 +149,8 @@ _PERSIST_OPTION = typer.Option(
),
)
-# Per-agent CLI flag for "run tools without prompting". opencode and openclaw have no
-# such flag (config only) and are handled in their config writers, so they are absent.
+# Per-agent CLI flag for "run tools without prompting". OpenCode (native --auto is
+# command-scoped, handled below) and OpenClaw (config-only) are absent from this prefix map.
_YOLO_COMMAND_FLAGS = {
"claude": ["--dangerously-skip-permissions"],
"codex": ["--dangerously-bypass-approvals-and-sandbox"],
@@ -166,6 +166,84 @@ def _yolo_command_flags(agent: str, yolo: bool) -> list:
return _YOLO_COMMAND_FLAGS.get(agent, []) if yolo else []
+# Subcommands that reject --auto (OpenCode exposes it only on the default TUI and `run`),
+# so `opencode serve --auto` is never emitted. Includes console/generate, hidden from
+# `opencode --help` but still registered. Unknown first positionals are TUI paths -> --auto.
+_OPENCODE_NON_AUTO_SUBCOMMANDS = frozenset(
+ "completion acp mcp attach debug providers auth agent upgrade uninstall serve web "
+ "models stats export import github pr session plugin plug db console generate".split()
+)
+_OPENCODE_GLOBAL_BOOLEAN_OPTIONS = frozenset(
+ "-h --help -v --version --print-logs --pure --mdns".split()
+)
+_OPENCODE_GLOBAL_VALUE_OPTIONS = frozenset(
+ "--log-level --port --hostname --mdns-domain --cors".split()
+)
+_OPENCODE_NATIVE_AUTO_MIN_VERSION = (1, 17, 12)
+
+
+def _opencode_supports_native_auto() -> bool:
+ executable = shutil.which("opencode")
+ if executable is None:
+ # No local binary: a --no-launch recipe may run elsewhere, and _run installs the
+ # current release on launch -- either way assume native --auto is available.
+ return True
+ try:
+ output = subprocess.check_output(
+ [executable, "--version"],
+ text = True,
+ timeout = 10,
+ stderr = subprocess.DEVNULL,
+ )
+ except Exception:
+ return False
+ match = re.search(r"(\d+)\.(\d+)\.(\d+)", output)
+ return bool(match) and tuple(int(part) for part in match.groups()) >= (
+ _OPENCODE_NATIVE_AUTO_MIN_VERSION
+ )
+
+
+def _opencode_subcommand(args: list[str]) -> Optional[str]:
+ """Return an explicit OpenCode subcommand after supported global options."""
+ index = 0
+ while index < len(args):
+ arg = args[index]
+ if arg == "--":
+ return None
+ if arg in _OPENCODE_GLOBAL_BOOLEAN_OPTIONS:
+ index += 1
+ continue
+ if arg in _OPENCODE_GLOBAL_VALUE_OPTIONS:
+ index += 2
+ continue
+ if any(arg.startswith(f"{option}=") for option in _OPENCODE_GLOBAL_VALUE_OPTIONS):
+ index += 1
+ continue
+ # A non-global option (e.g. --session) is a TUI flag; stop before its value is
+ # mistaken for a subcommand.
+ if arg.startswith("-"):
+ return None
+ return arg
+ return None
+
+
+def _opencode_native_auto_args(args: list[str], yolo: bool) -> tuple[list[str], bool]:
+ """Add OpenCode's native --auto when the selected command supports it."""
+ routed = list(args)
+ if not yolo:
+ return routed, False
+ if _opencode_subcommand(routed) in _OPENCODE_NON_AUTO_SUBCOMMANDS:
+ return routed, False
+ separator = routed.index("--") if "--" in routed else len(routed)
+ # --mini's runMini TUI forces auto=false and never forwards --auto, so appending it is
+ # useless; fall back to the config permission block so --yolo still auto-approves.
+ if any(arg == "--mini" or arg.startswith("--mini=") for arg in routed[:separator]):
+ return routed, False
+ if "--auto" not in routed[:separator]:
+ routed.insert(separator, "--auto")
+ return routed, True
+
+
def _hermes_install_hint() -> str:
return _HERMES_WINDOWS_INSTALL_HINT if os.name == "nt" else _HERMES_POSIX_INSTALL_HINT
@@ -1465,10 +1543,10 @@ def write_opencode_config(
compaction["reserved"] = max(1, window // 10)
tools = ("edit", "bash", "webfetch")
if yolo:
- # OpenCode has no --yolo flag; auto-approve is the config `permission` block
- # (singular). Allow the prompting tools and paths outside the launch directory so
- # tool calls don't block on the TUI. This rides inline (OPENCODE_CONFIG_CONTENT) so
- # --yolo works even over a project config.
+ # Fallback for commands without native --auto and for the append-safe bare
+ # --no-launch command (subcommand unknown yet). Rides inline (OPENCODE_CONFIG_CONTENT)
+ # so it wins over a project config. TUI and `run` launches use --auto and call here
+ # with yolo=False, letting OpenCode preserve explicit deny rules.
session_permission = {t: "allow" for t in tools}
session_permission["external_directory"] = {"*": "allow"}
config["permission"] = dict(session_permission)
@@ -1807,11 +1885,20 @@ def opencode(
# --no-launch, where the printed command is consumed by drivers that append a
# subcommand such as `run `; a leading --model would land before that
# subcommand and break it. Those paths rely on the inline pin instead.
+ native_auto = False
+ route_native_auto = yolo and _opencode_supports_native_auto()
if ctx.args:
- command = ["opencode", *ctx.args]
+ opencode_args, native_auto = _opencode_native_auto_args(list(ctx.args), route_native_auto)
+ command = ["opencode", *opencode_args]
elif launch:
- command = ["opencode", "--model", opencode_model]
+ opencode_args, native_auto = _opencode_native_auto_args(
+ ["--model", opencode_model],
+ route_native_auto,
+ )
+ command = ["opencode", *opencode_args]
else:
+ # Append-safe base: `opencode --auto run ...` parses as the TUI with a project
+ # "run", not the run subcommand. Command unknown here, so keep the config fallback.
command = ["opencode"]
# opencode keeps sessions in ~/.local/share/opencode (never relocated), so resume
# already survives exit; reopen the last one by passing `opencode --continue` through.
@@ -1820,12 +1907,18 @@ def opencode(
# OPENCODE_CONFIG is an overlay (loaded between the user's global and project
# configs), so this adds the Unsloth provider/model for the session without
# changing the user's default model. Key lives in the config, not the env.
- session_permission = write_opencode_config(base, key, entry, config_path, yolo = yolo)
+ session_permission = write_opencode_config(
+ base,
+ key,
+ entry,
+ config_path,
+ yolo = yolo and not native_auto,
+ )
# A project's own opencode.json outranks OPENCODE_CONFIG, so the session model pin
# would silently lose to a repo config. Carry it in OPENCODE_CONFIG_CONTENT, which
# outranks project config; the API key stays in the private file, never the env.
- # Only --yolo carries a permission here (its allow must win over a project config);
- # a non-yolo session returns no permission, so the project's own rules are honored.
+ # Only the config fallback carries a permission. Native --auto omits it (auto-approve
+ # asks, keep explicit denies); a non-yolo session omits it too, honoring project rules.
# opencode filters every provider (a config-defined custom one included) through
# its enabled_providers allowlist and disabled_providers denylist, and a model pin
# does not bypass that gate -- a filtered provider resolves to ModelNotFoundError.
diff --git a/unsloth_cli/tests/test_start.py b/unsloth_cli/tests/test_start.py
index 2405ba0480..5b2806be12 100644
--- a/unsloth_cli/tests/test_start.py
+++ b/unsloth_cli/tests/test_start.py
@@ -2190,8 +2190,18 @@ def test_yolo_aliases_are_interchangeable(fake_studio, alias):
assert "--dangerously-bypass-approvals-and-sandbox" in codex.output
assert "--dangerously-skip-permissions" not in codex.output
+ opencode = CliRunner().invoke(
+ start.start_app,
+ ["opencode", alias, "--no-launch", "run", "hello"],
+ )
+ assert opencode.exit_code == 0, opencode.output
+ assert _launch_command(opencode.output) == ["opencode", "run", "hello", "--auto"]
+ assert "permission" not in _opencode_inline_config(opencode.output)
-def test_yolo_opencode_writes_permission_block(fake_studio, tmp_path):
+
+def test_yolo_opencode_bare_no_launch_uses_permission_fallback(fake_studio, tmp_path):
+ # A bare --no-launch recipe stays append-safe (callers add a subcommand later);
+ # `opencode --auto run ...` would select the TUI, not `run`, so keep the config fallback.
result = CliRunner().invoke(start.start_app, ["opencode", "--yolo", "--no-launch"])
assert result.exit_code == 0, result.output
config = json.loads((tmp_path / "agents" / "opencode" / "opencode.json").read_text())
@@ -2203,6 +2213,172 @@ def test_yolo_opencode_writes_permission_block(fake_studio, tmp_path):
}
+def test_yolo_opencode_run_uses_native_auto(fake_studio):
+ result = CliRunner().invoke(
+ start.start_app,
+ ["opencode", "--yolo", "--no-launch", "run", "hello"],
+ )
+ assert result.exit_code == 0, result.output
+ command = _launch_command(result.output)
+ assert command == ["opencode", "run", "hello", "--auto"]
+ assert "permission" not in _opencode_inline_config(result.output)
+
+
+def test_yolo_opencode_tui_resume_uses_native_auto(fake_studio):
+ result = CliRunner().invoke(
+ start.start_app,
+ ["opencode", "--yolo", "--no-launch", "--session", "sid"],
+ )
+ assert result.exit_code == 0, result.output
+ command = _launch_command(result.output)
+ assert command == ["opencode", "--session", "sid", "--auto"]
+ assert "permission" not in _opencode_inline_config(result.output)
+
+
+def test_no_yolo_opencode_run_omits_native_auto(fake_studio):
+ result = CliRunner().invoke(
+ start.start_app,
+ ["opencode", "--no-launch", "run", "hello"],
+ )
+ assert result.exit_code == 0, result.output
+ assert _launch_command(result.output) == ["opencode", "run", "hello"]
+ assert "permission" not in _opencode_inline_config(result.output)
+
+
+def test_yolo_opencode_bare_launch_uses_native_auto(fake_studio, monkeypatch):
+ monkeypatch.setattr(start.shutil, "which", lambda _: "/usr/local/bin/opencode")
+ monkeypatch.setattr(start, "_opencode_supports_native_auto", lambda: True)
+ captured = _capture_launch(monkeypatch, ["opencode", "--yolo"])
+ assert captured["command"][1:] == [
+ "--model",
+ f"{start._OPENCODE_PROVIDER}/{MODEL['id']}",
+ "--auto",
+ ]
+ assert "permission" not in json.loads(captured["env"]["OPENCODE_CONFIG_CONTENT"])
+
+
+def test_yolo_opencode_native_auto_clears_prior_config_fallback(fake_studio, tmp_path):
+ fallback = CliRunner().invoke(
+ start.start_app,
+ ["opencode", "--yolo", "--no-launch"],
+ )
+ assert fallback.exit_code == 0, fallback.output
+
+ native = CliRunner().invoke(
+ start.start_app,
+ ["opencode", "--yolo", "--no-launch", "run", "hello"],
+ )
+ assert native.exit_code == 0, native.output
+ assert _launch_command(native.output) == ["opencode", "run", "hello", "--auto"]
+ assert "permission" not in _opencode_inline_config(native.output)
+ config = json.loads((tmp_path / "agents" / "opencode" / "opencode.json").read_text())
+ assert config["permission"] == {
+ "edit": "ask",
+ "bash": "ask",
+ "webfetch": "ask",
+ "external_directory": {"*": "ask"},
+ }
+
+
+@pytest.mark.parametrize(
+ ("version", "expected"),
+ [
+ ("1.17.11", False),
+ ("1.17.12", True),
+ ("opencode 1.18.2", True),
+ ("development build", False),
+ ],
+)
+def test_opencode_native_auto_version_gate(monkeypatch, version, expected):
+ monkeypatch.setattr(start.shutil, "which", lambda _: "/usr/local/bin/opencode")
+ monkeypatch.setattr(start.subprocess, "check_output", lambda *args, **kwargs: version)
+ assert start._opencode_supports_native_auto() is expected
+
+
+def test_opencode_native_auto_assumes_current_without_local_binary(monkeypatch):
+ monkeypatch.setattr(start.shutil, "which", lambda _: None)
+ assert start._opencode_supports_native_auto() is True
+
+
+def test_yolo_opencode_old_version_uses_config_fallback(fake_studio, monkeypatch):
+ monkeypatch.setattr(start.shutil, "which", lambda _: "/usr/local/bin/opencode")
+ monkeypatch.setattr(start.subprocess, "check_output", lambda *args, **kwargs: "1.17.11")
+ result = CliRunner().invoke(
+ start.start_app,
+ ["opencode", "--yolo", "--no-launch", "run", "hello"],
+ )
+ assert result.exit_code == 0, result.output
+ assert _launch_command(result.output) == ["opencode", "run", "hello"]
+ assert _opencode_inline_config(result.output)["permission"] == {
+ "edit": "allow",
+ "bash": "allow",
+ "webfetch": "allow",
+ "external_directory": {"*": "allow"},
+ }
+
+
+@pytest.mark.parametrize(
+ ("args", "expected", "native"),
+ [
+ ([], ["--auto"], True),
+ (["run", "hello"], ["run", "hello", "--auto"], True),
+ (
+ ["run", "hello", "--", "--literal"],
+ ["run", "hello", "--auto", "--", "--literal"],
+ True,
+ ),
+ (["--print-logs", "run", "hello"], ["--print-logs", "run", "hello", "--auto"], True),
+ (["--session", "serve"], ["--session", "serve", "--auto"], True),
+ (["serve"], ["serve"], False),
+ (["--print-logs", "serve"], ["--print-logs", "serve"], False),
+ (["run", "--auto", "hello"], ["run", "--auto", "hello"], True),
+ # Hidden commands that reject --auto fall back like the visible utility ones.
+ (["generate"], ["generate"], False),
+ (["console", "login"], ["console", "login"], False),
+ # --mini ignores --auto (runMini forces auto=false), so use the config fallback.
+ (["--mini"], ["--mini"], False),
+ (["--session", "sid", "--mini"], ["--session", "sid", "--mini"], False),
+ ],
+)
+def test_opencode_native_auto_args(args, expected, native):
+ assert start._opencode_native_auto_args(args, True) == (expected, native)
+ assert start._opencode_native_auto_args(args, False) == (args, False)
+
+
+def test_yolo_opencode_non_agent_subcommand_uses_config_fallback(fake_studio):
+ result = CliRunner().invoke(
+ start.start_app,
+ ["opencode", "--yolo", "--no-launch", "serve"],
+ )
+ assert result.exit_code == 0, result.output
+ command = _launch_command(result.output)
+ assert command == ["opencode", "serve"]
+ assert _opencode_inline_config(result.output)["permission"] == {
+ "edit": "allow",
+ "bash": "allow",
+ "webfetch": "allow",
+ "external_directory": {"*": "allow"},
+ }
+
+
+@pytest.mark.parametrize("passthrough", (["generate"], ["console", "login"], ["--mini"]))
+def test_yolo_opencode_no_auto_command_uses_config_fallback(fake_studio, passthrough):
+ # generate/console are hidden and reject --auto, --mini ignores it: none get --auto,
+ # all keep the config permission fallback.
+ result = CliRunner().invoke(
+ start.start_app,
+ ["opencode", "--yolo", "--no-launch", *passthrough],
+ )
+ assert result.exit_code == 0, result.output
+ assert _launch_command(result.output) == ["opencode", *passthrough]
+ assert _opencode_inline_config(result.output)["permission"] == {
+ "edit": "allow",
+ "bash": "allow",
+ "webfetch": "allow",
+ "external_directory": {"*": "allow"},
+ }
+
+
def test_no_yolo_opencode_has_no_permission_block(fake_studio, tmp_path):
result = CliRunner().invoke(start.start_app, ["opencode", "--no-launch"])
assert result.exit_code == 0, result.output
@@ -2549,15 +2725,16 @@ def test_openclaw_non_yolo_preserves_full_mode(tmp_path):
def test_yolo_command_flags_unmapped_agent_is_empty():
- # Config-based agents (and any typo) must yield no flag, not a KeyError.
+ # Placement-aware/config-based agents (and any typo) must yield no prefix flag.
assert start._yolo_command_flags("opencode", True) == []
assert start._yolo_command_flags("openclaw", True) == []
assert start._yolo_command_flags("claude", True) == ["--dangerously-skip-permissions"]
assert start._yolo_command_flags("claude", False) == []
-def test_yolo_config_agents_add_no_command_flag(fake_studio):
- # opencode/openclaw auto-approve is config-only; nothing should leak onto argv.
+def test_yolo_config_fallbacks_add_no_legacy_command_flag(fake_studio):
+ # OpenClaw is config-only; OpenCode's append-safe bare recipe uses its config fallback.
+ # Neither should leak a legacy yolo/dangerous alias onto argv.
for agent in ("opencode", "openclaw"):
result = CliRunner().invoke(start.start_app, [agent, "--yolo", "--no-launch"])
assert result.exit_code == 0, result.output
From e0132b6d6c414cece2bced7eaf164eeebe088dd1 Mon Sep 17 00:00:00 2001
From: Lee Jackson <130007945+Imagineer99@users.noreply.github.com>
Date: Mon, 20 Jul 2026 05:04:28 +0100
Subject: [PATCH 08/41] Pin the Hermes remote installer and harden consent
(#7179)
Pin the fetched Hermes install.sh/install.ps1 and the checkout they perform to an immutable upstream commit, and distinguish pinned from unpinned sources in the consent warning.
---
unsloth_cli/commands/start.py | 54 +++++++++++++++++++++++++++------
unsloth_cli/tests/test_start.py | 38 ++++++++++++++++++-----
2 files changed, 75 insertions(+), 17 deletions(-)
diff --git a/unsloth_cli/commands/start.py b/unsloth_cli/commands/start.py
index 8128447a02..6da31229f8 100644
--- a/unsloth_cli/commands/start.py
+++ b/unsloth_cli/commands/start.py
@@ -49,13 +49,22 @@ _HERMES_PROVIDER = "unsloth"
# the wizard's global API-key/model prompts would block the launch and point the
# user at a different (global) provider than the one Unsloth just configured.
# Both installers expose a skip flag: `-SkipSetup` (PowerShell) and
-# `--skip-setup` (POSIX; passed to the piped script via `bash -s --`).
+# `--skip-setup` (POSIX; passed to the piped script via `bash -s --`). Pin both
+# the fetched script and the repository checkout it performs to the same full
+# commit so a later change to either upstream branch cannot silently replace
+# code that Unsloth executes with the user's privileges.
+_HERMES_INSTALL_COMMIT = "f1af945f6c576eccb126fa955edc9be258b33020"
+_HERMES_INSTALL_BASE = (
+ "https://raw.githubusercontent.com/NousResearch/hermes-agent/"
+ f"{_HERMES_INSTALL_COMMIT}/scripts"
+)
_HERMES_WINDOWS_INSTALL_HINT = (
- "& ([scriptblock]::Create((irm https://hermes-agent.nousresearch.com/install.ps1))) -SkipSetup"
+ f"& ([scriptblock]::Create((irm {_HERMES_INSTALL_BASE}/install.ps1)))"
+ f" -SkipSetup -Commit {_HERMES_INSTALL_COMMIT}"
)
_HERMES_POSIX_INSTALL_HINT = (
- "curl -fsSL https://raw.githubusercontent.com/NousResearch/hermes-agent"
- "/main/scripts/install.sh | bash -s -- --skip-setup"
+ f"curl -fsSL {_HERMES_INSTALL_BASE}/install.sh | bash -s --"
+ f" --skip-setup --commit {_HERMES_INSTALL_COMMIT}"
)
# Hermes refuses to initialize when the model window is under 64,000 tokens; its
# error message points at the model.context_length / auxiliary.compression
@@ -1199,6 +1208,16 @@ def _install_source(install_hint: str) -> Optional[str]:
return match.group(0) if match else None
+def _pinned_raw_github_commit(source: str) -> Optional[str]:
+ """Return the immutable full commit in a raw GitHub URL, if present."""
+ match = re.match(
+ r"^https://raw\.githubusercontent\.com/[^/]+/[^/]+/([0-9a-f]{40})/",
+ source,
+ flags = re.IGNORECASE,
+ )
+ return match.group(1).lower() if match else None
+
+
def _install_agent(name: str, install_hint: str) -> Optional[str]:
# Missing agent under --launch: offer to run its documented install command, then
# re-resolve it on PATH. Consent-based (we never auto-run a remote install script
@@ -1212,12 +1231,27 @@ def _install_agent(name: str, install_hint: str) -> Optional[str]:
# and nothing checks a signature or hash on the fetched content. Naming the source
# turns a blind "yes" into informed consent.
source = _install_source(install_hint)
- warning = (
- f"This will download and RUN a script from {source} with your privileges"
- if source
- else f"This will RUN `{install_hint}` with your privileges"
- )
- typer.secho(f"{warning}; there is no signature or hash check.", fg = "yellow", err = True)
+ if source:
+ pinned_commit = _pinned_raw_github_commit(source)
+ if pinned_commit:
+ warning = (
+ "Security warning: This will download and execute a third-party script "
+ f"from {source} with your privileges. Unsloth pins this content to "
+ f"immutable upstream commit {pinned_commit}, but does not independently "
+ "verify or sandbox it. Continue only if you trust this source and commit."
+ )
+ else:
+ warning = (
+ "Security warning: This will download and execute an unverified third-party "
+ f"script from {source} with your privileges. Unsloth does not pin or verify "
+ "the downloaded content. Continue only if you trust this source."
+ )
+ else:
+ warning = (
+ f"This will RUN `{install_hint}` with your privileges; "
+ "there is no signature or hash check."
+ )
+ typer.secho(warning, fg = "yellow", err = True)
if not typer.confirm(f"Install `{name}` now with `{install_hint}`?", default = False):
return None
# Run each hint through the shell it is written for: PowerShell (irm | iex, or npm)
diff --git a/unsloth_cli/tests/test_start.py b/unsloth_cli/tests/test_start.py
index 5b2806be12..98bd9f8157 100644
--- a/unsloth_cli/tests/test_start.py
+++ b/unsloth_cli/tests/test_start.py
@@ -7,6 +7,7 @@ from __future__ import annotations
import json
import os
+import re
import shlex
import sys
import urllib.error
@@ -128,7 +129,7 @@ def test_install_agent_uses_powershell_on_windows(monkeypatch):
assert ran == [["powershell", "-NoProfile", "-Command", install_hint]]
-def test_install_agent_warns_and_names_remote_source(monkeypatch, capsys):
+def test_install_agent_warns_remote_installer_is_unverified_third_party(monkeypatch, capsys):
# Before the confirm, a remote installer must name the URL it fetches so the
# user consents to a specific source rather than blindly accepting.
monkeypatch.setattr(start.os, "name", "nt")
@@ -137,9 +138,23 @@ def test_install_agent_warns_and_names_remote_source(monkeypatch, capsys):
hint = "& ([scriptblock]::Create((irm https://hermes-agent.nousresearch.com/install.ps1))) -SkipSetup"
assert start._install_agent("hermes", hint) is None
err = capsys.readouterr().err
+ assert "Security warning" in err
+ assert "unverified third-party script" in err
assert "https://hermes-agent.nousresearch.com/install.ps1" in err
- assert "download and RUN" in err
- assert "signature or hash" in err
+ assert "Unsloth does not pin or verify the downloaded content" in err
+ assert "Continue only if you trust this source" in err
+
+
+def test_install_agent_reports_immutable_remote_installer_pin(monkeypatch, capsys):
+ monkeypatch.setattr(start.os, "name", "posix")
+ monkeypatch.setattr(start.sys, "stdin", SimpleNamespace(isatty = lambda: True))
+ monkeypatch.setattr(start.typer, "confirm", lambda *a, **k: False)
+ assert start._install_agent("hermes", start._HERMES_POSIX_INSTALL_HINT) is None
+ err = capsys.readouterr().err
+ assert start._HERMES_INSTALL_COMMIT in err
+ assert "immutable upstream commit" in err
+ assert "does not independently verify or sandbox it" in err
+ assert "does not pin or verify" not in err
def test_install_agent_warns_for_package_installer(monkeypatch, capsys):
@@ -160,8 +175,8 @@ def test_hermes_install_hint_is_windows_native_on_windows(monkeypatch):
# Scriptblock form so `-SkipSetup` reaches the installer and the interactive
# setup wizard is skipped during the unattended `unsloth start hermes` run.
assert start._hermes_install_hint() == (
- "& ([scriptblock]::Create((irm https://hermes-agent.nousresearch.com/install.ps1)))"
- " -SkipSetup"
+ f"& ([scriptblock]::Create((irm {start._HERMES_INSTALL_BASE}/install.ps1)))"
+ f" -SkipSetup -Commit {start._HERMES_INSTALL_COMMIT}"
)
@@ -170,11 +185,20 @@ def test_hermes_install_hint_is_bash_on_posix(monkeypatch):
# `bash -s -- --skip-setup` forwards the skip flag to the piped installer.
assert start._hermes_install_hint() == (
- "curl -fsSL https://raw.githubusercontent.com/NousResearch/hermes-agent"
- "/main/scripts/install.sh | bash -s -- --skip-setup"
+ f"curl -fsSL {start._HERMES_INSTALL_BASE}/install.sh | bash -s --"
+ f" --skip-setup --commit {start._HERMES_INSTALL_COMMIT}"
)
+def test_hermes_install_hints_pin_script_and_checkout_to_full_commit():
+ commit = start._HERMES_INSTALL_COMMIT
+ assert re.fullmatch(r"[0-9a-f]{40}", commit)
+ for hint in (start._HERMES_WINDOWS_INSTALL_HINT, start._HERMES_POSIX_INSTALL_HINT):
+ assert hint.count(commit) == 2
+ assert "/main/" not in hint
+ assert "hermes-agent.nousresearch.com" not in hint
+
+
def test_refresh_windows_path_noop_off_windows(monkeypatch):
monkeypatch.setattr(start.os, "name", "posix")
before = os.environ.get("PATH", "")
From 39497e6516bdc7d7edc2b09493b222c1dc2ce49c Mon Sep 17 00:00:00 2001
From: Lee Jackson <130007945+Imagineer99@users.noreply.github.com>
Date: Mon, 20 Jul 2026 05:05:22 +0100
Subject: [PATCH 09/41] Translate PWD for WSL-launched Windows agents (#7111)
Bridge PWD through WSLENV /p when launching a Windows npm shim from WSL so project-root discovery uses the live cwd. The no-launch recipe adds PWD/p without freezing PWD; the concrete cwd override applies only on direct launch.
---
unsloth_cli/commands/start.py | 16 ++++++++++++++--
unsloth_cli/tests/test_start.py | 16 ++++++++++++++--
2 files changed, 28 insertions(+), 4 deletions(-)
diff --git a/unsloth_cli/commands/start.py b/unsloth_cli/commands/start.py
index 6da31229f8..707c2b3c90 100644
--- a/unsloth_cli/commands/start.py
+++ b/unsloth_cli/commands/start.py
@@ -1274,6 +1274,16 @@ def _install_agent(name: str, install_hint: str) -> Optional[str]:
return executable
+def _wsl_shim_env(command: list, env: dict, unset_env: tuple) -> tuple[dict, tuple]:
+ wsl_env_bridge = _wsl_bridge_names(env, unset_env) if _wsl_windows_executable(command) else ()
+ if not wsl_env_bridge:
+ return env, wsl_env_bridge
+ # Bridge PWD via WSLENV (PWD/p) so the Windows shim finds its project root from the
+ # live cwd, not a stale inherited Linux PWD. Don't freeze env["PWD"]: a --no-launch
+ # recipe must translate the live PWD when run, not when generated; _launch overrides it.
+ return env, (*wsl_env_bridge, "PWD/p")
+
+
def _launch(
command: list,
env: dict,
@@ -1283,9 +1293,11 @@ def _launch(
executable = shutil.which(command[0]) or _install_agent(command[0], install_hint)
if executable is None:
_fail(f"`{command[0]}` not found on PATH. Install it with: {install_hint}")
- wsl_env_bridge = _wsl_bridge_names(env, unset_env) if _wsl_windows_executable(command) else ()
+ env, wsl_env_bridge = _wsl_shim_env(command, env, unset_env)
child_env = dict(os.environ)
if wsl_env_bridge:
+ # Override stale inherited PWD with the real cwd so the shim resolves the project root.
+ env = {**env, "PWD": os.getcwd()}
child_env["WSLENV"] = _merge_wslenv(child_env.get("WSLENV", ""), wsl_env_bridge)
for name in unset_env:
child_env[name] = ""
@@ -1353,8 +1365,8 @@ def _run(
if launch and clear_screen:
click.clear()
typer.echo(f"Unsloth {base} · model {entry['id']}")
- wsl_env_bridge = _wsl_bridge_names(env, unset_env) if _wsl_windows_executable(command) else ()
if not launch:
+ env, wsl_env_bridge = _wsl_shim_env(command, env, unset_env)
_print_env(env, command, unset_env = unset_env, wsl_env_bridge = wsl_env_bridge)
return
try:
diff --git a/unsloth_cli/tests/test_start.py b/unsloth_cli/tests/test_start.py
index 98bd9f8157..227918f63e 100644
--- a/unsloth_cli/tests/test_start.py
+++ b/unsloth_cli/tests/test_start.py
@@ -476,8 +476,10 @@ def test_connect_claude_launch_scrubs_conflicting_auth_env(fake_studio, monkeypa
reason = "WSL-from-Linux scenario (calling a Windows agent .exe from inside WSL); "
"os.name is 'posix' under WSL, so this path can't run on a native Windows runner.",
)
-def test_connect_claude_windows_shim_from_wsl_bridges_env(fake_studio, monkeypatch):
+def test_connect_claude_windows_shim_from_wsl_bridges_env(fake_studio, monkeypatch, tmp_path):
captured = {}
+ monkeypatch.chdir(tmp_path)
+ monkeypatch.setenv("PWD", "/stale/outer/repo")
monkeypatch.setenv("WSL_DISTRO_NAME", "Ubuntu")
monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-anthropic-stale")
monkeypatch.setenv("CLAUDE_CODE_OAUTH_TOKEN", "oauth-stale")
@@ -505,6 +507,8 @@ def test_connect_claude_windows_shim_from_wsl_bridges_env(fake_studio, monkeypat
assert captured["env"]["ANTHROPIC_AUTH_TOKEN"] == "sk-unsloth-feedfacefeedface"
assert captured["env"]["ANTHROPIC_BASE_URL"] == BASE
assert captured["env"]["ANTHROPIC_MODEL"] == MODEL["id"]
+ assert captured["env"]["PWD"] == str(tmp_path)
+ assert "PWD/p" in captured["env"]["WSLENV"].split(":")
for name in (
"ANTHROPIC_AUTH_TOKEN",
"ANTHROPIC_BASE_URL",
@@ -520,7 +524,11 @@ def test_connect_claude_windows_shim_from_wsl_bridges_env(fake_studio, monkeypat
reason = "WSL-from-Linux scenario (calling a Windows agent .exe from inside WSL); "
"os.name is 'posix' under WSL, so this path can't run on a native Windows runner.",
)
-def test_connect_claude_no_launch_windows_shim_from_wsl_prints_wslenv(fake_studio, monkeypatch):
+def test_connect_claude_no_launch_windows_shim_from_wsl_prints_wslenv(
+ fake_studio, monkeypatch, tmp_path
+):
+ monkeypatch.chdir(tmp_path)
+ monkeypatch.setenv("PWD", "/stale/outer/repo")
monkeypatch.setenv("WSL_DISTRO_NAME", "Ubuntu")
monkeypatch.setattr(
start.shutil, "which", lambda _: "/mnt/c/Users/samle/AppData/Roaming/npm/claude"
@@ -532,6 +540,10 @@ def test_connect_claude_no_launch_windows_shim_from_wsl_prints_wslenv(fake_studi
assert "export ANTHROPIC_API_KEY=" in result.output
assert "export CLAUDE_CODE_OAUTH_TOKEN=" in result.output
assert "export WSLENV=" in result.output
+ # PWD must NOT be frozen into the recipe (no `export PWD=`): WSLENV PWD/p translates the
+ # shell's live PWD at run time, so a recipe reused from another dir resolves the project root.
+ assert "export PWD=" not in result.output
+ assert "PWD/p" in result.output
assert "ANTHROPIC_AUTH_TOKEN" in result.output
assert "CLAUDE_CODE_OAUTH_TOKEN" in result.output
From 95d9970233ff3f248c9a5f89a987084e088623e9 Mon Sep 17 00:00:00 2001
From: Nilay <118994073+NilayYadav@users.noreply.github.com>
Date: Mon, 20 Jul 2026 12:42:42 +0530
Subject: [PATCH 10/41] persist llama.cpp KV cache across idle auto-unload
(slot save/restore) (#7204)
* Studio: persist llama.cpp KV cache across idle auto-unload (slot save/restore)
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Studio: address KV persistence review feedback
* Studio: guard KV restore on launch config
* Studio: fix KV resume purge race, fingerprint requested ctx, purge on disable
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Studio: re-check idle/keep-KV settings after slot save, ns file identity
* Studio: shard-aware KV guard, honor user --no-cache-prompt, early save cap
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Studio: honor LLAMA_ARG_CACHE_PROMPT env in slot-save guard
* Studio: derive prompt-cache state from final argv for slot saves
* Studio: stat LoRA/control-vector sidecars in KV restore fingerprint
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Studio: parse csv and FNAME:SCALE sidecar syntax in KV fingerprint
* Studio: address codex review on idle-unload KV resume
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Studio: harden slot-save cleanup, cap accounting, stale-KV guard, save timeout
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Studio: treat unavailable KV estimate as full-cap for slot-save disk check
---------
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Daniel Han
---
studio/backend/core/inference/llama_cpp.py | 293 +++++++++
.../backend/core/inference/llama_keepwarm.py | 138 +++-
.../core/inference/llama_server_args.py | 2 +
studio/backend/main.py | 3 +-
studio/backend/routes/inference.py | 2 +-
studio/backend/routes/settings.py | 21 +-
.../tests/test_llama_cpp_mtp_detection.py | 19 +
.../tests/test_llama_cpp_slot_resume.py | 494 +++++++++++++++
.../backend/tests/test_llama_server_args.py | 12 +
.../backend/tests/test_openai_auto_switch.py | 588 +++++++++++++++++-
.../utils/openai_auto_switch_settings.py | 61 +-
studio/backend/utils/paths/storage_roots.py | 5 +
.../settings/api/openai-auto-switch.ts | 19 +-
.../components/model-auto-switch-section.tsx | 26 +-
studio/frontend/src/i18n/locales/en.ts | 3 +
15 files changed, 1644 insertions(+), 42 deletions(-)
create mode 100644 studio/backend/tests/test_llama_cpp_slot_resume.py
diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py
index d7c7eed518..d4cab81bdf 100644
--- a/studio/backend/core/inference/llama_cpp.py
+++ b/studio/backend/core/inference/llama_cpp.py
@@ -23,6 +23,7 @@ import subprocess
import sys
import threading
import time
+import uuid
from pathlib import Path
from typing import (
Callable,
@@ -453,6 +454,23 @@ def _hf_offline_if_dns_dead():
os.environ.pop("TRANSFORMERS_OFFLINE", None)
+try:
+ _SLOT_SAVE_MAX_BYTES = int(os.environ.get("UNSLOTH_SLOT_SAVE_MAX_BYTES") or (10 << 30))
+except ValueError:
+ _SLOT_SAVE_MAX_BYTES = 10 << 30
+
+# The idle loop holds the lifecycle gate across a slot save, so a newly arriving
+# request waits on the in-flight save's HTTP call. Bound it (was 120s) so a slow
+# or stuck save can't stall the next request for minutes; best-effort save just
+# falls back to a plain unload. Override with UNSLOTH_SLOT_SAVE_TIMEOUT (seconds).
+try:
+ _SLOT_SAVE_HTTP_TIMEOUT = float(os.environ.get("UNSLOTH_SLOT_SAVE_TIMEOUT") or 30.0)
+except ValueError:
+ _SLOT_SAVE_HTTP_TIMEOUT = 30.0
+if _SLOT_SAVE_HTTP_TIMEOUT <= 0:
+ _SLOT_SAVE_HTTP_TIMEOUT = 30.0
+
+
def _swa_cache_path() -> Path:
home = os.environ.get("UNSLOTH_STUDIO_HOME") or os.environ.get("STUDIO_HOME")
base = Path(home) if home else Path.home() / ".unsloth" / "studio"
@@ -2000,6 +2018,12 @@ class LlamaCppBackend:
self._llama_log_path: Optional[Path] = None
self._cancel_event = threading.Event()
self._api_key: Optional[str] = None
+ self._slot_save_dir: Optional[str] = None
+ self._slot_save_binary: Optional[tuple[str, int]] = None
+ # (gguf_identity, launch_fingerprint) snapshotted at load, so a later slot
+ # save can tell whether the model files were swapped on disk since load.
+ self._slot_loaded_identity: Optional[tuple] = None
+ self._prompt_cache_disabled: bool = False
# True once a probe has completed; cleared on transient failure.
self._is_audio: bool = False
self._audio_type: Optional[str] = None
@@ -2638,6 +2662,7 @@ class LlamaCppBackend:
"supports_ctx_checkpoints": False,
"supports_no_cache_prompt": False,
"supports_metrics": False,
+ "supports_slot_save": False,
}
try:
mtime = int(Path(bin_path).stat().st_mtime)
@@ -2658,6 +2683,7 @@ class LlamaCppBackend:
supports_ctx_checkpoints = False
supports_no_cache_prompt = False
supports_metrics = False
+ supports_slot_save = False
try:
probe_env = cls._llama_server_env_for_binary(bin_path)
result = subprocess.run(
@@ -2756,6 +2782,7 @@ class LlamaCppBackend:
supports_ctx_checkpoints = _is_real("--ctx-checkpoints")
supports_no_cache_prompt = _is_real("--no-cache-prompt")
supports_metrics = _is_real("--metrics")
+ supports_slot_save = _is_real("--slot-save-path")
except (OSError, subprocess.SubprocessError) as exc:
logger.debug(f"llama-server --help probe failed: {exc}")
@@ -2773,6 +2800,7 @@ class LlamaCppBackend:
"supports_ctx_checkpoints": supports_ctx_checkpoints,
"supports_no_cache_prompt": supports_no_cache_prompt,
"supports_metrics": supports_metrics,
+ "supports_slot_save": supports_slot_save,
}
cls._capability_cache[cache_key] = info
return info
@@ -7332,6 +7360,26 @@ class LlamaCppBackend:
# when the binary advertises it (older/custom binaries may not).
if server_caps.get("supports_metrics"):
cmd.append("--metrics")
+ self._slot_save_dir = None
+ self._slot_save_binary = None
+ self._prompt_cache_disabled = False
+ if server_caps.get("supports_slot_save"):
+ try:
+ from utils.paths.storage_roots import ( # noqa: WPS433
+ llama_slot_cache_root,
+ )
+
+ slot_dir = llama_slot_cache_root()
+ slot_dir.mkdir(parents = True, exist_ok = True)
+ # Saved KV encodes chat content; keep it from other local users.
+ with contextlib.suppress(OSError):
+ os.chmod(slot_dir, 0o700)
+ cmd.extend(["--slot-save-path", str(slot_dir)])
+ self._slot_save_dir = str(slot_dir)
+ self._slot_save_binary = (binary, Path(binary).stat().st_mtime_ns)
+ except OSError:
+ self._slot_save_dir = None
+ self._slot_save_binary = None
cmd.extend(
self._ctx_integrity_flags(
n_parallel,
@@ -7529,6 +7577,7 @@ class LlamaCppBackend:
unsupported_cache_flags.append("--ctx-checkpoints")
if server_caps.get("supports_no_cache_prompt"):
cmd.append("--no-cache-prompt")
+ self._prompt_cache_disabled = True
else:
unsupported_cache_flags.append("--no-cache-prompt")
if unsupported_cache_flags:
@@ -8105,6 +8154,15 @@ class LlamaCppBackend:
if not self._healthy:
return False
+ # Snapshot the files the server actually loaded. If a GGUF shard or a
+ # LoRA/control-vector sidecar is swapped on disk afterwards while the
+ # old weights stay mapped, save_slots_for_resume() compares against
+ # this and refuses to persist KV that a reload could misapply.
+ if self._slot_save_dir:
+ self._slot_loaded_identity = (
+ self._gguf_file_identity(self._gguf_path),
+ self._slot_launch_fingerprint(),
+ )
return True
def _build_speculative_flags(
@@ -8690,6 +8748,10 @@ class LlamaCppBackend:
self._effective_context_length = None
self._max_context_length = None
self._reset_effective_parallel_slots()
+ self._slot_save_dir = None
+ self._slot_save_binary = None
+ self._slot_loaded_identity = None
+ self._prompt_cache_disabled = False
self._chat_template = None
self._chat_template_override = None
self._supports_reasoning = False
@@ -9216,6 +9278,237 @@ class LlamaCppBackend:
return False
return True
+ def _slot_launch_fingerprint(self) -> tuple:
+ # KV validity keys on extra args, stat'd sidecar weights, effective ctx.
+ sidecars = []
+ for path in self._sidecar_weight_files():
+ try:
+ st = os.stat(path)
+ sidecars.append((path, st.st_size, st.st_mtime_ns))
+ except OSError:
+ sidecars.append((path, None, None))
+ return (
+ tuple(self._extra_args or ()),
+ tuple(sidecars),
+ self._requested_n_ctx,
+ self._effective_context_length,
+ getattr(self, "_cache_type_kv", None),
+ self.effective_parallel_slots,
+ )
+
+ def _gguf_file_identity(self, path) -> Optional[tuple]:
+ # (size, mtime_ns) per shard: a split GGUF keys KV validity on every sibling.
+ p = Path(path)
+ paths = [p]
+ m = _SHARD_FULL_RE.match(p.name)
+ if m:
+ prefix, _first, total = m.groups()
+ paths = [
+ p.with_name(f"{prefix}-{i:05d}-of-{total}{p.suffix}")
+ for i in range(1, int(total) + 1)
+ ]
+ try:
+ return tuple((sp.stat().st_size, sp.stat().st_mtime_ns) for sp in paths)
+ except OSError:
+ return None
+
+ _SIDECAR_WEIGHT_FLAGS = (
+ "--lora",
+ "--lora-scaled",
+ "--control-vector",
+ "--control-vector-scaled",
+ )
+
+ def _sidecar_weight_files(self) -> list[str]:
+ # llama.cpp: comma-separated paths, FNAME:SCALE on -scaled (older builds: FNAME SCALE).
+ args = [str(a).strip() for a in (self._extra_args or ())]
+ files: list[str] = []
+ for i, arg in enumerate(args):
+ flag, sep, inline = arg.partition("=")
+ if flag not in self._SIDECAR_WEIGHT_FLAGS:
+ continue
+ operand = inline if sep else (args[i + 1] if i + 1 < len(args) else "")
+ if not operand:
+ continue
+ candidates = [operand]
+ pieces = [p for p in operand.split(",") if p]
+ if len(pieces) > 1:
+ candidates.extend(pieces)
+ if flag.endswith("-scaled"):
+ for item in list(candidates):
+ # ":" tail is a scale; rpartition spares drive letters.
+ head, colon, tail = item.rpartition(":")
+ if not (colon and head):
+ continue
+ try:
+ float(tail)
+ except ValueError:
+ continue
+ candidates.append(head)
+ for cand in candidates:
+ if cand not in files:
+ files.append(cand)
+ return files
+
+ def _prompt_cache_off(self) -> bool:
+ # Caching off makes restores useless; last prompt-cache flag wins, env only when unset.
+ last = None
+ for arg in self._extra_args or ():
+ flag = arg.strip().split("=", 1)[0]
+ if flag in ("--cache-prompt", "--no-cache-prompt"):
+ last = flag
+ if last is not None:
+ return last == "--no-cache-prompt"
+ if self._prompt_cache_disabled:
+ return True
+ if os.environ.get("LLAMA_ARG_NO_CACHE_PROMPT") is not None:
+ return True
+ env = (os.environ.get("LLAMA_ARG_CACHE_PROMPT") or "").strip().lower()
+ return env in {"off", "disabled", "false", "0"}
+
+ def save_slots_for_resume(
+ self, should_abort: Optional[Callable[[], bool]] = None
+ ) -> Optional[dict]:
+ if (
+ not self.is_loaded
+ or not self._slot_save_dir
+ or not self._gguf_path
+ or self._prompt_cache_off()
+ ):
+ return None
+ save_dir = Path(self._slot_save_dir)
+ gguf_stat = self._gguf_file_identity(self._gguf_path)
+ if gguf_stat is None:
+ return None
+ launch = self._slot_launch_fingerprint()
+ # If the GGUF or a sidecar was swapped on disk while the original weights
+ # stayed mapped, the live KV belongs to the old weights but a reload would
+ # load the new file. Persisting it would let restore misapply stale KV.
+ if self._slot_loaded_identity is not None and self._slot_loaded_identity != (
+ gguf_stat,
+ launch,
+ ):
+ logger.debug("Skipping slot save: model files changed on disk since load")
+ return None
+ try:
+ estimate = self._estimate_kv_cache_bytes(
+ self._effective_context_length or self._context_length or 0,
+ self._cache_type_kv,
+ n_parallel = self.effective_parallel_slots,
+ )
+ # Skip before writing anything when the estimate alone blows the cap,
+ # rather than fully writing a slot and discarding it afterwards.
+ if estimate > _SLOT_SAVE_MAX_BYTES:
+ logger.debug(
+ "Skipping slot save: estimated %d bytes exceeds cap %d",
+ estimate,
+ _SLOT_SAVE_MAX_BYTES,
+ )
+ return None
+ # A 0 estimate means metadata was insufficient, not a zero-byte cache:
+ # a slot can still be many GiB, so demand room for the whole cap before
+ # trusting the post-write check.
+ required = (estimate if estimate > 0 else _SLOT_SAVE_MAX_BYTES) + (1 << 30)
+ if shutil.disk_usage(save_dir).free < required:
+ logger.debug("Skipping slot save: insufficient free disk")
+ return None
+ except Exception:
+ pass
+ token = uuid.uuid4().hex[:8]
+ entries: list[dict] = []
+ total_bytes = 0
+ for slot in range(self.effective_parallel_slots):
+ # A request pending mid-save waits on the gate; stop wasting its time.
+ if should_abort is not None and should_abort():
+ break
+ filename = f"resume-{token}-slot{slot}.bin"
+ path = save_dir / filename
+ try:
+ resp = httpx.post(
+ f"{self.base_url}/slots/{slot}",
+ params = {"action": "save"},
+ json = {"filename": filename},
+ headers = self._auth_headers,
+ timeout = _SLOT_SAVE_HTTP_TIMEOUT,
+ trust_env = False,
+ )
+ except Exception as e:
+ logger.debug(f"slot {slot} save failed: {e}")
+ with contextlib.suppress(OSError):
+ path.unlink()
+ break
+ if resp.status_code != 200:
+ logger.debug(f"slot {slot} save returned HTTP {resp.status_code}")
+ with contextlib.suppress(OSError):
+ path.unlink()
+ continue
+ try:
+ body = resp.json()
+ if not isinstance(body, dict):
+ raise ValueError("slot save response was not a JSON object")
+ n_saved = int(body.get("n_saved") or 0)
+ except Exception as e:
+ # A 200 that still wrote a file but returns a malformed body must
+ # clean up like the transport/HTTP error paths above, or the file
+ # (which holds chat KV) is orphaned until the next startup sweep.
+ logger.debug(f"slot {slot} save returned an invalid response: {e}")
+ with contextlib.suppress(OSError):
+ path.unlink()
+ continue
+ if n_saved <= 0:
+ with contextlib.suppress(OSError):
+ path.unlink()
+ continue
+ # Account by the bytes actually on disk, not the server-reported
+ # count, so the cap holds even if a custom binary under-reports.
+ try:
+ n_written = path.stat().st_size
+ except OSError:
+ n_written = 0
+ total_bytes += n_written
+ entries.append({"id": slot, "filename": filename, "n_saved": n_saved})
+ if total_bytes > _SLOT_SAVE_MAX_BYTES:
+ break # already over the cap; the discard below cleans up
+ if not entries:
+ return None
+ if total_bytes > _SLOT_SAVE_MAX_BYTES:
+ logger.debug(
+ "Discarding slot save: %d bytes exceeds cap %d",
+ total_bytes,
+ _SLOT_SAVE_MAX_BYTES,
+ )
+ for entry in entries:
+ with contextlib.suppress(OSError):
+ (save_dir / entry["filename"]).unlink()
+ return None
+ return {
+ "dir": self._slot_save_dir,
+ "binary": self._slot_save_binary,
+ "gguf": str(self._gguf_path),
+ "gguf_stat": gguf_stat,
+ "launch": launch,
+ "slots": entries,
+ }
+
+ def restore_slots_for_resume(self, manifest: dict) -> None:
+ if not self.is_loaded or not self._slot_save_dir:
+ return
+ for entry in manifest.get("slots") or []:
+ try:
+ resp = httpx.post(
+ f"{self.base_url}/slots/{int(entry['id'])}",
+ params = {"action": "restore"},
+ json = {"filename": str(entry["filename"])},
+ headers = self._auth_headers,
+ timeout = _SLOT_SAVE_HTTP_TIMEOUT,
+ trust_env = False,
+ )
+ except Exception as e:
+ logger.debug(f"slot restore failed: {e}")
+ break
+ if resp.status_code != 200:
+ logger.debug(f"slot {entry.get('id')} restore returned HTTP {resp.status_code}")
+
def _maybe_recover_from_mtp_crash(self, exc: Optional[BaseException] = None) -> bool:
"""Schedule one background reload without MTP after a mid-generation death.
diff --git a/studio/backend/core/inference/llama_keepwarm.py b/studio/backend/core/inference/llama_keepwarm.py
index 86a8c8a404..3380ebf5f5 100644
--- a/studio/backend/core/inference/llama_keepwarm.py
+++ b/studio/backend/core/inference/llama_keepwarm.py
@@ -15,6 +15,7 @@ import asyncio
import contextlib
import threading
import time
+from pathlib import Path
from loggers import get_logger
@@ -30,6 +31,8 @@ _last_active = time.monotonic()
# otherwise 503 against an empty backend can reload it (set on unload, cleared on
# reload). Storing the quant means the reload restores the exact freed variant.
_last_unloaded_model = None
+# Slot KV manifest saved by the idle unload; whoever pops it owns deleting its files.
+_kv_resume = None
# Guards inflight bumps against the idle-check-then-unload race, and blocks new
# inference from starting mid-swap. Process-wide, not per-loop: the backend slot is
# shared across every event loop in the process, so a per-loop gate would let a
@@ -161,11 +164,17 @@ def inference_lifecycle_gate():
return _unload_gate()
-def note_model_loaded() -> None:
- """Record a successful GGUF load: stamp activity and drop any reload stash so
- a manual load clears it synchronously, not only on the next idle poll."""
+def note_model_loaded(backend = None) -> None:
+ """Stamp activity and synchronously drop any reload stash."""
_note_activity()
+ resume = take_kv_resume()
_set_last_unloaded(None)
+ if resume is None:
+ return
+ if backend is not None:
+ restore_kv_resume(backend, resume)
+ else:
+ _delete_resume_files(resume)
def note_model_unloaded() -> None:
@@ -182,9 +191,81 @@ def get_last_unloaded_model():
def _set_last_unloaded(value) -> None:
- global _last_unloaded_model
+ global _last_unloaded_model, _kv_resume
+ stale = None
with _lock:
_last_unloaded_model = value
+ if value is None and _kv_resume is not None:
+ stale, _kv_resume = _kv_resume, None
+ if stale:
+ _delete_resume_files(stale)
+
+
+def _delete_resume_files(manifest) -> None:
+ try:
+ base = Path(manifest.get("dir") or "")
+ for entry in manifest.get("slots") or []:
+ with contextlib.suppress(OSError):
+ (base / str(entry.get("filename"))).unlink()
+ except Exception:
+ pass
+
+
+def _set_kv_resume(value) -> None:
+ global _kv_resume
+ stale = None
+ with _lock:
+ if _kv_resume is not None and _kv_resume is not value:
+ stale = _kv_resume
+ _kv_resume = value
+ if stale:
+ _delete_resume_files(stale)
+
+
+def take_kv_resume():
+ global _kv_resume
+ with _lock:
+ manifest, _kv_resume = _kv_resume, None
+ return manifest
+
+
+def purge_kv_resume() -> None:
+ resume = take_kv_resume()
+ if resume:
+ _delete_resume_files(resume)
+
+
+def restore_kv_resume(backend, manifest) -> None:
+ try:
+ gguf = manifest.get("gguf")
+ binary = manifest.get("binary")
+ current = getattr(backend, "_gguf_path", None)
+ same_gguf = bool(gguf and current) and Path(current).resolve() == Path(gguf).resolve()
+ if same_gguf:
+ # Same path is not enough: shards may have been rewritten meanwhile.
+ identity = getattr(backend, "_gguf_file_identity", None)
+ same_gguf = callable(identity) and identity(current) == manifest.get("gguf_stat")
+ if same_gguf:
+ # Nor the same file: launch overrides can invalidate KV numerics.
+ fingerprint = getattr(backend, "_slot_launch_fingerprint", None)
+ same_gguf = callable(fingerprint) and manifest.get("launch") == fingerprint()
+ if same_gguf and binary and binary == getattr(backend, "_slot_save_binary", None):
+ logger.info("Restoring saved slot KV onto the reloaded model")
+ backend.restore_slots_for_resume(manifest)
+ except Exception as exc:
+ logger.debug("slot restore after reload failed: %s", exc)
+ finally:
+ _delete_resume_files(manifest)
+
+
+def sweep_slot_save_dir() -> None:
+ try:
+ from utils.paths.storage_roots import llama_slot_cache_root
+ for path in llama_slot_cache_root().glob("resume-*.bin"):
+ with contextlib.suppress(OSError):
+ path.unlink()
+ except Exception:
+ pass
class LlamaKeepWarmMiddleware:
@@ -266,7 +347,10 @@ def _loaded_identity(backend):
async def idle_unload_loop(poll_seconds: float = 15.0) -> None:
"""Unload the loaded GGUF once idle past the configured TTL. Inert when off."""
- from utils.openai_auto_switch_settings import get_auto_unload_idle_seconds
+ from utils.openai_auto_switch_settings import (
+ get_auto_unload_idle_seconds,
+ get_auto_unload_keep_kv,
+ )
seen_model = None
while True:
@@ -281,17 +365,47 @@ async def idle_unload_loop(poll_seconds: float = 15.0) -> None:
# Track by (id, variant): a (re)loaded model -- including the same repo
# at a different quant -- counts as activity so it survives one TTL
# before its first request (loads bypass the activity middleware).
- current = _loaded_identity(backend)
- if current != seen_model:
- seen_model = current
- if current is not None:
- _note_activity()
- _set_last_unloaded(None) # a model is loaded; drop stale stash
async with _unload_gate():
+ # Purging the stash mid-reload would race the restore.
+ current = _loaded_identity(backend)
+ if current != seen_model:
+ seen_model = current
+ if current is not None:
+ _note_activity()
+ _set_last_unloaded(None) # a model is loaded; drop stale stash
if backend.is_loaded and _is_idle(ttl):
freed = _loaded_identity(backend)
- await asyncio.to_thread(backend.unload_model)
+ manifest = None
+ if get_auto_unload_keep_kv():
+ try:
+ manifest = await asyncio.to_thread(
+ backend.save_slots_for_resume,
+ lambda: not _is_idle(ttl),
+ )
+ except Exception as exc:
+ logger.debug("slot save before idle unload failed: %s", exc)
+ # Re-read settings: the save can outlive a settings change.
+ ttl = get_auto_unload_idle_seconds()
+ if ttl <= 0 or not _is_idle(ttl):
+ if manifest:
+ _delete_resume_files(manifest)
+ continue
+ if manifest and not get_auto_unload_keep_kv():
+ _delete_resume_files(manifest)
+ manifest = None
+ try:
+ await asyncio.to_thread(backend.unload_model)
+ except Exception:
+ # Failed unload means nothing will stash the manifest.
+ if manifest:
+ _delete_resume_files(manifest)
+ raise
_set_last_unloaded(freed) # let an alias request reload it
+ if manifest and freed:
+ _set_kv_resume({"identity": freed, **manifest})
+ logger.info("Idle auto-unload: saved slot KV for restore on reload")
+ elif manifest:
+ _delete_resume_files(manifest)
logger.info("Idle auto-unload: freed GGUF after %ss idle", ttl)
seen_model = None
except Exception as exc:
diff --git a/studio/backend/core/inference/llama_server_args.py b/studio/backend/core/inference/llama_server_args.py
index e72e10e071..7b42d2f40d 100644
--- a/studio/backend/core/inference/llama_server_args.py
+++ b/studio/backend/core/inference/llama_server_args.py
@@ -70,6 +70,8 @@ _DENYLIST_GROUPS: tuple[frozenset[str], ...] = (
# llama-server's own built-in tools flag would silently stack on top of
# Unsloth's --enable-tools / --disable-tools policy resolver.
frozenset({"--tools"}),
+ # Slot-state dir: Studio owns it for KV persistence across idle unload.
+ frozenset({"--slot-save-path"}),
)
_DENYLIST: frozenset[str] = frozenset().union(*_DENYLIST_GROUPS)
diff --git a/studio/backend/main.py b/studio/backend/main.py
index 81d4c16e52..a1ff4d60da 100644
--- a/studio/backend/main.py
+++ b/studio/backend/main.py
@@ -547,8 +547,9 @@ async def lifespan(app: FastAPI):
threading.Thread(target = _warm_rag_embedder, daemon = True, name = "rag-embedder-warm").start()
# Idle auto-unload loop (no-op unless the OpenAI auto-unload TTL is set).
- from core.inference.llama_keepwarm import idle_unload_loop
+ from core.inference.llama_keepwarm import idle_unload_loop, sweep_slot_save_dir
+ sweep_slot_save_dir()
app.state.idle_unload_task = asyncio.create_task(idle_unload_loop())
# Initialize RSA key pair for API key encryption (external providers).
diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py
index 136e4f7645..9c08ea4b79 100644
--- a/studio/backend/routes/inference.py
+++ b/studio/backend/routes/inference.py
@@ -4710,7 +4710,7 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
# Clear any idle-unload reload stash now, not only on the next poll.
from core.inference.llama_keepwarm import note_model_loaded
- note_model_loaded()
+ await asyncio.to_thread(note_model_loaded, llama_backend)
# A plain load advertises its own identifier; auto-switch overwrites
# this with the repo id right after _load_model_impl returns.
llama_backend._openai_advertised_id = None
diff --git a/studio/backend/routes/settings.py b/studio/backend/routes/settings.py
index ab0fd2fd99..17e64df918 100644
--- a/studio/backend/routes/settings.py
+++ b/studio/backend/routes/settings.py
@@ -36,9 +36,10 @@ from utils.helper_precache_settings import (
)
from utils.coding_agents import CODING_AGENTS, detect_installed_coding_agents
from utils.openai_auto_switch_settings import (
- DEFAULT_AUTO_UNLOAD_IDLE_SECONDS,
+ DEFAULT_AUTO_UNLOAD_KEEP_KV,
DEFAULT_OPENAI_AUTO_SWITCH_ENABLED,
get_auto_unload_idle_seconds,
+ get_auto_unload_keep_kv,
get_model_overrides,
get_openai_auto_switch_enabled,
get_stored_auto_unload_idle_seconds,
@@ -90,7 +91,9 @@ class HelperPrecacheResponse(BaseModel):
class OpenAIAutoSwitchPayload(BaseModel):
enabled: bool
- auto_unload_idle_seconds: int = Field(default = DEFAULT_AUTO_UNLOAD_IDLE_SECONDS, ge = 0)
+ # None leaves the stored value untouched (partial updates can't clobber it).
+ auto_unload_idle_seconds: Optional[int] = Field(default = None, ge = 0)
+ auto_unload_keep_kv: Optional[bool] = None
class OpenAIAutoSwitchResponse(BaseModel):
@@ -101,6 +104,7 @@ class OpenAIAutoSwitchResponse(BaseModel):
# UNSLOTH_MODEL_IDLE_TTL set and nothing stored, this is true even while enabled
# is false, so the UI can show idle-unload as active instead of "needs enable".
idle_unload_active: bool = False
+ auto_unload_keep_kv: bool = DEFAULT_AUTO_UNLOAD_KEEP_KV
class ModelOverridePayload(BaseModel):
@@ -198,6 +202,7 @@ def get_openai_auto_switch(
enabled = get_openai_auto_switch_enabled(),
auto_unload_idle_seconds = get_stored_auto_unload_idle_seconds(),
idle_unload_active = get_auto_unload_idle_seconds() > 0,
+ auto_unload_keep_kv = get_auto_unload_keep_kv(),
)
@@ -206,8 +211,8 @@ def update_openai_auto_switch(
payload: OpenAIAutoSwitchPayload, current_subject: str = Depends(get_current_subject)
) -> OpenAIAutoSwitchResponse:
try:
- enabled, idle_seconds = set_openai_auto_switch(
- payload.enabled, payload.auto_unload_idle_seconds
+ enabled, idle_seconds, keep_kv = set_openai_auto_switch(
+ payload.enabled, payload.auto_unload_idle_seconds, payload.auto_unload_keep_kv
)
except ValueError as exc:
raise log_and_http_error(
@@ -217,10 +222,16 @@ def update_openai_auto_switch(
event = "settings.update_openai_auto_switch_failed",
log = logger,
) from exc
+ idle_unload_active = get_auto_unload_idle_seconds() > 0
+ if not keep_kv or not idle_unload_active:
+ # Keep-KV off or idle unload disabled: drop already-saved chat context too.
+ from core.inference.llama_keepwarm import purge_kv_resume
+ purge_kv_resume()
return OpenAIAutoSwitchResponse(
enabled = enabled,
auto_unload_idle_seconds = idle_seconds,
- idle_unload_active = get_auto_unload_idle_seconds() > 0,
+ idle_unload_active = idle_unload_active,
+ auto_unload_keep_kv = keep_kv,
)
diff --git a/studio/backend/tests/test_llama_cpp_mtp_detection.py b/studio/backend/tests/test_llama_cpp_mtp_detection.py
index 8fe04c0e39..68b706ebf9 100644
--- a/studio/backend/tests/test_llama_cpp_mtp_detection.py
+++ b/studio/backend/tests/test_llama_cpp_mtp_detection.py
@@ -741,6 +741,25 @@ def test_probe_reports_windows_cache_flags_absent_for_older_binary(tmp_path):
assert caps["supports_no_cache_prompt"] is False
+@_NEEDS_BASH
+def test_probe_detects_slot_save_path(tmp_path):
+ fake = _make_fake_llama_server(
+ tmp_path / "llama-server",
+ "--slot-save-path PATH path to save slot kv cache\n--threads N\n",
+ )
+ _clear_caps_cache()
+ caps = LlamaCppBackend.probe_server_capabilities(str(fake))
+ assert caps["supports_slot_save"] is True
+
+
+@_NEEDS_BASH
+def test_probe_reports_slot_save_absent_for_older_binary(tmp_path):
+ fake = _make_fake_llama_server(tmp_path / "llama-server", "--threads N\n")
+ _clear_caps_cache()
+ caps = LlamaCppBackend.probe_server_capabilities(str(fake))
+ assert caps["supports_slot_save"] is False
+
+
def test_build_ngram_mod_flags_new():
flags = _build_ngram_mod_flags({"ngram_mod_flavor": "new"})
assert flags == [
diff --git a/studio/backend/tests/test_llama_cpp_slot_resume.py b/studio/backend/tests/test_llama_cpp_slot_resume.py
new file mode 100644
index 0000000000..8b20c952c4
--- /dev/null
+++ b/studio/backend/tests/test_llama_cpp_slot_resume.py
@@ -0,0 +1,494 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+import os
+from types import SimpleNamespace
+
+import core.inference.llama_cpp as llama_cpp
+from core.inference.llama_cpp import LlamaCppBackend
+
+
+def _resume_backend(tmp_path, n_slots = 1):
+ backend = LlamaCppBackend()
+ backend._healthy = True
+ # No-op lifecycle methods so the atexit cleanup can kill the fake quietly.
+ backend._process = SimpleNamespace(
+ poll = lambda: None,
+ terminate = lambda: None,
+ wait = lambda *a, **k: 0,
+ kill = lambda: None,
+ pid = 0,
+ )
+ backend._port = 8081
+ backend._slot_save_dir = str(tmp_path)
+ backend._slot_save_binary = ("/bin/llama-server", 1)
+ (tmp_path / "model.gguf").write_bytes(b"gguf")
+ backend._gguf_path = str(tmp_path / "model.gguf")
+ backend._effective_parallel_slots = n_slots
+ backend._estimate_kv_cache_bytes = lambda *a, **k: 0
+ return backend
+
+
+def _fake_disk(monkeypatch, free = 1 << 40):
+ monkeypatch.setattr(llama_cpp.shutil, "disk_usage", lambda _p: SimpleNamespace(free = free))
+
+
+class _Resp:
+ def __init__(
+ self,
+ status_code = 200,
+ body = None,
+ ):
+ self.status_code = status_code
+ self._body = body or {}
+
+ def json(self):
+ return self._body
+
+
+def test_save_returns_none_when_slot_save_disabled(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path)
+ backend._slot_save_dir = None
+ monkeypatch.setattr(
+ llama_cpp.httpx,
+ "post",
+ lambda *a, **k: (_ for _ in ()).throw(AssertionError),
+ raising = False,
+ )
+ assert backend.save_slots_for_resume() is None
+
+
+def test_save_skipped_when_prompt_cache_disabled(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path)
+ backend._prompt_cache_disabled = True
+ monkeypatch.setattr(
+ llama_cpp.httpx,
+ "post",
+ lambda *a, **k: (_ for _ in ()).throw(AssertionError),
+ raising = False,
+ )
+ assert backend.save_slots_for_resume() is None
+
+
+def test_save_skipped_when_insufficient_free_disk(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path)
+ backend._estimate_kv_cache_bytes = lambda *a, **k: 1 << 40
+ _fake_disk(monkeypatch, free = 1 << 20)
+ monkeypatch.setattr(
+ llama_cpp.httpx,
+ "post",
+ lambda *a, **k: (_ for _ in ()).throw(AssertionError),
+ raising = False,
+ )
+ assert backend.save_slots_for_resume() is None
+
+
+def test_save_collects_manifest_across_slots(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path, n_slots = 2)
+ _fake_disk(monkeypatch)
+ calls = []
+
+ def fake_post(url, **kwargs):
+ calls.append((url, kwargs["params"], kwargs["json"]))
+ return _Resp(200, {"n_saved": 40, "n_written": 100})
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ manifest = backend.save_slots_for_resume()
+ assert manifest is not None
+ assert manifest["dir"] == str(tmp_path)
+ assert manifest["binary"] == ("/bin/llama-server", 1)
+ assert manifest["gguf"] == str(tmp_path / "model.gguf")
+ st = os.stat(manifest["gguf"])
+ assert manifest["gguf_stat"] == ((st.st_size, st.st_mtime_ns),)
+ assert manifest["launch"] == backend._slot_launch_fingerprint()
+ assert [e["id"] for e in manifest["slots"]] == [0, 1]
+ assert all(e["n_saved"] == 40 for e in manifest["slots"])
+ assert [c[1] for c in calls] == [{"action": "save"}] * 2
+ assert "/slots/0" in calls[0][0] and "/slots/1" in calls[1][0]
+
+
+def test_save_unlinks_empty_slot_and_returns_none(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path)
+ _fake_disk(monkeypatch)
+
+ def fake_post(url, **kwargs):
+ (tmp_path / kwargs["json"]["filename"]).write_bytes(b"")
+ return _Resp(200, {"n_saved": 0, "n_written": 0})
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ assert backend.save_slots_for_resume() is None
+ assert list(tmp_path.glob("resume-*.bin")) == [] # empty-slot file removed
+
+
+def test_save_cap_breach_discards_all_files(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path, n_slots = 2)
+ _fake_disk(monkeypatch)
+ monkeypatch.setattr(llama_cpp, "_SLOT_SAVE_MAX_BYTES", 150)
+
+ def fake_post(url, **kwargs):
+ (tmp_path / kwargs["json"]["filename"]).write_bytes(b"x" * 100)
+ return _Resp(200, {"n_saved": 40, "n_written": 100})
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ assert backend.save_slots_for_resume() is None # 200 bytes > 150 cap
+ assert list(tmp_path.glob("resume-*.bin")) == []
+
+
+def test_save_transport_error_aborts_remaining_slots(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path, n_slots = 3)
+ _fake_disk(monkeypatch)
+ calls = []
+
+ def fake_post(url, **kwargs):
+ calls.append(url)
+ raise OSError("connection refused")
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ assert backend.save_slots_for_resume() is None
+ assert len(calls) == 1 # no retries against a dead server
+
+
+def test_save_transport_error_unlinks_partial_file(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path)
+ _fake_disk(monkeypatch)
+
+ def fake_post(url, **kwargs):
+ (tmp_path / kwargs["json"]["filename"]).write_bytes(b"partial")
+ raise OSError("timed out")
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ assert backend.save_slots_for_resume() is None
+ assert list(tmp_path.glob("resume-*.bin")) == []
+
+
+def test_fingerprint_tracks_lora_sidecar_rewrite(tmp_path):
+ backend = _resume_backend(tmp_path)
+ adapter = tmp_path / "adapter.gguf"
+ adapter.write_bytes(b"v1")
+ backend._extra_args = ["--lora", str(adapter)]
+
+ before = backend._slot_launch_fingerprint()
+ adapter.write_bytes(b"v2-different") # re-exported adapter, same path
+ assert backend._slot_launch_fingerprint() != before
+
+ backend._extra_args = [f"--lora={adapter}"]
+ assert backend._sidecar_weight_files() == [str(adapter)]
+ backend._extra_args = ["--lora-scaled", str(adapter), "0.5"]
+ assert backend._sidecar_weight_files() == [str(adapter)]
+ backend._extra_args = ["--control-vector", str(adapter), "--threads", "4"]
+ assert backend._sidecar_weight_files() == [str(adapter)]
+
+
+def test_sidecar_files_parse_csv_and_colon_scale(tmp_path):
+ backend = _resume_backend(tmp_path)
+ a, b = tmp_path / "a.gguf", tmp_path / "b.gguf"
+
+ backend._extra_args = ["--lora", f"{a},{b}"]
+ files = backend._sidecar_weight_files()
+ assert str(a) in files and str(b) in files
+
+ backend._extra_args = ["--lora-scaled", f"{a}:0.5"]
+ assert str(a) in backend._sidecar_weight_files()
+
+ backend._extra_args = ["--control-vector-scaled", f"{a}:1.0,{b}:2.0"]
+ files = backend._sidecar_weight_files()
+ assert str(a) in files and str(b) in files
+
+ # Windows drive letter must not be mistaken for a scale separator.
+ backend._extra_args = ["--lora-scaled", "C:\\adapters\\a.gguf:0.75"]
+ assert "C:\\adapters\\a.gguf" in backend._sidecar_weight_files()
+ backend._extra_args = ["--lora", "C:\\adapters\\a.gguf"]
+ assert backend._sidecar_weight_files() == ["C:\\adapters\\a.gguf"]
+
+
+def test_fingerprint_tracks_colon_scaled_adapter_rewrite(tmp_path):
+ backend = _resume_backend(tmp_path)
+ adapter = tmp_path / "adapter.gguf"
+ adapter.write_bytes(b"v1")
+ backend._extra_args = ["--lora-scaled", f"{adapter}:0.5"]
+
+ before = backend._slot_launch_fingerprint()
+ adapter.write_bytes(b"v2-different") # re-exported adapter, same path
+ assert backend._slot_launch_fingerprint() != before
+
+
+def test_fingerprint_tracks_effective_context_length(tmp_path):
+ backend = _resume_backend(tmp_path)
+ backend._effective_context_length = 8192
+
+ before = backend._slot_launch_fingerprint()
+ backend._effective_context_length = 4096 # auto-fit landed smaller on reload
+ assert backend._slot_launch_fingerprint() != before
+
+
+def test_gguf_file_identity_covers_split_shards(tmp_path):
+ backend = _resume_backend(tmp_path)
+ first = tmp_path / "m-00001-of-00002.gguf"
+ second = tmp_path / "m-00002-of-00002.gguf"
+ first.write_bytes(b"a")
+ second.write_bytes(b"bb")
+
+ before = backend._gguf_file_identity(str(first))
+ st1, st2 = os.stat(first), os.stat(second)
+ assert before == ((st1.st_size, st1.st_mtime_ns), (st2.st_size, st2.st_mtime_ns))
+
+ second.write_bytes(b"rewritten") # sibling changes, primary untouched
+ after = backend._gguf_file_identity(str(first))
+ assert after is not None and after != before
+ assert after[0] == before[0] # primary shard unchanged
+
+ second.unlink()
+ assert backend._gguf_file_identity(str(first)) is None # missing shard
+
+
+def test_save_skipped_when_user_disabled_prompt_cache(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path)
+ backend._extra_args = ["--no-cache-prompt"]
+ monkeypatch.setattr(
+ llama_cpp.httpx,
+ "post",
+ lambda *a, **k: (_ for _ in ()).throw(AssertionError),
+ raising = False,
+ )
+ assert backend.save_slots_for_resume() is None
+
+
+def test_save_skipped_when_env_disables_prompt_cache(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path)
+ monkeypatch.setenv("LLAMA_ARG_CACHE_PROMPT", "0")
+ monkeypatch.setattr(
+ llama_cpp.httpx,
+ "post",
+ lambda *a, **k: (_ for _ in ()).throw(AssertionError),
+ raising = False,
+ )
+ assert backend.save_slots_for_resume() is None
+ monkeypatch.delenv("LLAMA_ARG_CACHE_PROMPT")
+ monkeypatch.setenv("LLAMA_ARG_NO_CACHE_PROMPT", "1") # legacy negative form
+ assert backend.save_slots_for_resume() is None
+
+
+def test_explicit_cache_prompt_flag_overrides_env(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path)
+ monkeypatch.setenv("LLAMA_ARG_CACHE_PROMPT", "0")
+ backend._extra_args = ["--cache-prompt"] # CLI wins over env in llama.cpp
+ _fake_disk(monkeypatch)
+ monkeypatch.setattr(
+ llama_cpp.httpx,
+ "post",
+ lambda *a, **k: _Resp(200, {"n_saved": 1, "n_written": 1}),
+ raising = False,
+ )
+ assert backend.save_slots_for_resume() is not None
+
+
+def test_user_cache_prompt_overrides_studio_no_cache_flag(monkeypatch, tmp_path):
+ # User extras follow Studio's flags, so an explicit --cache-prompt wins.
+ backend = _resume_backend(tmp_path)
+ backend._prompt_cache_disabled = True
+ backend._extra_args = ["--cache-prompt"]
+ _fake_disk(monkeypatch)
+ monkeypatch.setattr(
+ llama_cpp.httpx,
+ "post",
+ lambda *a, **k: _Resp(200, {"n_saved": 1, "n_written": 1}),
+ raising = False,
+ )
+ assert backend.save_slots_for_resume() is not None
+ # Last flag wins when both appear in extras.
+ backend._extra_args = ["--cache-prompt", "--no-cache-prompt"]
+ assert backend.save_slots_for_resume() is None
+
+
+def test_save_stops_writing_once_cap_exceeded(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path, n_slots = 3)
+ _fake_disk(monkeypatch)
+ monkeypatch.setattr(llama_cpp, "_SLOT_SAVE_MAX_BYTES", 150)
+ calls = []
+
+ def fake_post(url, **kwargs):
+ calls.append(url)
+ (tmp_path / kwargs["json"]["filename"]).write_bytes(b"x" * 100)
+ return _Resp(200, {"n_saved": 1, "n_written": 100})
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ assert backend.save_slots_for_resume() is None
+ assert len(calls) == 2 # cap blown after slot 1; slot 2 never attempted
+ assert list(tmp_path.glob("resume-*.bin")) == []
+
+
+def test_save_aborts_between_slots_when_no_longer_idle(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path, n_slots = 3)
+ _fake_disk(monkeypatch)
+ calls = []
+
+ def fake_post(url, **kwargs):
+ calls.append(url)
+ return _Resp(200, {"n_saved": 5, "n_written": 10})
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ aborts = iter([False, True, True])
+ manifest = backend.save_slots_for_resume(should_abort = lambda: next(aborts))
+ assert len(calls) == 1 # slots 1 and 2 skipped
+ assert manifest is not None
+ assert [e["id"] for e in manifest["slots"]] == [0]
+
+
+def test_save_non_200_slot_is_skipped_but_others_kept(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path, n_slots = 2)
+ _fake_disk(monkeypatch)
+
+ def fake_post(url, **kwargs):
+ if "/slots/0" in url:
+ return _Resp(500)
+ return _Resp(200, {"n_saved": 5, "n_written": 10})
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ manifest = backend.save_slots_for_resume()
+ assert manifest is not None
+ assert [e["id"] for e in manifest["slots"]] == [1]
+
+
+def test_restore_posts_each_slot_and_tolerates_failures(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path)
+ calls = []
+
+ def fake_post(url, **kwargs):
+ calls.append((url, kwargs["params"], kwargs["json"]))
+ return _Resp(500 if "/slots/0" in url else 200, {"n_restored": 5})
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ backend.restore_slots_for_resume(
+ {
+ "slots": [
+ {"id": 0, "filename": "resume-a-slot0.bin", "n_saved": 5},
+ {"id": 1, "filename": "resume-a-slot1.bin", "n_saved": 5},
+ ]
+ }
+ )
+ assert [c[1] for c in calls] == [{"action": "restore"}] * 2
+ assert calls[0][2] == {"filename": "resume-a-slot0.bin"}
+
+
+def test_restore_transport_error_stops_early(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path)
+ calls = []
+
+ def fake_post(url, **kwargs):
+ calls.append(url)
+ raise OSError("connection refused")
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ backend.restore_slots_for_resume(
+ {"slots": [{"id": 0, "filename": "a.bin"}, {"id": 1, "filename": "b.bin"}]}
+ )
+ assert len(calls) == 1
+
+
+def test_save_deletes_orphan_on_malformed_response(monkeypatch, tmp_path):
+ # A 200 that writes a file but returns a non-numeric counter must be cleaned
+ # up like any other save failure, not left orphaned holding chat KV.
+ backend = _resume_backend(tmp_path)
+ _fake_disk(monkeypatch)
+
+ def fake_post(url, **kwargs):
+ (tmp_path / kwargs["json"]["filename"]).write_bytes(b"chat-kv")
+ return _Resp(200, {"n_saved": "not-an-int"})
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ assert backend.save_slots_for_resume() is None
+ assert list(tmp_path.glob("resume-*.bin")) == []
+
+
+def test_save_deletes_orphan_on_non_dict_response(monkeypatch, tmp_path):
+ backend = _resume_backend(tmp_path)
+ _fake_disk(monkeypatch)
+
+ def fake_post(url, **kwargs):
+ (tmp_path / kwargs["json"]["filename"]).write_bytes(b"chat-kv")
+ return _Resp(200, ["unexpected", "list"])
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ assert backend.save_slots_for_resume() is None
+ assert list(tmp_path.glob("resume-*.bin")) == []
+
+
+def test_save_cap_uses_actual_file_size_not_reported_bytes(monkeypatch, tmp_path):
+ # A binary under-reporting n_written must not slip past the disk cap: the
+ # cap is enforced against the bytes actually on disk.
+ backend = _resume_backend(tmp_path)
+ _fake_disk(monkeypatch)
+ monkeypatch.setattr(llama_cpp, "_SLOT_SAVE_MAX_BYTES", 150)
+
+ def fake_post(url, **kwargs):
+ (tmp_path / kwargs["json"]["filename"]).write_bytes(b"x" * 200)
+ return _Resp(200, {"n_saved": 5, "n_written": 1}) # under-reported
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ assert backend.save_slots_for_resume() is None # 200 real bytes > 150 cap
+ assert list(tmp_path.glob("resume-*.bin")) == []
+
+
+def test_save_skipped_when_estimate_exceeds_cap(monkeypatch, tmp_path):
+ # An estimate over the cap skips before writing any slot at all.
+ backend = _resume_backend(tmp_path)
+ backend._estimate_kv_cache_bytes = lambda *a, **k: 1 << 40
+ monkeypatch.setattr(llama_cpp, "_SLOT_SAVE_MAX_BYTES", 1 << 20)
+ _fake_disk(monkeypatch)
+ monkeypatch.setattr(
+ llama_cpp.httpx,
+ "post",
+ lambda *a, **k: (_ for _ in ()).throw(AssertionError),
+ raising = False,
+ )
+ assert backend.save_slots_for_resume() is None
+
+
+def test_save_skipped_when_model_file_changed_since_load(monkeypatch, tmp_path):
+ # The GGUF/sidecars were swapped on disk after the server loaded them, so the
+ # live KV belongs to the old weights: refuse to persist it (no POST at all).
+ backend = _resume_backend(tmp_path)
+ backend._slot_loaded_identity = ((("stale", 0),), ()) # != current identity
+ _fake_disk(monkeypatch)
+ monkeypatch.setattr(
+ llama_cpp.httpx,
+ "post",
+ lambda *a, **k: (_ for _ in ()).throw(AssertionError),
+ raising = False,
+ )
+ assert backend.save_slots_for_resume() is None
+
+
+def test_save_proceeds_when_load_identity_matches(monkeypatch, tmp_path):
+ # Matching load-time snapshot: the save runs normally.
+ backend = _resume_backend(tmp_path)
+ backend._slot_loaded_identity = (
+ backend._gguf_file_identity(backend._gguf_path),
+ backend._slot_launch_fingerprint(),
+ )
+ _fake_disk(monkeypatch)
+
+ def fake_post(url, **kwargs):
+ (tmp_path / kwargs["json"]["filename"]).write_bytes(b"kv")
+ return _Resp(200, {"n_saved": 5, "n_written": 2})
+
+ monkeypatch.setattr(llama_cpp.httpx, "post", fake_post, raising = False)
+ manifest = backend.save_slots_for_resume()
+ assert manifest is not None
+ assert [e["id"] for e in manifest["slots"]] == [0]
+
+
+def test_save_skipped_when_estimate_unavailable_and_low_disk(monkeypatch, tmp_path):
+ # A 0 estimate means metadata was insufficient, not a zero-byte cache: the save
+ # must demand room for the whole cap, not just 1 GiB, on a low-disk host.
+ backend = _resume_backend(tmp_path)
+ backend._estimate_kv_cache_bytes = lambda *a, **k: 0 # metadata unavailable
+ monkeypatch.setattr(llama_cpp, "_SLOT_SAVE_MAX_BYTES", 8 << 30) # 8 GiB cap
+ _fake_disk(monkeypatch, free = 2 << 30) # 2 GiB free < 8 + 1 GiB required
+ monkeypatch.setattr(
+ llama_cpp.httpx,
+ "post",
+ lambda *a, **k: (_ for _ in ()).throw(AssertionError),
+ raising = False,
+ )
+ assert backend.save_slots_for_resume() is None
diff --git a/studio/backend/tests/test_llama_server_args.py b/studio/backend/tests/test_llama_server_args.py
index c6d16363f8..fa4ba71791 100644
--- a/studio/backend/tests/test_llama_server_args.py
+++ b/studio/backend/tests/test_llama_server_args.py
@@ -183,6 +183,8 @@ def test_non_flag_token_passes_through():
"--reranking",
# llama-server's own --tools clashes with Unsloth's tool policy.
"--tools",
+ # Slot-state dir: Studio owns it for KV persistence across idle unload.
+ "--slot-save-path",
],
)
def test_denylist_rejects_all_aliases(denied):
@@ -224,6 +226,16 @@ def test_denylist_rejects_equals_form():
validate_extra_args(["--port=9000"])
+def test_slot_save_path_is_managed_in_all_forms():
+ for args in (["--slot-save-path", "/tmp/x"], ["--slot-save-path=/tmp/x"], ["--slot-save-path"]):
+ with pytest.raises(ValueError, match = "--slot-save-path"):
+ validate_extra_args(args)
+ assert is_managed_flag("--slot-save-path") is True
+ assert is_managed_flag("--slot-save-path=/tmp/x") is True
+ # --slots (read-only diagnostics endpoint) stays a user choice.
+ assert is_managed_flag("--slots") is False
+
+
@pytest.mark.parametrize(
"padded",
[" --parallel", "--parallel ", "\t--parallel", " -np", "-np \n", "-np\t"],
diff --git a/studio/backend/tests/test_openai_auto_switch.py b/studio/backend/tests/test_openai_auto_switch.py
index c4c0ce15c9..1ee9ef36d3 100644
--- a/studio/backend/tests/test_openai_auto_switch.py
+++ b/studio/backend/tests/test_openai_auto_switch.py
@@ -8,6 +8,7 @@ tests/test_gguf_completion_usage.py.
"""
import asyncio
+import os
import pytest
@@ -18,6 +19,10 @@ from utils import openai_auto_switch_settings as settings
class _FakeBackend:
+ effective_parallel_slots = 1
+ _slot_save_binary = None
+ _gguf_path = None
+
def __init__(
self,
loaded_id = None,
@@ -29,6 +34,22 @@ class _FakeBackend:
self.hf_variant = hf_variant
self._openai_advertised_id = advertised_id
+ def save_slots_for_resume(self, should_abort = None):
+ return None
+
+ def restore_slots_for_resume(self, manifest):
+ return None
+
+ def _slot_launch_fingerprint(self):
+ return ((), None, None, 1)
+
+ def _gguf_file_identity(self, path):
+ try:
+ st = os.stat(path)
+ except OSError:
+ return None
+ return ((st.st_size, st.st_mtime_ns),)
+
class _LoadRecorder:
"""Stand-in for the load route: records calls and simulates a load."""
@@ -53,10 +74,15 @@ class _LoadRecorder:
from fastapi import HTTPException
raise HTTPException(status_code = 503, detail = "load failed")
self.backend.model_identifier = request.model_path
+ self.backend.hf_variant = getattr(request, "gguf_variant", None)
+ self.backend._gguf_path = request.model_path
self.backend.is_loaded = True
# Mirror _load_model_impl: a load advertises its own id until the
# auto-switch caller overwrites it with the repo id.
self.backend._openai_advertised_id = None
+ from core.inference import llama_keepwarm as kw
+
+ kw.note_model_loaded(self.backend)
return None
@@ -446,6 +472,75 @@ def test_idle_loop_unloads_after_ttl_and_stashes_for_reload(monkeypatch):
assert stash is not None and stash[0] == "unsloth/Idle-GGUF" and stash[1] == "Q4_K_M"
+def test_idle_loop_deletes_saved_kv_when_unload_fails(monkeypatch, tmp_path):
+ import time
+ from core.inference import llama_keepwarm as kw
+
+ monkeypatch.setattr(settings, "get_auto_unload_idle_seconds", lambda: 0.005)
+ monkeypatch.setattr(settings, "get_auto_unload_keep_kv", lambda: True)
+ kw._inflight = 0
+ kw._pending = 0
+ kw._last_active = time.monotonic() - 3600
+ kw._last_unloaded_model = None
+ kw._kv_resume = None
+
+ saved = tmp_path / "resume-abc-slot0.bin"
+ backend = _FakeBackend("unsloth/Idle-GGUF")
+ manifests = []
+
+ def _save(should_abort = None):
+ if manifests:
+ return None
+ saved.write_bytes(b"kv")
+ manifest = {"dir": str(tmp_path), "slots": [{"id": 0, "filename": saved.name}]}
+ manifests.append(manifest)
+ return manifest
+
+ def _unload():
+ raise RuntimeError("cuda teardown failed")
+
+ backend.save_slots_for_resume = _save
+ backend.unload_model = _unload
+ monkeypatch.setattr(inference_route, "get_llama_cpp_backend", lambda: backend)
+
+ async def _drive():
+ task = asyncio.create_task(kw.idle_unload_loop(poll_seconds = 0.01))
+ for _ in range(200):
+ await asyncio.sleep(0.01)
+ if manifests and not saved.exists():
+ break
+ task.cancel()
+ try:
+ await task
+ except asyncio.CancelledError:
+ pass
+
+ asyncio.run(_drive())
+ assert manifests and not saved.exists()
+ assert kw._kv_resume is None
+
+
+def test_disabling_idle_unload_purges_saved_kv(monkeypatch, tmp_path):
+ # PUT leaves keep-KV on but makes idle unload inactive: saved KV must go too.
+ import routes.settings as settings_route
+ from core.inference import llama_keepwarm as kw
+
+ saved = tmp_path / "resume-abc-slot0.bin"
+ saved.write_bytes(b"kv")
+ kw._kv_resume = {
+ "identity": ("m", None, "m"),
+ "dir": str(tmp_path),
+ "slots": [{"id": 0, "filename": saved.name}],
+ }
+ monkeypatch.setattr(settings_route, "set_openai_auto_switch", lambda *a: (False, 300, True))
+ monkeypatch.setattr(settings_route, "get_auto_unload_idle_seconds", lambda: 0)
+
+ payload = settings_route.OpenAIAutoSwitchPayload(enabled = False)
+ resp = settings_route.update_openai_auto_switch(payload, "tester")
+ assert resp.idle_unload_active is False and resp.auto_unload_keep_kv is True
+ assert kw._kv_resume is None and not saved.exists()
+
+
def test_audio_generate_is_tracked_as_inference_path():
# Direct GGUF TTS uses the llama backend and can outlive the idle TTL, so
# the keep-warm middleware must count it as in-flight inference.
@@ -2912,8 +3007,10 @@ def test_non_gguf_load_clears_reload_stash():
# A non-GGUF (Transformers/Unsloth) load must clear the stash like the GGUF
# branch, so it never lingers until the idle poll (or forever, idle-unload off).
import inspect
+
src = inspect.getsource(inference_route._load_model_impl)
- assert src.count("note_model_loaded()") >= 2
+ assert src.count("note_model_loaded()") >= 1 # non-GGUF branch
+ assert "to_thread(note_model_loaded, llama_backend)" in src # GGUF branch
def test_chat_rejects_malformed_tool_choice_before_switch(monkeypatch):
@@ -3121,6 +3218,495 @@ def test_responses_stream_hint_matches_toggle_regardless_of_active_model(monkeyp
assert "Model auto-switch" in non_gguf_loaded
+# ── idle-unload KV persistence (slot save/restore) ──────────────────
+
+
+def _seed_kv_manifest(
+ tmp_path,
+ identity = ("unsloth/A-GGUF", "Q4_K_M", "unsloth/A-GGUF"),
+ gguf = None,
+):
+ if gguf is None:
+ gguf_file = tmp_path / "model.gguf"
+ gguf_file.write_bytes(b"gguf")
+ gguf = str(gguf_file)
+ st = os.stat(gguf)
+ state_file = tmp_path / "resume-abc-slot0.bin"
+ state_file.write_bytes(b"kv")
+ return state_file, {
+ "identity": identity,
+ "dir": str(tmp_path),
+ "binary": ("/bin/llama-server", 111),
+ "gguf": gguf,
+ "gguf_stat": ((st.st_size, st.st_mtime_ns),),
+ "launch": ((), None, None, 1),
+ "slots": [{"id": 0, "filename": state_file.name, "n_saved": 42}],
+ }
+
+
+def _drive_idle_loop(
+ kw,
+ poll_seconds = 0.02,
+ run_for = 0.2,
+):
+ async def _drive():
+ task = asyncio.create_task(kw.idle_unload_loop(poll_seconds = poll_seconds))
+ await asyncio.sleep(run_for)
+ task.cancel()
+ try:
+ await task
+ except asyncio.CancelledError:
+ pass
+
+ asyncio.run(_drive())
+
+
+def test_idle_unload_saves_slots_before_unload_and_stashes_manifest(monkeypatch, tmp_path):
+ import time
+ from core.inference import llama_keepwarm as kw
+
+ monkeypatch.setattr(settings, "get_auto_unload_idle_seconds", lambda: 0.005)
+ monkeypatch.setattr(settings, "get_auto_unload_keep_kv", lambda: True)
+ kw._inflight = 0
+ kw._pending = 0
+ kw._last_active = time.monotonic() - 3600
+ kw._last_unloaded_model = None
+ kw._kv_resume = None
+
+ events = []
+ backend = _FakeBackend("unsloth/Idle-GGUF", hf_variant = "Q4_K_M")
+ manifest = {
+ "dir": str(tmp_path),
+ "binary": ("bin", 1),
+ "slots": [{"id": 0, "filename": "f.bin", "n_saved": 42}],
+ }
+
+ def _save(should_abort = None):
+ events.append("save")
+ return manifest
+
+ def _unload():
+ events.append("unload")
+ backend.is_loaded = False
+
+ backend.save_slots_for_resume = _save
+ backend.unload_model = _unload
+ monkeypatch.setattr(inference_route, "get_llama_cpp_backend", lambda: backend)
+
+ _drive_idle_loop(kw)
+ # KV must be saved while the server is still alive, then exactly one unload.
+ assert events == ["save", "unload"]
+ assert kw.get_last_unloaded_model()[:2] == ("unsloth/Idle-GGUF", "Q4_K_M")
+ resume = kw.take_kv_resume()
+ assert resume is not None
+ assert resume["identity"][:2] == ("unsloth/Idle-GGUF", "Q4_K_M")
+ assert resume["slots"][0]["filename"] == "f.bin"
+
+
+def test_idle_save_failure_still_unloads_plain(monkeypatch):
+ import time
+ from core.inference import llama_keepwarm as kw
+
+ monkeypatch.setattr(settings, "get_auto_unload_idle_seconds", lambda: 0.005)
+ monkeypatch.setattr(settings, "get_auto_unload_keep_kv", lambda: True)
+ kw._inflight = 0
+ kw._pending = 0
+ kw._last_active = time.monotonic() - 3600
+ kw._last_unloaded_model = None
+ kw._kv_resume = None
+
+ unloads = []
+ backend = _FakeBackend("unsloth/Idle-GGUF", hf_variant = "Q4_K_M")
+
+ def _save(should_abort = None):
+ raise RuntimeError("slot save exploded")
+
+ def _unload():
+ unloads.append(1)
+ backend.is_loaded = False
+
+ backend.save_slots_for_resume = _save
+ backend.unload_model = _unload
+ monkeypatch.setattr(inference_route, "get_llama_cpp_backend", lambda: backend)
+
+ _drive_idle_loop(kw)
+ assert unloads == [1] # the save failure must not skip the unload
+ assert kw.get_last_unloaded_model() is not None
+ assert kw.take_kv_resume() is None
+
+
+def test_keep_kv_setting_off_skips_save(monkeypatch):
+ import time
+ from core.inference import llama_keepwarm as kw
+
+ monkeypatch.setattr(settings, "get_auto_unload_idle_seconds", lambda: 0.005)
+ monkeypatch.setattr(settings, "get_auto_unload_keep_kv", lambda: False)
+ kw._inflight = 0
+ kw._pending = 0
+ kw._last_active = time.monotonic() - 3600
+ kw._last_unloaded_model = None
+ kw._kv_resume = None
+
+ saves, unloads = [], []
+ backend = _FakeBackend("unsloth/Idle-GGUF")
+
+ def _unload():
+ unloads.append(1)
+ backend.is_loaded = False
+
+ backend.save_slots_for_resume = lambda *a, **k: saves.append(1)
+ backend.unload_model = _unload
+ monkeypatch.setattr(inference_route, "get_llama_cpp_backend", lambda: backend)
+
+ _drive_idle_loop(kw)
+ assert saves == []
+ assert unloads == [1]
+ assert kw.take_kv_resume() is None
+
+
+def test_keep_kv_disabled_mid_save_discards_manifest(monkeypatch, tmp_path):
+ import time
+ from core.inference import llama_keepwarm as kw
+
+ keep = {"on": True}
+ monkeypatch.setattr(settings, "get_auto_unload_idle_seconds", lambda: 0.005)
+ monkeypatch.setattr(settings, "get_auto_unload_keep_kv", lambda: keep["on"])
+ kw._inflight = 0
+ kw._pending = 0
+ kw._last_active = time.monotonic() - 3600
+ kw._last_unloaded_model = None
+ kw._kv_resume = None
+
+ unloads = []
+ backend = _FakeBackend("unsloth/Idle-GGUF", hf_variant = "Q4_K_M")
+ state_file = tmp_path / "resume-mid-slot0.bin"
+ state_file.write_bytes(b"kv")
+ manifest = {
+ "dir": str(tmp_path),
+ "binary": ("bin", 1),
+ "slots": [{"id": 0, "filename": state_file.name, "n_saved": 1}],
+ }
+
+ def _save(should_abort = None):
+ keep["on"] = False # user flips the toggle while the save runs
+ return manifest
+
+ def _unload():
+ unloads.append(1)
+ backend.is_loaded = False
+
+ backend.save_slots_for_resume = _save
+ backend.unload_model = _unload
+ monkeypatch.setattr(inference_route, "get_llama_cpp_backend", lambda: backend)
+
+ _drive_idle_loop(kw)
+ assert unloads == [1] # still unloads; only the stash is dropped
+ assert kw.take_kv_resume() is None
+ assert not state_file.exists()
+
+
+def test_idle_ttl_disabled_mid_save_skips_unload(monkeypatch, tmp_path):
+ import time
+ from core.inference import llama_keepwarm as kw
+
+ ttl = {"v": 0.005}
+ monkeypatch.setattr(settings, "get_auto_unload_idle_seconds", lambda: ttl["v"])
+ monkeypatch.setattr(settings, "get_auto_unload_keep_kv", lambda: True)
+ kw._inflight = 0
+ kw._pending = 0
+ kw._last_active = time.monotonic() - 3600
+ kw._last_unloaded_model = None
+ kw._kv_resume = None
+
+ unloads = []
+ backend = _FakeBackend("unsloth/Idle-GGUF", hf_variant = "Q4_K_M")
+ state_file = tmp_path / "resume-mid-slot0.bin"
+ state_file.write_bytes(b"kv")
+ manifest = {
+ "dir": str(tmp_path),
+ "binary": ("bin", 1),
+ "slots": [{"id": 0, "filename": state_file.name, "n_saved": 1}],
+ }
+
+ def _save(should_abort = None):
+ ttl["v"] = 0 # user turns idle unload off while the save runs
+ return manifest
+
+ backend.save_slots_for_resume = _save
+ backend.unload_model = lambda: unloads.append(1)
+ monkeypatch.setattr(inference_route, "get_llama_cpp_backend", lambda: backend)
+
+ _drive_idle_loop(kw)
+ assert unloads == [] # the unload was cancelled by the setting change
+ assert kw.take_kv_resume() is None
+ assert not state_file.exists()
+
+
+def test_alias_reload_restores_slots_and_deletes_files(monkeypatch, tmp_path):
+ from core.inference import llama_keepwarm as kw
+
+ backend = _FakeBackend(None) # idle-unload emptied the backend
+ backend._slot_save_binary = ("/bin/llama-server", 111)
+ restored = []
+ backend.restore_slots_for_resume = lambda manifest: restored.append(manifest)
+
+ rec = _LoadRecorder(backend)
+ _wire(monkeypatch, enabled = True, resolves_to = None, backend = backend, recorder = rec)
+ monkeypatch.setattr(kw, "_inflight", 0)
+ state_file, manifest = _seed_kv_manifest(tmp_path)
+ monkeypatch.setattr(kw, "_last_unloaded_model", (manifest["gguf"], "Q4_K_M"))
+ monkeypatch.setattr(kw, "_kv_resume", manifest)
+
+ _run_hook("gpt-4o-mini")
+ assert len(rec.calls) == 1
+ assert len(restored) == 1 # same model + binary: restore ran
+ assert not state_file.exists() # state file deleted after the restore
+ assert kw._kv_resume is None
+
+
+def test_no_restore_when_different_model_loads(monkeypatch, tmp_path):
+ from core.inference import llama_keepwarm as kw
+
+ backend = _FakeBackend(None)
+ backend._slot_save_binary = ("/bin/llama-server", 111)
+ restored = []
+ backend.restore_slots_for_resume = lambda manifest: restored.append(manifest)
+ rec = _LoadRecorder(backend)
+ _wire(
+ monkeypatch,
+ enabled = True,
+ resolves_to = ("unsloth/B-GGUF", None, "unsloth/B-GGUF"),
+ backend = backend,
+ recorder = rec,
+ )
+ monkeypatch.setattr(kw, "_inflight", 0)
+ state_file, manifest = _seed_kv_manifest(tmp_path) # manifest is for model A
+ monkeypatch.setattr(kw, "_kv_resume", manifest)
+
+ _run_hook("unsloth/B-GGUF")
+ assert len(rec.calls) == 1
+ assert restored == [] # different model: never restored
+ assert not state_file.exists() # but the stale files are gone
+ assert kw._kv_resume is None
+
+
+def test_restore_skipped_when_binary_changed(monkeypatch, tmp_path):
+ from core.inference import llama_keepwarm as kw
+
+ state_file, manifest = _seed_kv_manifest(tmp_path)
+ backend = _FakeBackend("unsloth/A-GGUF", hf_variant = "Q4_K_M")
+ backend._gguf_path = manifest["gguf"]
+ backend._slot_save_binary = ("/bin/llama-server", 222) # newer mtime
+ restored = []
+ backend.restore_slots_for_resume = lambda manifest: restored.append(manifest)
+
+ kw.restore_kv_resume(backend, manifest)
+ assert restored == []
+ assert not state_file.exists()
+
+
+def test_restore_skipped_when_launch_config_changed(tmp_path):
+ from core.inference import llama_keepwarm as kw
+
+ state_file, manifest = _seed_kv_manifest(tmp_path)
+ backend = _FakeBackend("unsloth/A-GGUF", hf_variant = "Q4_K_M")
+ backend._gguf_path = manifest["gguf"]
+ backend._slot_save_binary = ("/bin/llama-server", 111)
+ backend._slot_launch_fingerprint = lambda: (("--rope-freq-scale", "0.5"), None, None, 1)
+ restored = []
+ backend.restore_slots_for_resume = lambda manifest: restored.append(manifest)
+
+ kw.restore_kv_resume(backend, manifest)
+ assert restored == []
+ assert not state_file.exists()
+
+
+def test_restore_skipped_when_gguf_rewritten_in_place(tmp_path):
+ from core.inference import llama_keepwarm as kw
+
+ state_file, manifest = _seed_kv_manifest(tmp_path)
+ with open(manifest["gguf"], "wb") as fh:
+ fh.write(b"different weights") # same path, new content
+ backend = _FakeBackend("unsloth/A-GGUF", hf_variant = "Q4_K_M")
+ backend._gguf_path = manifest["gguf"]
+ backend._slot_save_binary = ("/bin/llama-server", 111)
+ restored = []
+ backend.restore_slots_for_resume = lambda manifest: restored.append(manifest)
+
+ kw.restore_kv_resume(backend, manifest)
+ assert restored == []
+ assert not state_file.exists()
+
+
+def test_note_model_unloaded_purges_manifest_and_files(tmp_path):
+ from core.inference import llama_keepwarm as kw
+
+ state_file, manifest = _seed_kv_manifest(tmp_path)
+ kw._set_last_unloaded(("org/A-GGUF", "Q4_K_M"))
+ kw._set_kv_resume(manifest)
+ kw.note_model_unloaded()
+ assert kw.get_last_unloaded_model() is None
+ assert kw.take_kv_resume() is None
+ assert not state_file.exists()
+
+
+def test_note_model_loaded_purges_manifest_and_files(tmp_path):
+ from core.inference import llama_keepwarm as kw
+
+ state_file, manifest = _seed_kv_manifest(tmp_path)
+ kw._set_last_unloaded(("org/A-GGUF", "Q4_K_M"))
+ kw._set_kv_resume(manifest)
+ kw.note_model_loaded()
+ assert kw.get_last_unloaded_model() is None
+ assert kw.take_kv_resume() is None
+ assert not state_file.exists()
+
+
+def test_new_idle_save_purges_previous_manifest_files(tmp_path):
+ from core.inference import llama_keepwarm as kw
+
+ old_file, old_manifest = _seed_kv_manifest(tmp_path)
+ kw._set_kv_resume(old_manifest)
+ new_file = tmp_path / "resume-def-slot0.bin"
+ new_file.write_bytes(b"kv2")
+ kw._set_kv_resume(
+ {
+ "identity": ("unsloth/B-GGUF", None, "unsloth/B-GGUF"),
+ "dir": str(tmp_path),
+ "binary": ("/bin/llama-server", 111),
+ "slots": [{"id": 0, "filename": new_file.name, "n_saved": 7}],
+ }
+ )
+ assert not old_file.exists() # replaced manifest's files purged
+ assert new_file.exists()
+ assert kw.take_kv_resume()["slots"][0]["filename"] == new_file.name
+
+
+def test_sweep_slot_save_dir_removes_only_resume_files(monkeypatch, tmp_path):
+ from core.inference import llama_keepwarm as kw
+ from utils.paths import storage_roots
+
+ monkeypatch.setattr(storage_roots, "llama_slot_cache_root", lambda: tmp_path)
+ stale = tmp_path / "resume-old-slot0.bin"
+ stale.write_bytes(b"kv")
+ other = tmp_path / "unrelated.txt"
+ other.write_text("keep")
+ kw.sweep_slot_save_dir()
+ assert not stale.exists()
+ assert other.exists()
+
+
+def test_keep_kv_setting_roundtrip_and_default(monkeypatch):
+ import storage.studio_db as db
+
+ store = {}
+ monkeypatch.setattr(db, "upsert_app_settings", lambda m: store.update(m))
+ monkeypatch.setattr(settings, "_cached_setting", lambda k, d = None: store.get(k, d))
+
+ assert settings.get_auto_unload_keep_kv() is True # default when never stored
+ assert settings.set_openai_auto_switch(True, 60, False)[2] is False
+ assert store[settings.AUTO_UNLOAD_KEEP_KV_SETTING_KEY] is False
+ assert settings.get_auto_unload_keep_kv() is False
+ # None leaves the stored value untouched (older clients can't reset it).
+ assert settings.set_openai_auto_switch(True, 60, None)[2] is False
+ assert store[settings.AUTO_UNLOAD_KEEP_KV_SETTING_KEY] is False
+ with pytest.raises(ValueError, match = "true or false"):
+ settings.set_openai_auto_switch(True, 60, "garbage")
+
+
+def test_stale_stash_cleanup_waits_for_lifecycle_gate(monkeypatch, tmp_path):
+ # The loop's stale-stash purge must wait on the gate a mid-reload holds.
+ import time
+ from core.inference import llama_keepwarm as kw
+
+ monkeypatch.setattr(settings, "get_auto_unload_idle_seconds", lambda: 3600)
+ kw._inflight = 0
+ kw._pending = 0
+ kw._last_active = time.monotonic()
+ backend = _FakeBackend("unsloth/New-GGUF")
+ monkeypatch.setattr(inference_route, "get_llama_cpp_backend", lambda: backend)
+ state_file, manifest = _seed_kv_manifest(tmp_path)
+ kw._kv_resume = manifest
+ kw._last_unloaded_model = ("unsloth/A-GGUF", "Q4_K_M")
+
+ assert kw._lifecycle_lock.acquire(blocking = False) # simulate in-flight reload
+ try:
+ _drive_idle_loop(kw)
+ assert kw._kv_resume is manifest # purge deferred while the gate is held
+ assert state_file.exists()
+ finally:
+ kw._lifecycle_lock.release()
+ _drive_idle_loop(kw)
+ assert kw._kv_resume is None # gate freed: genuinely stale stash purged
+ assert not state_file.exists()
+
+
+def test_put_route_disabling_keep_kv_purges_saved_state(monkeypatch, tmp_path):
+ import routes.settings as settings_route
+ import storage.studio_db as db
+ from core.inference import llama_keepwarm as kw
+
+ store = {}
+ monkeypatch.setattr(db, "upsert_app_settings", lambda m: store.update(m))
+ monkeypatch.setattr(settings, "_cached_setting", lambda k, d = None: store.get(k, d))
+ state_file, manifest = _seed_kv_manifest(tmp_path)
+ monkeypatch.setattr(kw, "_kv_resume", manifest)
+
+ payload = settings_route.OpenAIAutoSwitchPayload(enabled = True, auto_unload_keep_kv = False)
+ resp = settings_route.update_openai_auto_switch(payload, "tester")
+ assert resp.auto_unload_keep_kv is False
+ assert kw._kv_resume is None
+ assert not state_file.exists()
+
+
+def test_keep_kv_only_update_leaves_env_idle_ttl_active(monkeypatch):
+ # A keep-KV-only update must not materialize the env TTL as a stored value.
+ import routes.settings as settings_route
+ import storage.studio_db as db
+
+ store = {}
+ monkeypatch.setattr(db, "upsert_app_settings", lambda m: store.update(m))
+ monkeypatch.setattr(settings, "_cached_setting", lambda k, d = None: store.get(k, d))
+ monkeypatch.setenv(settings.MODEL_IDLE_TTL_ENV_VAR, "600")
+
+ assert settings_route.OpenAIAutoSwitchPayload(enabled = False).auto_unload_idle_seconds is None
+ enabled, idle, keep_kv = settings.set_openai_auto_switch(False, None, False)
+ assert settings.AUTO_UNLOAD_IDLE_SETTING_KEY not in store # idle untouched
+ assert settings.get_auto_unload_idle_seconds() == 600 # env TTL still active
+ assert (enabled, idle, keep_kv) == (False, 600, False)
+
+
+def test_load_impl_notes_loaded_with_backend_off_loop():
+ import inspect
+ src = inspect.getsource(inference_route._load_model_impl)
+ assert "to_thread(note_model_loaded, llama_backend)" in src
+
+
+def test_restore_matches_gguf_realpath_across_naming(tmp_path):
+ from core.inference import llama_keepwarm as kw
+
+ blob = tmp_path / "blob.gguf"
+ blob.write_bytes(b"gguf")
+ link = tmp_path / "snapshot.gguf"
+ try:
+ link.symlink_to(blob)
+ except OSError:
+ pytest.skip("symlinks unsupported on this host")
+
+ backend = _FakeBackend("/hf/snapshots/d7f5", hf_variant = None)
+ backend._gguf_path = str(link) # reload resolved the symlink spelling
+ backend._slot_save_binary = ("/bin/llama-server", 111)
+ restored = []
+ backend.restore_slots_for_resume = lambda manifest: restored.append(manifest)
+ state_file, manifest = _seed_kv_manifest(
+ tmp_path, identity = ("unsloth/A-GGUF", None, "unsloth/A-GGUF"), gguf = str(blob)
+ )
+
+ kw.restore_kv_resume(backend, manifest)
+ assert len(restored) == 1 # names differ, file identical: restore ran
+ assert not state_file.exists()
+
+
def test_setter_rejects_idle_below_floor(monkeypatch):
import storage.studio_db as db
diff --git a/studio/backend/utils/openai_auto_switch_settings.py b/studio/backend/utils/openai_auto_switch_settings.py
index 462435e5d5..7007440f4c 100644
--- a/studio/backend/utils/openai_auto_switch_settings.py
+++ b/studio/backend/utils/openai_auto_switch_settings.py
@@ -30,11 +30,13 @@ from typing import Any, Optional
OPENAI_AUTO_SWITCH_SETTING_KEY = "openai_api_auto_switch_model"
AUTO_UNLOAD_IDLE_SETTING_KEY = "openai_api_auto_unload_idle_seconds"
+AUTO_UNLOAD_KEEP_KV_SETTING_KEY = "openai_api_auto_unload_keep_kv"
MODEL_OVERRIDES_SETTING_KEY = "openai_api_auto_switch_overrides"
MODEL_IDLE_TTL_ENV_VAR = "UNSLOTH_MODEL_IDLE_TTL"
DEFAULT_OPENAI_AUTO_SWITCH_ENABLED = False
DEFAULT_AUTO_UNLOAD_IDLE_SECONDS = 0
+DEFAULT_AUTO_UNLOAD_KEEP_KV = True
MIN_AUTO_UNLOAD_IDLE_SECONDS = 60
_CACHE_TTL_S = 2.0
@@ -158,29 +160,54 @@ def get_auto_unload_idle_seconds() -> int:
return env if env is not None else 0
-def set_openai_auto_switch(enabled: Any, idle_seconds: Any) -> tuple[bool, int]:
- """Set both auto-switch flags in one transaction so a settings PUT can't leave
- one key updated and the other stale. Both values are coerced before any write,
- so an invalid value raises without persisting either."""
+def get_auto_unload_keep_kv() -> bool:
+ """Whether the idle unload persists slot KV to disk for restore on reload."""
+ parsed = _coerce_bool(_cached_setting(AUTO_UNLOAD_KEEP_KV_SETTING_KEY, None))
+ return parsed if parsed is not None else DEFAULT_AUTO_UNLOAD_KEEP_KV
+
+
+def set_openai_auto_switch(
+ enabled: Any,
+ idle_seconds: Any,
+ keep_kv: Any = None,
+) -> tuple[bool, int, bool]:
+ """One-transaction write; ``None`` leaves a stored value untouched."""
parsed_enabled = _coerce_bool(enabled)
if parsed_enabled is None:
raise ValueError("OpenAI auto-switch must be true or false.")
- parsed_idle = _coerce_int(idle_seconds)
- if parsed_idle is None:
- raise ValueError("Auto-unload idle seconds must be a non-negative integer.")
- if 0 < parsed_idle < MIN_AUTO_UNLOAD_IDLE_SECONDS:
- raise ValueError(
- f"Auto-unload idle seconds must be 0 (off) or at least "
- f"{MIN_AUTO_UNLOAD_IDLE_SECONDS}."
- )
+ parsed_idle = None
+ if idle_seconds is not None:
+ parsed_idle = _coerce_int(idle_seconds)
+ if parsed_idle is None:
+ raise ValueError("Auto-unload idle seconds must be a non-negative integer.")
+ if 0 < parsed_idle < MIN_AUTO_UNLOAD_IDLE_SECONDS:
+ raise ValueError(
+ f"Auto-unload idle seconds must be 0 (off) or at least "
+ f"{MIN_AUTO_UNLOAD_IDLE_SECONDS}."
+ )
+ parsed_keep_kv = None
+ if keep_kv is not None:
+ parsed_keep_kv = _coerce_bool(keep_kv)
+ if parsed_keep_kv is None:
+ raise ValueError("Keep KV on idle unload must be true or false.")
from storage.studio_db import upsert_app_settings
- upsert_app_settings(
- {OPENAI_AUTO_SWITCH_SETTING_KEY: parsed_enabled, AUTO_UNLOAD_IDLE_SETTING_KEY: parsed_idle}
- )
+ updates: dict[str, Any] = {OPENAI_AUTO_SWITCH_SETTING_KEY: parsed_enabled}
+ if parsed_idle is not None:
+ updates[AUTO_UNLOAD_IDLE_SETTING_KEY] = parsed_idle
+ if parsed_keep_kv is not None:
+ updates[AUTO_UNLOAD_KEEP_KV_SETTING_KEY] = parsed_keep_kv
+ upsert_app_settings(updates)
_invalidate(OPENAI_AUTO_SWITCH_SETTING_KEY)
- _invalidate(AUTO_UNLOAD_IDLE_SETTING_KEY)
- return parsed_enabled, parsed_idle
+ if parsed_idle is not None:
+ _invalidate(AUTO_UNLOAD_IDLE_SETTING_KEY)
+ if parsed_keep_kv is not None:
+ _invalidate(AUTO_UNLOAD_KEEP_KV_SETTING_KEY)
+ return (
+ parsed_enabled,
+ parsed_idle if parsed_idle is not None else get_stored_auto_unload_idle_seconds(),
+ parsed_keep_kv if parsed_keep_kv is not None else get_auto_unload_keep_kv(),
+ )
def get_model_overrides() -> dict[str, dict]:
diff --git a/studio/backend/utils/paths/storage_roots.py b/studio/backend/utils/paths/storage_roots.py
index 1faa2b1281..35b8c57e9b 100644
--- a/studio/backend/utils/paths/storage_roots.py
+++ b/studio/backend/utils/paths/storage_roots.py
@@ -61,6 +61,11 @@ def cache_root() -> Path:
return studio_root() / "cache"
+def llama_slot_cache_root() -> Path:
+ """Dir llama-server saves/restores slot KV state in across idle unloads."""
+ return cache_root() / "llama-slots"
+
+
def studio_bin_root() -> Path:
"""Dir for Unsloth-managed executables (the `unsloth` shim, downloaded tools like cloudflared)."""
return studio_root() / "bin"
diff --git a/studio/frontend/src/features/settings/api/openai-auto-switch.ts b/studio/frontend/src/features/settings/api/openai-auto-switch.ts
index 80ffc084d0..47bad56eab 100644
--- a/studio/frontend/src/features/settings/api/openai-auto-switch.ts
+++ b/studio/frontend/src/features/settings/api/openai-auto-switch.ts
@@ -11,6 +11,8 @@ export type OpenAIAutoSwitchSettings = {
// True when the idle-unload loop will actually unload (e.g. enabled via the
// UNSLOTH_MODEL_IDLE_TTL env var even while the toggle is off).
idleUnloadActive: boolean;
+ // Persist the KV cache to disk on idle unload and restore it on reload.
+ autoUnloadKeepKv: boolean;
};
type ApiOpenAIAutoSwitchSettings = {
@@ -21,6 +23,8 @@ type ApiOpenAIAutoSwitchSettings = {
default_enabled: boolean;
// biome-ignore lint/style/useNamingConvention: API schema
idle_unload_active?: boolean;
+ // biome-ignore lint/style/useNamingConvention: API schema
+ auto_unload_keep_kv?: boolean;
};
let cachedSettings: OpenAIAutoSwitchSettings | null = null;
@@ -34,6 +38,7 @@ function fromApi(
autoUnloadIdleSeconds: settings.auto_unload_idle_seconds,
defaultEnabled: settings.default_enabled,
idleUnloadActive: settings.idle_unload_active ?? false,
+ autoUnloadKeepKv: settings.auto_unload_keep_kv ?? true,
};
}
@@ -66,15 +71,23 @@ export async function loadOpenAIAutoSwitchSettings() {
export async function updateOpenAIAutoSwitchSettings(
enabled: boolean,
- autoUnloadIdleSeconds: number,
+ autoUnloadIdleSeconds?: number,
+ autoUnloadKeepKv?: boolean,
): Promise {
const res = await authFetch("/api/settings/openai-auto-switch", {
method: "PUT",
headers: { "Content-Type": "application/json" },
body: JSON.stringify({
enabled,
- // biome-ignore lint/style/useNamingConvention: API schema
- auto_unload_idle_seconds: autoUnloadIdleSeconds,
+ // Omitted fields keep their stored value.
+ ...(autoUnloadIdleSeconds === undefined
+ ? {}
+ : // biome-ignore lint/style/useNamingConvention: API schema
+ { auto_unload_idle_seconds: autoUnloadIdleSeconds }),
+ ...(autoUnloadKeepKv === undefined
+ ? {}
+ : // biome-ignore lint/style/useNamingConvention: API schema
+ { auto_unload_keep_kv: autoUnloadKeepKv }),
}),
});
if (!res.ok) {
diff --git a/studio/frontend/src/features/settings/components/model-auto-switch-section.tsx b/studio/frontend/src/features/settings/components/model-auto-switch-section.tsx
index 32b3e53a2c..aa6857cff5 100644
--- a/studio/frontend/src/features/settings/components/model-auto-switch-section.tsx
+++ b/studio/frontend/src/features/settings/components/model-auto-switch-section.tsx
@@ -62,13 +62,18 @@ export function ModelAutoSwitchSection() {
const persist = async (
enabled: boolean,
- idleSeconds: number,
+ idleSeconds: number | undefined,
syncDraft = true,
+ keepKv?: boolean,
) => {
setIsSaving(true);
setError(null);
try {
- const saved = await updateOpenAIAutoSwitchSettings(enabled, idleSeconds);
+ const saved = await updateOpenAIAutoSwitchSettings(
+ enabled,
+ idleSeconds,
+ keepKv,
+ );
setSettings(saved);
if (syncDraft) {
setDraftIdleSeconds(String(saved.autoUnloadIdleSeconds));
@@ -107,6 +112,11 @@ export function ModelAutoSwitchSection() {
void persist(true, idleSeconds);
};
+ const handleKeepKvToggle = (keepKv: boolean) => {
+ if (!settings) return;
+ void persist(settings.enabled, undefined, false, keepKv);
+ };
+
return (
+ {settings?.idleUnloadActive ? (
+
+
+
+ ) : null}
);
}
diff --git a/studio/frontend/src/i18n/locales/en.ts b/studio/frontend/src/i18n/locales/en.ts
index de8ac17c29..cbddc9f0c2 100644
--- a/studio/frontend/src/i18n/locales/en.ts
+++ b/studio/frontend/src/i18n/locales/en.ts
@@ -232,6 +232,9 @@ export const en = {
loadError: "Failed to load model auto-switch settings.",
saveError: "Failed to save model auto-switch settings.",
idleError: "Enter 0 to keep the model loaded, or at least 60 seconds.",
+ keepKv: "Keep chat context across idle unload",
+ keepKvDescription:
+ "Save the model's KV cache to disk before an idle unload and restore it on reload, so resumed chats skip re-reading their history. Chat context is written to disk (up to 10 GB) until it is restored or cleaned up.",
},
previewSharing: {
sectionTitle: "Preview sharing",
From 9e334d552c77de8cc4b52ff1d2b13a93891d1366 Mon Sep 17 00:00:00 2001
From: alkinun
Date: Mon, 20 Jul 2026 10:23:37 +0300
Subject: [PATCH 11/41] Fix text-only VLM CPT packing truncation (#7211)
* Fix text-only VLM CPT packing truncation
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Handle streaming vision datasets in packing
* Harden multimodal packing detection
* Preserve safe packing boundaries
* Scope stream packing checks to VLMs
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Narrow VLM packing detection
* Align packing mode and eval safety
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Add qwen3_5/qwen3_next to PADDING_FREE_BLOCKLIST to avoid packed-sequence contamination
* Detect hybrid linear-attention models structurally instead of by name for packing guard
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Install wrapped-packing setup at the signature, not the Zoo license comment
The _unsloth_wrapped_packing / _inspect setup block was injected by matching the
exact 'All Unsloth Zoo code licensed under LGPLv3' comment line in the sourced
sft_prepare_dataset. The unsloth_zoo dependency is only lower-bounded, so a newer
Zoo that moves or drops that header made the setup a silent no-op while the
truncation and pack_dataset rewrites still emitted references to those names,
raising NameError on every SFT dataset preparation.
Anchor the setup on the function signature instead (a structural location that
always exists) and fail loudly if it cannot be found, so the helper variables are
always defined before they are referenced across Zoo versions.
Adds a regression test that patches in a Zoo source without the license header.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
---------
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Etherl <61019402+Etherll@users.noreply.github.com>
Co-authored-by: danielhanchen
---
studio/backend/core/training/trainer.py | 12 +-
tests/utils/test_packing.py | 371 +++++++++++++++++++++++-
unsloth/models/rl_replacements.py | 88 ++++--
unsloth/trainer.py | 177 ++++++++++-
4 files changed, 610 insertions(+), 38 deletions(-)
diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py
index 26720865f4..8e419849cb 100644
--- a/studio/backend/core/training/trainer.py
+++ b/studio/backend/core/training/trainer.py
@@ -3425,15 +3425,19 @@ class UnslothTrainer:
logger.info(
f"CPT: using UnslothTrainer with embedding_learning_rate={embedding_lr}\n"
)
+ cpt_args = _UnslothTrainingArguments(
+ embedding_learning_rate = embedding_lr,
+ **config_args,
+ )
+ if config_args.get("packing", False):
+ cpt_args.packing_strategy = "wrapped"
+ logger.info("CPT packing strategy: wrapped\n")
trainer_kwargs = {
"model": self.model,
"tokenizer": sft_tokenizer,
"train_dataset": dataset["dataset"],
"data_collator": data_collator,
- "args": _UnslothTrainingArguments(
- embedding_learning_rate = embedding_lr,
- **config_args,
- ),
+ "args": cpt_args,
}
if eval_dataset is not None:
trainer_kwargs["eval_dataset"] = eval_dataset
diff --git a/tests/utils/test_packing.py b/tests/utils/test_packing.py
index a8557d8533..98c29d9f0f 100644
--- a/tests/utils/test_packing.py
+++ b/tests/utils/test_packing.py
@@ -14,6 +14,7 @@
# along with this program. If not, see .
from unsloth import FastLanguageModel
+import unsloth.trainer as trainer_module
from unsloth.utils import attention_dispatch as attention_dispatch_utils
from unsloth.utils.packing import (
configure_padding_free,
@@ -29,7 +30,7 @@ from unittest.mock import patch
import pytest
import torch
-from datasets import Dataset
+from datasets import Dataset, IterableDataset
from trl import SFTConfig, SFTTrainer
from trl.trainer.sft_trainer import DataCollatorForLanguageModeling
@@ -160,6 +161,374 @@ def test_configure_padding_free():
assert config.remove_unused_columns is False
+def _patch_fake_sft_trainer():
+ class FakeSFTTrainer:
+ def __init__(self, *args, **kwargs):
+ self.model = args[0] if len(args) >= 1 else kwargs["model"]
+ self.args = args[1] if len(args) >= 2 else kwargs["args"]
+ self.data_collator = args[2] if len(args) >= 3 else kwargs.get("data_collator")
+
+ trainer_module._patch_sft_trainer_auto_packing(SimpleNamespace(SFTTrainer = FakeSFTTrainer))
+ return FakeSFTTrainer
+
+
+def _vlm_model():
+ return SimpleNamespace(
+ config = SimpleNamespace(
+ architectures = ["Gemma4ForConditionalGeneration"],
+ model_type = "gemma4",
+ vision_config = SimpleNamespace(),
+ ),
+ max_seq_length = 16,
+ )
+
+
+def _text_model():
+ return SimpleNamespace(
+ config = SimpleNamespace(
+ architectures = ["LlamaForCausalLM"],
+ model_type = "llama",
+ ),
+ max_seq_length = 16,
+ )
+
+
+class _CharacterTokenizer:
+ bos_token = None
+ eos_token = None
+ chat_template = None
+
+ def __call__(self, texts, **kwargs):
+ is_batched = isinstance(texts, list)
+ if not is_batched:
+ texts = [texts]
+ input_ids = [[ord(char) for char in text] for text in texts]
+ if kwargs.get("truncation") and kwargs.get("max_length") is not None:
+ input_ids = [ids[: kwargs["max_length"]] for ids in input_ids]
+ return {"input_ids": input_ids if is_batched else input_ids[0]}
+
+
+def test_vlm_text_dataset_allows_explicit_packing():
+ fake_trainer = _patch_fake_sft_trainer()
+ config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
+
+ trainer = fake_trainer(
+ model = _vlm_model(),
+ args = config,
+ processing_class = object(),
+ train_dataset = Dataset.from_dict({"text": ["text-only CPT sample"]}),
+ )
+
+ assert config.packing is True
+ assert config.padding_free is True
+ assert trainer.model._unsloth_allow_packed_overlength is True
+
+
+def test_vlm_without_processing_class_still_disables_packing():
+ fake_trainer = _patch_fake_sft_trainer()
+ config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
+
+ fake_trainer(
+ _vlm_model(),
+ config,
+ None,
+ Dataset.from_dict({"text": ["text-only sample"]}),
+ )
+
+ assert config.packing is False
+ assert config.padding_free is False
+
+
+@pytest.mark.parametrize(
+ ("model_type", "architecture"),
+ (
+ ("t5", "T5ForConditionalGeneration"),
+ ("bart", "BartForConditionalGeneration"),
+ ("whisper", "WhisperForConditionalGeneration"),
+ ("csm", "CsmForConditionalGeneration"),
+ ),
+)
+def test_nonvision_conditional_generation_keeps_packing(model_type, architecture):
+ fake_trainer = _patch_fake_sft_trainer()
+ config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
+ model = SimpleNamespace(
+ config = SimpleNamespace(model_type = model_type, architectures = [architecture]),
+ max_seq_length = 16,
+ )
+
+ trainer = fake_trainer(
+ model,
+ config,
+ None,
+ Dataset.from_dict({"text": ["text-only sample"]}),
+ )
+
+ assert config.packing is True
+ assert config.padding_free is True
+ assert trainer.model._unsloth_allow_packed_overlength is True
+
+
+def test_vlm_vision_dataset_still_disables_packing():
+ fake_trainer = _patch_fake_sft_trainer()
+ config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
+
+ fake_trainer(
+ _vlm_model(),
+ config,
+ None,
+ Dataset.from_dict({"images": [None], "text": ["multimodal sample"]}),
+ None,
+ object(),
+ )
+
+ assert config.packing is False
+ assert config.padding_free is False
+
+
+@pytest.mark.parametrize(
+ "vision_column",
+ ("pixel_values", "pixel_attention_mask", "image_grid_thw"),
+)
+def test_vlm_preprocessed_vision_dataset_disables_packing(vision_column):
+ fake_trainer = _patch_fake_sft_trainer()
+ config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
+
+ fake_trainer(
+ model = _vlm_model(),
+ args = config,
+ processing_class = object(),
+ train_dataset = Dataset.from_dict({"input_ids": [[1]], vision_column: [None]}),
+ )
+
+ assert config.packing is False
+ assert config.padding_free is False
+
+
+@pytest.mark.parametrize("dict_eval", (False, True))
+def test_vlm_vision_eval_dataset_disables_packing(dict_eval):
+ fake_trainer = _patch_fake_sft_trainer()
+ config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
+ eval_dataset = Dataset.from_dict({"input_ids": [[1]], "pixel_values": [None]})
+ if dict_eval:
+ eval_dataset = {"vision": eval_dataset}
+
+ fake_trainer(
+ model = _vlm_model(),
+ args = config,
+ processing_class = object(),
+ train_dataset = Dataset.from_dict({"text": ["text-only training sample"]}),
+ eval_dataset = eval_dataset,
+ )
+
+ assert config.packing is False
+ assert config.padding_free is False
+
+
+def test_vlm_streaming_vision_dataset_without_metadata_disables_packing():
+ fake_trainer = _patch_fake_sft_trainer()
+ config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
+ dataset = IterableDataset.from_generator(
+ lambda: iter([{"images": [None], "text": "multimodal sample"}])
+ )
+ assert dataset.column_names is None
+
+ fake_trainer(
+ model = _vlm_model(),
+ args = config,
+ processing_class = object(),
+ train_dataset = dataset,
+ )
+
+ assert config.packing is False
+ assert config.padding_free is False
+ assert next(iter(dataset))["text"] == "multimodal sample"
+
+
+@pytest.mark.parametrize("data_collator", (None, object()))
+def test_stateful_stream_is_not_consumed_during_detection(data_collator):
+ class StatefulDataset:
+ def __init__(self):
+ self.rows = iter([{"text": "first"}, {"text": "second"}])
+
+ def __iter__(self):
+ return (row for row in self.rows)
+
+ fake_trainer = _patch_fake_sft_trainer()
+ config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
+ dataset = StatefulDataset()
+
+ fake_trainer(
+ model = _vlm_model(),
+ args = config,
+ processing_class = object(),
+ data_collator = data_collator,
+ train_dataset = dataset,
+ )
+
+ assert config.packing is False
+ assert config.padding_free is False
+ assert next(iter(dataset))["text"] == "first"
+
+
+def test_text_model_stream_without_metadata_keeps_packing():
+ class StatefulDataset:
+ def __init__(self):
+ self.rows = iter([{"text": "first"}, {"text": "second"}])
+
+ def __iter__(self):
+ return (row for row in self.rows)
+
+ fake_trainer = _patch_fake_sft_trainer()
+ config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
+ dataset = StatefulDataset()
+
+ trainer = fake_trainer(
+ model = _text_model(),
+ args = config,
+ processing_class = object(),
+ train_dataset = dataset,
+ )
+
+ assert config.packing is True
+ assert config.padding_free is True
+ assert trainer.model._unsloth_allow_packed_overlength is True
+ assert next(iter(dataset))["text"] == "first"
+
+
+def test_bfd_packing_truncates_before_packing(monkeypatch):
+ args = SimpleNamespace(
+ dataset_num_proc = 1,
+ dataset_text_field = "text",
+ max_length = 4,
+ packing_strategy = "bfd",
+ )
+ trainer = SimpleNamespace(model = None)
+ dataset = Dataset.from_dict({"prompt": ["abc"], "completion": ["defghij"]})
+ prepare_globals = SFTTrainer._prepare_dataset.__globals__
+
+ def passthrough_pack_dataset(dataset, seq_length, strategy, map_kwargs):
+ return dataset
+
+ monkeypatch.setitem(prepare_globals, "pack_dataset", passthrough_pack_dataset)
+ packed = SFTTrainer._prepare_dataset(
+ trainer,
+ dataset,
+ _CharacterTokenizer(),
+ args,
+ True,
+ None,
+ "train",
+ )
+
+ assert len(packed["input_ids"][0]) == args.max_length
+
+
+def test_wrapped_strategy_without_packing_still_truncates():
+ args = SimpleNamespace(
+ dataset_num_proc = 1,
+ dataset_text_field = "text",
+ max_length = 4,
+ packing_strategy = "wrapped",
+ )
+ trainer = SimpleNamespace(model = None)
+ dataset = Dataset.from_dict({"text": ["abcdefghi"]})
+
+ prepared = SFTTrainer._prepare_dataset(
+ trainer,
+ dataset,
+ _CharacterTokenizer(),
+ args,
+ False,
+ None,
+ "train",
+ )
+
+ assert len(prepared["input_ids"][0]) == args.max_length
+
+
+@pytest.mark.parametrize("legacy_api", (False, True))
+def test_wrapped_packing_preserves_overlength_tokens(monkeypatch, legacy_api):
+ args_kwargs = {
+ "dataset_num_proc": 1,
+ "dataset_text_field": "text",
+ "max_length": 4,
+ }
+ if not legacy_api:
+ args_kwargs["packing_strategy"] = "wrapped"
+ args = SimpleNamespace(**args_kwargs)
+ trainer = SimpleNamespace(model = None)
+ dataset = Dataset.from_dict({"text": ["abcdefghi"]})
+ prepare_globals = SFTTrainer._prepare_dataset.__globals__
+ pack_dataset = prepare_globals["pack_dataset"]
+
+ def legacy_pack_dataset(
+ dataset,
+ seq_length,
+ map_kwargs = None,
+ ):
+ return pack_dataset(dataset, seq_length, "wrapped", map_kwargs)
+
+ if legacy_api:
+ monkeypatch.setitem(prepare_globals, "pack_dataset", legacy_pack_dataset)
+
+ packed = SFTTrainer._prepare_dataset(
+ trainer,
+ dataset,
+ _CharacterTokenizer(),
+ args,
+ True,
+ None,
+ "train",
+ )
+
+ packed_ids = packed["input_ids"]
+ assert sum(len(input_ids) for input_ids in packed_ids) == 9
+ assert all(len(input_ids) <= args.max_length for input_ids in packed_ids)
+
+
+# Named to match the unsloth_zoo helper: sft_trainer_prepare_dataset sources it by
+# name and renames "def sft_prepare_dataset" -> "def _prepare_dataset". This fixture
+# deliberately omits the "All Unsloth Zoo code licensed under LGPLv3" header to emulate
+# a newer, compatible Zoo whose header moved (the dependency is only lower-bounded).
+def sft_prepare_dataset(
+ self, dataset, processing_class, args, packing, formatting_func, dataset_text_field
+):
+ do_truncation = True
+ # Mirror the Zoo call so the "truncation = do_truncation," injection anchor
+ # survives formatting (a bare tuple assignment gets rewritten to a paren form).
+ dataset = processing_class(
+ dataset,
+ truncation = do_truncation,
+ )
+ return dataset
+
+
+def test_wrapped_packing_setup_survives_missing_zoo_header(monkeypatch):
+ # Regression: the wrapped-packing setup used to anchor on the Zoo license comment,
+ # so a header change made it a no-op while the truncation reference still landed,
+ # NameError-ing every SFT dataset preparation. It must now install via the
+ # signature and always precede the reference.
+ import ast
+ import textwrap
+ import unsloth.models.rl_replacements as rlr
+
+ monkeypatch.setitem(rlr.RL_REPLACEMENTS, "sft_prepare_dataset", sft_prepare_dataset)
+
+ source = (
+ "def _prepare_dataset(self, dataset, processing_class, args, packing, "
+ "formatting_func, dataset_text_field):\n return dataset\n"
+ )
+ patched = rlr.sft_trainer_prepare_dataset("_prepare_dataset", source)
+
+ assert "_unsloth_wrapped_packing = packing" in patched
+ assert "import inspect as _inspect" in patched
+ assert "not _unsloth_wrapped_packing" in patched
+ assert patched.index("_unsloth_wrapped_packing = packing") < patched.index(
+ "truncation = do_truncation and not _unsloth_wrapped_packing"
+ )
+ ast.parse(textwrap.dedent(patched))
+
+
class _DummyChild(torch.nn.Module):
def __init__(self):
super().__init__()
diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py
index ffb845b04f..b0709f7376 100644
--- a/unsloth/models/rl_replacements.py
+++ b/unsloth/models/rl_replacements.py
@@ -276,17 +276,13 @@ def dpo_trainer_vision_signature_columns(function_name, function):
_extra_columns = "".join(f' "{_k}",\n' for _k in _DPO_VISION_KEYS)
new_function = function.replace(
' "image_sizes",\n "token_type_ids",\n',
- f' "image_sizes",\n'
- f"{_extra_columns}"
- f' "token_type_ids",\n',
+ f' "image_sizes",\n{_extra_columns} "token_type_ids",\n',
)
if new_function != function:
return new_function
return function.replace(
' "image_sizes",\n "ref_chosen_logps",\n',
- f' "image_sizes",\n'
- f"{_extra_columns}"
- f' "ref_chosen_logps",\n',
+ f' "image_sizes",\n{_extra_columns} "ref_chosen_logps",\n',
)
@@ -458,6 +454,60 @@ def sft_trainer_prepare_dataset(function_name, function):
if matched:
# Use fast version!
function = inspect.getsource(fast_sft_prepare_dataset)
+ # why: install the wrapped-packing setup (and the `_inspect` import the
+ # truncation / pack_dataset rewrites below depend on) at the function
+ # signature, a structural anchor that always exists, rather than the
+ # unsloth_zoo license-comment line. That header is only lower-bounded, so a
+ # newer Zoo may move or drop it; anchoring there let the setup silently
+ # no-op while the references still landed, NameError-ing every SFT dataset
+ # preparation. Fail loudly if even the signature cannot be located.
+ _wrapped_packing_setup = (
+ " import inspect as _inspect\n"
+ " try:\n"
+ ' _unsloth_pack_has_strategy = "strategy" in _inspect.signature(pack_dataset).parameters\n'
+ " except Exception:\n"
+ " _unsloth_pack_has_strategy = True\n"
+ " _unsloth_wrapped_packing = packing and (\n"
+ ' getattr(args, "packing_strategy", None) == "wrapped"\n'
+ " or not _unsloth_pack_has_strategy\n"
+ " )\n"
+ )
+ function, _n_setup = re.subn(
+ r"(def sft_prepare_dataset\s*\(.*?\)\s*(?:->[^:\n]*)?:[ \t]*\n)",
+ lambda match: match.group(1) + _wrapped_packing_setup,
+ function,
+ count = 1,
+ flags = re.DOTALL,
+ )
+ if _n_setup != 1:
+ raise RuntimeError(
+ "Unsloth: failed to install wrapped-packing support into "
+ "sft_prepare_dataset (signature not found); please file a bug report."
+ )
+ function = function.replace(
+ "truncation = do_truncation,",
+ "truncation = do_truncation and not _unsloth_wrapped_packing,",
+ )
+ function = function.replace(
+ "if do_truncation and max_seq_length > 0:",
+ "if do_truncation and not _unsloth_wrapped_packing and max_seq_length > 0:",
+ )
+ function = function.replace(
+ """dataset = pack_dataset(
+ dataset.select_columns(used_column_names),
+ max_seq_length,
+ getattr(args, "packing_strategy", "bfd"),
+ map_kwargs,
+ )""",
+ """_pack_kwargs = {"map_kwargs": map_kwargs}
+ if "strategy" in _inspect.signature(pack_dataset).parameters:
+ _pack_kwargs["strategy"] = getattr(args, "packing_strategy", "bfd")
+ dataset = pack_dataset(
+ dataset.select_columns(used_column_names),
+ max_seq_length,
+ **_pack_kwargs,
+ )""",
+ )
function = function.split("\n")
function = "\n".join(" " * 4 + x for x in function)
function = function.replace("def sft_prepare_dataset", "def _prepare_dataset")
@@ -2120,19 +2170,21 @@ def grpo_trainer_compute_loss(function_name, function):
logits_to_keep,
batch_size = None,
compute_entropy = False,
- compute_efficient = False: self._get_per_token_logps(
- model, input_ids, attention_mask, logits_to_keep, compute_efficient
+ compute_efficient = False: (
+ self._get_per_token_logps(
+ model, input_ids, attention_mask, logits_to_keep, compute_efficient
+ )
+ if hasattr(self, "_get_per_token_logps")
+ else self._get_per_token_logps_and_entropies(
+ model,
+ input_ids,
+ attention_mask,
+ logits_to_keep,
+ batch_size,
+ compute_entropy,
+ compute_efficient,
+ )[0]
)
- if hasattr(self, "_get_per_token_logps")
- else self._get_per_token_logps_and_entropies(
- model,
- input_ids,
- attention_mask,
- logits_to_keep,
- batch_size,
- compute_entropy,
- compute_efficient,
- )[0]
) # logps
per_token_logps = get_logps_func(
diff --git a/unsloth/trainer.py b/unsloth/trainer.py
index 83cb1758f0..61d41aad21 100644
--- a/unsloth/trainer.py
+++ b/unsloth/trainer.py
@@ -100,6 +100,10 @@ PADDING_FREE_BLOCKLIST = {
"gemma2", # - gemma2: Uses slow_attention_softcapping which has torch.compile issues
"gpt_oss", # - gpt_oss: Uses Flex Attention which doesn't handle padding_free correctly
}
+# Hybrid linear-attention / state-space models (Qwen3.5, Qwen3-Next, ...) carry a
+# recurrent gated-delta state plus a causal conv1d. Sample packing / padding-free
+# flattens the batch, so those ops leak state across sequence boundaries. Detected
+# structurally by _is_hybrid_linear_attention_model rather than by model name.
def _should_pack(config) -> bool:
@@ -137,6 +141,132 @@ def _should_skip_auto_packing_error(exc: Exception) -> bool:
return any(msg in message for msg in _AUTO_PACK_SKIP_MESSAGES)
+_VISION_DATASET_KEYS = frozenset(
+ {
+ "image",
+ "images",
+ "image_grid_thw",
+ "image_position_ids",
+ "image_sizes",
+ "mm_token_type_ids",
+ "pixel_attention_mask",
+ "pixel_position_ids",
+ "pixel_values",
+ "pixel_values_videos",
+ "video",
+ "videos",
+ "video_grid_thw",
+ }
+)
+
+
+def _is_vlm_config(config, model_types = ()) -> bool:
+ if any(
+ hasattr(config, attr)
+ for attr in ("vision_config", "img_processor", "image_token_index", "projector_config")
+ ):
+ return True
+
+ architectures = getattr(config, "architectures", None) or ()
+ try:
+ from transformers.models.auto import modeling_auto
+
+ mappings = (
+ getattr(modeling_auto, "MODEL_FOR_IMAGE_TEXT_TO_TEXT_MAPPING_NAMES", {}) or {},
+ getattr(modeling_auto, "MODEL_FOR_VISION_2_SEQ_MAPPING_NAMES", {}) or {},
+ )
+ registry_types = set().union(*(mapping.keys() for mapping in mappings))
+ registry_classes = set().union(*(mapping.values() for mapping in mappings))
+ config_types = set(model_types or ())
+ model_type = getattr(config, "model_type", None)
+ if model_type is not None:
+ config_types.add(model_type)
+ if not config_types.isdisjoint(registry_types) or any(
+ architecture in registry_classes for architecture in architectures
+ ):
+ return True
+ except Exception:
+ pass
+ return any(
+ isinstance(architecture, str) and architecture.endswith("ForVisionText2Text")
+ for architecture in architectures
+ )
+
+
+def _is_vision_dataset(dataset, *, unknown_is_vision = False) -> bool:
+ if dataset is None:
+ return False
+ column_names = getattr(dataset, "column_names", None)
+ if column_names is not None:
+ return not _VISION_DATASET_KEYS.isdisjoint(column_names)
+ # Unknown-schema streams cannot be safely probed without potentially dropping a sample.
+ return unknown_is_vision
+
+
+def _is_vision_eval_dataset(dataset, *, unknown_is_vision = False) -> bool:
+ if isinstance(dataset, dict):
+ return any(
+ _is_vision_dataset(split, unknown_is_vision = unknown_is_vision)
+ for split in dataset.values()
+ )
+ return _is_vision_dataset(dataset, unknown_is_vision = unknown_is_vision)
+
+
+_HYBRID_CONFIG_MARKERS = (
+ "linear_conv_kernel_dim",
+ "linear_key_head_dim",
+ "linear_value_head_dim",
+ "full_attention_interval",
+)
+
+
+def _is_hybrid_linear_attention_model(model) -> bool:
+ """Detect models mixing linear-attention / state-space mixers (gated-delta,
+ Mamba-style) with a causal conv1d, e.g. Qwen3.5 / Qwen3-Next. Packing and
+ padding-free flatten the batch, and those recurrent + conv ops leak state
+ across sequence boundaries, so they must not be packed. Uses composite
+ structural evidence rather than a model-name match."""
+ if model is None:
+ return False
+
+ # Config-level: explicit hybrid layer schedule or linear-attn markers.
+ for config in (
+ getattr(model, "config", None),
+ getattr(getattr(model, "config", None), "text_config", None),
+ ):
+ if config is None:
+ continue
+ layer_types = getattr(config, "layer_types", None)
+ if isinstance(layer_types, (list, tuple)) and any(
+ isinstance(t, str) and "linear_attention" in t for t in layer_types
+ ):
+ return True
+ if any(hasattr(config, marker) for marker in _HYBRID_CONFIG_MARKERS):
+ return True
+
+ # Module-level: a mixer carrying a recurrent gated-delta op plus a conv1d.
+ named_modules = getattr(model, "named_modules", None)
+ if named_modules is None:
+ return False
+ seen = set()
+ for _, module in named_modules():
+ if id(module) in seen:
+ continue
+ seen.add(id(module))
+ cls = type(module).__name__
+ if not (
+ cls.endswith("GatedDeltaNet") or "LinearAttention" in cls or cls.endswith("Mamba2Mixer")
+ ):
+ continue
+ has_recurrent = any(
+ hasattr(module, attr)
+ for attr in ("chunk_gated_delta_rule", "recurrent_gated_delta_rule", "A_log")
+ )
+ if has_recurrent and hasattr(module, "conv1d"):
+ return True
+ return False
+
+
# Unsloth gradient accumulation fix:
from transformers import __version__ as transformers_version, ProcessorMixin
@@ -498,30 +628,43 @@ def _patch_sft_trainer_auto_packing(trl_module):
else:
config_arg = kwargs.get("args")
- model = kwargs.get("model")
- is_unsupported_model = False
+ model = args[0] if len(args) >= 1 else kwargs.get("model")
is_vlm = False
+ is_unsupported_model = False
+ is_hybrid = False
if model is not None:
model_config = getattr(model, "config", None)
if model_config is not None:
model_types = get_transformers_model_type(model_config)
is_unsupported_model = any(x in PADDING_FREE_BLOCKLIST for x in model_types)
+ is_vlm = _is_vlm_config(model_config, model_types)
+ is_hybrid = _is_hybrid_linear_attention_model(model)
- architectures = getattr(model_config, "architectures", None)
- if architectures is None:
- architectures = []
- is_vlm = any(x.endswith("ForConditionalGeneration") for x in architectures)
- is_vlm = is_vlm or hasattr(model_config, "vision_config")
-
- processing_class = kwargs.get("processing_class") or kwargs.get("tokenizer")
- data_collator = kwargs.get("data_collator")
+ processing_class = (
+ args[5] if len(args) >= 6 else kwargs.get("processing_class") or kwargs.get("tokenizer")
+ )
+ data_collator = args[2] if len(args) >= 3 else kwargs.get("data_collator")
+ train_dataset = args[3] if len(args) >= 4 else kwargs.get("train_dataset")
+ eval_dataset = args[4] if len(args) >= 5 else kwargs.get("eval_dataset")
+ is_processor = isinstance(processing_class, ProcessorMixin)
+ is_auto_processor_vlm = is_vlm and processing_class is None
+ is_vision_dataset = (
+ data_collator is None
+ and not is_processor
+ and (
+ _is_vision_dataset(train_dataset, unknown_is_vision = is_vlm)
+ or _is_vision_eval_dataset(eval_dataset, unknown_is_vision = is_vlm)
+ )
+ )
# Disable padding-free for VLMs / custom collators / blocklisted models
blocked = (
(data_collator is not None)
- or isinstance(processing_class, ProcessorMixin)
- or is_vlm
+ or is_processor
+ or is_auto_processor_vlm
+ or is_vision_dataset
or is_unsupported_model
+ or is_hybrid
or (
os.environ.get("UNSLOTH_RETURN_LOGITS", "0") == "1"
) # Disable padding free on forced logits
@@ -535,10 +678,14 @@ def _patch_sft_trainer_auto_packing(trl_module):
if blocked and requested_pack:
reason = "custom data collator"
- if data_collator is None and isinstance(processing_class, ProcessorMixin):
+ if data_collator is None and is_processor:
reason = "processor-based model"
- elif is_vlm:
- reason = "vision-language model"
+ elif is_auto_processor_vlm:
+ reason = "vision-language model with auto processor"
+ elif is_vision_dataset:
+ reason = "vision dataset"
+ elif is_hybrid:
+ reason = "hybrid linear-attention model"
elif is_unsupported_model:
reason = f"unsupported model type(s): {', '.join(model_types)}"
message = f"Unsloth: Sample packing skipped ({reason} detected)."
From cf912cbd881190f411147b8c93294408efe5f90c Mon Sep 17 00:00:00 2001
From: Naitik Pal
Date: Mon, 20 Jul 2026 13:03:56 +0530
Subject: [PATCH 12/41] feat(studio): add UNSLOTH_LLAMA_CPP_BACKEND env var to
force CPU fallback #7213 (#7228)
* test(studio): add e2e test for cpu-fallback overriding vulkan
* feat(studio): add UNSLOTH_LLAMA_CPP_BACKEND env var
* feat(studio): add UNSLOTH_LLAMA_CPP_BACKEND env var
* Preserve UNSLOTH_LLAMA_CPP_BACKEND=cpu across llama.cpp updates for PR #7228
The in-app updater rebuilt the installer command without --cpu-fallback and
only re-asserted Vulkan, so accepting a llama.cpp update after forcing CPU on
an Intel iGPU host re-ran host detection and routed back to the crashing Vulkan
bundle (#7213). Record install_kind in the prebuilt marker and re-assert
--cpu-fallback on update when the installed bundle is CPU.
Also make setup.sh's UNSLOTH_LLAMA_CPP_BACKEND check case-insensitive to match
setup.ps1, and add tests for the updater CPU preservation and the setup.sh flag
plumbing.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Trim and validate UNSLOTH_LLAMA_CPP_BACKEND, warn on unknown values for PR #7228
Trim surrounding whitespace and lowercase the value in both setup.sh and
setup.ps1, so values like ' cpu ' or 'CPU' still force the CPU-only prebuilt.
An unrecognized value (e.g. 'gpu') now prints a warning instead of silently
falling back to auto. Extend test_setup_llama_cpp_backend.py to cover both
scripts, including trimmed, empty and unknown values.
* Preserve arm64 CPU installs on update and honor CPU override in Windows prune for PR #7228
The update-path CPU preservation only matched install_kind ending in -cpu, so
arm64 CPU bundles (linux-arm64, windows-arm64) were re-routed to a GPU or source
build on update. Match the full set of CPU-only kinds instead.
Persisting install_kind also activated the previously inert Windows
mismatch-prune in setup.ps1: on a GPU host with UNSLOTH_LLAMA_CPP_BACKEND=cpu it
saw the windows-cpu marker as mismatched and deleted it every rerun. Normalize
the override once and make CPU expected so a deliberate CPU install is kept.
Extend the tests to cover both.
* Document legacy llama.cpp markers keep heal-to-GPU on update for PR #7228
Legacy prebuilt markers written before install_kind was persisted intentionally
do not force --cpu-fallback on update: the in-app updater lets them re-resolve
(heal to a GPU bundle) per the existing behavior from #6097, and only markers
that explicitly record a CPU install_kind are pinned to CPU. Add a comment and a
regression case documenting the boundary.
* Tighten llama.cpp CPU-fallback comments for PR #7228
* Fix Windows install-prune to keep valid Intel/fallback bundles for PR #7228
Persisting install_kind activated the setup.ps1 mismatch-prune, whose
expectedKinds was incomplete: the non-NVIDIA/non-AMD branch omitted
windows-vulkan (the Intel auto-route) and the GPU branches omitted the
windows-cpu/windows-arm64 fallback the installer uses when a GPU prebuilt is
missing. That made every setup rerun delete and re-download a valid Intel Vulkan
(or CPU-fallback) install. List all kinds the installer can produce per host so
only a bundle the host cannot run is pruned. Cover the full matrix in tests.
* Persist force_cpu marker flag so only forced CPU installs re-assert on update for PR #7228
* Add --force-cpu for deliberate CPU installs and warn on macOS for PR #7228
* Record force_cpu when reusing a matching CPU bundle for PR #7228
* Accept force_cpu keyword in installer test validator fakes for PR #7228
---------
Co-authored-by: danielhanchen
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Daniel Han
---
.../tests/test_install_resolve_prebuilt.py | 95 +++++++++++
studio/backend/tests/test_llama_cpp_update.py | 58 ++++++-
.../tests/test_setup_llama_cpp_backend.py | 154 ++++++++++++++++++
studio/backend/utils/llama_cpp_update.py | 12 +-
studio/install_llama_prebuilt.py | 69 +++++++-
studio/setup.ps1 | 13 ++
studio/setup.sh | 20 +++
.../test_install_llama_prebuilt_logic.py | 4 +
8 files changed, 410 insertions(+), 15 deletions(-)
create mode 100644 studio/backend/tests/test_setup_llama_cpp_backend.py
diff --git a/studio/backend/tests/test_install_resolve_prebuilt.py b/studio/backend/tests/test_install_resolve_prebuilt.py
index e97ca47717..3ebad861ad 100644
--- a/studio/backend/tests/test_install_resolve_prebuilt.py
+++ b/studio/backend/tests/test_install_resolve_prebuilt.py
@@ -445,6 +445,101 @@ def test_route_to_vulkan_prebuilt_cpu_fallback_wins():
assert routed is host
+@pytest.mark.parametrize("cpu_flag", ["--cpu-fallback", "--force-cpu"])
+def test_resolve_prebuilt_cpu_fallback_overrides_intel_vulkan(monkeypatch, capsys, cpu_flag):
+ """Either CPU flag via CLI must suppress Vulkan even on an Intel GPU host: both
+ drop GPU detection (--force-cpu additionally persists, on the install path)."""
+ monkeypatch.setattr(
+ ilp,
+ "detect_host",
+ lambda: _host(is_linux = True, is_x86_64 = True, has_intel_gpu = True),
+ )
+ seen = {}
+
+ def _resolver(tag, host, repo, published_release_tag):
+ seen["host"] = host
+ seen["repo"] = repo
+ raise ilp.PrebuiltFallback("no asset")
+
+ monkeypatch.setattr(ilp, "resolve_simple_install_release_plans", _resolver)
+ monkeypatch.setattr(
+ sys,
+ "argv",
+ [
+ "install_llama_prebuilt.py",
+ "--resolve-prebuilt",
+ "latest",
+ cpu_flag,
+ "--output-format",
+ "json",
+ ],
+ )
+ assert ilp.main() == ilp.EXIT_SUCCESS
+ # The CPU flag must suppress Intel GPU, route to fork (not upstream Vulkan)
+ assert seen["host"].has_intel_gpu is False
+ assert seen["repo"] == FORK
+
+
+@pytest.mark.parametrize(
+ "flags, expect_force, expect_persist",
+ [
+ ([], False, False),
+ # Automatic/transient last resort (arm64 GPU-build recovery): drops GPU but
+ # does NOT persist, so a later update heals to a GPU bundle (#6097).
+ (["--cpu-fallback"], True, False),
+ # Deliberate CPU-only (UNSLOTH_LLAMA_CPP_BACKEND=cpu): drops GPU AND persists so
+ # the updater re-asserts it and never revives the Intel iGPU crash (#7213).
+ (["--force-cpu"], True, True),
+ (["--cpu-fallback", "--force-cpu"], True, True),
+ ],
+)
+def test_cli_cpu_flags_thread_force_and_persist(
+ monkeypatch, tmp_path, flags, expect_force, expect_persist
+):
+ captured = {}
+ monkeypatch.setattr(ilp, "install_prebuilt", lambda **kw: captured.update(kw))
+ monkeypatch.setattr(
+ sys,
+ "argv",
+ ["install_llama_prebuilt.py", "--install-dir", str(tmp_path / "llama.cpp"), *flags],
+ )
+ assert ilp.main() == ilp.EXIT_SUCCESS
+ assert captured["force_cpu"] is expect_force
+ assert captured["persist_force_cpu"] is expect_persist
+
+
+@pytest.mark.parametrize(
+ "existing, requested, expected",
+ [
+ # A deliberate --force-cpu on top of a naturally-installed CPU bundle (same
+ # asset, install skipped) must still flip the marker to true (#7213).
+ (False, True, True),
+ (None, True, True),
+ # No spurious writes when already in sync, and a released force syncs down.
+ (True, True, True),
+ (False, False, False),
+ (True, False, False),
+ ],
+)
+def test_sync_marker_force_cpu(tmp_path, existing, requested, expected):
+ marker = {"tag": "b9585", "asset": "llama-b9585-bin-ubuntu-x64.tar.gz"}
+ if existing is not None:
+ marker["force_cpu"] = existing
+ marker_path = tmp_path / "UNSLOTH_PREBUILT_INFO.json"
+ marker_path.write_text(json.dumps(marker))
+ ilp.sync_marker_force_cpu(tmp_path, requested)
+ written = json.loads(marker_path.read_text())
+ assert written["force_cpu"] is expected
+ # Unrelated fields are preserved.
+ assert written["asset"] == "llama-b9585-bin-ubuntu-x64.tar.gz"
+
+
+def test_sync_marker_force_cpu_missing_marker_is_noop(tmp_path):
+ # No marker (or unreadable) must not crash the reuse path.
+ ilp.sync_marker_force_cpu(tmp_path, True)
+ assert not (tmp_path / "UNSLOTH_PREBUILT_INFO.json").exists()
+
+
def test_route_to_vulkan_prebuilt_hidden_nvidia_not_rerouted():
# A mixed NVIDIA+Intel host that hid NVIDIA (CUDA_VISIBLE_DEVICES=""/-1):
# physical NVIDIA present but not usable. Must NOT auto-route to Vulkan, or
diff --git a/studio/backend/tests/test_llama_cpp_update.py b/studio/backend/tests/test_llama_cpp_update.py
index 83ea07a066..f12384231f 100644
--- a/studio/backend/tests/test_llama_cpp_update.py
+++ b/studio/backend/tests/test_llama_cpp_update.py
@@ -83,6 +83,7 @@ def _write_install(
repo: str = "unslothai/llama.cpp",
asset: str | None = None,
release_tag: str | None = None,
+ force_cpu: bool | None = None,
) -> str:
"""Create a fake prebuilt install and return the llama-server path."""
bin_dir = dir_ / "build" / "bin"
@@ -99,6 +100,8 @@ def _write_install(
}
if asset is not None:
marker["asset"] = asset
+ if force_cpu is not None:
+ marker["force_cpu"] = force_cpu
(dir_ / MARKER).write_text(json.dumps(marker))
return str(binary)
@@ -493,6 +496,47 @@ def test_start_update_preserves_vulkan_via_env(monkeypatch, tmp_path):
assert popen_kwargs["env"]["UNSLOTH_FORCE_VULKAN"] == "1"
+@pytest.mark.parametrize(
+ "force_cpu, expect_flag",
+ [
+ # A deliberate CPU install (marker force_cpu=True) re-asserts --force-cpu on
+ # update so detect_host on a GPU host cannot re-route and revive the crash
+ # (#7213); --force-cpu also re-persists the flag for the next update.
+ (True, True),
+ # A transient fallback (or a legacy marker without the flag) stays free to
+ # heal to a GPU bundle (#6097).
+ (False, False),
+ (None, False),
+ ],
+)
+def test_start_update_cpu_fallback_preserved_by_flag(monkeypatch, tmp_path, force_cpu, expect_flag):
+ asset = "llama-b9493-bin-ubuntu-x64.tar.gz"
+ install_dir = tmp_path / "llama.cpp"
+ binary = _write_install(install_dir, "b9493", asset = asset, force_cpu = force_cpu)
+ monkeypatch.setattr(upd, "_find_binary", lambda: binary)
+ monkeypatch.setattr(upd, "_installer_script", lambda: tmp_path / "install_llama_prebuilt.py")
+ monkeypatch.setattr(freshness, "_fetch_latest_release_tag", lambda repo, timeout = 5.0: "b9518")
+
+ captured: dict = {}
+
+ def _on_start(cmd):
+ captured["cmd"] = cmd
+ _write_install(install_dir, "b9518", asset = asset, force_cpu = force_cpu)
+
+ _patch_installer_popen(monkeypatch, lines = ["installed\n"], on_start = _on_start)
+
+ assert upd.start_update()["started"] is True
+ deadline = time.time() + 10
+ while time.time() < deadline:
+ job = upd.get_update_status()["job"]
+ if job["state"] in ("success", "error"):
+ break
+ time.sleep(0.05)
+ assert job["state"] == "success", job
+ assert ("--force-cpu" in captured["cmd"]) is expect_flag
+ assert "--cpu-fallback" not in captured["cmd"]
+
+
def test_start_update_reports_full_release_tag(monkeypatch, tmp_path):
install_dir = tmp_path / "llama.cpp"
binary = _write_install(install_dir, "b9595")
@@ -676,7 +720,7 @@ def test_install_cmd_rocm_marker_forwards_gfx(monkeypatch, tmp_path):
assert "--rocm-gfx" in cmd
assert cmd[cmd.index("--rocm-gfx") + 1] == "gfx110x"
assert "--has-rocm" not in cmd
- assert "--cpu-fallback" not in cmd
+ assert "--force-cpu" not in cmd
assert "--simple-policy" not in cmd
assert "--published-repo" in cmd and "unslothai/llama.cpp" in cmd
@@ -690,17 +734,17 @@ def test_install_cmd_fork_rocm_marker_forwards_has_rocm(monkeypatch, tmp_path):
def test_install_cmd_ggml_cpu_marker_has_no_cpu_fallback(monkeypatch, tmp_path):
- # Legacy CPU installs recorded a ggml-org marker (new installs use the fork).
- # Re-running into the same install-dir/repo reproduces the same CPU bundle;
- # --cpu-fallback (which force-drops GPU detection) is reserved for setup.sh's
- # arm64 rescue and must not appear here.
+ # Legacy CPU installs recorded a ggml-org marker (new installs use the fork) with
+ # no force_cpu field. Re-running into the same install-dir/repo reproduces the same
+ # CPU bundle; --force-cpu (the persisted-CPU re-assert) must not appear for a marker
+ # that never recorded a deliberate CPU choice, so it can still heal to GPU (#6097).
cmd = _capture_install_cmd(
monkeypatch,
tmp_path,
repo = "ggml-org/llama.cpp",
asset = "llama-b9334-bin-ubuntu-x64.tar.gz",
)
- assert "--cpu-fallback" not in cmd
+ assert "--force-cpu" not in cmd
assert "--rocm-gfx" not in cmd
assert "--has-rocm" not in cmd
assert "--simple-policy" not in cmd
@@ -714,7 +758,7 @@ def test_install_cmd_cuda_marker_minimal_and_backward_compatible(monkeypatch, tm
assert "--simple-policy" not in cmd
assert "--rocm-gfx" not in cmd
assert "--has-rocm" not in cmd
- assert "--cpu-fallback" not in cmd
+ assert "--force-cpu" not in cmd
def test_install_cmd_pins_offered_release_tag(monkeypatch, tmp_path):
diff --git a/studio/backend/tests/test_setup_llama_cpp_backend.py b/studio/backend/tests/test_setup_llama_cpp_backend.py
new file mode 100644
index 0000000000..36928c680c
--- /dev/null
+++ b/studio/backend/tests/test_setup_llama_cpp_backend.py
@@ -0,0 +1,154 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""setup.sh and setup.ps1 must map UNSLOTH_LLAMA_CPP_BACKEND=cpu to
+install_llama_prebuilt.py's --force-cpu so users can force the CPU-only prebuilt
+on GPU hosts (#7213). The match is case-insensitive and whitespace-trimmed, an
+unrecognized value warns instead of silently falling back, and macOS warns (no
+CPU-only bundle). Runs the real block extracted from each script so the tests
+track the shipped logic.
+"""
+
+import os
+import re
+import shutil
+import subprocess
+from pathlib import Path
+
+import pytest
+
+_STUDIO = Path(__file__).resolve().parents[2]
+_SETUP_SH = _STUDIO / "setup.sh"
+_SETUP_PS1 = _STUDIO / "setup.ps1"
+_SKIP_NO_BASH = pytest.mark.skipif(shutil.which("bash") is None, reason = "bash unavailable")
+_SKIP_NO_PWSH = pytest.mark.skipif(shutil.which("pwsh") is None, reason = "pwsh unavailable")
+
+
+def _backend_block() -> str:
+ text = _SETUP_SH.read_text(encoding = "utf-8")
+ m = re.search(r"_llama_backend=.*?esac", text, re.DOTALL)
+ assert m, "UNSLOTH_LLAMA_CPP_BACKEND block not found in setup.sh"
+ return m.group(0)
+
+
+def _run(value: str | None, system: str = "Linux") -> tuple[list[str], str]:
+ # Pass the value through env (not the script text) so whitespace survives, and
+ # stub the setup.sh logging helpers the unknown-value branch calls. system sets
+ # _HOST_SYSTEM so the macOS (Darwin) no-op branch can be exercised.
+ env = {k: v for k, v in os.environ.items() if k != "UNSLOTH_LLAMA_CPP_BACKEND"}
+ if value is not None:
+ env["UNSLOTH_LLAMA_CPP_BACKEND"] = value
+ harness = (
+ f'_PREBUILT_CMD=()\nC_WARN=""\n_HOST_SYSTEM="{system}"\n'
+ 'step() { printf "STEP: %s\\n" "$*" >&2; }\n'
+ f"{_backend_block()}\n"
+ 'printf "%s\\n" "${_PREBUILT_CMD[@]}"'
+ )
+ out = subprocess.run(
+ ["bash", "-c", harness], capture_output = True, text = True, env = env, check = True
+ )
+ return out.stdout.split(), out.stderr
+
+
+@_SKIP_NO_BASH
+@pytest.mark.parametrize("value", ["cpu", "CPU", "Cpu", " cpu ", "CPU\t"])
+def test_backend_cpu_appends_flag(value):
+ # A deliberate CPU choice persists, so it uses --force-cpu (not the transient
+ # --cpu-fallback the arm64 GPU-build recovery uses).
+ args, stderr = _run(value)
+ assert "--force-cpu" in args
+ assert "--cpu-fallback" not in args
+ assert "Ignoring" not in stderr
+
+
+@_SKIP_NO_BASH
+@pytest.mark.parametrize("value", ["cpu", "CPU", " cpu "])
+def test_backend_cpu_macos_warns_no_flag(value):
+ # macOS has no CPU-only bundle (the universal build already runs on CPU), so the
+ # override warns instead of writing a misleading forced-CPU marker.
+ args, stderr = _run(value, system = "Darwin")
+ assert "--force-cpu" not in args
+ assert "--cpu-fallback" not in args
+ assert "macOS" in stderr
+
+
+@_SKIP_NO_BASH
+@pytest.mark.parametrize("value", [None, "", "auto", "AUTO", " "])
+def test_backend_auto_no_flag_no_warn(value):
+ args, stderr = _run(value)
+ assert "--force-cpu" not in args
+ assert "Ignoring" not in stderr
+
+
+@_SKIP_NO_BASH
+@pytest.mark.parametrize("value", ["vulkan", "gpu", "cuda"])
+def test_backend_unknown_warns_and_no_flag(value):
+ args, stderr = _run(value)
+ assert "--force-cpu" not in args
+ assert "Ignoring" in stderr
+
+
+@_SKIP_NO_BASH
+def test_arm64_recovery_uses_transient_cpu_fallback():
+ # The arm64 Linux GPU-build recovery must stay transient (--cpu-fallback), never
+ # the persisted --force-cpu, so a later update can still heal to a GPU bundle (#6097).
+ text = _SETUP_SH.read_text(encoding = "utf-8")
+ m = re.search(r"_ARM64_CPU_CMD=\((.*?)\)", text, re.DOTALL)
+ assert m, "arm64 CPU recovery command not found in setup.sh"
+ block = m.group(1)
+ assert "--cpu-fallback" in block
+ assert "--force-cpu" not in block
+
+
+def _ps1_search(pattern: str, flags = 0) -> str:
+ m = re.search(pattern, _SETUP_PS1.read_text(encoding = "utf-8"), flags)
+ assert m, f"setup.ps1 block not found: {pattern}"
+ return m.group(0)
+
+
+def _run_ps1(value: str | None) -> str:
+ # The override is normalized (assign + warn) at the top of the prebuilt block and
+ # applied to $prebuiltArgs lower down; compose both real snippets.
+ normalize = _ps1_search(
+ r'\$llamaBackend = "\$\(\$env:UNSLOTH_LLAMA_CPP_BACKEND\)".*?Write-Host.*?\n\s*\}',
+ re.DOTALL,
+ )
+ apply_flag = _ps1_search(
+ r'if \(\$llamaBackend -eq "cpu"\) \{\s*\$prebuiltArgs \+= "--force-cpu"\s*\}'
+ )
+ env = {k: v for k, v in os.environ.items() if k != "UNSLOTH_LLAMA_CPP_BACKEND"}
+ if value is not None:
+ env["UNSLOTH_LLAMA_CPP_BACKEND"] = value
+ harness = f'$prebuiltArgs = @()\n{normalize}\n{apply_flag}\n"ARGS:" + ($prebuiltArgs -join ",")'
+ out = subprocess.run(
+ ["pwsh", "-NoProfile", "-Command", harness],
+ capture_output = True,
+ text = True,
+ env = env,
+ check = True,
+ )
+ return out.stdout
+
+
+@_SKIP_NO_PWSH
+@pytest.mark.parametrize("value", ["cpu", "CPU", "Cpu", " cpu ", "CPU\t"])
+def test_ps1_backend_cpu_appends_flag(value):
+ out = _run_ps1(value)
+ assert "--force-cpu" in out
+ assert "Ignoring" not in out
+
+
+@_SKIP_NO_PWSH
+@pytest.mark.parametrize("value", [None, "", "auto", "AUTO", " "])
+def test_ps1_backend_auto_no_flag_no_warn(value):
+ out = _run_ps1(value)
+ assert "--force-cpu" not in out
+ assert "Ignoring" not in out
+
+
+@_SKIP_NO_PWSH
+@pytest.mark.parametrize("value", ["vulkan", "gpu", "cuda"])
+def test_ps1_backend_unknown_warns_and_no_flag(value):
+ out = _run_ps1(value)
+ assert "--force-cpu" not in out
+ assert "Ignoring" in out
diff --git a/studio/backend/utils/llama_cpp_update.py b/studio/backend/utils/llama_cpp_update.py
index 31dbda63ea..67733bde35 100644
--- a/studio/backend/utils/llama_cpp_update.py
+++ b/studio/backend/utils/llama_cpp_update.py
@@ -479,6 +479,7 @@ def _run_update(
asset: Optional[str],
script: Path,
pin_release_tag: Optional[str] = None,
+ force_cpu: bool = False,
) -> None:
"""Worker: put the backend into a maintenance state, run the installer for
the latest prebuilt, then refresh caches so the next load uses the new build.
@@ -522,6 +523,12 @@ def _run_update(
if pin_release_tag:
cmd.extend(["--published-release-tag", pin_release_tag])
cmd.extend(_rocm_install_args(asset))
+ # Re-assert a deliberate CPU install (--force-cpu) so detect_host on a GPU host
+ # does not re-route to a GPU/Vulkan bundle and revive the crash (#7213). --force-cpu
+ # (not --cpu-fallback) also re-persists force_cpu, keeping the choice across future
+ # updates. A natural fallback (or a legacy marker without the flag) heals to GPU (#6097).
+ if force_cpu:
+ cmd.append("--force-cpu")
logger.info("llama update: installing", cmd = " ".join(cmd))
# Stream progress lines into job["progress"].
env = dict(os.environ, UNSLOTH_PROGRESS_PERCENT_STEP = "5")
@@ -671,6 +678,7 @@ def start_update() -> dict:
repo = marker.get("published_repo") or DEFAULT_PUBLISHED_REPO
from_tag = marker.get("tag") or marker.get("release_tag")
asset = marker.get("asset")
+ force_cpu = bool(marker.get("force_cpu"))
# Install exactly the release the banner offered: the installer's own
# "latest" is commit-date ordered and can lag the published_at pick
# above, reinstalling the current build in a loop (the #6219 class).
@@ -705,6 +713,8 @@ def start_update() -> dict:
repo = (res or {}).get("repo") or DEFAULT_PUBLISHED_REPO
from_tag = None
asset = (res or {}).get("asset")
+ # Source builds carry no forced-CPU marker, so nothing to preserve here.
+ force_cpu = False
# No pin: source-build detection resolves via --resolve-prebuilt latest,
# the same resolver the unpinned apply uses, so the two already agree.
pin_release_tag = None
@@ -735,7 +745,7 @@ def start_update() -> dict:
thread = threading.Thread(
target = _run_update,
- args = (install_dir, repo, asset, script, pin_release_tag),
+ args = (install_dir, repo, asset, script, pin_release_tag, force_cpu),
name = "llama-cpp-update",
daemon = True,
)
diff --git a/studio/install_llama_prebuilt.py b/studio/install_llama_prebuilt.py
index 9bbd0cb8be..b8182a534b 100644
--- a/studio/install_llama_prebuilt.py
+++ b/studio/install_llama_prebuilt.py
@@ -6450,6 +6450,7 @@ def write_prebuilt_metadata(
choice: AssetChoice,
approved_checksums: ApprovedReleaseChecksums,
prebuilt_fallback_used: bool,
+ force_cpu: bool = False,
) -> None:
source_asset_name, source_sha256 = selected_source_archive_metadata(
approved_checksums,
@@ -6474,6 +6475,10 @@ def write_prebuilt_metadata(
"release_tag": release_tag,
"published_repo": approved_checksums.repo,
"asset": choice.name,
+ # True only for a deliberate CPU choice (--force-cpu). The updater re-asserts it
+ # so a forced CPU install is not re-routed to a GPU bundle (#7213). An automatic
+ # --cpu-fallback (e.g. arm64 GPU-build recovery) stays False so it can heal to GPU.
+ "force_cpu": force_cpu,
"asset_sha256": choice.expected_sha256,
"source": choice.source_label,
# Binary-side repo/tag for non-fork sources (e.g. the ggml-org upstream
@@ -6501,6 +6506,24 @@ def write_prebuilt_metadata(
(install_dir / "UNSLOTH_PREBUILT_INFO.json").write_text(json.dumps(metadata, indent = 2) + "\n")
+def sync_marker_force_cpu(install_dir: Path, persist_force_cpu: bool) -> None:
+ """Sync only the force_cpu flag of an existing marker when the resolved bundle is
+ unchanged, so the install is skipped without a full metadata rewrite. A deliberate
+ --force-cpu on top of a naturally installed CPU bundle (same asset) must still be
+ recorded, else the updater will not re-assert it and can re-route the install to a
+ GPU/Vulkan bundle that revives the crash (#7213)."""
+ marker_path = install_dir / "UNSLOTH_PREBUILT_INFO.json"
+ try:
+ marker = json.loads(marker_path.read_text())
+ except (OSError, ValueError):
+ return
+ if not isinstance(marker, dict) or bool(marker.get("force_cpu")) == persist_force_cpu:
+ return
+ marker["force_cpu"] = persist_force_cpu
+ marker_path.write_text(json.dumps(marker, indent = 2) + "\n")
+ log(f"existing install reused; recorded force_cpu={persist_force_cpu} from this run")
+
+
def expected_install_fingerprint(
*,
llama_tag: str,
@@ -6746,6 +6769,7 @@ def validate_prebuilt_choice(
approved_checksums: ApprovedReleaseChecksums,
prebuilt_fallback_used: bool,
quantized_path: Path,
+ force_cpu: bool = False,
) -> tuple[Path, Path]:
source_repo, source_ref, source_archive, exact_source = preferred_source_archive(
approved_checksums, llama_tag
@@ -6786,6 +6810,7 @@ def validate_prebuilt_choice(
choice = choice,
approved_checksums = approved_checksums,
prebuilt_fallback_used = prebuilt_fallback_used,
+ force_cpu = force_cpu,
)
# Hashless external prebuilts are not in the approved-sha256
# manifest and rely on the functional smoke test as their only integrity gate,
@@ -6828,6 +6853,7 @@ def validate_prebuilt_attempts(
approved_checksums: ApprovedReleaseChecksums,
initial_fallback_used: bool = False,
existing_install_dir: Path | None = None,
+ force_cpu: bool = False,
) -> tuple[AssetChoice, Path, bool]:
attempt_list = list(attempts)
if not attempt_list:
@@ -6880,6 +6906,7 @@ def validate_prebuilt_attempts(
approved_checksums = approved_checksums,
prebuilt_fallback_used = tried_fallback,
quantized_path = quantized_path,
+ force_cpu = force_cpu,
)
except Exception as exc:
remove_tree(staging_dir)
@@ -6939,8 +6966,8 @@ def _route_to_vulkan_prebuilt(
"""Point a Vulkan-capable host at the upstream ggml-org Vulkan prebuilt.
The unsloth published repo ships only CUDA/ROCm/CPU assets, so Vulkan comes
- from UPSTREAM_REPO. Two triggers route here, both suppressed under
- --cpu-fallback (the explicit "give me CPU" last resort wins):
+ from UPSTREAM_REPO. Two triggers route here, both suppressed when a CPU flag
+ (--cpu-fallback or --force-cpu, folded into force_cpu) wins:
* UNSLOTH_FORCE_VULKAN forces Vulkan over the detected CUDA/ROCm backend;
* an auto-detected Intel GPU with NO physical NVIDIA/ROCm -- the purpose
of the has_intel_gpu probe, since the fork manifest ships no Vulkan asset.
@@ -7020,8 +7047,11 @@ def install_prebuilt(
override_has_rocm: bool = False,
override_rocm_gfx: str | None = None,
force_cpu: bool = False,
+ persist_force_cpu: bool = False,
instruction_cleanup_root: Path | None = None,
) -> None:
+ # force_cpu drops GPU detection (mechanism, both --cpu-fallback and --force-cpu);
+ # persist_force_cpu records the deliberate choice so the updater re-asserts it.
host = detect_host()
host = _apply_host_overrides(
host,
@@ -7072,6 +7102,9 @@ def install_prebuilt(
"existing llama.cpp install already matches selected release "
f"{current.release_tag} upstream_tag={current.llama_tag}; skipping download and install"
)
+ # Reused bundle is unchanged, but a fresh --force-cpu still must be
+ # recorded so the updater re-asserts it (#7213).
+ sync_marker_force_cpu(install_dir, persist_force_cpu)
return
with tempfile.TemporaryDirectory(prefix = "unsloth-llama-prebuilt-") as tmp:
work_dir = Path(tmp)
@@ -7092,6 +7125,7 @@ def install_prebuilt(
"existing llama.cpp install already matches fallback release "
f"{plan.release_tag} upstream_tag={plan.llama_tag}; skipping reinstall"
)
+ sync_marker_force_cpu(install_dir, persist_force_cpu)
return
log(
"selected "
@@ -7112,6 +7146,8 @@ def install_prebuilt(
initial_fallback_used = release_index > 0,
# Skip is gated per-attempt inside, so pass the dir always.
existing_install_dir = install_dir,
+ # Persist only the deliberate choice, not a transient fallback.
+ force_cpu = persist_force_cpu,
)
except ExistingInstallSatisfied:
return
@@ -7209,8 +7245,21 @@ def parse_args() -> argparse.Namespace:
default = False,
help = (
"Select the CPU prebuilt for this OS/arch even when a GPU is present. "
- "setup.sh uses this as a last resort for arm64 Linux GPU hosts whose "
- "source build failed (no arm64 CUDA prebuilt exists anywhere)."
+ "Automatic/transient: setup.sh uses this as a last resort for arm64 Linux "
+ "GPU hosts whose source build failed. Does NOT persist, so a later update "
+ "heals back to a GPU bundle once one is available (#6097). Use --force-cpu "
+ "for a deliberate CPU-only choice that survives updates."
+ ),
+ )
+ parser.add_argument(
+ "--force-cpu",
+ action = "store_true",
+ default = False,
+ help = (
+ "Deliberate CPU-only install (UNSLOTH_LLAMA_CPP_BACKEND=cpu). Drops GPU "
+ "detection like --cpu-fallback but also records force_cpu in the marker, so "
+ "the in-app updater re-asserts CPU and never re-routes to a GPU/Vulkan "
+ "bundle that would revive the Intel iGPU crash (#7213)."
),
)
resolve_group = parser.add_mutually_exclusive_group()
@@ -7333,16 +7382,19 @@ def main() -> int:
# Host-aware "is a prebuilt available" probe, no download. Every host now
# plans against the fork (args.published_repo defaults to it); an explicit
# --published-repo overrides. PrebuiltFallback == source build.
+ # Both flags drop GPU detection; --force-cpu additionally persists (install
+ # path only). The probe only needs the mechanism, so OR them.
+ _cpu_mechanism = args.cpu_fallback or args.force_cpu
host = _apply_host_overrides(
detect_host(),
override_has_rocm = args.has_rocm,
override_rocm_gfx = args.rocm_gfx,
- force_cpu = args.cpu_fallback,
+ force_cpu = _cpu_mechanism,
)
# Same Vulkan routing the install path applies, so the probe's answer
# matches what would install (an Intel/forced-Vulkan host -> upstream).
host, repo, release_tag = _route_to_vulkan_prebuilt(
- host, args.published_repo, args.published_release_tag or "", force_cpu = args.cpu_fallback
+ host, args.published_repo, args.published_release_tag or "", force_cpu = _cpu_mechanism
)
try:
_requested, plans = resolve_simple_install_release_plans(
@@ -7380,7 +7432,10 @@ def main() -> int:
published_release_tag = args.published_release_tag or "",
override_has_rocm = args.has_rocm,
override_rocm_gfx = args.rocm_gfx,
- force_cpu = args.cpu_fallback,
+ # Both drop GPU detection; only --force-cpu (deliberate) is recorded so the
+ # updater re-asserts it. --cpu-fallback stays transient and heals to GPU.
+ force_cpu = args.cpu_fallback or args.force_cpu,
+ persist_force_cpu = args.force_cpu,
instruction_cleanup_root = install_arg.absolute(),
)
return EXIT_SUCCESS
diff --git a/studio/setup.ps1 b/studio/setup.ps1
index 98e801cd3c..f7d33a1142 100644
--- a/studio/setup.ps1
+++ b/studio/setup.ps1
@@ -32,6 +32,10 @@ $PackageDir = Split-Path -Parent $ScriptDir
# (no matching GitHub release), forces a source build, and causes HTTP 422
# errors. Only use "master" temporarily when the latest release is missing
# support for a new model architecture.
+#
+# UNSLOTH_LLAMA_CPP_BACKEND : "auto" (default) or "cpu". When "cpu", forces
+# the CPU-only prebuilt bundle on GPU hosts. Fixes Intel iGPU Vulkan
+# crashes (#7213).
$DefaultLlamaPrForce = ""
$DefaultLlamaSource = "https://github.com/ggml-org/llama.cpp"
$DefaultLlamaTag = "latest"
@@ -3367,6 +3371,15 @@ if ($LocalLlamaCppLinked) {
if ($env:UNSLOTH_LLAMA_RELEASE_TAG) {
$prebuiltArgs += @("--published-release-tag", $env:UNSLOTH_LLAMA_RELEASE_TAG)
}
+ # UNSLOTH_LLAMA_CPP_BACKEND=cpu (case-insensitive, whitespace-trimmed) forces the
+ # CPU-only prebuilt via --force-cpu (persisted so updates keep it). Fixes Intel
+ # iGPU Vulkan crash (#7213).
+ $llamaBackend = "$($env:UNSLOTH_LLAMA_CPP_BACKEND)".Trim().ToLowerInvariant()
+ if ($llamaBackend -eq "cpu") {
+ $prebuiltArgs += "--force-cpu"
+ } elseif ($llamaBackend -and $llamaBackend -ne "auto") {
+ Write-Host "[WARN] Ignoring UNSLOTH_LLAMA_CPP_BACKEND='$($env:UNSLOTH_LLAMA_CPP_BACKEND)' (expected 'auto' or 'cpu')" -ForegroundColor Yellow
+ }
$prevEAPPrebuilt = $ErrorActionPreference
$ErrorActionPreference = "Continue"
$previousNativeErrorPreference = $null
diff --git a/studio/setup.sh b/studio/setup.sh
index 8d47eecfda..df7178c662 100755
--- a/studio/setup.sh
+++ b/studio/setup.sh
@@ -36,6 +36,10 @@ fi
# forces a source build, and causes HTTP 422 errors.
# Only use "master" temporarily when the latest release
# is missing support for a new model architecture.
+#
+# UNSLOTH_LLAMA_CPP_BACKEND : "auto" (default) or "cpu". When "cpu", forces
+# the CPU-only prebuilt bundle on GPU hosts.
+# Fixes Intel iGPU Vulkan crashes (#7213).
# ──────────────────────────────────────────────────────────────────────────
_DEFAULT_LLAMA_PR_FORCE=""
_DEFAULT_LLAMA_SOURCE="https://github.com/ggml-org/llama.cpp"
@@ -1359,6 +1363,22 @@ else
# present so it can still attempt a prebuilt. Mirrors setup.ps1 behaviour.
_PREBUILT_CMD+=(--has-rocm)
fi
+ # UNSLOTH_LLAMA_CPP_BACKEND=cpu (case-insensitive, trimmed) forces the CPU-only
+ # prebuilt via --force-cpu, bypassing Vulkan/CUDA/ROCm. Fixes Intel iGPU crash (#7213).
+ # No effect on macOS: the universal bundle already runs on CPU (Metal is a runtime
+ # -ngl choice), so warn instead of writing a misleading forced-CPU marker.
+ _llama_backend="$(printf '%s' "${UNSLOTH_LLAMA_CPP_BACKEND:-auto}" | awk '{$1=$1; print tolower($0)}')"
+ case "$_llama_backend" in
+ cpu)
+ if [ "$_HOST_SYSTEM" = "Darwin" ]; then
+ step "llama.cpp" "UNSLOTH_LLAMA_CPP_BACKEND=cpu has no effect on macOS (universal build; use -ngl 0 at runtime for CPU-only)" "$C_WARN" >&2
+ else
+ _PREBUILT_CMD+=(--force-cpu)
+ fi
+ ;;
+ ""|auto) ;;
+ *) step "llama.cpp" "Ignoring UNSLOTH_LLAMA_CPP_BACKEND='$UNSLOTH_LLAMA_CPP_BACKEND' (expected 'auto' or 'cpu')" "$C_WARN" >&2 ;;
+ esac
_PREBUILT_LOG="$(mktemp)"
set +e
if _is_verbose; then
diff --git a/tests/studio/install/test_install_llama_prebuilt_logic.py b/tests/studio/install/test_install_llama_prebuilt_logic.py
index e995e5033e..9a094ddc0e 100644
--- a/tests/studio/install/test_install_llama_prebuilt_logic.py
+++ b/tests/studio/install/test_install_llama_prebuilt_logic.py
@@ -1348,6 +1348,7 @@ def test_install_prebuilt_falls_back_to_older_release_plan(
approved_checksums,
initial_fallback_used = False,
existing_install_dir = None,
+ force_cpu = False,
):
call_log.append((llama_tag, initial_fallback_used))
if llama_tag == "b9002":
@@ -2551,6 +2552,7 @@ def test_install_prebuilt_skips_when_older_release_fallback_matches_existing_ins
approved_checksums,
initial_fallback_used = False,
existing_install_dir = None,
+ force_cpu = False,
):
call_log.append(llama_tag)
raise PrebuiltFallback("validation failed for latest release")
@@ -2698,6 +2700,7 @@ def test_install_prebuilt_skips_same_release_fallback_attempt_when_installed(
approved_checksums,
prebuilt_fallback_used,
quantized_path,
+ force_cpu = False,
):
attempted_names.append(choice.name)
if choice.name == first_choice.name:
@@ -2824,6 +2827,7 @@ def test_install_prebuilt_same_tag_upstream_failure_uses_older_unsloth_release_p
approved_checksums,
initial_fallback_used = False,
existing_install_dir = None,
+ force_cpu = False,
):
attempted.append((llama_tag, release_tag, attempts[0].source_label))
if llama_tag == "b9002":
From 07272b9278eaa2813c30f3b12f712276ff97fa01 Mon Sep 17 00:00:00 2001
From: Daniel Han
Date: Mon, 20 Jul 2026 00:57:02 -0700
Subject: [PATCH 13/41] Experimental: correct varlen sample packing for hybrid
linear-attention models (#7249)
* Fix text-only VLM CPT packing truncation
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Handle streaming vision datasets in packing
* Harden multimodal packing detection
* Preserve safe packing boundaries
* Scope stream packing checks to VLMs
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Narrow VLM packing detection
* Align packing mode and eval safety
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Add qwen3_5/qwen3_next to PADDING_FREE_BLOCKLIST to avoid packed-sequence contamination
* Detect hybrid linear-attention models structurally instead of by name for packing guard
* Add experimental varlen packing for hybrid linear-attention models
Feed seq_idx to the causal conv and cu_seqlens to the gated-delta scan so
sample packing / padding-free reset state at sequence boundaries for hybrid
linear-attention models (Qwen3.5, Qwen3-Next). Gated behind
UNSLOTH_EXPERIMENTAL_HYBRID_PACKING and fail-closed: when the flag is off or
the accelerated kernels (causal_conv1d + fla) are unavailable, the guard keeps
these models on the padded path.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Harden hybrid linear-attention varlen packing shim
Make patch_hybrid_linear_attention_varlen robust across transformers 4.57.6
through 5.x and TRL 0.22.2 through 1.x, following the import_fixes.py style:
- Read UNSLOTH_EXPERIMENTAL_HYBRID_PACKING at call time so the flag takes effect
when set after importing unsloth.
- Idempotent: repeat calls on a patched model return True without re-validating
the wrappers or double-wrapping; signatures are checked on captured originals.
- Prefer the authoritative packed_seq_lengths (via get_packed_info_from_kwargs)
over position_ids resets, handling pad_to_multiple_of trailing tokens.
- Suppress injection for cached forwards (use_cache / past_key_values) so
generation and eval are left on the untouched decode path.
- Validate every gated-delta module before mutating any (transactional).
- Bind position_ids / use_cache from both positional and keyword args.
- Verify dispatch at runtime (Unsloth wraps each module forward, so the mixer
source is not statically inspectable) and warn once if the shim is never hit.
- Emit one deduped diagnostic on each fail-closed path.
Add CPU unit tests covering the hybrid guard detection, the boundary builders,
and the shim (fail-closed, active, idempotent, cached no-op, runtime handshake).
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Abort hybrid packing when the varlen shim is not fully dispatched
The runtime handshake used a single per-module hit flag written by both the conv
and scan wrappers, so a partial dispatch (only one kernel routed through
self.) passed the any() check and trained on contaminated data, and a
missing dispatch only logged a warning. Track conv and scan dispatch separately,
require both on every gated-delta module on the first packed forward, and raise
before loss/backward when either is missing (the batch is already flattened, so
there is no padded recovery at that point). Also skip an empty packed_seq_lengths
before it reaches max(), and document the position_ids fallback's left-pad
assumption.
Add tests for no-dispatch and partial (conv-only / scan-only) abort, the
packed_seq_lengths preference over a competing position_ids, MRoPE 3D position
ids, and the pad_to_multiple_of trailing-segment path through the metadata builder.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Import the hybrid packing patch from its submodule to satisfy the import-hoist lint
* Fail closed for hybrid packing on encoder-decoder, chunked-loss, and string-name models
The varlen shim only helps decoder-only hybrid models that run their mixer
through self. on a live nn.Module forward. Three cases slipped past
the guard:
- Encoder-decoder configs (is_encoder_decoder) reached the packing path even
though flattening a cross-attention batch is unsound. Block them explicitly.
- TRL's chunked_nll loss (the 1.x default) calls the backbone directly and
bypasses model.forward, so the per-instance forward wrapper that refreshes
the varlen stash never runs. Detect that path and keep the model padded.
- A string model_name reaches the trainer before the module exists, so the
instance shim has nothing to patch. Resolve the config up front and keep
string hybrids on the padded path.
Adds encoder-decoder / decoder-only / chunked-loss / string-model tests.
* Harden the SFT source-injection replacements and forward auth args for string models
The wrapped-packing injection rewrote the sourced unsloth_zoo sft_prepare_dataset
with str.replace anchored on the exact 'All Unsloth Zoo code licensed under
LGPLv3' comment. str.replace never raises on a missing anchor, so a supported
newer unsloth_zoo (the dependency is only lower-bounded) that moved that header
would silently drop the setup while the truncation and pack_dataset edits still
referenced _unsloth_wrapped_packing / _inspect, raising NameError on every SFT
dataset preparation.
- Install the setup at the sft_prepare_dataset signature via re.subn (a structural
anchor that always exists) and raise if even that is missing.
- Route the remaining edits through a _require_replace helper that fails loudly on a
missing required anchor (or warns once for an optional one), formalizing the
verify-then-replace idiom the DPO patchers in this file already use.
- Reuse the guarded _unsloth_pack_has_strategy at the pack_dataset call instead of
re-calling inspect.signature(pack_dataset) unguarded, so a non-introspectable
pack_dataset cannot crash there after the setup already handled it.
- _resolve_string_model_config now forwards token / use_auth_token / cache_dir /
code_revision, so a private hybrid resolves its config instead of falling through
as non-hybrid and enabling packing without the varlen shim.
Adds regression tests for the drift-resistant injection, the helper, and the
string-model auth forwarding.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Honor top-level SFTConfig.trust_remote_code when resolving a string model
TRL merges the top-level args.trust_remote_code into the load via
model_init_kwargs.setdefault("trust_remote_code", args.trust_remote_code) before
create_model_from_path, so a remote-code hybrid is commonly set with
SFTConfig(trust_remote_code=True) rather than inside model_init_kwargs. The config
probe only read model_init_kwargs, so AutoConfig could fail for such a model, leave
model_config None, and let the guard treat it as non-hybrid, enabling packing
without the varlen shim. Mirror TRL's setdefault (model_init_kwargs wins).
* Tighten hybrid-packing comments for concision
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
---------
Co-authored-by: alkinun
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Etherl <61019402+Etherll@users.noreply.github.com>
---
tests/utils/test_packing.py | 579 +++++++++++++++++++++++++++---
unsloth/models/rl_replacements.py | 90 +++--
unsloth/trainer.py | 93 ++++-
unsloth/utils/packing.py | 302 ++++++++++++++++
4 files changed, 984 insertions(+), 80 deletions(-)
diff --git a/tests/utils/test_packing.py b/tests/utils/test_packing.py
index 98c29d9f0f..1b8bb65058 100644
--- a/tests/utils/test_packing.py
+++ b/tests/utils/test_packing.py
@@ -15,6 +15,7 @@
from unsloth import FastLanguageModel
import unsloth.trainer as trainer_module
+import unsloth.utils.packing as packing_module
from unsloth.utils import attention_dispatch as attention_dispatch_utils
from unsloth.utils.packing import (
configure_padding_free,
@@ -22,6 +23,7 @@ from unsloth.utils.packing import (
enable_padding_free_metadata,
enable_sample_packing,
mask_packed_sequence_boundaries,
+ patch_hybrid_linear_attention_varlen,
)
from contextlib import ExitStack
@@ -161,6 +163,327 @@ def test_configure_padding_free():
assert config.remove_unused_columns is False
+# --- Hybrid linear-attention guard + varlen shim (PR #7211 / #7249) ---------------
+
+
+def _hybrid_config_model():
+ # Qwen3.5 / Qwen3-Next style: explicit linear_attention layer schedule.
+ return SimpleNamespace(
+ config = SimpleNamespace(layer_types = ["linear_attention", "full_attention"])
+ )
+
+
+def _gemma3_model():
+ # Has layer_types but no linear_attention -> must NOT be flagged as hybrid.
+ return SimpleNamespace(
+ config = SimpleNamespace(
+ model_type = "gemma3", layer_types = ["sliding_attention", "full_attention"]
+ ),
+ )
+
+
+def _dense_qwen3_model():
+ return SimpleNamespace(
+ config = SimpleNamespace(model_type = "qwen3", architectures = ["Qwen3ForCausalLM"])
+ )
+
+
+class _FakeGatedDeltaNet(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.conv1d = torch.nn.Conv1d(4, 4, 3, groups = 4)
+ self.A_log = torch.nn.Parameter(torch.zeros(4))
+
+ def forward(self, hidden_states, **kwargs): # dispatch through self.
+ return self.chunk_gated_delta_rule(self.causal_conv1d_fn(hidden_states))
+
+
+class _FakeHybridModel(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.config = SimpleNamespace() # no markers -> forces module-level detection
+ self.linear_attn = _FakeGatedDeltaNet()
+
+
+def test_is_hybrid_linear_attention_detects_and_excludes():
+ is_hybrid = trainer_module._is_hybrid_linear_attention_model
+ assert is_hybrid(_hybrid_config_model()) is True
+ assert is_hybrid(_FakeHybridModel()) is True # module-structural evidence
+ assert is_hybrid(_text_model()) is False # Llama
+ assert is_hybrid(_gemma3_model()) is False # layer_types without linear_attention
+ assert is_hybrid(_dense_qwen3_model()) is False # dense Qwen3
+ assert is_hybrid(None) is False
+
+
+def test_varlen_from_position_ids():
+ cu, seq_idx = packing_module._varlen_from_position_ids(torch.tensor([[0, 1, 0, 0, 1, 2]]))
+ assert cu.tolist() == [0, 2, 3, 6]
+ assert seq_idx.tolist() == [[0, 0, 1, 2, 2, 2]]
+ assert (
+ packing_module._varlen_from_position_ids(torch.tensor([[0, 1, 2, 3]])) is None
+ ) # single sequence
+ assert packing_module._varlen_from_position_ids(torch.tensor([[1, 2, 3]])) is None # first != 0
+ assert (
+ packing_module._varlen_from_position_ids(torch.tensor([[0, 1], [0, 1]])) is None
+ ) # normal 2-row batch
+ assert packing_module._varlen_from_position_ids(None) is None
+
+
+def test_seq_idx_from_cu_seqlens_handles_trailing_pad():
+ cu = torch.tensor([0, 2, 5], dtype = torch.int32)
+ boundaries, seq_idx = packing_module._seq_idx_from_cu_seqlens(cu, total = 8) # pad_to_multiple_of
+ assert boundaries.tolist() == [0, 2, 5, 8]
+ assert seq_idx.tolist() == [[0, 0, 1, 1, 1, 2, 2, 2]]
+ boundaries2, _ = packing_module._seq_idx_from_cu_seqlens(cu, total = 5) # exact fit
+ assert boundaries2.tolist() == [0, 2, 5]
+ assert (
+ packing_module._seq_idx_from_cu_seqlens(torch.tensor([1, 2], dtype = torch.int32), total = 2)
+ is None
+ )
+ assert packing_module._seq_idx_from_cu_seqlens(cu, total = 3) is None # boundaries exceed total
+
+
+def test_hybrid_varlen_metadata_prefers_packed_seq_lengths():
+ # A competing position_ids would segment [0, 3, 6]; packed_seq_lengths must win.
+ kwargs = {
+ "input_ids": torch.zeros(1, 6, dtype = torch.long),
+ "packed_seq_lengths": torch.tensor([2, 1, 3], dtype = torch.int32),
+ "position_ids": torch.tensor([[0, 1, 2, 0, 1, 2]]),
+ }
+ cu, seq_idx = packing_module._hybrid_varlen_metadata(kwargs)
+ assert cu.tolist() == [0, 2, 3, 6]
+ assert seq_idx.tolist() == [[0, 0, 1, 2, 2, 2]]
+
+
+def test_hybrid_varlen_metadata_suppressed_when_cached():
+ base = {
+ "input_ids": torch.zeros(1, 6, dtype = torch.long),
+ "packed_seq_lengths": torch.tensor([2, 1, 3], dtype = torch.int32),
+ }
+ assert packing_module._hybrid_varlen_metadata({**base, "use_cache": True}) is None
+ assert packing_module._hybrid_varlen_metadata({**base, "past_key_values": object()}) is None
+
+
+def test_hybrid_varlen_metadata_none_for_plain_batch():
+ kwargs = {
+ "input_ids": torch.zeros(1, 4, dtype = torch.long),
+ "position_ids": torch.tensor([[0, 1, 2, 3]]),
+ }
+ assert packing_module._hybrid_varlen_metadata(kwargs) is None
+
+
+def _make_fake_kernels():
+ def causal_conv1d_fn(
+ x,
+ weight = None,
+ bias = None,
+ activation = None,
+ seq_idx = None,
+ ):
+ causal_conv1d_fn.calls.append(seq_idx)
+ return x
+
+ causal_conv1d_fn.calls = []
+
+ def chunk_gated_delta_rule(
+ q,
+ k = None,
+ v = None,
+ cu_seqlens = None,
+ **kw,
+ ):
+ chunk_gated_delta_rule.calls.append(cu_seqlens)
+ return q
+
+ chunk_gated_delta_rule.calls = []
+ return causal_conv1d_fn, chunk_gated_delta_rule
+
+
+class _ShimGatedDeltaNet(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.conv1d = torch.nn.Conv1d(4, 4, 3, groups = 4)
+ self.causal_conv1d_fn, self.chunk_gated_delta_rule = _make_fake_kernels()
+
+ def forward(self, hidden_states, **kwargs):
+ return self.chunk_gated_delta_rule(self.causal_conv1d_fn(hidden_states))
+
+
+class _ShimHybridModel(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.config = SimpleNamespace(layer_types = ["linear_attention", "full_attention"])
+ self.linear_attn = _ShimGatedDeltaNet()
+
+ def forward(
+ self,
+ input_ids = None,
+ position_ids = None,
+ packed_seq_lengths = None,
+ use_cache = None,
+ **kwargs,
+ ):
+ return self.linear_attn(input_ids.float())
+
+
+def test_patch_hybrid_varlen_flag_off(monkeypatch):
+ monkeypatch.delenv("UNSLOTH_EXPERIMENTAL_HYBRID_PACKING", raising = False)
+ model = _ShimHybridModel()
+ assert patch_hybrid_linear_attention_varlen(model) is False
+ assert not getattr(model, "_unsloth_varlen_forward_wrapped", False)
+
+
+def test_patch_hybrid_varlen_active_and_idempotent(monkeypatch):
+ monkeypatch.setenv("UNSLOTH_EXPERIMENTAL_HYBRID_PACKING", "1")
+ model = _ShimHybridModel()
+ conv_orig, scan_orig = (
+ model.linear_attn.causal_conv1d_fn,
+ model.linear_attn.chunk_gated_delta_rule,
+ )
+
+ assert patch_hybrid_linear_attention_varlen(model) is True
+ assert model._unsloth_varlen_forward_wrapped is True
+ assert model.linear_attn._unsloth_varlen_wrapped is True
+ assert patch_hybrid_linear_attention_varlen(model) is True # idempotent, no double-wrap
+
+ conv_orig.calls.clear()
+ scan_orig.calls.clear()
+ packing_module._HYBRID_WARNED.clear()
+ ids = torch.zeros(1, 6, dtype = torch.long)
+ model(
+ input_ids = ids,
+ packed_seq_lengths = torch.tensor([2, 1, 3], dtype = torch.int32),
+ use_cache = False,
+ )
+ assert conv_orig.calls[-1] is not None # seq_idx injected
+ assert scan_orig.calls[-1].tolist() == [0, 2, 3, 6] # cu_seqlens injected
+ assert not packing_module._HYBRID_WARNED # handshake passed, no rejection
+
+ conv_orig.calls.clear()
+ scan_orig.calls.clear()
+ model(
+ input_ids = ids, packed_seq_lengths = torch.tensor([2, 1, 3], dtype = torch.int32), use_cache = True
+ )
+ assert conv_orig.calls[-1] is None # cached forward -> no injection
+ assert scan_orig.calls[-1] is None
+
+
+def test_patch_hybrid_varlen_torch_fallback_fail_closed(monkeypatch):
+ monkeypatch.setenv("UNSLOTH_EXPERIMENTAL_HYBRID_PACKING", "1")
+ model = _ShimHybridModel()
+
+ def torch_chunk_gated_delta_rule(
+ q,
+ cu_seqlens = None,
+ **kw,
+ ):
+ return q
+
+ model.linear_attn.chunk_gated_delta_rule = torch_chunk_gated_delta_rule
+ assert patch_hybrid_linear_attention_varlen(model) is False
+ assert not getattr(model, "_unsloth_varlen_forward_wrapped", False)
+
+
+def test_patch_hybrid_varlen_bad_signature_fail_closed(monkeypatch):
+ monkeypatch.setenv("UNSLOTH_EXPERIMENTAL_HYBRID_PACKING", "1")
+ model = _ShimHybridModel()
+
+ def scan_no_cu(q, **kw): # missing cu_seqlens
+ return q
+
+ model.linear_attn.chunk_gated_delta_rule = scan_no_cu
+ assert patch_hybrid_linear_attention_varlen(model) is False
+
+
+def _hybrid_model_with_gdn(gdn_forward):
+ # Build a fake hybrid model whose gated-delta mixer forward is `gdn_forward`.
+ class _GatedDeltaNet(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.conv1d = torch.nn.Conv1d(4, 4, 3, groups = 4)
+ self.causal_conv1d_fn, self.chunk_gated_delta_rule = _make_fake_kernels()
+
+ forward = gdn_forward
+
+ class _Model(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.config = SimpleNamespace(layer_types = ["linear_attention", "full_attention"])
+ self.linear_attn = _GatedDeltaNet()
+
+ def forward(
+ self,
+ input_ids = None,
+ packed_seq_lengths = None,
+ use_cache = None,
+ **kwargs,
+ ):
+ return self.linear_attn(input_ids.float())
+
+ return _Model()
+
+
+def test_patch_hybrid_varlen_no_dispatch_aborts(monkeypatch):
+ # Dispatch is verified at runtime, not statically. A mixer that never calls
+ # self. installs the shim, but the first packed forward aborts (both
+ # boundary kernels are load-bearing).
+ monkeypatch.setenv("UNSLOTH_EXPERIMENTAL_HYBRID_PACKING", "1")
+ model = _hybrid_model_with_gdn(lambda self, hidden_states, **kw: hidden_states)
+ assert patch_hybrid_linear_attention_varlen(model) is True # kernels valid -> installs
+ with pytest.raises(RuntimeError, match = "both invoked"):
+ model(
+ input_ids = torch.zeros(1, 6),
+ packed_seq_lengths = torch.tensor([2, 1, 3], dtype = torch.int32),
+ use_cache = False,
+ )
+
+
+def test_patch_hybrid_varlen_partial_dispatch_aborts(monkeypatch):
+ # Only the conv fires; the scan would leak state. Both must be invoked, so abort.
+ monkeypatch.setenv("UNSLOTH_EXPERIMENTAL_HYBRID_PACKING", "1")
+ conv_only = _hybrid_model_with_gdn(
+ lambda self, hidden_states, **kw: self.causal_conv1d_fn(hidden_states)
+ )
+ assert patch_hybrid_linear_attention_varlen(conv_only) is True
+ with pytest.raises(RuntimeError, match = "both invoked"):
+ conv_only(
+ input_ids = torch.zeros(1, 6),
+ packed_seq_lengths = torch.tensor([2, 1, 3], dtype = torch.int32),
+ use_cache = False,
+ )
+
+ scan_only = _hybrid_model_with_gdn(
+ lambda self, hidden_states, **kw: self.chunk_gated_delta_rule(hidden_states)
+ )
+ assert patch_hybrid_linear_attention_varlen(scan_only) is True
+ with pytest.raises(RuntimeError, match = "both invoked"):
+ scan_only(
+ input_ids = torch.zeros(1, 6),
+ packed_seq_lengths = torch.tensor([2, 1, 3], dtype = torch.int32),
+ use_cache = False,
+ )
+
+
+def test_varlen_from_position_ids_mrope_3d():
+ pos = (
+ torch.tensor([[0, 1, 0, 0, 1, 2]]).unsqueeze(0).expand(3, 1, 6).clone()
+ ) # [3,1,T] text plane
+ cu, seq_idx = packing_module._varlen_from_position_ids(pos)
+ assert cu.tolist() == [0, 2, 3, 6]
+ assert seq_idx.tolist() == [[0, 0, 1, 2, 2, 2]]
+
+
+def test_hybrid_varlen_metadata_trailing_pad():
+ # packed_seq_lengths sum to 6 but the flattened input is 8 (pad_to_multiple_of).
+ kwargs = {
+ "input_ids": torch.zeros(1, 8, dtype = torch.long),
+ "packed_seq_lengths": torch.tensor([2, 1, 3], dtype = torch.int32),
+ }
+ cu, seq_idx = packing_module._hybrid_varlen_metadata(kwargs)
+ assert cu.tolist() == [0, 2, 3, 6, 8]
+ assert seq_idx.tolist() == [[0, 0, 1, 2, 2, 2, 3, 3]]
+
+
def _patch_fake_sft_trainer():
class FakeSFTTrainer:
def __init__(self, *args, **kwargs):
@@ -245,29 +568,101 @@ def test_vlm_without_processing_class_still_disables_packing():
("t5", "T5ForConditionalGeneration"),
("bart", "BartForConditionalGeneration"),
("whisper", "WhisperForConditionalGeneration"),
- ("csm", "CsmForConditionalGeneration"),
),
)
-def test_nonvision_conditional_generation_keeps_packing(model_type, architecture):
+def test_encoder_decoder_disables_packing(model_type, architecture):
+ # Text-only encoder-decoder models are not VLMs, but their bidirectional encoder
+ # attends across concatenated samples once padding-free drops attention_mask.
fake_trainer = _patch_fake_sft_trainer()
config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
model = SimpleNamespace(
- config = SimpleNamespace(model_type = model_type, architectures = [architecture]),
+ config = SimpleNamespace(
+ model_type = model_type,
+ architectures = [architecture],
+ is_encoder_decoder = True,
+ ),
max_seq_length = 16,
)
- trainer = fake_trainer(
- model,
- config,
- None,
- Dataset.from_dict({"text": ["text-only sample"]}),
+ trainer = fake_trainer(model, config, None, Dataset.from_dict({"text": ["text-only sample"]}))
+
+ assert config.packing is False
+ assert config.padding_free is False
+
+
+def test_decoder_only_conditional_generation_keeps_packing():
+ # CSM is decoder-only despite the ForConditionalGeneration name -> packing stays on.
+ fake_trainer = _patch_fake_sft_trainer()
+ config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
+ model = SimpleNamespace(
+ config = SimpleNamespace(
+ model_type = "csm",
+ architectures = ["CsmForConditionalGeneration"],
+ is_encoder_decoder = False,
+ ),
+ max_seq_length = 16,
)
+ trainer = fake_trainer(model, config, None, Dataset.from_dict({"text": ["text-only sample"]}))
+
assert config.packing is True
assert config.padding_free is True
assert trainer.model._unsloth_allow_packed_overlength is True
+def _hybrid_trainer_model():
+ return SimpleNamespace(
+ config = SimpleNamespace(
+ model_type = "qwen3_next",
+ architectures = ["Qwen3NextForCausalLM"],
+ layer_types = ["linear_attention", "full_attention"],
+ ),
+ max_seq_length = 16,
+ )
+
+
+def test_hybrid_varlen_active_enables_packing(monkeypatch):
+ # Baseline: shim active + no forward bypass -> hybrid packing is allowed.
+ monkeypatch.setattr(trainer_module, "_chunked_loss_bypasses_forward", lambda config: False)
+ monkeypatch.setattr(trainer_module, "patch_hybrid_linear_attention_varlen", lambda model: True)
+ fake_trainer = _patch_fake_sft_trainer()
+ config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
+ fake_trainer(_hybrid_trainer_model(), config, None, Dataset.from_dict({"text": ["x"]}))
+ assert config.packing is True
+ assert config.padding_free is True
+
+
+def test_hybrid_chunked_loss_stays_on_padded_path(monkeypatch):
+ # TRL's chunked-loss forward bypass leaves the varlen shim off -> block packing.
+ monkeypatch.setattr(trainer_module, "_chunked_loss_bypasses_forward", lambda config: True)
+ monkeypatch.setattr(trainer_module, "patch_hybrid_linear_attention_varlen", lambda model: True)
+ fake_trainer = _patch_fake_sft_trainer()
+ config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
+ fake_trainer(_hybrid_trainer_model(), config, None, Dataset.from_dict({"text": ["x"]}))
+ assert config.packing is False
+ assert config.padding_free is False
+
+
+def test_string_hybrid_model_disables_packing(monkeypatch):
+ # A string model= is materialized after init; a hybrid string is blocked because the
+ # shim cannot patch a not-yet-built model.
+ monkeypatch.setattr(
+ trainer_module,
+ "_resolve_string_model_config",
+ lambda name, cfg: SimpleNamespace(
+ model_type = "qwen3_next",
+ architectures = ["Qwen3NextForCausalLM"],
+ layer_types = ["linear_attention", "full_attention"],
+ ),
+ )
+ monkeypatch.setattr(trainer_module, "patch_hybrid_linear_attention_varlen", lambda model: True)
+ fake_trainer = _patch_fake_sft_trainer()
+ config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
+ fake_trainer("Qwen/Qwen3-Next-80B-A3B", config, None, Dataset.from_dict({"text": ["x"]}))
+ assert config.packing is False
+ assert config.padding_free is False
+
+
def test_vlm_vision_dataset_still_disables_packing():
fake_trainer = _patch_fake_sft_trainer()
config = SimpleNamespace(packing = True, padding_free = None, remove_unused_columns = True)
@@ -486,49 +881,6 @@ def test_wrapped_packing_preserves_overlength_tokens(monkeypatch, legacy_api):
assert all(len(input_ids) <= args.max_length for input_ids in packed_ids)
-# Named to match the unsloth_zoo helper: sft_trainer_prepare_dataset sources it by
-# name and renames "def sft_prepare_dataset" -> "def _prepare_dataset". This fixture
-# deliberately omits the "All Unsloth Zoo code licensed under LGPLv3" header to emulate
-# a newer, compatible Zoo whose header moved (the dependency is only lower-bounded).
-def sft_prepare_dataset(
- self, dataset, processing_class, args, packing, formatting_func, dataset_text_field
-):
- do_truncation = True
- # Mirror the Zoo call so the "truncation = do_truncation," injection anchor
- # survives formatting (a bare tuple assignment gets rewritten to a paren form).
- dataset = processing_class(
- dataset,
- truncation = do_truncation,
- )
- return dataset
-
-
-def test_wrapped_packing_setup_survives_missing_zoo_header(monkeypatch):
- # Regression: the wrapped-packing setup used to anchor on the Zoo license comment,
- # so a header change made it a no-op while the truncation reference still landed,
- # NameError-ing every SFT dataset preparation. It must now install via the
- # signature and always precede the reference.
- import ast
- import textwrap
- import unsloth.models.rl_replacements as rlr
-
- monkeypatch.setitem(rlr.RL_REPLACEMENTS, "sft_prepare_dataset", sft_prepare_dataset)
-
- source = (
- "def _prepare_dataset(self, dataset, processing_class, args, packing, "
- "formatting_func, dataset_text_field):\n return dataset\n"
- )
- patched = rlr.sft_trainer_prepare_dataset("_prepare_dataset", source)
-
- assert "_unsloth_wrapped_packing = packing" in patched
- assert "import inspect as _inspect" in patched
- assert "not _unsloth_wrapped_packing" in patched
- assert patched.index("_unsloth_wrapped_packing = packing") < patched.index(
- "truncation = do_truncation and not _unsloth_wrapped_packing"
- )
- ast.parse(textwrap.dedent(patched))
-
-
class _DummyChild(torch.nn.Module):
def __init__(self):
super().__init__()
@@ -759,3 +1111,128 @@ def test_packing_sdpa(tmp_path):
if hasattr(trainer, "accelerator"):
trainer.accelerator.free_memory()
+
+
+# --- wrapped-packing source-injection robustness (reviewer.py / fork findings) --------
+
+
+# fmt: off
+# Named to match the unsloth_zoo helper (sourced by name, "def sft_prepare_dataset" ->
+# "def _prepare_dataset"). Deliberately OMITS the "licensed under LGPLv3" header to
+# emulate a newer Zoo whose header moved (dependency is only lower-bounded). Source only.
+def sft_prepare_dataset(
+ self, dataset, processing_class, args, packing, formatting_func, dataset_text_field
+):
+ do_truncation = True
+ max_seq_length = 4
+ used_column_names = ["text"]
+ map_kwargs = {}
+ dataset = processing_class(dataset, truncation = do_truncation,)
+ if do_truncation and max_seq_length > 0:
+ pass
+ if packing:
+ dataset = pack_dataset(
+ dataset.select_columns(used_column_names),
+ max_seq_length,
+ getattr(args, "packing_strategy", "bfd"),
+ map_kwargs,
+ )
+ return dataset
+# fmt: on
+
+
+def test_wrapped_packing_injection_is_drift_resistant(monkeypatch):
+ # Regression: the setup used to anchor on the Zoo license comment, so a header
+ # change silently no-op'd it while the truncation/pack edits still referenced its
+ # variables -> NameError on every SFT prep. It must now install via the signature
+ # before those references, and the pack edit must reuse the guarded
+ # _unsloth_pack_has_strategy instead of re-calling _inspect.signature(pack_dataset).
+ import ast
+ import textwrap
+ import unsloth.models.rl_replacements as rlr
+
+ monkeypatch.setitem(rlr.RL_REPLACEMENTS, "sft_prepare_dataset", sft_prepare_dataset)
+
+ source = (
+ "def _prepare_dataset(self, dataset, processing_class, args, packing, "
+ "formatting_func, dataset_text_field):\n return dataset\n"
+ )
+ patched = rlr.sft_trainer_prepare_dataset("_prepare_dataset", source)
+
+ # setup installed despite the missing header, and before it is referenced
+ assert "_unsloth_wrapped_packing = packing" in patched
+ assert "import inspect as _inspect" in patched
+ assert patched.index("_unsloth_wrapped_packing = packing") < patched.index(
+ "truncation = do_truncation and not _unsloth_wrapped_packing"
+ )
+ # the pack edit reuses the guarded flag (signature inspected exactly once, in setup)
+ assert "if _unsloth_pack_has_strategy:" in patched
+ assert patched.count("_inspect.signature(pack_dataset)") == 1
+ ast.parse(textwrap.dedent(patched))
+
+
+def test_require_replace_raises_on_missing_anchor():
+ from unsloth.models.rl_replacements import _require_replace
+
+ assert _require_replace("abc", "b", "B") == "aBc"
+ with pytest.raises(RuntimeError):
+ _require_replace("abc", "z", "Z", where = "unit test")
+ # an optional edit warns once and returns the source unchanged (no dangling ref)
+ assert _require_replace("abc", "z", "Z", required = False, where = "optional") == "abc"
+
+
+def test_resolve_string_model_config_forwards_token(monkeypatch):
+ import transformers
+
+ captured = {}
+
+ class _FakeAutoConfig:
+ @staticmethod
+ def from_pretrained(name, **kwargs):
+ captured.update(kwargs)
+ return SimpleNamespace(is_encoder_decoder = False)
+
+ monkeypatch.setattr(transformers, "AutoConfig", _FakeAutoConfig)
+
+ config_arg = SimpleNamespace(
+ model_init_kwargs = {
+ "token": "hf_secret",
+ "trust_remote_code": True,
+ "cache_dir": "/tmp/cache",
+ "torch_dtype": "bfloat16", # not a config arg -> must NOT be forwarded
+ }
+ )
+ result = trainer_module._resolve_string_model_config("org/private-hybrid", config_arg)
+
+ assert result is not None
+ assert captured.get("token") == "hf_secret"
+ assert captured.get("trust_remote_code") is True
+ assert captured.get("cache_dir") == "/tmp/cache"
+ assert "torch_dtype" not in captured
+
+
+def test_resolve_string_model_config_merges_top_level_trust_remote_code(monkeypatch):
+ import transformers
+
+ captured = {}
+
+ class _FakeAutoConfig:
+ @staticmethod
+ def from_pretrained(name, **kwargs):
+ captured.update(kwargs)
+ return SimpleNamespace(is_encoder_decoder = False)
+
+ monkeypatch.setattr(transformers, "AutoConfig", _FakeAutoConfig)
+
+ # SFTConfig(trust_remote_code=True) with no model_init_kwargs entry is honored
+ config_arg = SimpleNamespace(model_init_kwargs = {}, trust_remote_code = True)
+ trainer_module._resolve_string_model_config("org/remote-hybrid", config_arg)
+ assert captured.get("trust_remote_code") is True
+
+ # model_init_kwargs wins over the top-level flag (mirrors TRL's setdefault)
+ captured.clear()
+ config_arg = SimpleNamespace(
+ model_init_kwargs = {"trust_remote_code": False}, trust_remote_code = True
+ )
+ trainer_module._resolve_string_model_config("org/remote-hybrid", config_arg)
+ assert captured.get("trust_remote_code") is False
diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py
index b0709f7376..4ef3af6add 100644
--- a/unsloth/models/rl_replacements.py
+++ b/unsloth/models/rl_replacements.py
@@ -437,6 +437,52 @@ RL_FUNCTIONS["dpo_trainer"].append(dpo_trainer_compute_loss_liger)
RL_EXTRA_ARGS["dpo_trainer"].append(dpo_trainer_data_collator_vision_keys)
+_WRAPPED_PACKING_SETUP = (
+ " import inspect as _inspect\n"
+ " try:\n"
+ ' _unsloth_pack_has_strategy = "strategy" in _inspect.signature(pack_dataset).parameters\n'
+ " except Exception:\n"
+ " _unsloth_pack_has_strategy = True\n"
+ " _unsloth_wrapped_packing = packing and (\n"
+ ' getattr(args, "packing_strategy", None) == "wrapped"\n'
+ " or not _unsloth_pack_has_strategy\n"
+ " )\n"
+)
+
+_WARNED_MISSING_ANCHORS = set()
+
+
+def _require_replace(
+ function,
+ old,
+ new,
+ *,
+ count = 1,
+ required = True,
+ where = "",
+):
+ """str.replace that never silently no-ops a load-bearing source edit.
+
+ Plain str.replace returns the source unchanged when the anchor is absent, so a
+ drifted anchor in a newer TRL / unsloth_zoo would skip the edit while later edits
+ still reference helper variables it should have introduced (NameError at runtime).
+ Fail loudly for a required edit, warn once and skip for an optional one, so a
+ drifted source can never corrupt the patched function silently.
+ """
+ if old not in function:
+ detail = f" ({where})" if where else ""
+ if required:
+ raise RuntimeError(
+ f"Unsloth: source anchor not found{detail}; the patched function is out "
+ "of sync with this TRL / unsloth_zoo version. Please file a bug report."
+ )
+ if where not in _WARNED_MISSING_ANCHORS:
+ _WARNED_MISSING_ANCHORS.add(where)
+ logger.warning(f"Unsloth: skipped an optional source edit{detail} (anchor not found).")
+ return function
+ return function.replace(old, new, count)
+
+
# Fix tokenizer double BOS
def sft_trainer_prepare_dataset(function_name, function):
if function_name != "_prepare_non_packed_dataloader" and function_name != "_prepare_dataset":
@@ -454,27 +500,14 @@ def sft_trainer_prepare_dataset(function_name, function):
if matched:
# Use fast version!
function = inspect.getsource(fast_sft_prepare_dataset)
- # why: install the wrapped-packing setup (and the `_inspect` import the
- # truncation / pack_dataset rewrites below depend on) at the function
- # signature, a structural anchor that always exists, rather than the
- # unsloth_zoo license-comment line. That header is only lower-bounded, so a
- # newer Zoo may move or drop it; anchoring there let the setup silently
- # no-op while the references still landed, NameError-ing every SFT dataset
- # preparation. Fail loudly if even the signature cannot be located.
- _wrapped_packing_setup = (
- " import inspect as _inspect\n"
- " try:\n"
- ' _unsloth_pack_has_strategy = "strategy" in _inspect.signature(pack_dataset).parameters\n'
- " except Exception:\n"
- " _unsloth_pack_has_strategy = True\n"
- " _unsloth_wrapped_packing = packing and (\n"
- ' getattr(args, "packing_strategy", None) == "wrapped"\n'
- " or not _unsloth_pack_has_strategy\n"
- " )\n"
- )
+ # why: anchor the wrapped-packing setup on the function signature -- a
+ # structural anchor that always exists -- not the unsloth_zoo license comment,
+ # which is only lower-bounded and a newer Zoo may move or drop. Anchoring there
+ # let the setup silently no-op while edits below referenced its variables,
+ # NameError-ing every SFT dataset prep. Fail loudly if the signature is missing.
function, _n_setup = re.subn(
r"(def sft_prepare_dataset\s*\(.*?\)\s*(?:->[^:\n]*)?:[ \t]*\n)",
- lambda match: match.group(1) + _wrapped_packing_setup,
+ lambda match: match.group(1) + _WRAPPED_PACKING_SETUP,
function,
count = 1,
flags = re.DOTALL,
@@ -484,15 +517,25 @@ def sft_trainer_prepare_dataset(function_name, function):
"Unsloth: failed to install wrapped-packing support into "
"sft_prepare_dataset (signature not found); please file a bug report."
)
- function = function.replace(
+ # why: route each edit through _require_replace so a drifted anchor fails
+ # loudly instead of leaving a dangling reference to the setup variables.
+ function = _require_replace(
+ function,
"truncation = do_truncation,",
"truncation = do_truncation and not _unsloth_wrapped_packing,",
+ where = "sft_prepare_dataset truncation flag",
)
- function = function.replace(
+ function = _require_replace(
+ function,
"if do_truncation and max_seq_length > 0:",
"if do_truncation and not _unsloth_wrapped_packing and max_seq_length > 0:",
+ where = "sft_prepare_dataset truncation guard",
)
- function = function.replace(
+ # why: reuse the guarded _unsloth_pack_has_strategy from the setup instead of
+ # re-calling _inspect.signature(pack_dataset) here -- the setup wraps that call
+ # in try/except, so a non-introspectable pack_dataset must not crash here.
+ function = _require_replace(
+ function,
"""dataset = pack_dataset(
dataset.select_columns(used_column_names),
max_seq_length,
@@ -500,13 +543,14 @@ def sft_trainer_prepare_dataset(function_name, function):
map_kwargs,
)""",
"""_pack_kwargs = {"map_kwargs": map_kwargs}
- if "strategy" in _inspect.signature(pack_dataset).parameters:
+ if _unsloth_pack_has_strategy:
_pack_kwargs["strategy"] = getattr(args, "packing_strategy", "bfd")
dataset = pack_dataset(
dataset.select_columns(used_column_names),
max_seq_length,
**_pack_kwargs,
)""",
+ where = "sft_prepare_dataset pack_dataset call",
)
function = function.split("\n")
function = "\n".join(" " * 4 + x for x in function)
diff --git a/unsloth/trainer.py b/unsloth/trainer.py
index 61d41aad21..1c30192301 100644
--- a/unsloth/trainer.py
+++ b/unsloth/trainer.py
@@ -17,6 +17,7 @@ import os
import psutil
import warnings
from dataclasses import dataclass, field
+from types import SimpleNamespace
from typing import Optional, List
from functools import wraps
@@ -32,6 +33,7 @@ from unsloth.utils import (
enable_padding_free_metadata,
enable_sample_packing,
)
+from unsloth.utils.packing import patch_hybrid_linear_attention_varlen
from unsloth_zoo.training_utils import (
unsloth_train as _unsloth_train,
)
@@ -101,9 +103,9 @@ PADDING_FREE_BLOCKLIST = {
"gpt_oss", # - gpt_oss: Uses Flex Attention which doesn't handle padding_free correctly
}
# Hybrid linear-attention / state-space models (Qwen3.5, Qwen3-Next, ...) carry a
-# recurrent gated-delta state plus a causal conv1d. Sample packing / padding-free
-# flattens the batch, so those ops leak state across sequence boundaries. Detected
-# structurally by _is_hybrid_linear_attention_model rather than by model name.
+# recurrent gated-delta state plus a causal conv1d that leak across sequence
+# boundaries once packing flattens the batch. Detected structurally by
+# _is_hybrid_linear_attention_model, not by model name.
def _should_pack(config) -> bool:
@@ -267,6 +269,57 @@ def _is_hybrid_linear_attention_model(model) -> bool:
return False
+def _resolve_string_model_config(model_name, config_arg):
+ """TRL materializes a string ``model=`` inside ``__init__``; resolve its config
+ up front so the packing guards run before the dataset is packed. Best-effort:
+ returns None if the config cannot be loaded."""
+ try:
+ from transformers import AutoConfig
+
+ init_kwargs = getattr(config_arg, "model_init_kwargs", None) or {}
+ # why: forward auth + cache args too. Dropping token/use_auth_token made a
+ # private hybrid fail to load (resolve as None) -> treated as non-hybrid ->
+ # packing enabled without the shim even though TRL later loads it with the token.
+ forward = {
+ key: init_kwargs[key]
+ for key in (
+ "trust_remote_code",
+ "revision",
+ "subfolder",
+ "token",
+ "use_auth_token",
+ "cache_dir",
+ "code_revision",
+ )
+ if key in init_kwargs
+ }
+ # why: TRL merges top-level args.trust_remote_code into the load via setdefault
+ # before create_model_from_path, so honor it here (model_init_kwargs wins), else
+ # a remote-code hybrid with SFTConfig(trust_remote_code=True) resolves as None
+ # and skips the guard.
+ top_level_trust_remote_code = getattr(config_arg, "trust_remote_code", None)
+ if top_level_trust_remote_code is not None:
+ forward.setdefault("trust_remote_code", top_level_trust_remote_code)
+ return AutoConfig.from_pretrained(model_name, **forward)
+ except Exception:
+ return None
+
+
+def _chunked_loss_bypasses_forward(config) -> bool:
+ """TRL's default ``loss_type="chunked_nll"`` patches the model forward and calls
+ the backbone directly, so a forward wrapper never runs. Detect it so hybrid
+ packing stays on the padded path instead of silently skipping the varlen shim."""
+ try:
+ import trl.trainer.sft_trainer as _sft_trainer
+ except Exception:
+ return False
+ if not hasattr(_sft_trainer, "_patch_chunked_ce_lm_head"):
+ return False # TRL has no chunked-CE path -> forward is not bypassed
+ if getattr(config, "use_liger_kernel", False):
+ return False # liger forces loss_type="nll" -> normal forward
+ return getattr(config, "loss_type", None) in (None, "chunked_nll")
+
+
# Unsloth gradient accumulation fix:
from transformers import __version__ as transformers_version, ProcessorMixin
@@ -632,13 +685,38 @@ def _patch_sft_trainer_auto_packing(trl_module):
is_vlm = False
is_unsupported_model = False
is_hybrid = False
+ is_encoder_decoder = False
+ hybrid_varlen_active = False
if model is not None:
model_config = getattr(model, "config", None)
+ if model_config is None and isinstance(model, str):
+ # TRL builds a string model inside __init__; resolve its config now.
+ model_config = _resolve_string_model_config(model, config_arg)
if model_config is not None:
model_types = get_transformers_model_type(model_config)
is_unsupported_model = any(x in PADDING_FREE_BLOCKLIST for x in model_types)
is_vlm = _is_vlm_config(model_config, model_types)
- is_hybrid = _is_hybrid_linear_attention_model(model)
+ is_encoder_decoder = bool(getattr(model_config, "is_encoder_decoder", False))
+ hybrid_target = (
+ SimpleNamespace(config = model_config)
+ if isinstance(model, str) and model_config is not None
+ else model
+ )
+ is_hybrid = _is_hybrid_linear_attention_model(hybrid_target)
+ # Hybrid models corrupt packed batches unless the gated-delta conv + scan
+ # reset at sequence boundaries. Enable the experimental varlen shim (flag +
+ # kernels) so packing stays correct, else keep them blocked. A string model
+ # (patched only after init) and TRL's chunked-loss forward bypass both leave
+ # the shim off, so hybrid packing falls back to the padded path.
+ if (
+ is_hybrid
+ and not isinstance(model, str)
+ and not _chunked_loss_bypasses_forward(config_arg)
+ ):
+ try:
+ hybrid_varlen_active = patch_hybrid_linear_attention_varlen(model)
+ except Exception:
+ hybrid_varlen_active = False
processing_class = (
args[5] if len(args) >= 6 else kwargs.get("processing_class") or kwargs.get("tokenizer")
@@ -664,7 +742,8 @@ def _patch_sft_trainer_auto_packing(trl_module):
or is_auto_processor_vlm
or is_vision_dataset
or is_unsupported_model
- or is_hybrid
+ or is_encoder_decoder
+ or (is_hybrid and not hybrid_varlen_active)
or (
os.environ.get("UNSLOTH_RETURN_LOGITS", "0") == "1"
) # Disable padding free on forced logits
@@ -684,7 +763,9 @@ def _patch_sft_trainer_auto_packing(trl_module):
reason = "vision-language model with auto processor"
elif is_vision_dataset:
reason = "vision dataset"
- elif is_hybrid:
+ elif is_encoder_decoder:
+ reason = "encoder-decoder model"
+ elif is_hybrid and not hybrid_varlen_active:
reason = "hybrid linear-attention model"
elif is_unsupported_model:
reason = f"unsupported model type(s): {', '.join(model_types)}"
diff --git a/unsloth/utils/packing.py b/unsloth/utils/packing.py
index dd0a1bfb62..f8d539fb93 100644
--- a/unsloth/utils/packing.py
+++ b/unsloth/utils/packing.py
@@ -17,8 +17,11 @@
from __future__ import annotations
+import inspect
import logging
+import os
from collections import OrderedDict
+from functools import wraps
from typing import Any, Iterable, Optional, Sequence, Tuple
import torch
@@ -218,6 +221,305 @@ def enable_padding_free_metadata(model, trainer):
collator._unsloth_padding_free_lengths_wrapped = True
+# --- Experimental: correct packing / padding-free for hybrid linear-attention ---
+# Qwen3.5 / Qwen3-Next mix a gated-delta recurrence with a causal conv1d. Packing
+# flattens the batch, and both ops leak state across sequence boundaries unless we
+# pass seq_idx (conv) and cu_seqlens (scan). Only the accelerated kernels accept
+# these, so we fail closed on the pure-torch fallbacks. Gated behind an env flag.
+#
+# Overrides only the per-module prefill kernels (causal_conv1d_fn /
+# chunk_gated_delta_rule), leaving decode untouched so generation is unaffected.
+# Recompute-safe under gradient checkpointing; never fires for cached forwards.
+# Feature-detect (never version-detect), fail closed, idempotent, one deduped
+# diagnostic when it declines to activate.
+_HYBRID_PACKING_ENV_VAR = "UNSLOTH_EXPERIMENTAL_HYBRID_PACKING"
+_HYBRID_LOGGER = logging.getLogger("unsloth.hybrid_packing")
+_HYBRID_WARNED: set = set()
+
+
+def _hybrid_packing_enabled() -> bool:
+ # Read at call time so setting the flag after `import unsloth` still takes effect.
+ return os.environ.get(_HYBRID_PACKING_ENV_VAR, "0").strip().lower() in {
+ "1",
+ "true",
+ "yes",
+ "on",
+ }
+
+
+def _hybrid_reject(reason: str) -> bool:
+ # One deduped diagnostic explaining why hybrid packing stayed on the padded path.
+ if reason not in _HYBRID_WARNED:
+ _HYBRID_WARNED.add(reason)
+ _HYBRID_LOGGER.warning(
+ "Unsloth: hybrid linear-attention packing disabled (padded path): %s.",
+ reason,
+ )
+ return False
+
+
+def _iter_gated_delta_modules(model):
+ modules, seen = [], set()
+ for module in model.modules():
+ if id(module) in seen:
+ continue
+ seen.add(id(module))
+ if type(module).__name__.endswith("GatedDeltaNet") and hasattr(module, "conv1d"):
+ modules.append(module)
+ return modules
+
+
+def _hybrid_varlen_kernels_available(gated_delta_modules) -> Optional[str]:
+ """None if every module can use the accelerated varlen path, else a short
+ reason string. All modules are validated before any are mutated; signatures
+ are read off the captured originals when already wrapped.
+
+ Dispatch (the mixer actually calling self.causal_conv1d_fn /
+ self.chunk_gated_delta_rule) is verified at RUNTIME by the forward-wrapper
+ handshake, not statically: Unsloth's compile-disable shim hides it from
+ inspect.getsource, and every supported transformers release dispatches
+ through the instance attribute."""
+ if not gated_delta_modules:
+ return "no gated-delta modules found"
+ for module in gated_delta_modules:
+ conv = getattr(module, "_unsloth_varlen_orig_conv", None) or getattr(
+ module,
+ "causal_conv1d_fn",
+ None,
+ )
+ scan = getattr(module, "_unsloth_varlen_orig_scan", None) or getattr(
+ module,
+ "chunk_gated_delta_rule",
+ None,
+ )
+ if conv is None or scan is None:
+ return "accelerated kernels missing (install causal_conv1d and fla)"
+ if getattr(scan, "__name__", "").startswith("torch_") or getattr(
+ conv,
+ "__name__",
+ "",
+ ).startswith("torch_"):
+ return "pure-torch kernel fallback in use"
+ try:
+ if "seq_idx" not in inspect.signature(conv).parameters:
+ return "conv kernel does not accept seq_idx"
+ if "cu_seqlens" not in inspect.signature(scan).parameters:
+ return "scan kernel does not accept cu_seqlens"
+ except (TypeError, ValueError):
+ return "kernel signature not introspectable"
+ return None
+
+
+def _varlen_from_position_ids(position_ids):
+ """(cu_seqlens int32[n+1], seq_idx int32[1,T]) for a flattened padding-free
+ batch, else None. Padding-free position_ids reset to 0 at each sequence start;
+ accepts only a validated single-row pack (normal batch or single sequence ->
+ None). Fallback used only when packed_seq_lengths is absent: it assumes
+ right-packed reset position_ids and would mis-segment a left-padded row, which
+ is why packed_seq_lengths is always preferred."""
+ if position_ids is None:
+ return None
+ pos = position_ids
+ if pos.dim() == 3: # MRoPE [n_planes, 1, T] -> text plane is index 0
+ pos = pos[0]
+ if pos.dim() != 2 or pos.shape[0] != 1:
+ return None
+ row = pos[0]
+ total = row.shape[0]
+ starts = (row == 0).nonzero(as_tuple = False).flatten()
+ if starts.numel() <= 1 or int(starts[0].item()) != 0:
+ return None
+ cu_seqlens = torch.cat(
+ [
+ starts.to(torch.int32),
+ torch.tensor([total], dtype = torch.int32, device = row.device),
+ ]
+ )
+ return _seq_idx_from_cu_seqlens(cu_seqlens, total)
+
+
+def _seq_idx_from_cu_seqlens(cu_seqlens, total):
+ """(cu_seqlens int32[n+1], seq_idx int32[1,total]) partitioning [0, total),
+ else None. Appends a trailing segment for pad_to_multiple_of zero tokens so the
+ boundaries always cover the full flattened length the kernels see."""
+ if cu_seqlens is None or cu_seqlens.numel() < 2 or int(cu_seqlens[0].item()) != 0:
+ return None
+ boundaries = cu_seqlens.to(torch.int32)
+ last = int(boundaries[-1].item())
+ if last > total:
+ return None
+ if last < total: # trailing pad tokens -> one final segment
+ boundaries = torch.cat(
+ [
+ boundaries,
+ torch.tensor([total], dtype = torch.int32, device = boundaries.device),
+ ]
+ )
+ lengths = boundaries[1:] - boundaries[:-1]
+ if not bool((lengths > 0).all()):
+ return None
+ seq_idx = torch.repeat_interleave(
+ torch.arange(lengths.numel(), dtype = torch.int32, device = boundaries.device),
+ lengths.to(torch.int64),
+ ).unsqueeze(0)
+ return boundaries, seq_idx
+
+
+def _hybrid_varlen_metadata(kwargs):
+ """Boundary metadata (cu_seqlens, seq_idx) for one flattened packed forward,
+ else None. Prefers the authoritative packed_seq_lengths, falls back to
+ reset-style position_ids. Returns None for cached forwards and non-packed
+ batches so decode / eval / normal batches are a strict no-op."""
+ if kwargs.get("use_cache"):
+ return None
+ if kwargs.get("past_key_values") is not None or kwargs.get("cache_params") is not None:
+ return None
+ total, device = None, None
+ for key in ("input_ids", "inputs_embeds", "position_ids"):
+ tensor = kwargs.get(key)
+ if tensor is not None and hasattr(tensor, "shape"):
+ total = tensor.shape[1] if key == "inputs_embeds" else tensor.shape[-1]
+ device = tensor.device
+ break
+ if total is None:
+ return None
+ psl = kwargs.get("packed_seq_lengths")
+ if psl is not None and getattr(psl, "numel", lambda: 1)() > 0: # skip empty (no max())
+ info = get_packed_info_from_kwargs(kwargs, device)
+ if info is not None:
+ _, cu_seqlens, _ = info
+ built = _seq_idx_from_cu_seqlens(cu_seqlens, total)
+ if built is not None:
+ return built
+ return _varlen_from_position_ids(kwargs.get("position_ids"))
+
+
+def patch_hybrid_linear_attention_varlen(model) -> bool:
+ """Feed seq_idx / cu_seqlens to the gated-delta conv + scan so packing and
+ padding-free reset state at sequence boundaries. Gated by
+ UNSLOTH_EXPERIMENTAL_HYBRID_PACKING and fail-closed. Returns True when the
+ varlen path is active, so the caller may allow packing for the model.
+ Idempotent: repeat calls on an already-patched model return True."""
+ if not _hybrid_packing_enabled():
+ return False
+ gated_delta_modules = _iter_gated_delta_modules(model)
+
+ # Idempotency: an already fully-patched model stays active without re-validation.
+ if (
+ getattr(model, "_unsloth_varlen_forward_wrapped", False)
+ and gated_delta_modules
+ and all(getattr(m, "_unsloth_varlen_wrapped", False) for m in gated_delta_modules)
+ ):
+ return True
+
+ reason = _hybrid_varlen_kernels_available(gated_delta_modules)
+ if reason is not None:
+ return _hybrid_reject(reason)
+
+ # Transactional: every module validated above, now wrap each and stash originals.
+ for module in gated_delta_modules:
+ if getattr(module, "_unsloth_varlen_wrapped", False):
+ continue
+ conv_orig, scan_orig = module.causal_conv1d_fn, module.chunk_gated_delta_rule
+ module._unsloth_varlen_orig_conv = conv_orig
+ module._unsloth_varlen_orig_scan = scan_orig
+
+ @wraps(conv_orig)
+ def conv_fn(
+ *args,
+ _orig = conv_orig,
+ _module = module,
+ **kwargs,
+ ):
+ varlen = getattr(_module, "_unsloth_varlen", None)
+ if varlen is not None:
+ _module._unsloth_varlen_conv_hit = True # runtime dispatch handshake
+ if kwargs.get("seq_idx") is None:
+ kwargs["seq_idx"] = varlen[1]
+ return _orig(*args, **kwargs)
+
+ @wraps(scan_orig)
+ def scan_fn(
+ *args,
+ _orig = scan_orig,
+ _module = module,
+ **kwargs,
+ ):
+ varlen = getattr(_module, "_unsloth_varlen", None)
+ if varlen is not None:
+ _module._unsloth_varlen_scan_hit = True
+ if kwargs.get("cu_seqlens") is None:
+ kwargs["cu_seqlens"] = varlen[0]
+ return _orig(*args, **kwargs)
+
+ module.causal_conv1d_fn = conv_fn
+ module.chunk_gated_delta_rule = scan_fn
+ module._unsloth_varlen = None
+ module._unsloth_varlen_wrapped = True
+
+ # Refresh the boundary stash on the outermost forward (once per step, outside
+ # gradient-checkpoint recompute, so it stays valid for recomputed inner
+ # forwards). Read from both positional and keyword args via the bound signature.
+ if not getattr(model, "_unsloth_varlen_forward_wrapped", False):
+ forward_orig = model.forward
+ try:
+ forward_sig = inspect.signature(forward_orig)
+ except (TypeError, ValueError):
+ forward_sig = None
+
+ @wraps(forward_orig)
+ def forward_with_varlen(*args, **kwargs):
+ try:
+ bound = dict(kwargs)
+ if forward_sig is not None and args:
+ bound.update(forward_sig.bind_partial(*args).arguments)
+ varlen = _hybrid_varlen_metadata(bound)
+ except Exception:
+ varlen = None
+ first_pack = varlen is not None and not getattr(
+ model,
+ "_unsloth_varlen_handshake_done",
+ False,
+ )
+ for module in gated_delta_modules:
+ module._unsloth_varlen = varlen
+ if first_pack:
+ module._unsloth_varlen_conv_hit = False
+ module._unsloth_varlen_scan_hit = False
+ out = forward_orig(*args, **kwargs)
+ # Runtime dispatch handshake: on the first packed forward, confirm BOTH
+ # boundary kernels ran for EVERY module. seq_idx (conv) and cu_seqlens
+ # (scan) are both load-bearing, so a partial/absent dispatch (a future
+ # version no longer routing through self.) leaves cross-sequence
+ # contamination. The batch is already flattened with no padded recovery,
+ # so abort before loss/backward rather than train on corrupted data.
+ if first_pack:
+ model._unsloth_varlen_handshake_done = True
+ missing = [
+ type(m).__name__
+ for m in gated_delta_modules
+ if not (
+ getattr(m, "_unsloth_varlen_conv_hit", False)
+ and getattr(m, "_unsloth_varlen_scan_hit", False)
+ )
+ ]
+ if missing:
+ for m in gated_delta_modules:
+ m._unsloth_varlen = None
+ _hybrid_reject("varlen conv/scan not both dispatched (dispatch changed?)")
+ raise RuntimeError(
+ "Unsloth: experimental hybrid packing cannot continue because the "
+ "varlen conv/scan wrappers were not both invoked for "
+ f"{sorted(set(missing))}. Unset UNSLOTH_EXPERIMENTAL_HYBRID_PACKING "
+ "to train these models on the padded path."
+ )
+ return out
+
+ model.forward = forward_with_varlen
+ model._unsloth_varlen_forward_wrapped = True
+ return True
+
+
def get_packed_info_from_kwargs(
kwargs: dict, device: torch.device
) -> Optional[Tuple[torch.Tensor, torch.Tensor, int]]:
From 3ab8dce97a95923b2b4e6741e9df9cba1a9baaca Mon Sep 17 00:00:00 2001
From: Daniel Han
Date: Mon, 20 Jul 2026 00:58:52 -0700
Subject: [PATCH 14/41] install: let UNSLOTH_TORCH_INDEX_FAMILY / _URL override
CUDA wheel detection (#6692)
* install: let UNSLOTH_TORCH_INDEX_FAMILY / _URL override CUDA wheel detection
get_torch_index_url (and the studio-update mirror _detect_cuda_torch_index_url)
chose the torch wheel family solely by probing the host GPU, with no override.
In a headless / container / CI build the host driver is visible via the
/proc/driver/nvidia/gpus fallback but nvidia-smi cannot report a CUDA version,
so the function fell back to its cu126 default and installed the wrong wheels
(e.g. a cu128 image got cu126 torch).
Add an explicit override checked before any probing, in both the shell installer
and the Python studio-update path:
- UNSLOTH_TORCH_INDEX_URL full index URL, used verbatim (wins)
- UNSLOTH_TORCH_INDEX_FAMILY family (cpu, cu128, rocm6.4, ...) appended to the
mirror base (UNSLOTH_PYTORCH_MIRROR still honoured)
This matches how the published GPU images select CUDA -- vLLM and SGLang take the
CUDA version from an explicit build ARG rather than detecting it, and the Unsloth
Docker base image already pins the cu128 index directly. Desktop installs are
unchanged: with no override set, detection runs exactly as before.
Adds test_get_torch_index_url.sh cases for the override (family, full URL,
precedence, mirror base, trailing-slash strip, empty-ignored).
* install: make the torch-index override authoritative across ROCm paths
Address review feedback on the override added in this PR so a pinned index is
honoured everywhere, not just in get_torch_index_url:
- Skip the WSL ROCm bootstrap (root privilege + large downloads, probes
/dev/dxg) when UNSLOTH_TORCH_INDEX_URL / _FAMILY is set; it previously ran
before the override was consulted.
- Skip the Radeon/Strix rerouting (which re-probes the GPU and overwrites the
resolved URL with repo.radeon.com / repo.amd.com) when the index is pinned, so
an explicit ROCm override (e.g. UNSLOTH_TORCH_INDEX_FAMILY=rocm6.4) is kept.
- install_python_stack.py: derive _TORCH_BACKEND from the override when
UNSLOTH_TORCH_BACKEND is unset (standalone studio update), so _ensure_rocm_torch
/ _ensure_cuda_torch repair to the requested family instead of re-detecting.
- Strip ALL leading/trailing slashes in the shell override to match the Python
side (avoids 404s on strict pip proxies).
Adds test cases for double-slash and leading/trailing-slash overrides.
* install: honor pinned torch index in CUDA/ROCm repair paths
Follow-up to the override work in this PR: the get_torch_index_url / install.sh
reroute already respect a pinned UNSLOTH_TORCH_INDEX_URL / _FAMILY, but the
Python repair helpers in install_python_stack.py still re-probed the GPU and
could overwrite the pinned family. Make the pin authoritative there too:
- _ensure_cuda_torch: an explicit cu* pin commits to CUDA wheels, so repair a
ROCm-poisoned venv even when no NVIDIA GPU is visible here (headless /
container / CI cross-install), instead of bailing on the GPU-presence gate.
- _ensure_rocm_torch: skip the AMD per-gfx (Strix) reroute when a ROCm index is
pinned, and in the generic reinstall path install from the pinned URL verbatim
rather than re-detecting the host ROCm version. gfx*/rocm7.2 indexes serve
torch 2.11+, so select the 2.11 package specs for a gfx leaf.
- install.sh: raise the torch constraint to 2.11 for */gfx* indexes too, matching
rocm7.2, so a pinned full-URL/family override that returns early keeps a valid
constraint.
Add _explicit_torch_index_url / _explicit_rocm_torch_index_url helpers and tests
covering the no-GPU CUDA pin repair and the explicit gfx index honored verbatim.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: honor torch-index override on the Windows installers too
The pinned-index work landed for install.sh and install_python_stack.py, but the
Windows installers still picked the wheel index from GPU probing. Extend the same
UNSLOTH_TORCH_INDEX_URL / _FAMILY contract so a pinned index wins on every platform:
- install.ps1: Get-TorchIndexUrl returns the pinned URL/family before nvidia-smi
probing; the AMD ROCm reroute is skipped when the index is pinned, so an explicit
cpu/cu* pin on an AMD host is not overwritten.
- studio/setup.ps1: add shared Get-PinnedTorchIndexUrl / Get-TorchIndexLeaf helpers;
the stale-venv check, the install selection and the AMD reroute all honor the pin,
and the CPU/CUDA install pulls from the resolved index URL.
- tests: parity test that all four installers read both override vars and the two
Windows installers gate the AMD reroute on the pinned flag.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: complete pinned-index handling for ROCm/Windows edge cases
Follow-ups to the override work flagged in review:
- install.ps1: a pinned gfx*/rocm>=7.2 index previously skipped the AMD reroute
that sets the torch>=2.11 floor, so the generic install used torch>=2.4,<2.11
and could resolve the known-bad _grouped_mm wheel. Route a pinned ROCm index
through the ROCm install path with the 2.11 floor + companions, and guard the
companion-spec lookup so a skipped reroute block cannot null-deref.
- studio/setup.ps1: the stale-venv check compared the installed flavor (cuXXX/cpu,
with +rocm misread as cpu) against the raw pinned leaf (gfx1151 / rocm6.4), so a
correct pinned ROCm venv was always marked stale. Classify +rocm wheels as the
generic 'rocm' flavor and normalize a pinned rocm*/gfx* leaf to 'rocm' before
comparing (cu* stays specific so cu126-vs-cu128 still rebuilds).
- install_python_stack.py: _ensure_cuda_torch now also reinstalls from a pinned
CUDA index when the venv carries a CPU wheel (headless CPU-venv-to-CUDA
cross-install via 'studio update'), not only when it finds a ROCm build.
- tests: parity assertions already cover all four installers honoring the override.
* install: finish pinned ROCm/CUDA edge cases on Windows + repair path
Follow-ups to the previous round:
- studio/setup.ps1: a pinned gfx*/rocm>=7.2 index now routes through the ROCm
install path with the 2.11 floor + companions (it previously fell through to the
CUDA branch with bare torch/torchvision/torchaudio against the ROCm index). The
CPU/CUDA fallback index is forced to the CPU wheel index when a ROCm index is
active, so a failed pinned-ROCm install does not retry the ROCm mirror.
- studio/setup.ps1: the stale-venv check no longer treats an unrecognized pinned
URL leaf (e.g. a PEP 503 mirror ending in /simple) as a torch flavor tag, which
was marking a correct venv stale; cu*/cpu/rocm/gfx leaves are still compared.
- install.ps1: the post-failure CPU fallback uses an explicit CPU index instead of
, which for a pinned ROCm index was the ROCm mirror itself (so the
'fallback' just retried the failing index and aborted the installer).
- install_python_stack.py: _ensure_cuda_torch now also reinstalls when the venv's
CUDA family differs from a pinned one (installed cu126 vs pinned cu128), not only
CPU->CUDA; the probe reports the installed cuXXX tag for the comparison.
* install: keep the ROCm to CPU fallback install inside the retry-helper window
The pinned-ROCm CPU fallback computes an explicit CPU index, but the comment
explaining why it cannot reuse $TorchIndexUrl pushed the actual
Invoke-InstallCommandRetry / --force-reinstall call more than 600 chars past the
"ROCm PyTorch install failed" message, so test_pr5940_followups's window check
no longer saw the retry helper. Move the CPU-index computation and its comment
above the failure substep so the retrying force-reinstall stays adjacent to the
message. No behavior change: same explicit CPU index, same retry, same
--force-reinstall.
* install: address #6692 review round 5 (ROCm/CPU pin edge cases)
setup.ps1:
- Stale-venv check: treat an AMD/ROCm host (HasROCm or a resolved gfx arch) with
no explicit pin as expecting "rocm", not "cpu", so a healthy +rocm venv is not
flagged stale (which made installer-managed setup exit and direct update rebuild).
- Pinned-ROCm install failure now routes into the force-reinstall CPU branch:
CuTag stays the rocm/gfx leaf on failure, so the condition also checks
ROCmCpuFallback; otherwise the CUDA branch installed from the CPU index without
--force-reinstall and kept the partial ROCm torch.
- Explicit ROCm pin compare no longer collapses gfx*/rocm* to a generic "rocm":
it compares the +rocmX.Y version (and the torch 2.11 line for gfx pins) so
changing the pinned family (e.g. rocm6.4 -> gfx1151) rebuilds and applies it.
install_python_stack.py:
- _ensure_rocm_torch: an explicit ROCm wheel-index pin now bypasses the
NVIDIA-present / no-AMD-GPU / unreadable-ROCm gates (headless/container/CI
cross-install), mirroring the explicit-CUDA-pin bypass in _ensure_cuda_torch.
- Add _ensure_cpu_torch: an explicit CPU pin (FAMILY=cpu or /cpu URL) now has a
repair path that reinstalls CPU torch over an existing CUDA/ROCm build on a
standalone update (which skips install.sh's flavor enforcement).
install.sh:
- Pin torchvision/torchaudio companions alongside torch for the rocm7.2 / per-gfx
index and the Strix reroute (those AMD indexes publish companions independently
and a bare name can resolve a torch-2.12-built wheel, an ABI mismatch).
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* torch-index override: classify CUDA pin by leaf; trim blank shell overrides
_ensure_cuda_torch only overrode the NVIDIA-presence gate for *any* pinned index,
so a non-CUDA mirror URL (or a ROCm/CPU pin) on a non-NVIDIA host with ROCm torch
could force a CUDA reinstall over a working ROCm venv. Add
_explicit_cuda_torch_index_url() (leaf cu*), matching the ROCm/CPU helpers, and
gate on it instead.
install.sh::get_torch_index_url treated a whitespace-only UNSLOTH_TORCH_INDEX_URL
/ _FAMILY as authoritative (yielding an invalid index), unlike the Python .strip()
and PowerShell IsNullOrWhiteSpace paths; trim leading/trailing whitespace first.
* install: honor pinned torch index over CVD/GPU gates and fix leaf-based ROCm classification
- install_python_stack.py: an explicit cu* pin now clears the CUDA_VISIBLE_DEVICES
empty/-1 hide gate as well as the NVIDIA-presence gate, so
CVD=-1 UNSLOTH_TORCH_INDEX_FAMILY=cu128 studio update repairs to CUDA wheels
(parity with install.sh's get_torch_index_url override, which skips all GPU
probing). Unpinned CVD=-1 still skips.
- install_python_stack.py: _ensure_cpu_torch installs the bounded _CPU_TORCH_PKG_SPEC
instead of a bare torch/torchvision/torchaudio trio; the /cpu index now also
serves torch 2.11+, which is outside the supported <2.11 range.
- install.sh: the torch>=2.11 constraint case matches the index leaf (rocm7.2|gfx*)
instead of the whole URL, so a mirror base path containing a gfx/rocm7.2 segment
with a cu*/cpu family is not false-matched onto the 2.11 line.
- setup.ps1: the stale-venv check expects rocm torch only for arches the install
path maps to a repo.amd.com wheel index; an unmapped/unreadable arch installs
CPU, so a correct CPU venv is no longer marked stale.
- Tests for each of the above.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: tighten pinned torch-index override edge cases
- install.sh: trim whitespace-only UNSLOTH_TORCH_INDEX_URL/_FAMILY before the
_torch_index_pinned guard, matching get_torch_index_url, so a blank override no
longer skips the WSL bootstrap and Radeon/Strix reroutes while detection still
picks the normal index.
- install.sh / install.ps1 / setup.ps1 / install_python_stack.py: force the torch
2.11 floor only for the gfx families with the <2.11 _grouped_mm bug (gfx120X-all,
gfx1151, gfx1150). A pinned override to gfx110X-all/gfx90a/gfx908 stays on the
default range, matching the automatic AMD path.
- install_python_stack.py _ensure_cuda_torch: treat an untagged CUDA build under a
CUDA pin as a family mismatch (reinstall), and match cuXXX pins narrowly (cu +
digits) so a custom/current mirror leaf no longer forces CUDA over a CPU/ROCm venv.
- install_python_stack.py _ensure_rocm_torch: reinstall when an explicit ROCm pin
names a different ROCm family than the already-installed ROCm torch (the ROCm
analogue of the CUDA cuXXX mismatch repair).
Adds tests for each case.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: fix second-order edge cases in pinned torch-index ROCm/CUDA handling
Parse the ROCm torch probe positionally so an empty HIP marker is kept:
CPU/CUDA torch no longer reads as HIP, so the ROCm reinstall is not skipped.
Emit one "|" line (like the CUDA probe) for a robust parse.
Limit the gfx torch 2.11 expectation to the install allowlist
(gfx120X-all/gfx1151/gfx1150). A pinned gfx110X-all/gfx90a/gfx908 index stays
on the default <2.11 specs, so a correct 2.10+rocm wheel is no longer judged a
mismatch and force-reinstalled every update.
Distinguish an AMD per-arch wheel (three-part +rocmA.B.C) from a generic
pytorch.org wheel (two-part +rocmA.B): a gfx per-arch pin over a generic 2.11
wheel now reinstalls the per-arch wheel, while an already-installed per-arch
wheel is not re-flagged (no reinstall loop).
Mirror all of the above in setup.ps1 via new Test-RocmGfx211Leaf /
Test-CudaFamilyLeaf / Get-RocmPinStaleTags helpers, reused by both the
install-spec path and the stale-venv check so they cannot diverge again.
Require a digit after "cu" (^cu[0-9]) in setup.ps1, install.ps1 and install.sh
so a mirror leaf like /custom or /current is not branded CUDA and does not
rebuild the venv every run.
Add tests: CPU/CUDA probe -> has_hip_torch False; gfx110X-all pin + 2.10 wheel
not stale; gfx1151 pin + generic 2.11 wheel stale; gfx1151 pin + per-arch wheel
not stale; /custom and /current not CUDA; plus cross-language allowlist and
cu-digit parity guards, and a PowerShell unit test for the new setup.ps1 helpers.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Fix ROCm/gfx pin case normalization, ROCm-tag requirement, and CUDA-leaf classification
Normalize torch-index leaves to lowercase before the gfx*/rocm*/cu* allowlist
matches so the canonical gfx120X-all (capital X) gets the torch 2.11 floor in
install.sh (leaf, flavor and repairable helpers). Require an installed +rocm
local tag before a rocmX.Y or non-2.11 gfx pin is judged satisfied in
setup.ps1 Get-RocmPinStaleTags and the Python _rocm_pin_family_mismatch, so an
untagged CPU/CUDA wheel never leaves the pin unapplied. Classify a leaf as CUDA
only via ^cu[0-9]: the Python _TORCH_BACKEND derivation now uses
_is_cuda_family_leaf, and install.sh brands cuda only on cu[0-9]* (unset on an
unknown /current /custom mirror leaf) so the stack probes the GPU instead of
skipping ROCm repair. Add bash, Python and PowerShell tests for capital
gfx120X-all floor, current/custom not-cuda, and untagged-wheel ROCm pins.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: converge torch-index pin detection via a per-venv marker
Introduce a torch-index MARKER that records the exact wheel --index-url used
after each successful torch install, so `unsloth studio update` / repair makes
the "did the pinned index change?" decision by an EXACT string compare rather
than inferring it from the wheel +rocm/+cu version tag. The tag cannot encode
the AMD per-arch gfx family (two 2.11 gfx indexes both install +rocm7.13.0), so
the tag heuristic missed a gfx1151 -> gfx120X-all switch and a custom-URL swap.
Marker path is per-venv (.unsloth-torch-index), one line = the resolved index
URL, written atomically (temp + rename). Path, format and normalization are
shared across all four installers (install.sh, install_python_stack.py,
setup.ps1, install.ps1).
- Reapply gfx pins on a per-arch target change: the marker's exact compare
reinstalls when the pinned index differs, even when both wheels share a tag.
- Honor custom ROCm URL pins during repair: an explicit index whose leaf is not
rocm/gfx/cu/cpu (e.g. simple, current) now reinstalls torch VERBATIM from the
pin when it differs from the marker ("URL wins verbatim").
- Align the KNOWN-2.11 rocm/gfx set to exactly rocm7.2 plus the gfx allowlist
gfx120x-all/gfx1151/gfx1150 in every language; stop treating an unknown newer
rocm (rocm7.3, which does not exist) as the 2.11 line speculatively.
Backward compatible: with no marker (old venvs, torch installed out-of-band) the
existing +rocm/version-tag heuristics still decide, and a matching marker never
reinstall-loops. A cu128 CUDA pin stays a CUDA pin; custom and current leaves are
not CUDA. Adds marker tests (py/sh/ps) plus cross-installer parity checks.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: keep the torch-index marker additive to flavor validation
Three narrow fixes in the marker-based stale-venv detection:
- setup.ps1: a matching marker no longer overwrites the detected installed
flavor. The marker compare is now an additional rebuild trigger, so a stale
wheel (torch swapped to a +cpu build while the marker still records a cuXXX
pin) is still caught by the flavor check instead of being masked as up to date.
- setup.ps1: a supported AMD arch carrying CPU torch is no longer marked stale
and wiped. The downstream AMD Windows ROCm override upgrades CPU torch to ROCm
in place, so wiping first would delete the venv and abort with "Virtual
environment not found". Only a genuinely wrong CUDA wheel still rebuilds.
- install.sh: the Radeon --find-links path records its repo.radeon.com base in
the marker instead of the generic pytorch.org ROCm fallback index, so a later
pin to that generic family correctly reinstalls rather than comparing equal.
Mirrors install.ps1/setup.ps1, which already record the real AMD index.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: honor custom pins and repair pinned venvs in place
Four follow-ups to the torch-index marker work:
- install_python_stack.py: _ensure_cuda_torch/_ensure_rocm_torch now bail when an
explicit custom-index pin names no known torch family, so a verbatim URL override
(a private/simple mirror) is not clobbered by auto-detected CUDA/ROCm wheels
before _ensure_verbatim_torch_index applies it.
- install_python_stack.py: the ROCm marker is additive, not a substitute -- a
matching marker still runs the family/version check so a wheel swapped after the
marker was written is caught. Mirrors setup.ps1.
- setup.ps1: a stale venv under an explicit pin, whose torch still imports, is
repaired in place (force-reinstall torch from the pin in the dependency pass)
instead of wiped. The wipe path only delegates to install.ps1, so on a direct
update it stranded the user at "Virtual environment not found" instead of
applying the new pin. A broken venv or unpinned drift still wipes/delegates.
- install.ps1: when a pinned ROCm install fails over to a CPU base, the marker now
records the CPU index actually used instead of the ROCm pin, so the next managed
setup does not see CPU torch under a ROCm pin and abort as stale.
* setup.ps1: keep the ROCm CPU-fallback force line the pr5940 test guards
5c93ffd4 folded the pin-change force-reinstall into the ROCm CPU-fallback
condition on one line, so the exact literal that test_pr5940_followups.py checks
(if ($ROCmCpuFallback) { $cpuForce = @("--force-reinstall") }) no longer appeared
and the test failed. Split the two conditions into separate if lines: the ROCm
fallback line is restored verbatim and the pin-change force is its own line. Both
still set $cpuForce to the array, so @splat passes one arg.
* install: honor exact CUDA/custom index URL pins in the torch-index marker
Address three Codex review findings on the torch-index marker mechanism:
- install.sh: after the ROCm CPU repair reinstalls torch from the generic
$TORCH_INDEX_URL, record that as the marker source. A Radeon --find-links
install set _TORCH_MARKER_INDEX_URL to its repo.radeon.com base earlier, so
leaving it made the marker misreport Radeon wheels and a later Radeon pin would
compare equal and skip a needed reinstall.
- install_python_stack.py: _ensure_cuda_torch now consults the exact-URL marker
(_marker_pin_mismatch) when the installed +cuXXX tag matches the pinned leaf,
so a same-leaf CUDA mirror change (official cu128 to an internal cu128 mirror)
is reinstalled and re-recorded instead of skipped.
- _normalize_index_url / _normalize_family_leaf (install.sh, setup.ps1,
install_python_stack.py): lowercase only KNOWN wheel-family leaves (rocm/gfx/
cpu/cuXXX) so gfx120X-all still matches gfx120x-all, while a custom
(unknown-family) leaf keeps its case so a verbatim URL pin like /Current does
not compare equal to /current. Tests updated to assert the refined behavior.
* install: fix 3 torch-index marker edge cases (CPU mirror pin, Radeon leaf, migrated venv)
Addresses three review findings on the torch-index override path:
1. CPU index URL change on an already-CPU venv. _ensure_cpu_torch returned
early whenever torch was already a CPU build, so a standalone update that
moved the pin (official /cpu -> a private UNSLOTH_PYTORCH_MIRROR /cpu, same
+cpu tag) never reinstalled. It now consults the exact-URL marker and
reinstalls only when _marker_pin_mismatch reports a different index,
mirroring the CUDA/ROCm same-family handling. A matching marker (or none)
still leaves CPU torch untouched, so there is no reinstall loop.
2. Radeon find-links directory misclassified as a pip ROCm family. A
repo.radeon.com/.../rocm-rel-7.2.1 leaf starts with "rocm" but is a
find-links listing, not a pip --index-url. The old startswith(("rocm",
"gfx")) test routed it into a --index-url reinstall that fails against
find-links. New _is_pip_rocm_family_leaf gates on ^rocm\d / gfx (matching
install.sh's rocm[0-9]* and setup.ps1's ^(rocm[0-9]|gfx)), so a Radeon URL
routes to the verbatim/marker path instead.
3. Migrated venv rewriting its marker to a pin it did not install. install.sh
and install.ps1 write the marker unconditionally, so a migration that
preserves existing torch recorded the newly requested pin and a later
update then found a matching marker and skipped the reinstall the pin
needs (e.g. a per-arch gfx1151 -> gfx120X-all switch, identical +rocm tag).
Both now track _TORCH_INSTALLED_THIS_RUN and write the marker only when
torch was actually installed or repaired this run.
Also add Get-NormalizedFamilyLeaf to the setup.ps1 helper-extraction list in
test_torch_index_marker.ps1 (it was added to setup.ps1 and the shell test in an
earlier round but missed here) and add two unit tests covering findings 1 and 2.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: keep pinned torch repairs on the pinned index
Two fixes for explicit index pins (UNSLOTH_TORCH_INDEX_FAMILY / _URL):
1. install_python_stack.py's repair paths ran uv without clearing the
inherited uv index env vars. uv resolves the default index (--index-url
or --default-index) at the LOWEST priority, so a UV_INDEX or
UV_EXTRA_INDEX_URL mirror in the environment won for any package it
served: a cu128-pinned repair could install torch from the mirror and
then record the cu128 marker it never used. Verified empirically: with
UV_EXTRA_INDEX_URL=.../cu126 exported, uv pip install torch
--index-url .../cu128 resolves torch 2.13.0+cu126. Strip the four uv
index env vars for pinned-index commands only, mirroring the gate
install.sh, install.ps1 and setup.ps1 already have; non-pinned installs
keep the user's mirror.
2. install.ps1 routed any pinned leaf matching rocm* through the ROCm
--default-index path, so a custom find-links leaf like rocm-rel-7.2.1
was treated as a PEP 503 ROCm index and could silently fall back to CPU
torch on resolution failure. Require a digit after rocm, matching
install.sh's rocm[0-9]* and install_python_stack.py's ^rocm\d.
Adds parity + unit tests for both (11 new tests).
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: keep pinned repairs off UV_TORCH_BACKEND and narrow setup.ps1's rocm pin match
Round 2 of the pinned-index hardening:
1. _build_uv_cmd converted UV_TORCH_BACKEND into --torch-backend before the
new env isolation could act, and uv's torch backend redirects torch
resolution to its own per-backend index even when --index-url is given
(verified: a cu128-pinned dry run with UV_TORCH_BACKEND=cpu resolves
torch 2.13.0+cpu). Pinned-index commands now never receive the flag and
UV_TORCH_BACKEND joins the stripped env vars, so uv cannot re-read it.
2. setup.ps1's pinned reroute had the same bare rocm* glob install.ps1 had:
a custom find-links leaf like rocm-rel-7.2.1 was routed through the ROCm
--index-url path instead of the verbatim unknown-pin path. Now requires
a digit after rocm, matching install.ps1, install.sh and
_is_pip_rocm_family_leaf.
3. The marker test's case-normalization checks used -eq, which is
case-insensitive in PowerShell, making them vacuous, and the unknown-leaf
expectation was written lowercased while the implementation deliberately
preserves custom-leaf case. Tightened to -ceq with the case-preserving
expected value.
Adds unit + parity tests for 1 and 2 (5 new tests).
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: extend the pinned-index guards to every remaining surface
Round 3 of the pinned-index hardening, closing the same holes on the
surfaces the earlier rounds missed:
1. install.sh's pinned-install env scrub now clears UV_TORCH_BACKEND (uv's
torch backend redirects torch resolution to its own per-backend index
even against --default-index), and both PowerShell wrappers clear it in
their pinned-install scrubs, matching install_python_stack.py.
2. setup.ps1's marker stale check still classified any rocm* leaf as a
PyTorch ROCm family while the install selection is digit-gated, so a
custom rocm-current / rocm-rel-7.2.1 pin stale-compared as
not-rocm vs rocm and force-reinstalled on every studio update. The
stale check now uses the same ^rocm\d gate.
3. install_python_stack.py's pinned-command scrub also strips
PIP_EXTRA_INDEX_URL for the pip fallback: pip adds the env extra index
in addition to --index-url, so an inherited mirror could satisfy torch
off the pin while the marker recorded the pinned URL. PIP_INDEX_URL
needs no strip since the explicit --index-url flag overrides it.
Parity + unit tests extended (4 new tests).
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: scrub find-links and carry the pinned scrub through pip fallbacks
Round 4 of the pinned-index hardening:
1. UV_FIND_LINKS joins every pinned-install scrub (install.sh, install.ps1,
setup.ps1, install_python_stack.py): uv's --find-links locations can
satisfy torch off the pinned index the same way an extra index does.
2. setup.ps1's Fast-Install restored the scrubbed vars in its finally
BEFORE the pip fallback ran, and never touched the pip env vars at all,
so a failed uv attempt fell back to python -m pip with an inherited
PIP_EXTRA_INDEX_URL / PIP_FIND_LINKS able to win over the pinned
--index-url. The scrub now wraps the whole function (uv attempt + pip
fallback) and includes the pip vars; restore happens after both.
3. install_python_stack.py's scrub also strips PIP_FIND_LINKS for its own
pip fallback, completing the PIP_EXTRA_INDEX_URL fix from round 3.
Parity tests extended (2 new tests).
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: digit-gate rocm leaves in marker normalization and ROCm side effects
Round 5 of the pinned-index hardening (three custom-rocm-leaf edge cases):
1. _normalize_family_leaf lowercased every leaf starting with rocm, so a
custom mirror leaf like rocm-Current compared equal to its lowercase form
and a case-only pin change was skipped. URL paths can be case-sensitive.
The rocm prefix is now digit-gated (rocm[0-9]*, matching
_is_pip_rocm_family_leaf) in install.sh, setup.ps1 and
install_python_stack.py, so only true family leaves (rocm7.2) are
lowercased; a custom rocm-* leaf keeps its case.
2. setup.ps1 Test-MarkerPinMismatch compared normalized URLs with -ne, which
is case-insensitive in PowerShell, so a case-only marker change (Simple
vs simple) was treated as matching and the reinstall skipped. Now -cne.
3. install.sh gated the AMD bitsandbytes install and the "repair ROCm torch"
--default-index reinstall on a bare whole-URL rocm glob, so a custom
CPU/CUDA/private index whose leaf merely starts with rocm (rocm-current)
was force-repaired from the wrong ROCm-only path whenever torch.version.hip
was empty. Both now gate on _torch_index_is_rocm_family, computed once from
the digit-gated leaf (rocm[0-9]*/gfx*).
Tests: 4 new parity assertions plus 2 case-sensitivity marker checks.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: apply an explicit custom torch-index pin on the first update
Round 6: an explicitly-set custom (unknown-family) UNSLOTH_TORCH_INDEX_URL
was silently ignored on the first `studio update` of a venv that predates
the marker feature, on both platforms, because the no-marker case was
treated as "do nothing" and the version-tag heuristics cannot judge an
unknown leaf.
1. install_python_stack.py _ensure_verbatim_torch_index now reinstalls
verbatim when the marker is ABSENT (None), not only when it differs, and
short-circuits only when the marker already records this exact pin. It
then writes the marker, so every later update is a no-op. A user who did
not set the override gets pin=None and is untouched, so an out-of-band
torch install is never clobbered.
2. setup.ps1: for an unknown-family pin on a marker-less venv the stale-venv
check now sets PinChangedForceReinstall so the torch block reinstalls in
place from the pin. It deliberately does NOT set shouldRebuild, which
would wipe the venv and strand a direct `studio update`.
3. setup.sh (the Linux `studio update` entry point) skipped
install_python_stack.py entirely when unsloth was already current, so the
marker-driven reinstall (both the verbatim custom pin and the cu/rocm
flavor and family-change repair, e.g. gfx1151 to gfx120X-all) never ran.
It now forces the dependency pass when a torch-index pin env var is set;
the pass is idempotent and no-ops when the marker already matches. This
mirrors setup.ps1's stale-venv pre-check.
Tests: 3 new parity assertions.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* test: expect first-update reinstall for a no-marker custom index pin
Follow-up to d671d8fb2: _ensure_verbatim_torch_index now applies an
explicit unknown-family URL pin verbatim on the first update when the
marker is absent (instead of no-op), so the old
test_verbatim_custom_url_no_marker_is_noop assertion was stale. Rewritten
as test_verbatim_custom_url_no_marker_reinstalls_once: asserts the one
verbatim reinstall from the pinned URL, that the marker is written, and
that a second call with the pin still set is idempotent (no reinstall
loop).
* install: gate the pinned update pass on the marker and record a pin baseline
Round 8, two follow-ups to the round-6 first-update pin fix:
1. setup.sh forced the full dependency pass on EVERY `studio update` while a
torch-index pin stayed exported, even after the marker already recorded the
same pin, turning quick updates into the expensive pass every time. It now
probes install_python_stack.py --torch-pin-needs-apply (which reuses the
exact marker normalization) and forces the pass only when the pin is not yet
applied (marker absent or different); an already-applied persistent pin keeps
the fast path. A probe error fails safe toward running the pass. setup.ps1
gets the same probe in its fast path for parity.
2. A known-family full-URL pin on a venv predating the marker (e.g. an installed
cu128 build and UNSLOTH_TORCH_INDEX_URL pointing at a same-family mirror) left
the marker absent forever: the _ensure_* helpers deliberately do not force a
multi-GB reinstall of identical-family wheels on an old venv, so nothing
recorded the pin and every update re-entered the pass. _record_torch_index_pin_baseline
now records the resolved pin as a baseline after the ensure sequence when the
family already matches and no marker exists, so the pin is tracked (a later
genuine change is detected and applied) and the update loop is broken, without
the redundant reinstall.
Tests: 3 new baseline unit tests, 4 new parity assertions, and the CLI probe.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* setup.sh: keep the pin probe's exit 1 from killing the update under set -e
The --torch-pin-needs-apply probe deliberately exits 1 for the common
steady-state answer (pin already recorded, keep the fast path), but it ran
as a bare command under set -euo pipefail, so the whole studio update
aborted before the exit code was even captured. Absorb the status with
|| _PIN_NEEDS_APPLY=$? and pre-seed 0 so all three outcomes route as
documented: 0 runs the pass, 1 keeps the fast path, anything else fails
safe into the pass. Parity test asserts the guard.
* install: strip pin credentials, disable uv config discovery, bound verbatim installs
Four verified fix groups from a 12-reviewer audit of the torch-index
override feature, each reproduced before fixing:
1. Credential persistence: all four marker writers stored the raw pin URL,
so an authenticated pin (https://user:token@mirror/simple) persisted its
credentials in .unsloth-torch-index (mode 0644 under a default POSIX
umask) and install_python_stack.py printed pin URLs verbatim in repair
messages. Userinfo is now stripped before persisting and in every
log/substep that interpolates a pin, via lockstep helpers
(_strip_index_url_credentials in install.sh / install_python_stack.py,
Remove-IndexUrlCredentials in install.ps1 / setup.ps1). The three
normalizers strip too, so an OLD marker that already carries credentials
still compares equal to the same pin: no reinstall loop on upgrade.
Query strings deliberately stay in the marker; two indexes distinguished
only by query must not compare equal.
2. uv configuration discovery beat the explicit pin: with a discovered
uv.toml declaring torch-backend = "cpu" or a [[index]] entry, uv 0.10.12
resolves torch 2.13.0+cpu against an explicit --index-url/.../cu126 pin;
UV_NO_CONFIG=1 restores +cu126 (reproduced both ways). The pinned-install
scrub in all four installers now sets UV_NO_CONFIG=1 and drops
UV_CONFIG_FILE.
3. The verbatim custom-index update path installed a bare, unconstrained
torch trio while fresh installs from the same unknown-leaf pin apply the
supported range; _ensure_verbatim_torch_index now installs the bounded
trio spec, closing the fresh-vs-update asymmetry.
4. Query-bearing pins (.../cu128?token=x) classified by raw leaf split and
force-reinstalled on every update (the installed cu128 never equals
cu128?token=x). Query/fragment are now stripped before leaf
classification in all four implementations; the marker comparison keeps
the query per (1).
Rejected after verification (no change): the pin-baseline record cannot
produce a wrong later decision (every pin change still mismatches and
reinstalls from the new pin); the venv temp-file symlink scenarios require
an attacker who already owns the environment; pathological inputs like
" / cu128 / " have no realistic caller and fail loudly.
Parity, stack, rocm-support, marker (sh + ps1), pin-stale, index-url and
flavor suites all pass (455 python + full shell/ps1 batteries).
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: harden custom-pin repair against clobber, broken torch, and pip config
Four follow-ups to the pinned-index audit fixes:
1. setup.ps1 routed an unknown-leaf custom pin through the CUDA branch with
a bare torch trio while install.ps1 (fresh) and the Python verbatim path
bound the supported range; the pinned unknown-leaf route now applies the
same torch>=2.4,<2.11.0 bound. Known cu* leaves and unpinned runs are
unchanged.
2. The final torch safety pass could not repair a clobbered unknown-family
pin: intermediate dependency steps can pull torch from PyPI (the pass
exists for exactly that reason), but the verbatim helper short-circuited
on marker==pin and no flavor tag exists to probe. The helper now keeps a
per-run snapshot of the installed trio (taken after a verbatim reinstall
or on the first matching-marker pass) and reinstalls from the pin when
the final pass sees the trio drifted. Probe failure skips the
comparison; a reinstall refreshes the snapshot, so no loop.
3. _record_torch_index_pin_baseline could freeze a known-family pin as
applied on a venv whose torch is missing or broken (every family helper
returns without reinstalling when its probe fails), making
--torch-pin-needs-apply report done forever. The baseline now probes the
installed flavor and records only on a match: a cuXXX pin requires the
matching +cuXXX tag, cpu requires a cpu build, rocm/gfx requires hip;
probe failure records nothing.
4. The pinned pip fallback stripped PIP_* env vars but user/site pip config
files still applied (a configured global.extra-index-url can satisfy
torch off the pin). PIP_CONFIG_FILE is now pointed at the null device
for pinned commands (pip loads no config files then), in
_install_env_for_cmd and setup.ps1's Fast-Install pinned scrub.
install.sh / install.ps1 have no pip fallback (uv-only), verified.
Tests: 7 new rocm_support tests (snapshot reset fixture), 1 stack test,
2 parity tests. Full battery green (464 python, sh and ps1 suites).
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: complete the pin-repair coverage across the fast path and platforms
Three cross-platform follow-ups to the round-2 pin-repair fixes:
1. The --torch-pin-needs-apply probe only compared marker==pin, so a torch
trio clobbered to the wrong family (a cpu wheel replacing cu128 via a
later pip install) with a still-matching marker reported "already
applied" and the _ensure_{cuda,rocm,cpu} repair never ran on the Linux
fast path. The probe is now a testable _torch_pin_needs_apply() that also
checks the installed flavor against a known-family pin (via a shared
_torch_flavor_matches_pin() helper, so the baseline and the probe cannot
drift). An unknown-family pin has no flavor to validate and a failed
probe cannot prove drift, so both keep the fast path.
2. macOS ARM (real CPU/MPS torch, not NO_TORCH) never applied an unknown-
family custom pin on update: both the verbatim path and the baseline
returned on IS_MACOS while fresh install.sh honors the pin, so the marker
was never written and setup.sh forced the dependency pass on every update
forever. The guards are now IS_MAC_INTEL (Intel mac is already NO_TORCH),
and the final pass applies the pin on macOS ARM.
3. The round-2 final verbatim repair sat in the step-13 sequence guarded
not IS_WINDOWS, so on Windows a dependency step that clobbered torch after
the pin was applied was masked by the matching marker (setup.ps1 does not
re-validate the main venv's torch after calling this script -- verified).
Step 13 now runs the verbatim snapshot-drift repair on Windows and macOS
ARM too; the Linux-oriented cuda/rocm/cpu family helpers stay Linux-only.
Tests: 13 new rocm_support cases (flavor drift, macOS ARM, Windows repair),
parity updates. Full battery green (475 python, sh and ps1 suites).
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: strip query tokens from the marker and tighten the pin-drift probe
Four follow-ups to the round-3 pin-repair fixes:
1. The credential stripper feeding the torch-index marker and the logged repair
messages dropped only user:pass@ userinfo, so a private feed that carries its
auth token in the query string (.../simple?token=SECRET) persisted the token
in the world-readable marker (mode 0644 under a default umask) and printed it
in substep output. All four strippers (install.sh, install.ps1,
studio/setup.ps1, install_python_stack.py) now drop the query and fragment
before building the sanitized URL. A query is not part of a PEP 503 index's
identity, so this also stops a rotated token from spuriously mismatching the
marker and forcing a needless reinstall.
2. The --torch-pin-needs-apply fast-path probe accepted an untagged CUDA build
(no +cuXXX local tag) under a specific cuXXX pin, but _ensure_cuda_torch
reinstalls exactly that build to enforce the pin. The probe was more lenient
than the repair, so the repair pass was skipped on the fast path.
_torch_flavor_matches_pin now reports a mismatch for an untagged build under a
cuXXX pin, forcing the pass.
3. The probe's ROCm branch accepted any HIP build for a rocm/gfx pin, while
_ensure_rocm_torch decides a reinstall with the per-arch
_rocm_pin_family_mismatch predicate (a generic +rocm7.2 wheel under a per-arch
gfx pin, or a wrong ROCm version, is a mismatch). The probe now reuses that
predicate, so it is as strict as the repair. This needs the installed torch
version, so _probe_torch_flavor now returns (marker, cutag, version) and
_torch_flavor_matches_pin takes the pin URL (extracting the leaf internally).
4. On Windows a known-family cu*/cpu pin is applied to the main venv by setup.ps1
before install_python_stack.py runs; a later dependency step can clobber it,
and the GPU-aware _ensure_{cuda,cpu}_torch self-skip on Windows while the
verbatim helper handles only unknown-family pins, so nothing repaired the
clobber (setup.ps1 does not re-validate the main venv's torch afterward,
verified). New _ensure_pinned_known_family_torch reinstalls a drifted cu*/cpu
pin in the step-13 Windows/macOS-ARM branch; rocm/gfx per-arch specs stay owned
by setup.ps1, unknown-family by the verbatim helper.
A speculative ROCm 2.11 floor was also raised but is unreachable: the rocm7.2
index publishes no 2.x wheel below 2.11.0, and an unknown newer rocm is not
floored speculatively.
Tests: query/fragment strip cases in the sh + ps1 marker suites and the Python
strip/marker tests; the tri-state helper and the probe/baseline harnesses moved
to the (marker, cutag, version) flavor with matching versions; new probe cases
(untagged CUDA, generic-rocm-under-gfx) and 8 _ensure_pinned_known_family_torch
tests; a four-way query-strip parity assertion. Full battery green (1150 python,
sh 26/26 marker, ps1 marker/flavor/pin-stale).
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: reinstall markerless gfx pins and cap custom-index updates at torch 2.11
Two follow-ups from the pin-marker audit:
1. A markerless venv with a gfx per-arch 2.11 pin trusted the wheel version
tag, which is byte-identical (+rocm7.13.0) across gfx120X-all / gfx1151 /
gfx1150. A pre-marker install holding one gfx arch's wheel that is now
pinned to a DIFFERENT gfx index was therefore never switched:
_rocm_pin_family_mismatch returns no-mismatch for any three-part +rocm
2.11 wheel, and _ensure_rocm_torch's absent-marker branch fell through to
that heuristic. _ensure_rocm_torch now forces a one-time reinstall when the
marker is absent AND the pin leaf is a 2.11 gfx per-arch index; the reinstall
writes the marker, so the next update compares exactly and does not loop
(the correctly-pinned no-reinstall guarantee then comes from the exact marker
compare, not the ambiguous tag). Non-gfx-2.11 pins (rocmX.Y, non-2.11 gfx)
stay on the tag heuristic -- their tags are distinguishable.
2. The verbatim custom-index update path used _CUDA_TORCH_PKG_SPEC (torch
<2.12.0) while a FRESH install of the same unknown leaf caps torch at
<2.11.0 (install.sh's default TORCH_CONSTRAINT, and setup.ps1's custom-pin
branch), so a private /simple mirror publishing torch 2.11 could upgrade a
`studio update` to a state the fresh installer never produces. Added
_CUSTOM_INDEX_TORCH_PKG_SPEC (torch>=2.4,<2.11.0), used only by the verbatim
path; companions stay pinned for the same exclusive --index-url ABI reason
as _CUDA_TORCH_PKG_SPEC (a bare name could pull a torch-2.12-built
torchvision). _CUDA_TORCH_PKG_SPEC is unchanged (known-family cu/cpu repair
correctly tracks install.sh's widened cu ceiling).
Tests: 2 new markerless-gfx cases (one-time reinstall + marker write + no-loop
second run, and the rocmX.Y absent-marker no-op), the pre-existing markerless
gfx no-reinstall test flipped to assert the one-time reinstall (it had encoded
the old tag-trusting behavior), and the custom-index bound assertions. 488
passed.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: a matching marker must not mask a broken, clobbered, or misclassified torch
Four round-6 follow-ups, all closing cases where a matching torch-index
marker wrongly vouched for a torch that is not actually the pinned one:
1. _is_cuda_family_leaf matched cu+digits by PREFIX (^cu[0-9]), so a custom
mirror leaf like cu128-private classified as CUDA family; the flavor check
then compared the installed cu128 tag to the whole leaf cu128-private and
forced a reinstall on EVERY update (never converging). The cu family is
now matched EXACTLY (re.fullmatch cu[0-9]+), so a cu-suffixed custom leaf
routes through the verbatim/unknown path with a stable marker. Mirrored in
install.sh (_normalize_family_leaf: strip cu, require an all-digit
remainder) and setup.ps1 / install.ps1 (^cu[0-9]+$).
2. _torch_pin_needs_apply returned False on a failed torch probe (missing or
unimportable) under a matching marker, so setup.sh kept the fast path and
a broken torch was never repaired. A failed probe now forces the pass: the
marker cannot vouch for a torch that does not import, forcing is idempotent,
and once torch imports again the probe succeeds and the forcing stops
(self-resolving). Reverses the round-4 conservative choice for this case.
3. _ensure_verbatim_torch_index snapshotted the installed trio on the first
pass with a matching marker and treated an unimportable torch (snapshot
None) as "no drift, skip", so a torch clobbered to a broken state before
the run was masked. A None snapshot now reapplies the pin. A torch
clobbered to a WORKING-but-wrong build under an unknown-family pin remains
undetectable from metadata (no flavor tag; reinstalling every update would
be the loop this avoids) and is documented as a known limitation.
4. The step-13 Windows final repair reran only the verbatim (unknown-family)
and known-family cu*/cpu paths, so a clobbered explicit rocm/gfx pin (the
wheel setup.ps1 installed from AMD's per-arch index) was left in place. The
branch now also runs _ensure_rocm_torch on Windows for an explicit rocm/gfx
pin; it has a Windows path and no-ops when torch already links HIP, so it
only reinstalls a genuinely clobbered ROCm venv (loop-safe).
Tests: the round-4 failed-probe-trusts-marker test flipped to force the pass;
new cases for the cu-suffix no-loop, the broken-torch verbatim reinstall, and
the Windows rocm final-repair structure; item-2 exact-cu parity assertions.
490 passed. sh/ps1 marker + flavor + pin-stale suites all green.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: repair Windows ROCm pins from the pinned URL and honor NO_TORCH
Four round-7 review items, two of them regressions in the round-6 work:
1. _torch_pin_needs_apply ignored UNSLOTH_NO_TORCH. With a torch-index env
var set and no marker, the failed-probe branch forced the dependency pass
on every `studio update`, and the pass (which also honors NO_TORCH) never
installs torch or writes a marker, so nothing could ever stop the forcing.
It now returns False immediately under NO_TORCH: the pin only matters once
torch is actually installed.
2. The step-13 Windows final repair (round-6) restored a clobbered explicit
rocm/gfx pin by calling _ensure_rocm_torch, whose Windows path reinstalls
from the arch AUTO-DETECTED via hipinfo, not from the pin. A user pinning a
different gfx family or a private mirror was restored from the wrong source
(and the wrong marker written), and a headless box was skipped entirely
(the arch probe returns nothing). The repair now goes through
_ensure_pinned_known_family_torch, which reinstalls from the PINNED url with
the same per-arch floor setup.ps1 uses (2.11-line gfx leaves) or a bare trio
(older arches, rocmN mirrors). It is gated on IS_WINDOWS since macOS ARM has
no ROCm, and the existing flavor check keeps it loop-safe (a matching HIP
wheel is left alone).
3. _ensure_verbatim_torch_index's broken-torch check (round-6) used
"_installed_trio_snapshot() is None", but that helper reports a REMOVED torch
as "torch==absent" (a non-None tuple) and a broken import as the stale
on-disk version, so a missing or unimportable torch under a matching marker
was read as "no drift" and skipped. The matching-marker path now confirms
torch health with an import probe (_probe_torch_flavor): a torch that does
not import reapplies the pin, while a healthy torch keeps the snapshot-based
intra-run drift detection.
4. A unit test for _ensure_cpu_torch did not pin NO_TORCH False like its
siblings, so a suite run with UNSLOTH_NO_TORCH=1 in the environment made the
guard return early and the reinstall assertions fail spuriously.
Tests: the round-6 broken-torch verbatim test re-encodes the non-None
"torch==absent" snapshot case (the exact state the old "is None" check missed);
new Windows-ROCm pinned-repair cases (reinstall from the pin, per-arch floor vs
bare spec, matching-wheel no-op, off-Windows no-op); a NO_TORCH fast-path probe
case; the parity test now asserts the Windows final branch does not auto-detect
the ROCm index and that the helper reinstalls from the explicit pin. 494 passed.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: floor the rocm7.2 index in the Windows pin repair; isolate marker tests
Three round-8 review items, two of them downstream of the round-7 changes:
1. _ensure_pinned_known_family_torch gave a rocm index leaf a bare
torch/torchvision/torchaudio trio while flooring only gfx* leaves, so a
Windows venv clobbered under an explicit rocm7.2 pin could reinstall an
unbounded or ABI-mismatched trio from that exclusive --index-url. It now
mirrors the spec the initial ROCm paths pin: the rocm7.2 floor for 2.11-line
gfx leaves and rocm leaves that serve torch 2.11, the <2.11 default for
older rocm versions, and a bare trio only for older gfx per-arch leaves
(which publish no floor), matching _ROCM_TORCH_PKG_SPECS / _ensure_rocm_torch.
2. test_verbatim_custom_url_no_marker_reinstalls_once called
_ensure_verbatim_torch_index twice; the second call now hits the
matching-marker health probe, and with pip_install mocked torch never becomes
importable, so in a no-torch environment _probe_torch_flavor returned None and
forced another reinstall, failing the idempotence assertion. The test now pins
a healthy flavor so the idempotence check is about the marker, not ambient
torch.
3. The TestEnsureRocmTorchMarker fixture patched os.environ per test but not
_TORCH_BACKEND, which install_python_stack.py computes once at import from
UNSLOTH_TORCH_BACKEND. A runner starting with a cuda/cpu backend made
_ensure_rocm_torch early-return and skip the mocked repair these tests
exercise. The fixture now neutralizes _TORCH_BACKEND so the marker tests are
independent of the caller's installer-pin environment.
Tests: the Windows floor-spec test now asserts a rocm7.2 mirror pin uses the
rocm7.2 floor (not bare), plus a new rocm7.1 case that must fall back to the
<2.11 default; the marker suite passes under a hostile
UNSLOTH_TORCH_BACKEND=cuda / UNSLOTH_TORCH_INDEX_URL env. 495 passed.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: apply same-flavor pin repoints, keep ROCm fallback nonfatal, bound custom companions
Four round-9 review items, two of them regressions in the round-7 pin helper:
1. _ensure_pinned_known_family_torch returned as satisfied whenever the installed
flavor matched the pin, so a same-flavor SOURCE change (one /cpu or /cu128
mirror to another, or a gfx1151 -> gfx120x-all per-arch switch, both carrying
the same wheel tag) was never applied, while _torch_pin_needs_apply kept forcing
the pass on the marker mismatch forever. It now also reinstalls when the marker
records a DIFFERENT index of the same flavor, rewriting the marker so the next
update matches (no loop), exactly as the Linux _ensure_{cuda,cpu}_torch helpers
do. An absent marker on an already-matching venv is still left to the baseline
recorder (no forced reinstall of a correct pre-marker venv).
2. That helper reinstalled a Windows ROCm pin with the FATAL pip_install, so when
setup.ps1 had taken its CPU fallback (the pinned AMD index unavailable), the
final repair re-hit the same missing index and aborted the whole install. The
ROCm reinstall is now nonfatal (pip_install_try): on failure it leaves the CPU
base in place and writes no ROCm marker, so the install completes -- matching
_ensure_rocm_torch's Windows path. cu*/cpu pins stay fatal (authoritative source).
3. install.sh left torchvision/torchaudio bare for a pinned custom/unknown-leaf
index (a private /simple mirror), unlike the Python update path's
_CUSTOM_INDEX_TORCH_PKG_SPEC, so a mirror also exposing newer companion wheels
could resolve a torch-2.12-built torchvision against the capped <2.11 torch. It
now bounds the companions (torchvision>=0.19,<0.26.0 / torchaudio>=2.4,<2.11.0)
for a custom leaf, gated on an empty _expected_torch_flavor_tag so known families
keep their curated bare/floored companions.
4. install.sh's _expected_torch_flavor_tag matched cu[0-9]* by prefix, so a custom
leaf like cu128-private classified as the cu128 family and force-reinstalled a
correct +cu128 wheel on every run. It now requires exact cu+digits (routing the
suffixed leaf to the custom path), matching the Python re.fullmatch(cu[0-9]+) and
PowerShell, and feeding item 3's custom-leaf detection.
Tests: new cases for the same-flavor marker-change reinstall, the nonfatal ROCm
fallback (no marker on failure), the rocm7.2/older-rocm floor selection now split
across the nonfatal path, cu-suffixed custom leaves in test_torch_flavor.sh, and the
custom-leaf companion bounds in test_torch_constraint.sh. 497 python + 143 shell
assertions pass; the marker suite still passes under a hostile
UNSLOTH_TORCH_BACKEND=cuda env.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: bound custom-pin companions on the Windows setup path; isolate pin-probe tests
Two round-10 review items:
1. setup.ps1's custom/unknown-leaf pin branch capped only torch ($cudaTorchSpec)
and still asked the exclusive index for bare torchvision/torchaudio, so a
private mirror that also serves newer companion wheels could install a
torch<2.11 wheel alongside a torchvision>=0.26 / torchaudio>=2.11 built for a
newer torch ABI, after which the marker records the pin as applied. It now
bounds the whole trio (torch>=2.4,<2.11.0 / torchvision>=0.19,<0.26.0 /
torchaudio>=2.4,<2.11.0) for a pinned non-cu-family leaf, matching install.sh,
install.ps1's fresh pinned install, and install_python_stack.py's
_CUSTOM_INDEX_TORCH_PKG_SPEC. This completes the companion-bounds fix across all
three installers; known cu* leaves keep bare specs (the family index bounds them).
2. The _torch_pin_needs_apply probe tests did not pin NO_TORCH False, so a test
process launched with UNSLOTH_NO_TORCH=1 short-circuited the probe (the round-7
guard) and returned False for cases that expect the pass to run. The _needs_apply
helper now patches NO_TORCH (default False) around the call, and the dedicated
no-torch case passes no_torch=True explicitly.
Tests: the cross-platform parity test now asserts setup.ps1 bounds the full trio
(not just torch) for a custom leaf; the pin-probe suite passes under a hostile
UNSLOTH_NO_TORCH=1 environment. setup.ps1 parses clean; 497 python + shell suites
green.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: bound custom rocm-* pins, redact diag tokens, snapshot custom pins before base update
Three round-11 review items, all reproduced before fixing:
1. install.sh's custom-index companion bounds gated on _expected_torch_flavor_tag
returning empty, but that helper returned "rocm" for ANY rocm* leaf, so a custom
mirror whose leaf starts with rocm but is not a pip family (a private rocm-current
mirror, a Radeon find-links rocm-rel-7.2.1) escaped the bounds and installed bare
torchvision/torchaudio. It now digit-gates rocm to rocm[0-9]* (matching the Python
_is_pip_rocm_family_leaf ^rocm\d), so those custom leaves return "" and the <2.11
companion caps apply; real rocm7.2 / gfx per-arch indexes still classify as rocm.
2. _tauri_torch_index_family classified by the raw last path segment, so a pinned URL
carrying auth in the query (.../rocm7.2?token=SECRET) had the token echoed verbatim
into the emitted [TAURI:DIAG] line. It now strips query/fragment before classifying
(mirroring the marker/log credential stripping), so no token reaches the diagnostic
output; as a side effect .../cu128?token=x now classifies as cu128 instead of auto.
3. On studio update, the core package step (a newer unsloth can require a torch the
custom pin does not satisfy, pulling a default PyPI trio) runs BEFORE the step-2b
verbatim check, which then recorded the already-clobbered trio as the baseline for a
matching marker and left the pin unapplied. A new _capture_verbatim_baseline() records
the pre-clobber trio before the core step, so the verbatim pass detects the drift and
reapplies the pin. Captures only for a matching custom pin with importable torch; a
mismatched/absent marker or broken torch is left to _ensure_verbatim_torch_index.
Tests: _expected_torch_flavor_tag rocm-current / rocm-rel cases; _tauri_torch_index_family
token/fragment redaction with a no-leak regression guard; _capture_verbatim_baseline
record/skip cases plus an end-to-end clobber-detection scenario; a structural guard that
the capture runs before the core step. 501 python + shell suites pass; install.sh bash -n
clean, shellcheck unchanged from base.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: match rocm family leaves exactly, enforce the rocm7.2 torch line, repair a broken pinned torch
A pinned index is a pip ROCm --index-url family only when its leaf is an exact
rocm / rocm. (rocm7.2) or a gfx* per-arch leaf. The prior
^rocm[0-9] prefix match also caught suffixed private-mirror leaves (rocm7.2-private,
rocm7-current), routing them through the ROCm/companion-family path instead of the
verbatim pin: the companion bounds were skipped and, on a pre-marker venv with a
compatible +rocm wheel, the pin was never applied. Match the family exactly through one
shared helper at every site:
- install_python_stack.py: _is_pip_rocm_family_leaf (re.fullmatch), plus the two other
loose gates it feeds (_normalize_family_leaf, _torch_flavor_matches_pin).
- install.sh: a new _is_pip_rocm_family_leaf routes _expected_torch_flavor_tag,
_torch_index_repairable, _normalize_family_leaf and the ROCm side-effect gate.
- setup.ps1: a new Test-PipRocmFamilyLeaf routes Get-NormalizedFamilyLeaf and both
pinned reroutes; install.ps1 anchors its reroute regex.
_rocm_pin_family_mismatch (and its setup.ps1 mirror Get-RocmPinStaleTags) compared only
the ROCm version, so a +rocm7.2 wheel whose torch release drifted off the 2.11 line
(2.12/2.13 from an out-of-band upgrade or a custom rocm7.2 mirror) satisfied the family
check while violating _ROCM_TORCH_PKG_SPECS['rocm7.2'] (torch>=2.11,<2.12). Flag it stale
so the repair reinstalls to floor; >=2.11 alone is not enough, so the release is compared
exactly against the 2.11 line for a KNOWN-2.11 rocm pin.
_ensure_pinned_known_family_torch returned on a failed import probe, but
_torch_pin_needs_apply forces the dependency pass on that same failed probe: a broken
torch under a known-family pin was left in place and the pass was forced on every update.
Treat an unimportable torch as drift and reinstall the pinned trio (the spec and marker
derive from the pinned leaf, not the absent flavor); once it lands the probe succeeds and
the fast path returns.
Tests: exact-match cases across test_torch_flavor.sh, test_rocm_support.py,
test_cross_platform_parity.py and the two .ps1 helper suites; the rocm7.2 release-line
and broken-probe-reinstall cases; extraction lists updated for the new helpers.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: anchor the PS pinned-ROCm floor gate and bound install.ps1 custom-pin companions
Round 12 made every family CLASSIFIER exact, but the Windows install-flow floor gate reads
$_pinRocm211 directly from the raw pinned leaf with an unanchored -match '^rocm(\d+)\.(\d+)'
BEFORE any exact classification runs. A suffixed custom leaf (rocm7.2-private) matches that
rocm7.2 prefix, so it takes the 2.11-floor branch and is force-routed through the ROCm
install path before the exact-match elseif can send it to the verbatim install. Anchor the
match ($) in both install.ps1 and setup.ps1 so only an exact rocmX.Y leaf is floored; a
suffixed or newer-suffix leaf falls through to the verbatim path. The Python floor
selection is already exact (dict lookups gated on _is_pip_rocm_family_leaf), so only the two
PS scripts needed this.
install.ps1's custom (non-cu-family) pinned-torch install bounded torch>=2.4,<2.11.0 but
left torchvision/torchaudio bare, so a private mirror serving newer companions could pull a
wheel built for a newer torch ABI while the marker records the pin as applied. Bound both
companions (torchvision>=0.19,<0.26.0 / torchaudio>=2.4,<2.11.0) when the leaf is not a
cu family index (a cu index bounds its own resolution), matching setup.ps1's
Test-CudaFamilyLeaf gate and _CUSTOM_INDEX_TORCH_PKG_SPEC.
Tests: parity guards for the anchored floor gate in both PS scripts and for install.ps1's
bounded custom-pin companions.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: tighten comments in the torch-index-override paths
Collapse the verbose comment and docstring blocks added across the installer
scripts and their tests to fewer, clearer lines without changing behaviour.
Remove a duplicated CUDA-spec comment block. Comments/docstrings only; no code
changes (AST-verified).
* install: repair a broken pinned torch on Linux, strip trailing slash in tauri family, count the final step
_ensure_cuda_torch / _ensure_cpu_torch returned on a failed import probe (torch present but
unimportable). With an explicit CUDA/CPU pin, _torch_pin_needs_apply forces the dependency
pass on that same failed probe, and the base package update does not force-reinstall an
already-installed torch distribution, so the broken torch was left in place and the pass
reran every update without repairing it. Treat a failed probe under a pin as drift and
reinstall from the pinned index (the reinstall rewrites the marker and the next probe
imports, so no loop). This is the Linux counterpart of the known-family repair fix.
_tauri_torch_index_family stripped the query/fragment before classifying but not a trailing
slash, so a token-authenticated pin like .../cu128/?token=x collapsed to .../cu128/ and fell
through the exact-suffix */cu128 and */cpu arms to "auto". Strip a trailing slash too,
mirroring _torch_index_url_leaf.
The Windows / macOS-ARM final torch-repair step (_ensure_pinned_known_family_torch) runs a
progress step that base_total never counted (the final-step increment was gated to Linux),
so _STEP ran one past _TOTAL on those platforms. Add the missing increment.
Tests: broken-probe reinstall for the CUDA (family and URL pins) and CPU paths; trailing
slash / slash+token cases for _tauri_torch_index_family; a full-flow progress-count guard
asserting _STEP == _TOTAL on Windows and Linux.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: tighten comments in the torch-index-override paths
* install: harden the torch-index pin across all four installers
Redact index-URL credentials from captured install logs before they print on
failure. uv/pip failure text embeds the failing --index-url verbatim, so a
user:token@ or ?token= secret could leak into the console. Add a shared
redaction pass (_redact_install_output / Redact-InstallOutput) wired into the
error-output dump in install.sh, install.ps1, setup.ps1 and
install_python_stack.py. Verbose mode still streams live uncaptured output, so
it is intentionally left unredacted (developer opt-in).
Trim trailing slashes on the PATH only for a verbatim UNSLOTH_TORCH_INDEX_URL
override, preserving a ?query/#fragment token. A whole-URL rstrip corrupted a
base64 token ending in "/", and a single-slash strip left .../cu128//
classifying as an empty leaf. Add _trim_index_path_slashes /
Trim-IndexPathSlashes and route the override through it; strip ALL trailing
slashes in the backend-branding leaf classifier so a double slash still yields
the real leaf.
Reject a trailing-dot ROCm leaf (rocm7.) in the bash family validator so it
matches Python re.fullmatch(rocm\d+(?:\.\d+)?) and the PowerShell regex: both the
major and the minor must be non-empty digits, so rocm7. is a custom verbatim pin,
not a pip ROCm family.
Scrub PIP_NO_INDEX and PIP_INDEX_URL for a pinned install in the two installers
that have a plain-pip fallback (install_python_stack.py, setup.ps1):
PIP_NO_INDEX=1 makes the fallback ignore every index including the pinned
--index-url, and PIP_INDEX_URL replaces it. install.sh and install.ps1 install
via uv --default-index (which ignores pip config/env), so they are unaffected.
Add unit tests (bash, Python, PowerShell) and cross-platform parity tests
covering credential redaction, path-only slash trimming, the rocm7. validator,
the double-slash leaf, and the PIP_NO_INDEX/PIP_INDEX_URL scrub.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: redact captured torch-install output and warn on a failed pinned ROCm repair
Close a redaction gap the earlier pass missed: setup.ps1's direct
`Fast-Install ... | Out-String` branches (ROCm from $ROCmIndexUrl, CPU/CUDA from
$TorchInstallIndexUrl, plus the Triton and T5 sub-venv installs) printed the
captured $output verbatim on failure, bypassing Redact-InstallOutput. A private
index carrying userinfo or a ?token= in the pin could leak into Windows Studio
setup logs. Route every `Write-Host $output` through Redact-InstallOutput.
Warn on a failed pinned Windows ROCm reinstall in
_ensure_pinned_known_family_torch: the branch printed "reinstalling from it" then
called pip_install_try, but had no else, so a failure continued silently and left
the user believing the pin was applied while the old CPU/wrong torch survived.
Mirror the auto-ROCm Windows path and warn, telling the user to retry.
* install: redact captured output on the pip fallback and optional-install failure paths
The uv install path already redacted its captured output, but pip_install's pip
fallback runs through run(), which printed result.stdout verbatim on failure, and
_print_optional_install_failure did the same. A pinned --index-url carrying
userinfo or a ?token= could still leak there when uv is unavailable or the pip
fallback also fails. Route both through _redact_install_output. The verbose
pip_install_try path stays raw (developer opt-in), matching the other installers.
* install: split the survive-updates marker subsystem into a follow-up
The torch-index override PR grew a persisted per-venv marker plus repair
machinery (stale-pin detection, verbatim re-apply, update-time reinstall
triggers) that roughly doubled it. That subsystem is orthogonal to the core
feature and is being reworked in a follow-up (versioned/hashed marker,
full-URL pin baseline), so it moves there wholesale instead of shipping
twice.
What this PR still does: UNSLOTH_TORCH_INDEX_URL / UNSLOTH_TORCH_INDEX_FAMILY
pick the torch wheel index at install time in all four installers, with the
exact rocm/gfx/cpu/cu leaf classification, the torch 2.11 floor for the
per-arch AMD indexes, bounded companions for custom leaves, credential
redaction of captured installer output, path-only slash trimming, and the
uv/pip index env scrubs. Flavor-based repair keeps honoring the pin: a wrong
family under an explicit pin still reinstalls from the pinned URL, and
setup.ps1 repairs a pinned stale venv in place instead of wiping it.
What moves to the follow-up: the .unsloth-torch-index marker file and its
writers/readers/normalizers, exact-URL pin-change detection on update
(same-tag gfx switches, custom-mirror repoints), the verbatim trio snapshot
and clobber re-apply, the pin-baseline recorder, and the
--torch-pin-needs-apply fast-path probe in setup.sh / setup.ps1. Their tests
(the marker sh/ps1 suites, the stale-pin suite, and the marker classes in the
rocm/cuda/parity suites) move with them; the removed code is preserved on a
local archive branch to seed that PR.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: re-apply a ROCm pin over an existing HIP wheel via the version tag
The subsystem split left an explicit ROCm/gfx pin unenforced on `studio
update` whenever the venv already imported ANY ROCm torch: the pinned
reinstall lived inside the `elif not has_hip_torch` branch, so a rocm6.4 to
rocm7.2 switch, a gfx1151 pin over a generic +rocm7.2 wheel, or a broken
2.12+rocm7.2 drift never re-applied the pin.
Restore the markerless half of that detection: _rocm_pin_family_mismatch
compares the pinned leaf against the installed wheel tag (exact rocmX.Y
compare, the 2.11 gfx per-arch allowlist, the untagged-wheel rule), the HIP
probe emits "|" again so the installed tag is available,
and _ensure_rocm_torch reinstalls from the pinned URL when the tag mismatches
even though HIP torch is present. setup.ps1 mirrors it: the stale-venv check
routes a pinned rocm/gfx leaf through Get-RocmPinStaleTags instead of
collapsing it to a generic "rocm" flavor, and the existing pinned in-place
repair (no wipe) applies the change.
What still waits for the follow-up marker PR, by design: pin changes the
wheel tag cannot see -- a per-arch switch between two 2.11 gfx indexes
(identical +rocm7.13.0 tag), a custom-mirror URL repoint under the same
family leaf, and unknown-family verbatim pins. Those need the persisted
index record.
Tests restored with the code: the _rocm_pin_family_mismatch table, the five
update-path cases (older-rocm reinstall, gfx-over-pre-2.11 reinstall,
matching-pin no-reinstall, non-2.11 gfx no-reinstall, gfx-over-generic-2.11
reinstall), the "|" probe-format guards, and the AST-extracted
Get-RocmPinStaleTags suite for setup.ps1.
* install: compare major-only rocm pins, redact URL fragments, bound pinned CPU trio
Three review fixes on the restored pin-repair path.
The family classifier accepts a major-only rocm leaf (rocm7), but the
mismatch comparators only parsed rocmX.Y, so a rocm7 pin fell through to the
2.11-line fallback and INVERTED both verdicts: an installed +rocm6.4 wheel
compared as satisfied (pin never re-applied) while a matching +rocm7.2 wheel
compared as stale (reinstall loop). Major-only pins now compare on the major
alone in _rocm_pin_family_mismatch and Get-RocmPinStaleTags: rocm6.x under a
rocm7 pin is a mismatch, any rocm7.x satisfies it, an untagged wheel never
does, and a bare +rocm tag with an unreadable version is accepted (matching
the existing lenient unreadable fallback).
The output redactors scrubbed userinfo and ?query= values but not #fragments,
so a pin like https://mirror/whl/cu128#token=secret leaked the secret in
captured uv/pip failure text -- inconsistent with the URL handling itself,
which already treats fragments as sensitive. All four redactors gain a
URL-anchored fragment rule (anchored so a bare "# comment" line in tool
output is never touched).
setup.ps1's CPU branch installed a bare torch/torchvision/torchaudio trio;
fine for the unpinned host default, but a PINNED cpu index routes through the
same branch and the /cpu index serves newer torch, so a fresh pinned CPU
install could land an unsupported trio that _ensure_cpu_torch then keeps
(it accepts any CPU build). Under a pin the branch now installs the bounded
trio mirroring _CPU_TORCH_PKG_SPEC (torch>=2.4,<2.12.0 and matching
companions); the unpinned path is unchanged.
Tests: major-only rows in the Python mismatch table and the AST-extracted
setup.ps1 suite; fragment + query-plus-fragment + bare-hash-comment cases in
all four redactor suites; a parity check that the pinned CPU trio bounds
exist, are gated on the pin, and mirror the Python repair spec.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* install: tighten comments in the torch index override paths
* tests: track the moved pass-through inheritance in the gguf order check
Main moved the llama_extra_args pass-through inheritance out of the
GGUF branch into _resolve_inherited_extra_args, which runs before it,
so the source-order assertion's "if request.llama_extra_args is None"
anchor no longer exists inside the branch and the check failed after
the main merge. The test now asserts the same property in the current
shape: inheritance before the GGUF branch (a carried --no-mmproj still
shapes the hub guard's companion requirement), and marker, hub guard,
unload in order within the branch. Full file passes (32 tests).
* tests: anchor the inheritance order check on the call, not the definition
source.index("_resolve_inherited_extra_args(") matched the function
definition, which always precedes the endpoint, so the ordering
assertion was vacuously true. Anchoring on "= _resolve_inherited_
extra_args(" pins the first call site inside the load endpoint (line
4505), which is the statement whose position relative to the GGUF
branch the test is meant to guard. 32 tests pass.
* tests: align the gguf order test with main
Main fixed the stale ordering assertion in PR 7252; adopting its
version verbatim removes this file from the branch diff entirely and
avoids a conflict on the next main merge. 32 tests pass.
* install: bound the companion constraints to torch's window everywhere
A full platform x vendor validation matrix over this branch surfaced a
real trio mismatch on the cpu/mac paths: torch is capped <2.11 (installs
2.10.0+cpu) but the bare torchaudio companion resolves 2.11.0+cpu,
because torchaudio 2.11 dropped its exact torch pin. Reproduced in a
sandboxed end to end cpu install. torchvision still exact-pins torch and
self-corrected.
The default companion constraints are now bounded to torch's window
(<0.26 / <2.11) and widen together with the cu* torch window (<0.27 /
<2.12), so every leaf resolves a paired trio. Verified with uv dry-runs
on the cpu, cu130, and rocm6.4 leaves (2.10.0/0.25.0/2.10.0,
2.11.0/0.26.0/2.11.0, 2.9.1/0.24.1/2.9.1) and a rerun of the sandboxed
cpu install, which now lands torch 2.10.0+cpu with torchaudio
2.10.0+cpu.
The Strix WSL reroute now also forwards UNSLOTH_TORCH_INDEX_URL and
UNSLOTH_TORCH_INDEX_FAMILY into the rerouted 24.04 distro; dropping
them silently reverted the child install to auto-detection, defeating
the pin this branch introduces.
test_torch_constraint.sh updated: the bounded companions must appear at
the defaults and the custom-leaf block, no bare companion may remain,
and the cu* widen must carry the companions with it.
* install: harden the override path against reroute drift and credential leaks
Review sweep focused on default-path idempotency found no defects on the
unset path; these fixes cover the override path and failure reporting.
install.sh:
- The early WSL Strix Halo distro reroute now honors an explicit index
pin (UNSLOTH_TORCH_INDEX_URL / _FAMILY): the pin is used in the current
distro instead of probing the GPU and re-entering another distribution,
matching the contract of the later Radeon and Strix guards. Whitespace
only values do not gate, in parity with get_torch_index_url.
- Verbose mode now streams installer output through the credential
redactor; it previously bypassed the redaction the quiet path applies.
The exit code survives the pipe via an rc file since the script runs
under plain sh with no pipefail.
- The kept-release fallback warning now strips credentials from the
index URL before printing it.
install.ps1:
- Bounded torchvision and torchaudio next to every capped torch install
(custom pin, ROCm CPU fallback, CUDA flavor repair). torchaudio 2.11
dropped its exact torch pin from the wheel metadata, so a bare
companion beside torch<2.11 can resolve a mismatched 2.11.0 build,
cu family indexes included. Mirrors the install.sh companion bounds.
studio/install_python_stack.py:
- The verbose failure path now redacts index URLs in pip and uv output
before printing, matching every other output site in the file.
All sh, ps1 and python installer test suites pass (the host-defaults
suite has a known pre-existing failure unrelated to this change).
* install: redact verbose Windows installer output and repair the parity tests
Follow-ups to the override-hardening commit, from review:
- install.ps1 Invoke-InstallCommand and setup.ps1 Invoke-SetupCommand now
pipe verbose output through Redact-InstallOutput per record, and the
three verbose Fast-Install torch call sites (ROCm, CPU, CUDA) do the
same: uv and pip echo the pinned index URL, credentials included, in
their errors, and verbose mode previously bypassed the redaction the
quiet paths apply. ForEach-Object and Out-Host leave $LASTEXITCODE
untouched, verified with a native command exiting 7 behind the pipe.
- test_cross_platform_parity.py: the install.ps1 companion-bounds
assertion now matches the implemented behavior (bounds on every index,
no cu-family exemption, since torchaudio 2.11 dropped its exact torch
pin) instead of requiring the removed $_pinCuLeaf gate.
- test_rocm_support.py: the WSL reroute guard test slices the whole
function body to its closing brace instead of a fixed 1200-character
window, which the new pin-gate preamble had outgrown.
428 tests pass across the parity, install stack and rocm support suites;
the sh and ps1 installer suites pass unchanged.
* install: tighten comments in the torch-index and ROCm/CUDA repair paths
* install: digit-gate the gfx family leaf and honor ROCm pins in the Windows repair
Two review follow-ups on the override path:
- The pip ROCm family predicate accepted ANY gfx-prefixed leaf, so a
custom verbatim pin like /gfx-private classified as a ROCm family and
enabled the ROCm-only side effects (AMD bitsandbytes, ROCm torch
repair) on a mirror that may serve CPU/CUDA wheels. gfx now requires a
following digit (gfx90a, gfx1151, gfx120X-all), consistently in
install.sh, install_python_stack.py, install.ps1 (family gate and
expected-flavor classifier) and setup.ps1, matching the strictness the
rocm side already had (rocm7.2-private stays verbatim). The broader
backend BRANDING globs are unchanged on purpose: radeon repo leaves
(rocm-rel-X.Y) must still brand the rocm backend without being
force-repaired as a family.
- The Windows branch of the ROCm torch repair always installed from the
public per-arch index, ignoring an explicit ROCm-family pin: after a
pinned setup.ps1 install failed to a CPU base, the repair retried
repo.amd.com instead of the pinned index. The branch now resolves
_explicit_rocm_torch_index_url() first, uses it as the install index
when set, and mirrors the Linux pin contract by skipping the NVIDIA
and gfx-detection gates a pin is documented to override.
Source-assertion tests updated to the tightened predicate and the new
repair label. 1165 tests pass across the parity, install stack and
studio install suites; the sh and ps1 suites pass; both PowerShell
installers parse clean.
* Remove scratch archives accidentally committed with the comment pass
The temp/ archive copies of installer and test files were working
scratch, not PR content, and inflated the diff by about nine thousand
lines.
---------
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
---
install.ps1 | 153 ++-
install.sh | 343 +++++--
studio/install_python_stack.py | 872 +++++++++++++-----
studio/setup.ps1 | 402 +++++++-
tests/python/test_cross_platform_parity.py | 606 ++++++++++++
tests/python/test_install_python_stack.py | 134 +++
tests/run_all.sh | 1 +
tests/sh/test_get_torch_index_url.sh | 57 ++
tests/sh/test_redact_install_output.sh | 89 ++
tests/sh/test_torch_constraint.sh | 74 ++
tests/sh/test_torch_flavor.sh | 89 +-
tests/studio/install/test_cuda_repair.py | 189 +++-
.../install/test_gpu_detection_followups.py | 131 ++-
tests/studio/install/test_pr5940_followups.py | 2 +-
tests/studio/install/test_rocm_support.py | 419 ++++++++-
tests/studio/test_setup_pin_stale.ps1 | 114 +++
tests/studio/test_torch_flavor.ps1 | 15 +-
.../studio/test_torch_index_pin_hardening.ps1 | 78 ++
18 files changed, 3382 insertions(+), 386 deletions(-)
create mode 100755 tests/sh/test_redact_install_output.sh
create mode 100644 tests/studio/test_setup_pin_stale.ps1
create mode 100644 tests/studio/test_torch_index_pin_hardening.ps1
diff --git a/install.ps1 b/install.ps1
index df49414620..6e059ee0dd 100644
--- a/install.ps1
+++ b/install.ps1
@@ -53,7 +53,8 @@ function Install-UnslothStudio {
param([string]$TorchIndexUrl)
if ($SkipTorch) { return "none" }
if ([string]::IsNullOrWhiteSpace($TorchIndexUrl)) { return "none" }
- $leaf = ($TorchIndexUrl.TrimEnd('/') -split '/')[-1].ToLowerInvariant()
+ # Drop query/fragment first so a token-authenticated pin classifies by family.
+ $leaf = (($TorchIndexUrl -split '[?#]', 2)[0].TrimEnd('/') -split '/')[-1].ToLowerInvariant()
if (@("cpu", "cu118", "cu124", "cu126", "cu128", "cu130") -contains $leaf) { return $leaf }
if ($leaf -match '^rocm[0-9]+\.[0-9]+$') { return $leaf }
return "auto"
@@ -62,7 +63,8 @@ function Install-UnslothStudio {
function Get-TauriGpuBranch {
param([string]$TorchIndexFamily)
if ($SkipTorch) { return "no_torch" }
- if ($TorchIndexFamily -like "cu*") { return "cuda" }
+ # Require a digit after "cu" so /current or /custom isn't branded CUDA (parity ^cu[0-9]).
+ if ($TorchIndexFamily -match '^cu[0-9]') { return "cuda" }
if ($TorchIndexFamily -like "rocm*") { return "rocm" }
if ($TorchIndexFamily -eq "cpu") { return "cpu" }
return "unknown"
@@ -467,22 +469,35 @@ function Install-UnslothStudio {
}
}
+ # Redact index-URL credentials (userinfo + ?query= + #fragment) from captured installer
+ # output before printing on failure; uv/pip errors echo the failing --index-url verbatim.
+ # Mirrors the other installers. Verbose mode streams uncaptured, so it isn't redacted.
+ function Redact-InstallOutput {
+ param([string]$Text)
+ if (-not $Text) { return $Text }
+ $Text = $Text -replace '(https?://)[^/@\s`]+@', '$1@'
+ $Text = $Text -replace '([?&][^=\s&`]+)=[^\s`]+', '$1='
+ # A #token=... fragment is as sensitive as a query; URL-anchored.
+ return $Text -replace '(https?://[^\s`#]+)#[^\s`]+', '$1#'
+ }
+
# Run native commands quietly by default to match install.sh behavior.
# Full command output is shown only when --verbose / UNSLOTH_VERBOSE=1.
function Invoke-InstallCommand {
param(
[Parameter(Mandatory = $true)][ScriptBlock]$Command
)
- # Installer-pinned index installs (torch) must beat an inherited uv mirror
- # (#6898): when the command pins an index, clear every uv index env var so
- # it wins, then restore in finally. Other installs keep the user's mirror.
+ # Installer-pinned index installs (torch) must beat an inherited uv mirror (#6898):
+ # for --default-index, clear the uv index env vars (restore in finally) and set
+ # UV_NO_CONFIG=1 so a uv.toml/pyproject index can't outrank the CLI pin (uv 0.10).
$savedUvIndex = $null
if ($Command.ToString() -match '--default-index') {
$savedUvIndex = @{}
- foreach ($n in 'UV_DEFAULT_INDEX', 'UV_INDEX_URL', 'UV_INDEX', 'UV_EXTRA_INDEX_URL') {
+ foreach ($n in 'UV_DEFAULT_INDEX', 'UV_INDEX_URL', 'UV_INDEX', 'UV_EXTRA_INDEX_URL', 'UV_TORCH_BACKEND', 'UV_FIND_LINKS', 'UV_CONFIG_FILE', 'UV_NO_CONFIG') {
$savedUvIndex[$n] = [Environment]::GetEnvironmentVariable($n)
Remove-Item "Env:$n" -ErrorAction SilentlyContinue
}
+ $env:UV_NO_CONFIG = '1'
}
$prevEap = $ErrorActionPreference
$ErrorActionPreference = "Continue"
@@ -493,17 +508,23 @@ function Install-UnslothStudio {
# Merge stderr into stdout so progress/warning output stays visible
# without flipping $? on successful native commands (PS 5.1 treats
# stderr records as errors that set $? = $false even on exit code 0).
- & $Command 2>&1 | Out-Host
+ # Redact per record: uv echoes index URLs (credentials and all) in
+ # its errors, and verbose mode must not bypass the quiet path's
+ # redaction. ForEach-Object/Out-Host leave $LASTEXITCODE untouched.
+ & $Command 2>&1 | ForEach-Object { Redact-InstallOutput "$_" } | Out-Host
} else {
$output = & $Command 2>&1 | Out-String
if ($LASTEXITCODE -ne 0) {
- Write-Host $output -ForegroundColor Red
+ Write-Host (Redact-InstallOutput $output) -ForegroundColor Red
}
}
return [int]$LASTEXITCODE
} finally {
$ErrorActionPreference = $prevEap
- if ($savedUvIndex) { foreach ($n in $savedUvIndex.Keys) { if ($null -ne $savedUvIndex[$n]) { Set-Item "Env:$n" $savedUvIndex[$n] } } }
+ if ($savedUvIndex) {
+ Remove-Item "Env:UV_NO_CONFIG" -ErrorAction SilentlyContinue
+ foreach ($n in $savedUvIndex.Keys) { if ($null -ne $savedUvIndex[$n]) { Set-Item "Env:$n" $savedUvIndex[$n] } }
+ }
}
}
@@ -1960,10 +1981,31 @@ exit 0
# On an AMD GPU (no NVIDIA), surface the optional WSL-ROCm driver hint.
if (-not $HasNvidiaSmi -and ($ROCmGfxArch -or $ROCmGpuLabel)) { Show-AmdWslDriverHint }
+ # Trim trailing slashes from the URL PATH only, preserving ?query / #fragment: a whole-URL
+ # TrimEnd corrupts a token ending in "/", a single strip leaves .../cu128// empty. Shared.
+ function Trim-IndexPathSlashes {
+ param([string]$Url)
+ $value = $Url.Trim()
+ $idx = $value.IndexOfAny([char[]]@('?', '#'))
+ if ($idx -lt 0) {
+ return $value.TrimEnd('/')
+ }
+ return $value.Substring(0, $idx).TrimEnd('/') + $value.Substring($idx)
+ }
+
# ── Choose the correct PyTorch index URL based on driver CUDA version ──
# Mirrors Get-PytorchCudaTag in setup.ps1.
function Get-TorchIndexUrl {
$baseUrl = if ($env:UNSLOTH_PYTORCH_MIRROR) { $env:UNSLOTH_PYTORCH_MIRROR.TrimEnd('/') } else { "https://download.pytorch.org/whl" }
+ # Explicit pin -- skip ALL GPU probing (headless / CI / cross-install).
+ # UNSLOTH_TORCH_INDEX_URL wins (full URL, verbatim); _FAMILY is the leaf appended
+ # to the mirror base. Matches install.sh / install_python_stack.py.
+ if (-not [string]::IsNullOrWhiteSpace($env:UNSLOTH_TORCH_INDEX_URL)) {
+ return (Trim-IndexPathSlashes $env:UNSLOTH_TORCH_INDEX_URL)
+ }
+ if (-not [string]::IsNullOrWhiteSpace($env:UNSLOTH_TORCH_INDEX_FAMILY)) {
+ return "$baseUrl/$($env:UNSLOTH_TORCH_INDEX_FAMILY.Trim().Trim('/'))"
+ }
if (-not $NvidiaSmiExe) { return "$baseUrl/cpu" }
try {
$output = Invoke-NvidiaSmiBounded $NvidiaSmiExe
@@ -1984,6 +2026,25 @@ exit 0
return "$baseUrl/cu126"
}
+ # Strip userinfo AND query/fragment so an authenticated pin never leaks. Shared with
+ # _strip_index_url_credentials (install.sh / py / setup.ps1).
+ function Remove-IndexUrlCredentials {
+ param([string]$Url)
+ $sep = $Url.IndexOf('://')
+ if ($sep -lt 0) { return $Url }
+ $scheme = $Url.Substring(0, $sep)
+ $rest = $Url.Substring($sep + 3)
+ # Drop query / fragment (may hold auth tokens).
+ $q = $rest.IndexOfAny([char[]]('?', '#'))
+ if ($q -ge 0) { $rest = $rest.Substring(0, $q) }
+ $slash = $rest.IndexOf('/')
+ $authority = if ($slash -ge 0) { $rest.Substring(0, $slash) } else { $rest }
+ $at = $authority.LastIndexOf('@')
+ $host_ = if ($at -ge 0) { $authority.Substring($at + 1) } else { $authority }
+ if ($slash -ge 0) { return "${scheme}://${host_}$($rest.Substring($slash))" }
+ return "${scheme}://${host_}"
+ }
+
# ── Torch flavor helpers (to repair a stale CPU / wrong-CUDA wheel) ──
# torch.__version__ -> flavor tag (cuXXX / rocm / cpu); untagged wheel = cpu,
# matching setup.ps1's stale-venv parse.
@@ -2002,11 +2063,13 @@ exit 0
param([string]$TorchIndexUrl, [string]$ROCmIndexUrl)
if (-not [string]::IsNullOrWhiteSpace($ROCmIndexUrl)) { return 'rocm' }
if ([string]::IsNullOrWhiteSpace($TorchIndexUrl)) { return $null }
- $leaf = ($TorchIndexUrl.TrimEnd('/') -split '/')[-1].ToLowerInvariant()
+ # Drop query/fragment first so .../cu128?token=x classifies as cu128 (else it reinstalls every run).
+ $leaf = (($TorchIndexUrl -split '[?#]', 2)[0].TrimEnd('/') -split '/')[-1].ToLowerInvariant()
if ($leaf -match '^cu\d+$') { return $leaf }
if ($leaf -eq 'cpu') { return 'cpu' }
if ($leaf -match '^rocm') { return 'rocm' }
- if ($leaf -match '^gfx') { return 'rocm' }
+ # gfx must be followed by a digit (an architecture leaf); gfx-private is custom.
+ if ($leaf -match '^gfx[0-9]') { return 'rocm' }
return $null
}
@@ -2041,6 +2104,10 @@ exit 0
} catch { return $null }
}
+ # An explicit pin is authoritative: the AMD ROCm reroute below must not rewrite it
+ # (e.g. a deliberate cpu pin on an AMD host).
+ $TorchIndexPinned = (-not [string]::IsNullOrWhiteSpace($env:UNSLOTH_TORCH_INDEX_URL)) -or `
+ (-not [string]::IsNullOrWhiteSpace($env:UNSLOTH_TORCH_INDEX_FAMILY))
$TorchIndexUrl = Get-TorchIndexUrl
# ── GPU arch → newest compatible Windows ROCm wheel release ──
@@ -2052,7 +2119,9 @@ exit 0
# Override with UNSLOTH_ROCM_WINDOWS_MIRROR for air-gapped / mirror installs.
$ROCmIndexUrl = $null
$ROCmTorchFloor = $null
- if (($HasROCm -or $ROCmGfxArch) -and $TorchIndexUrl -like "*/cpu" -and -not $SkipTorch) {
+ $PinnedRocmVisionSpec = $null
+ $PinnedRocmAudioSpec = $null
+ if (-not $TorchIndexPinned -and ($HasROCm -or $ROCmGfxArch) -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
@@ -2102,6 +2171,32 @@ exit 0
}
}
+ # A gfx*/rocm pin skips the auto-reroute above, but the generic CPU/CUDA install below
+ # would use torch>=2.4,<2.11 and pull a known-bad wheel on the gfx115x/gfx120x/rocm>=7.2
+ # indexes (the _grouped_mm bug). Route a pinned ROCm index through the ROCm path.
+ if ($TorchIndexPinned -and -not $ROCmIndexUrl -and -not $SkipTorch) {
+ $_pinLeaf = (($TorchIndexUrl -split '[?#]', 2)[0].TrimEnd('/') -split '/')[-1].ToLower()
+ $_pinRocm211 = $false
+ # Anchor ($) so a suffixed custom leaf (rocm7.2-private) falls through to verbatim.
+ if ($_pinLeaf -match '^rocm(\d+)\.(\d+)$') {
+ # Only KNOWN-2.11 rocm (rocm7.2) gets the floor. Matches Test-RocmKnown211Version.
+ $_pinRocm211 = ([int]$Matches[1] -eq 7 -and [int]$Matches[2] -eq 2)
+ }
+ # Only the 2.11-allowlist gfx arches need the floor; others publish <2.11 and stay bare.
+ $_pinGfx211 = @('gfx120x-all', 'gfx1151', 'gfx1150') -contains $_pinLeaf
+ if ($_pinGfx211 -or $_pinRocm211) {
+ $ROCmIndexUrl = $TorchIndexUrl
+ $ROCmTorchFloor = "torch>=2.11.0,<2.12.0"
+ $PinnedRocmVisionSpec = "torchvision>=0.26.0,<0.27.0"
+ $PinnedRocmAudioSpec = "torchaudio>=2.11.0,<2.12.0"
+ substep "pinned ROCm index ($_pinLeaf) -- enforcing $ROCmTorchFloor" "Cyan"
+ } elseif ($_pinLeaf -match '^gfx[0-9]' -or $_pinLeaf -match '^rocm[0-9]+(\.[0-9]+)?$') {
+ # Other gfx / older rocm (<=7.1) ship torch <2.11; route via the ROCm path with
+ # bare specs. Only EXACT rocm/gfx* are families; a suffixed leaf is verbatim.
+ $ROCmIndexUrl = $TorchIndexUrl
+ }
+ }
+
if ($ROCmIndexUrl) {
$TorchIndexFamily = "rocm"
} else {
@@ -2164,8 +2259,8 @@ exit 0
}
if ($_Migrated) {
- # Migrated env: force-reinstall unsloth+unsloth-zoo to ensure clean state
- # in the new venv location, while preserving existing torch/CUDA
+ # Migrated env: force-reinstall unsloth+unsloth-zoo for a clean state, preserving
+ # existing torch/CUDA unless the flavor repair below re-lands it.
Write-TauriLog "STEP" "Installing unsloth"
substep "upgrading unsloth in migrated environment..."
if ($SkipTorch) {
@@ -2210,22 +2305,24 @@ exit 0
substep "skipping PyTorch (--no-torch flag set)." "Yellow"
} elseif ($ROCmIndexUrl) {
Write-TauriLog "STEP" "Installing PyTorch (AMD ROCm Windows)"
- substep "installing PyTorch from $ROCmIndexUrl..."
+ substep "installing PyTorch from $(Remove-IndexUrlCredentials $ROCmIndexUrl)..."
$torchSpec = if ($ROCmTorchFloor) { $ROCmTorchFloor } else { "torch" }
# Pin the companions to match $torchSpec; bare names can resolve an
# ABI-incompatible torchvision/torchaudio on AMD's per-arch index.
- $visionSpec = if ($ROCmGfxArch -and $torchvisionFloorMap.ContainsKey($ROCmGfxArch)) { $torchvisionFloorMap[$ROCmGfxArch] } else { "torchvision" }
- $audioSpec = if ($ROCmGfxArch -and $torchaudioFloorMap.ContainsKey($ROCmGfxArch)) { $torchaudioFloorMap[$ROCmGfxArch] } else { "torchaudio" }
+ $visionSpec = if ($PinnedRocmVisionSpec) { $PinnedRocmVisionSpec } elseif ($ROCmGfxArch -and $torchvisionFloorMap -and $torchvisionFloorMap.ContainsKey($ROCmGfxArch)) { $torchvisionFloorMap[$ROCmGfxArch] } else { "torchvision" }
+ $audioSpec = if ($PinnedRocmAudioSpec) { $PinnedRocmAudioSpec } elseif ($ROCmGfxArch -and $torchaudioFloorMap -and $torchaudioFloorMap.ContainsKey($ROCmGfxArch)) { $torchaudioFloorMap[$ROCmGfxArch] } else { "torchaudio" }
$torchInstallExit = Invoke-InstallCommandRetry -Label "install PyTorch (AMD ROCm)" { uv pip install --python $VenvPython --force-reinstall --default-index $ROCmIndexUrl $torchSpec $visionSpec $audioSpec }
if ($torchInstallExit -ne 0) {
- # Transient AMD-index failure: fall back to a CPU base so the install
- # still completes; Unsloth setup retries ROCm afterwards.
+ # Transient AMD-index failure: fall back to a CPU base (Unsloth setup retries
+ # ROCm). Use an explicit CPU index -- for a pinned ROCm index $TorchIndexUrl IS
+ # the ROCm mirror, so reusing it would just retry it.
+ $CpuFallbackIndexUrl = if ($env:UNSLOTH_PYTORCH_MIRROR) { "$($env:UNSLOTH_PYTORCH_MIRROR.TrimEnd('/'))/cpu" } else { "https://download.pytorch.org/whl/cpu" }
substep "ROCm PyTorch install failed (exit $torchInstallExit); using a CPU base, Unsloth setup retries ROCm." "Yellow"
# --force-reinstall: a failed ROCm install can leave an unpinned ROCm
# torch (e.g. 2.10.0+rocm on gfx110X/gfx90a) that still satisfies the CPU
# torch>= range, so without it uv would keep the ROCm build and only swap
# the companions -- a mismatched venv the flavor-repair block won't fix.
- $torchInstallExit = Invoke-InstallCommandRetry -Label "install PyTorch (CPU fallback)" { uv pip install --python $VenvPython --force-reinstall "torch>=2.4,<2.11.0" torchvision torchaudio --default-index $TorchIndexUrl }
+ $torchInstallExit = Invoke-InstallCommandRetry -Label "install PyTorch (CPU fallback)" { uv pip install --python $VenvPython --force-reinstall "torch>=2.4,<2.11.0" "torchvision>=0.19,<0.26.0" "torchaudio>=2.4,<2.11.0" --default-index $CpuFallbackIndexUrl }
if ($torchInstallExit -ne 0) {
Write-Host "[ERROR] Failed to install PyTorch (ROCm and CPU base both failed, exit code $torchInstallExit)" -ForegroundColor Red
return (Exit-InstallFailure "Failed to install PyTorch (exit code $torchInstallExit)" $torchInstallExit)
@@ -2238,8 +2335,14 @@ exit 0
}
} else {
Write-TauriLog "STEP" "Installing PyTorch"
- substep "installing PyTorch ($TorchIndexUrl)..."
- $torchInstallExit = Invoke-InstallCommandRetry -Label "install PyTorch" { uv pip install --python $VenvPython "torch>=2.4,<2.11.0" torchvision torchaudio --default-index $TorchIndexUrl }
+ substep "installing PyTorch ($(Remove-IndexUrlCredentials $TorchIndexUrl))..."
+ # Bound the companions to the capped torch on EVERY index, cu
+ # families included: torchaudio 2.11 dropped its exact torch pin from
+ # the wheel metadata, so a bare companion next to torch<2.11 can
+ # resolve a mismatched 2.11.0 build. Mirrors install.sh.
+ $_pinVisionSpec = "torchvision>=0.19,<0.26.0"
+ $_pinAudioSpec = "torchaudio>=2.4,<2.11.0"
+ $torchInstallExit = Invoke-InstallCommandRetry -Label "install PyTorch" { uv pip install --python $VenvPython "torch>=2.4,<2.11.0" $_pinVisionSpec $_pinAudioSpec --default-index $TorchIndexUrl }
if ($torchInstallExit -ne 0) {
Write-Host "[ERROR] Failed to install PyTorch (exit code $torchInstallExit)" -ForegroundColor Red
return (Exit-InstallFailure "Failed to install PyTorch (exit code $torchInstallExit)" $torchInstallExit)
@@ -2335,8 +2438,8 @@ exit 0
$rocmSpec = if ($ROCmTorchFloor) { $ROCmTorchFloor } else { "torch" }
# Pin companions like the fresh ROCm path (bare names can pull an
# ABI-incompatible torchvision/torchaudio from the per-arch index).
- $visionSpec = if ($ROCmGfxArch -and $torchvisionFloorMap.ContainsKey($ROCmGfxArch)) { $torchvisionFloorMap[$ROCmGfxArch] } else { "torchvision" }
- $audioSpec = if ($ROCmGfxArch -and $torchaudioFloorMap.ContainsKey($ROCmGfxArch)) { $torchaudioFloorMap[$ROCmGfxArch] } else { "torchaudio" }
+ $visionSpec = if ($PinnedRocmVisionSpec) { $PinnedRocmVisionSpec } elseif ($ROCmGfxArch -and $torchvisionFloorMap -and $torchvisionFloorMap.ContainsKey($ROCmGfxArch)) { $torchvisionFloorMap[$ROCmGfxArch] } else { "torchvision" }
+ $audioSpec = if ($PinnedRocmAudioSpec) { $PinnedRocmAudioSpec } elseif ($ROCmGfxArch -and $torchaudioFloorMap -and $torchaudioFloorMap.ContainsKey($ROCmGfxArch)) { $torchaudioFloorMap[$ROCmGfxArch] } else { "torchaudio" }
substep "PyTorch flavor mismatch (installed $installedTorchTag, need ROCm) -- reinstalling correct build..." "Yellow"
$torchFixExit = Invoke-InstallCommand { uv pip install --python $VenvPython --force-reinstall --default-index $ROCmIndexUrl $rocmSpec $visionSpec $audioSpec }
if ($torchFixExit -ne 0) {
@@ -2347,7 +2450,7 @@ exit 0
} elseif ($expectedTorchTag -ne 'rocm') {
# CUDA: stale +cpu (or wrong cuXXX) against a CUDA index -> reinstall triplet.
substep "PyTorch flavor mismatch (installed $installedTorchTag, need $expectedTorchTag) -- reinstalling correct build..." "Yellow"
- $torchFixExit = Invoke-InstallCommand { uv pip install --python $VenvPython "torch>=2.4,<2.11.0" torchvision torchaudio --default-index $TorchIndexUrl --reinstall-package torch --reinstall-package torchvision --reinstall-package torchaudio }
+ $torchFixExit = Invoke-InstallCommand { uv pip install --python $VenvPython "torch>=2.4,<2.11.0" "torchvision>=0.19,<0.26.0" "torchaudio>=2.4,<2.11.0" --default-index $TorchIndexUrl --reinstall-package torch --reinstall-package torchvision --reinstall-package torchaudio }
if ($torchFixExit -ne 0) {
Write-Host "[ERROR] Failed to reinstall PyTorch with the correct CUDA build (exit code $torchFixExit)" -ForegroundColor Red
return (Exit-InstallFailure "Failed to reinstall PyTorch ($expectedTorchTag) (exit code $torchFixExit)" $torchFixExit)
diff --git a/install.sh b/install.sh
index 7918a2bd23..c02552628f 100755
--- a/install.sh
+++ b/install.sh
@@ -159,18 +159,58 @@ run_maybe_quiet() {
fi
}
+# Trim trailing slashes from the URL PATH only, preserving ?query / #fragment: a whole-URL
+# strip corrupts a token ending in "/", a single strip leaves .../cu128// empty. Shared.
+_trim_index_path_slashes() {
+ _tips_v="$1"
+ case "$_tips_v" in
+ *[?#]*)
+ _tips_head="${_tips_v%%[?#]*}"
+ _tips_tail="${_tips_v#"$_tips_head"}"
+ ;;
+ *)
+ _tips_head="$_tips_v"
+ _tips_tail=""
+ ;;
+ esac
+ while [ -n "$_tips_head" ] && [ "${_tips_head%/}" != "$_tips_head" ]; do
+ _tips_head="${_tips_head%/}"
+ done
+ printf '%s%s' "$_tips_head" "$_tips_tail"
+}
+
+# Redact index-URL credentials (userinfo + ?query= + #fragment) from captured installer
+# output before printing on failure; uv/pip errors echo the failing --index-url verbatim.
+# Mirrors the other installers. Verbose mode streams uncaptured, so it isn't redacted.
+_redact_install_output() {
+ sed -E \
+ -e 's#(https?://)[^/@[:space:]`]+@#\1@#g' \
+ -e 's#([?&][^=[:space:]&`]+)=[^[:space:]`]+#\1=#g' \
+ -e 's|(https?://[^[:space:]`#]+)#[^[:space:]`]+|\1#|g' \
+ "$@"
+}
+
run_install_cmd() {
_label="$1"
shift
- # Installer-pinned index installs (torch) must beat an inherited uv mirror
- # (#6898): when we pass --default-index, neutralize every uv index env var so
- # the pinned index wins. Other installs keep the user's mirror.
+ # Installer-pinned index installs (torch) must beat an inherited uv mirror (#6898):
+ # for --default-index, neutralize the uv index/backend/config vars (UV_TORCH_BACKEND
+ # redirects torch; UV_NO_CONFIG=1 + dropping UV_CONFIG_FILE stops a uv.toml/pyproject
+ # index outranking the CLI pin, uv 0.10).
case " $* " in
- *" --default-index "*) set -- env -u UV_DEFAULT_INDEX -u UV_INDEX_URL -u UV_INDEX -u UV_EXTRA_INDEX_URL "$@" ;;
+ *" --default-index "*) set -- env -u UV_DEFAULT_INDEX -u UV_INDEX_URL -u UV_INDEX -u UV_EXTRA_INDEX_URL -u UV_TORCH_BACKEND -u UV_FIND_LINKS -u UV_CONFIG_FILE UV_NO_CONFIG=1 "$@" ;;
esac
if _is_verbose; then
- "$@" && return 0
- _rc=$?
+ # Stream through the redactor: uv echoes index URLs (credentials and
+ # all) in its errors, and verbose mode previously bypassed the
+ # redaction the quiet path applies. The rc file preserves the
+ # command's exit code across the pipe without relying on pipefail
+ # (this script runs under plain sh).
+ _rcf=$(mktemp)
+ { "$@" 2>&1; printf '%s' "$?" > "$_rcf"; } | _redact_install_output
+ _rc=$(cat "$_rcf" 2>/dev/null || echo 1)
+ rm -f "$_rcf"
+ [ "${_rc:-1}" -eq 0 ] 2>/dev/null && return 0
step "error" "$_label failed (exit code $_rc)" "$C_ERR" >&2
return "$_rc"
fi
@@ -178,7 +218,7 @@ run_install_cmd() {
"$@" >"$_log" 2>&1 && { rm -f "$_log"; return 0; }
_rc=$?
step "error" "$_label failed (exit code $_rc)" "$C_ERR" >&2
- cat "$_log" >&2
+ _redact_install_output "$_log" >&2
rm -f "$_log"
return $_rc
}
@@ -257,7 +297,7 @@ _install_bnb_rocm() {
fi
_bnb_rc=$?
if _is_verbose; then
- cat "$_bnb_log" >&2
+ _redact_install_output "$_bnb_log" >&2
fi
rm -f "$_bnb_log"
step "warning" "$_label (pre-release) failed (exit code $_bnb_rc)" "$C_WARN" >&2
@@ -310,6 +350,11 @@ _tauri_torch_index_family() {
return
fi
_diag_url="${1:-}"
+ # Strip query/fragment AND a trailing slash before classifying (like _torch_index_url_leaf):
+ # a token isn't echoed into [TAURI:DIAG], and .../cu128/?token=x still classifies as cu128.
+ _diag_url="${_diag_url%%\?*}"
+ _diag_url="${_diag_url%%#*}"
+ _diag_url="${_diag_url%/}"
case "$_diag_url" in
*/cu118) echo "cu118" ;;
*/cu124) echo "cu124" ;;
@@ -343,7 +388,8 @@ _tauri_gpu_branch() {
return
fi
case "$_diag_family" in
- cu*) echo "cuda" ;;
+ # Require a digit after cu so /current or /custom isn't branded CUDA (parity ^cu[0-9]).
+ cu[0-9]*) echo "cuda" ;;
rocm*)
if [ "$_diag_radeon" = true ]; then
echo "rocm_radeon"
@@ -1575,6 +1621,12 @@ _has_usable_nvidia_gpu() {
# the STUDIO_HOME mkdir/venv so the origin distro is untouched.
_maybe_reroute_strixhalo_to_2404() {
[ "${OS:-}" = "wsl" ] || return 0
+ # An explicit index pin skips every GPU-driven reroute (same contract as
+ # the later Radeon/Strix guard): the pin is honored in THIS distro rather
+ # than probing the GPU and switching distributions. Whitespace-only
+ # overrides do not gate (parity with get_torch_index_url).
+ _rr_pin=$(printf '%s' "${UNSLOTH_TORCH_INDEX_URL:-}${UNSLOTH_TORCH_INDEX_FAMILY:-}" | tr -d '[:space:]')
+ [ -n "$_rr_pin" ] && return 0
[ "${SKIP_TORCH:-false}" = "false" ] || return 0
[ "${UNSLOTH_SKIP_ROCM_WSL_SETUP:-0}" = "1" ] && return 0
[ "${UNSLOTH_WSL_REROUTED:-0}" = "1" ] && return 0
@@ -1636,6 +1688,10 @@ _maybe_reroute_strixhalo_to_2404() {
# Forward explicit ROCm-bootstrap consent (e.g. Tauri) so the child auto-enables the
# GPU instead of falling back to the desktop-app prompt path.
[ "${UNSLOTH_ROCM_WSL_AUTO:-0}" = "1" ] && _rr_exports="$_rr_exports; export UNSLOTH_ROCM_WSL_AUTO=1"
+ # Forward a pinned torch index into the rerouted distro; dropping it would
+ # silently revert the child install to auto-detection.
+ [ -n "${UNSLOTH_TORCH_INDEX_URL:-}" ] && _rr_exports="$_rr_exports; export UNSLOTH_TORCH_INDEX_URL=$(_rr_q "$UNSLOTH_TORCH_INDEX_URL")"
+ [ -n "${UNSLOTH_TORCH_INDEX_FAMILY:-}" ] && _rr_exports="$_rr_exports; export UNSLOTH_TORCH_INDEX_FAMILY=$(_rr_q "$UNSLOTH_TORCH_INDEX_FAMILY")"
[ "$_SKIP_AUTOSTART" = true ] && _rr_exports="$_rr_exports; export UNSLOTH_SKIP_AUTOSTART=1"
_rr_args=""
[ "$PACKAGE_NAME" != "unsloth" ] && _rr_args="$_rr_args --package $(_rr_q "$PACKAGE_NAME")"
@@ -2001,6 +2057,15 @@ if [ "$SKIP_TORCH" = false ] && [ "$OS" = "macos" ] && [ "$_ARCH" = "arm64" ]; t
TORCH_CONSTRAINT="torch>=2.6,<2.11.0"
fi
fi
+# Companion (torchvision/torchaudio) constraints, bounded to torch's window.
+# torchaudio 2.11 dropped its exact torch pin, so a bare companion next to a
+# <2.11-capped torch resolves torchaudio 2.11 (verified: cpu leaf installed
+# torch 2.10.0+cpu with torchaudio 2.11.0+cpu). torchvision still exact-pins
+# torch and self-corrects, but is bounded for symmetry. Widened alongside the
+# cu* torch window below; the torch-2.11 AMD paths (rocm7.2 / per-gfx / Strix)
+# pin their own trio.
+TORCHVISION_CONSTRAINT="torchvision>=0.19,<0.26.0"
+TORCHAUDIO_CONSTRAINT="torchaudio>=2.4,<2.11.0"
# ── Resolve repo root (for --local installs) ──
_REPO_ROOT="$(cd "$(dirname "$0" 2>/dev/null || echo ".")" && pwd)"
@@ -2069,6 +2134,24 @@ _has_amd_rocm_gpu() {
get_torch_index_url() {
_base="${UNSLOTH_PYTORCH_MIRROR:-https://download.pytorch.org/whl}"
_base="${_base%/}"
+ # Explicit override -- skip ALL GPU probing (headless / container / CI / cross-install).
+ # UNSLOTH_TORCH_INDEX_URL wins (full URL, verbatim); _FAMILY is the leaf (cpu, cu128, ...)
+ # appended to the mirror base. Trim whitespace so a whitespace-only value is unset.
+ _url="${UNSLOTH_TORCH_INDEX_URL:-}"
+ _url="${_url#"${_url%%[![:space:]]*}"}"; _url="${_url%"${_url##*[![:space:]]}"}"
+ if [ -n "$_url" ]; then
+ # Trim trailing PATH slashes (a multi-slash path 404s on strict pip proxies) while
+ # preserving a ?query/#fragment token (a whole-URL strip would eat a "/"-ending token).
+ _url=$(_trim_index_path_slashes "$_url")
+ echo "$_url"; return
+ fi
+ _family="${UNSLOTH_TORCH_INDEX_FAMILY:-}"
+ _family="${_family#"${_family%%[![:space:]]*}"}"; _family="${_family%"${_family##*[![:space:]]}"}"
+ if [ -n "$_family" ]; then
+ 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
# Try nvidia-smi -- require the binary to actually list a usable GPU.
@@ -2197,6 +2280,45 @@ _torch_flavor_tag() {
esac
}
+# Final path segment of a wheel index URL ($1), lowercased, query/fragment stripped first
+# so a token-authenticated pin (.../cu128?token=x) classifies as cu128 (else it reinstalls
+# every update). Classification only. Shared with the py / ps1 leaf extractors.
+_torch_index_url_leaf() {
+ _tl_u="${1%%\?*}"
+ _tl_u="${_tl_u%%#*}"
+ # Strip ALL trailing slashes, not one: .../rocm7.2// must yield rocm7.2, not an empty leaf.
+ while [ -n "$_tl_u" ] && [ "${_tl_u%/}" != "$_tl_u" ]; do
+ _tl_u="${_tl_u%/}"
+ done
+ printf '%s' "${_tl_u##*/}" | tr '[:upper:]' '[:lower:]'
+}
+
+# True (exit 0) when a lowercased leaf is an EXACT pip ROCm family: rocm[.]
+# or a gfx ARCHITECTURE leaf (gfx followed by a digit: gfx90a, gfx1151, gfx120x-all). A leaf
+# that merely starts with rocm/gfx (rocm7.2-private, gfx-private) is a custom verbatim pin.
+# Matches the py / ps1 sides.
+_is_pip_rocm_family_leaf() {
+ case "$1" in
+ gfx[0-9]*) return 0 ;;
+ rocm[0-9]*)
+ # Exact rocm[.]: both major and minor must be non-empty all-digits
+ # (rocm7., rocm7.2.1, rocm7.2-private are all custom pins, not a family).
+ _rocm_rest="${1#rocm}"
+ case "$_rocm_rest" in
+ *.*.*) return 1 ;;
+ *.*)
+ _rocm_minor="${_rocm_rest#*.}"
+ case "${_rocm_rest%%.*}" in "" | *[!0-9]*) return 1 ;; esac
+ case "$_rocm_minor" in "" | *[!0-9]*) return 1 ;; esac
+ ;;
+ *[!0-9]*) return 1 ;;
+ esac
+ return 0
+ ;;
+ *) return 1 ;;
+ esac
+}
+
# Whether release base $1 (X.Y[.Z...]) falls inside constraint window $2
# ("torch>=A.B[.C],
# rocm). Empty on an unknown leaf (odd mirror) so the repair safely no-ops.
_expected_torch_flavor_tag() {
- _u="${1%/}"
- _leaf="${_u##*/}"
+ _leaf=$(_torch_index_url_leaf "$1")
case "$_leaf" in
- cu[0-9]*) echo "$_leaf" ;;
- cpu) echo "cpu" ;;
- rocm*|gfx*) echo "rocm" ;;
- *) echo "" ;;
+ cu[0-9]*)
+ # Exact cu + digits only; a cu*-suffixed leaf (cu128-private) -> "" (custom),
+ # else a correct +cu128 wheel is force-reinstalled every run.
+ case "${_leaf#cu}" in
+ *[!0-9]*) echo "" ;;
+ *) echo "$_leaf" ;;
+ esac
+ ;;
+ cpu) echo "cpu" ;;
+ # Exact rocm/gfx families only; a custom rocm*-suffixed leaf -> "" (custom).
+ *)
+ if _is_pip_rocm_family_leaf "$_leaf"; then echo "rocm"; else echo ""; fi
+ ;;
esac
}
@@ -2308,14 +2438,42 @@ _expected_torch_flavor_tag() {
# fresh-install paths above already use -- so a stale wheel is auto-repairable.
# Unknown/odd-mirror leaves -> no, so we warn rather than risk a wrong reinstall.
_torch_index_repairable() {
- _u="${1%/}"
- _leaf="${_u##*/}"
+ _leaf=$(_torch_index_url_leaf "$1")
case "$_leaf" in
- cu[0-9]*|rocm[0-9]*|gfx*) echo "yes" ;;
- *) echo "no" ;;
+ cu[0-9]*) echo "yes" ;;
+ # Only EXACT rocm/gfx families resolve via --default-index; a suffixed leaf is verbatim.
+ *)
+ if _is_pip_rocm_family_leaf "$_leaf"; then echo "yes"; else echo "no"; fi
+ ;;
esac
}
+# Remove credentials from a wheel index URL ($1) so an authenticated pin never leaks:
+# drops userinfo AND query/fragment; scheme/host/path stay exact. Shared with py / ps1.
+_strip_index_url_credentials() {
+ _sic_url="$1"
+ case "$_sic_url" in
+ *://*) ;;
+ *) printf '%s' "$_sic_url"; return ;;
+ esac
+ _sic_scheme="${_sic_url%%://*}"
+ _sic_rest="${_sic_url#*://}"
+ # Drop query / fragment (may hold auth tokens).
+ _sic_rest="${_sic_rest%%\?*}"
+ _sic_rest="${_sic_rest%%#*}"
+ _sic_auth="${_sic_rest%%/*}"
+ # Drop user:pass@ userinfo if present.
+ case "$_sic_auth" in
+ *@*) _sic_host="${_sic_auth##*@}" ;;
+ *) _sic_host="$_sic_auth" ;;
+ esac
+ if [ "$_sic_auth" = "$_sic_rest" ]; then
+ printf '%s://%s' "$_sic_scheme" "$_sic_host"
+ else
+ printf '%s://%s/%s' "$_sic_scheme" "$_sic_host" "${_sic_rest#*/}"
+ fi
+}
+
get_radeon_wheel_url() {
# Only meaningful on Linux. Picks a repo.radeon.com base URL whose listing
# contains torch wheels. Tries paths like rocm-rel-7.2.1/, rocm-rel-7.2/,
@@ -2561,7 +2719,19 @@ _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: skip the WSL ROCm bootstrap and the Radeon/Strix reroute below (which would
+# re-probe the GPU and overwrite the pin). Trim whitespace first (parity with
+# get_torch_index_url): a whitespace-only override is unset there, so must not flip this true.
+_torch_index_pinned=false
+_ti_url_trim="${UNSLOTH_TORCH_INDEX_URL:-}"
+_ti_url_trim="${_ti_url_trim#"${_ti_url_trim%%[![:space:]]*}"}"; _ti_url_trim="${_ti_url_trim%"${_ti_url_trim##*[![:space:]]}"}"
+_ti_family_trim="${UNSLOTH_TORCH_INDEX_FAMILY:-}"
+_ti_family_trim="${_ti_family_trim#"${_ti_family_trim%%[![:space:]]*}"}"; _ti_family_trim="${_ti_family_trim%"${_ti_family_trim##*[![:space:]]}"}"
+if [ -n "$_ti_url_trim" ] || [ -n "$_ti_family_trim" ]; then
+ _torch_index_pinned=true
+fi
+[ "$_torch_index_pinned" = true ] || _maybe_bootstrap_rocm_wsl || true
TORCH_INDEX_URL=$(get_torch_index_url)
@@ -2572,29 +2742,74 @@ TORCH_INDEX_URL=$(get_torch_index_url)
# whose base path happens to contain "rocm" or "gfx" must not mislabel a
# cu*/cpu index as ROCm (radeon repo URLs end in rocm-rel-X.Y/, Strix
# overrides in gfxNNNN/, so the trailing slash is stripped first).
-_torch_index_leaf="${TORCH_INDEX_URL%/}"
+# Lowercase the leaf so every gfx*/rocm*/cu* arm matches regardless of case (canonical AMD
+# RDNA4 leaf is gfx120X-all). CUDA is branded only on a real cu[0-9]* leaf, so a mirror
+# leaf (/current) does NOT commit a CUDA backend; an unknown leaf leaves the var unset so
+# the stack probes the GPU. Query/fragment dropped first, then ALL trailing slashes (in
+# lockstep with the shared _torch_index_url_leaf extractor).
+_torch_index_leaf="${TORCH_INDEX_URL%%\?*}"
+_torch_index_leaf="${_torch_index_leaf%%#*}"
+# Strip ALL trailing slashes, not one: .../cu128// must yield cu128, not an empty leaf.
+while [ -n "$_torch_index_leaf" ] && [ "${_torch_index_leaf%/}" != "$_torch_index_leaf" ]; do
+ _torch_index_leaf="${_torch_index_leaf%/}"
+done
_torch_index_leaf="${_torch_index_leaf##*/}"
+_torch_index_leaf=$(printf '%s' "$_torch_index_leaf" | tr '[:upper:]' '[:lower:]')
case "$_torch_index_leaf" in
rocm*|gfx*) export UNSLOTH_TORCH_BACKEND="rocm" ;;
cpu) export UNSLOTH_TORCH_BACKEND="cpu" ;;
- *) export UNSLOTH_TORCH_BACKEND="cuda" ;;
+ cu[0-9]*) export UNSLOTH_TORCH_BACKEND="cuda" ;;
+ # Unknown leaf (odd mirror, /current): unset so a stale inherited value can't leak and
+ # the stack probes the GPU.
+ *) unset UNSLOTH_TORCH_BACKEND ;;
esac
-# rocm7.2 and the CUDA cu12x/cu13x indexes now ship torch 2.11.x, so widen the
-# ceiling to <2.12.0 (matches the base image and _CUDA_TORCH_PKG_SPEC in
-# studio/install_python_stack.py). Keep the >=2.4 floor so an older CUDA index
-# (e.g. cu118) still resolves. Match on _torch_index_leaf, not the full URL, so
-# a mirror whose base path contains cu*/rocm7.2 but resolves to a cpu/older-rocm
-# leaf keeps the default <2.11.0.
+# Whether TORCH_INDEX_URL names an actual pip ROCm family (rocm* / gfx*), gating the
+# ROCm-only side effects below (AMD bitsandbytes, ROCm-torch repair). Digit-gated so a leaf
+# merely STARTING with "rocm" isn't force-repaired from the wrong path.
+if _is_pip_rocm_family_leaf "$_torch_index_leaf"; then
+ _torch_index_is_rocm_family=true
+else
+ _torch_index_is_rocm_family=false
+fi
+
+# rocm7.2 and the per-gfx indexes with the _grouped_mm <2.11 bug (gfx120X-all, gfx1151,
+# gfx1150) ship torch 2.11.0 -- raise the floor (also covers a pinned override that skipped
+# the Strix reroute). Pin the companions too: the per-gfx index publishes them independently
+# and a bare name can resolve a 2.12 ABI-mismatched wheel. Match on the FINAL leaf so a
+# custom mirror with a gfx/rocm7.2 path segment but a cu*/cpu family isn't forced.
case "$_torch_index_leaf" in
- rocm7.2) TORCH_CONSTRAINT="torch>=2.11.0,<2.12.0" ;;
- cu[0-9]*) TORCH_CONSTRAINT="torch>=2.4,<2.12.0" ;;
+ rocm7.2|gfx120x-all|gfx1151|gfx1150)
+ TORCH_CONSTRAINT="torch>=2.11.0,<2.12.0"
+ TORCHVISION_CONSTRAINT="torchvision>=0.26.0,<0.27.0"
+ TORCHAUDIO_CONSTRAINT="torchaudio>=2.11.0,<2.12.0"
+ ;;
+ # CUDA cu12x/cu13x indexes ship torch 2.11.x: widen the ceiling to <2.12.0 (matches
+ # _CUDA_TORCH_PKG_SPEC) and widen the companions with it so the trio stays paired.
+ cu[0-9]*)
+ TORCH_CONSTRAINT="torch>=2.4,<2.12.0"
+ TORCHVISION_CONSTRAINT="torchvision>=0.19,<0.27.0"
+ TORCHAUDIO_CONSTRAINT="torchaudio>=2.4,<2.12.0"
+ ;;
esac
+# A pinned custom/unknown-leaf index (/simple, /current, /cu128-private) has no curated
+# companion set, so bound torchvision/torchaudio to the same <2.11 range the Python path pins
+# (else a mirror with newer companions resolves a 2.12 ABI-mismatched wheel). Known families
+# keep their curated companions above (_expected_torch_flavor_tag returns "" only for custom).
+if [ "$_torch_index_pinned" = true ] && \
+ [ -z "$(_expected_torch_flavor_tag "$TORCH_INDEX_URL")" ]; then
+ TORCHVISION_CONSTRAINT="torchvision>=0.19,<0.26.0"
+ TORCHAUDIO_CONSTRAINT="torchaudio>=2.4,<2.11.0"
+fi
+
# 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 when the index is pinned: an explicit override 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 && \
@@ -2671,10 +2886,14 @@ case "$TORCH_INDEX_URL" in
done
TORCH_INDEX_URL="${_amd_strix_base}/${_strix_gfx}/"
TORCH_CONSTRAINT="torch>=2.11.0,<2.12.0"
+ # Pin companions to 2.11 (per-gfx index publishes them independently).
+ TORCHVISION_CONSTRAINT="torchvision>=0.26.0,<0.27.0"
+ TORCHAUDIO_CONSTRAINT="torchaudio>=2.11.0,<2.12.0"
_amd_gpu_radeon=false
fi
;;
esac
+fi # _torch_index_pinned guard (Radeon + Strix reroute)
# Re-run over an existing install: keep the previous venv's torch RELEASE; the fresh
# index above supplies the right flavor for this machine. Evaluated HERE, after every
# index/constraint decision including the Strix reroute, so the window checked is the
@@ -2821,7 +3040,7 @@ case "$TORCH_INDEX_URL" in
if [ "$_amd_gpu_radeon" = true ]; then
substep "wheels: repo.radeon.com (Radeon)"
else
- substep "wheels: $TORCH_INDEX_URL"
+ substep "wheels: $(_strip_index_url_credentials "$TORCH_INDEX_URL")"
fi
;;
esac
@@ -2867,8 +3086,8 @@ for _p in ('torch', 'torchvision', 'torchaudio'):
}
if [ "$_MIGRATED" = true ]; then
- # Migrated env: force-reinstall unsloth+unsloth-zoo to ensure clean state
- # in the new venv location, while preserving existing torch/CUDA
+ # Migrated env: force-reinstall unsloth+unsloth-zoo for a clean state, preserving
+ # existing torch/CUDA unless the ROCm repair below fires.
substep "upgrading unsloth in migrated environment..."
if [ "$SKIP_TORCH" = true ]; then
# No-torch: install unsloth + unsloth-zoo with --no-deps (current
@@ -2909,18 +3128,14 @@ if [ "$_MIGRATED" = true ]; then
# AMD ROCm: install bitsandbytes even in migrated environments so
# existing ROCm installs gain the AMD bitsandbytes build without a
# fresh reinstall.
- if [ "$SKIP_TORCH" = false ]; then
- case "$TORCH_INDEX_URL" in
- */rocm*|*/gfx*)
- _install_bnb_rocm "install bitsandbytes (AMD)" "$_VENV_PY"
- # Repair ROCm torch if overwritten during migrated install
- _has_hip=$("$_VENV_PY" -c "import torch; print(getattr(torch.version,'hip','') or '')" 2>/dev/null || true)
- if [ -z "$_has_hip" ]; then
- substep "repairing ROCm torch (overwritten by dependency resolution)..."
- _install_torch_default_index --force-reinstall
- fi
- ;;
- esac
+ if [ "$SKIP_TORCH" = false ] && [ "$_torch_index_is_rocm_family" = true ]; then
+ _install_bnb_rocm "install bitsandbytes (AMD)" "$_VENV_PY"
+ # Repair ROCm torch if overwritten during migrated install
+ _has_hip=$("$_VENV_PY" -c "import torch; print(getattr(torch.version,'hip','') or '')" 2>/dev/null || true)
+ if [ -z "$_has_hip" ]; then
+ substep "repairing ROCm torch (overwritten by dependency resolution)..."
+ _install_torch_default_index --force-reinstall
+ fi
fi
elif [ -n "$TORCH_INDEX_URL" ]; then
# Fresh: Step 1 - install torch from explicit index (skip when --no-torch or Intel Mac)
@@ -3074,7 +3289,7 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
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"
+ substep "[WARN] Radeon repo lacks a compatible wheel set for this Python; falling back to ROCm index ($(_strip_index_url_credentials "$TORCH_INDEX_URL"))" "$C_WARN"
_install_torch_default_index
else
substep "installing PyTorch from Radeon repo (${_RADEON_BASE_URL})..."
@@ -3095,7 +3310,7 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
fi
fi
else
- substep "[WARN] Radeon repo unavailable; falling back to ROCm index ($TORCH_INDEX_URL)" "$C_WARN"
+ substep "[WARN] Radeon repo unavailable; falling back to ROCm index ($(_strip_index_url_credentials "$TORCH_INDEX_URL"))" "$C_WARN"
_install_torch_default_index
fi
else
@@ -3103,19 +3318,15 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
_install_torch_default_index
fi
else
- substep "installing PyTorch ($TORCH_INDEX_URL)..."
+ substep "installing PyTorch ($(_strip_index_url_credentials "$TORCH_INDEX_URL"))..."
_install_torch_default_index
fi
# AMD ROCm: install bitsandbytes (once, after torch, for all ROCm paths).
# Gate on SKIP_TORCH=false so a user running with --no-torch on a ROCm
# host stays in GGUF-only mode rather than pulling in bitsandbytes,
# which is only useful once torch is present for training.
- if [ "$SKIP_TORCH" = false ]; then
- case "$TORCH_INDEX_URL" in
- */rocm*|*/gfx*)
- _install_bnb_rocm "install bitsandbytes (AMD)" "$_VENV_PY"
- ;;
- esac
+ if [ "$SKIP_TORCH" = false ] && [ "$_torch_index_is_rocm_family" = true ]; then
+ _install_bnb_rocm "install bitsandbytes (AMD)" "$_VENV_PY"
fi
# Fresh: Step 2 - install unsloth, preserving the torch Step 1 installed
tauri_log "STEP" "Installing Unsloth"
@@ -3161,16 +3372,12 @@ elif [ -n "$TORCH_INDEX_URL" ]; then
_UNSLOTH_TORCH_OVERRIDES=""
# AMD ROCm: repair torch if the unsloth/unsloth-zoo install pulled in
# CUDA torch from PyPI, overwriting the ROCm wheels installed in Step 1.
- if [ "$SKIP_TORCH" = false ]; then
- case "$TORCH_INDEX_URL" in
- */rocm*|*/gfx*)
- _has_hip=$("$_VENV_PY" -c "import torch; print(getattr(torch.version,'hip','') or '')" 2>/dev/null || true)
- if [ -z "$_has_hip" ]; then
- substep "repairing ROCm torch (overwritten by dependency resolution)..."
- _install_torch_default_index --force-reinstall
- fi
- ;;
- esac
+ if [ "$SKIP_TORCH" = false ] && [ "$_torch_index_is_rocm_family" = true ]; then
+ _has_hip=$("$_VENV_PY" -c "import torch; print(getattr(torch.version,'hip','') or '')" 2>/dev/null || true)
+ if [ -z "$_has_hip" ]; then
+ substep "repairing ROCm torch (overwritten by dependency resolution)..."
+ _install_torch_default_index --force-reinstall
+ fi
fi
else
# Fallback: GPU detection failed to produce a URL -- let uv resolve torch
@@ -3217,7 +3424,7 @@ if [ "$SKIP_TORCH" = false ] && [ -n "${TORCH_INDEX_URL:-}" ]; then
substep "[WARN] PyTorch is CPU-only but a $_expected_torch_tag GPU build was expected for this machine." "$C_WARN"
substep "[WARN] Training and GPU inference will run on CPU until this is fixed." "$C_WARN"
substep "[WARN] Re-run this installer, or reinstall the GPU build manually:" "$C_WARN"
- substep "[WARN] uv pip install --python \"$_VENV_PY\" \"$TORCH_CONSTRAINT\" torchvision torchaudio --default-index $TORCH_INDEX_URL --reinstall-package torch --reinstall-package torchvision --reinstall-package torchaudio" "$C_WARN"
+ substep "[WARN] uv pip install --python \"$_VENV_PY\" \"$TORCH_CONSTRAINT\" \"$TORCHVISION_CONSTRAINT\" \"$TORCHAUDIO_CONSTRAINT\" --default-index $(_strip_index_url_credentials "$TORCH_INDEX_URL") --reinstall-package torch --reinstall-package torchvision --reinstall-package torchaudio" "$C_WARN"
fi
fi
fi
diff --git a/studio/install_python_stack.py b/studio/install_python_stack.py
index 95c9356d4a..9921b83543 100644
--- a/studio/install_python_stack.py
+++ b/studio/install_python_stack.py
@@ -44,11 +44,10 @@ IS_MAC_INTEL = IS_MACOS and platform.machine() == "x86_64"
IS_MAC_ARM = IS_MACOS and platform.machine() == "arm64"
IS_LINUX = sys.platform.startswith("linux")
-# DiskPart-prompt suppression: amd-smi auto-elevates on Windows, popping a
-# UAC/DiskPart prompt mid-install. This installer only spawns probes and pip/uv
-# (none need elevation), so set __COMPAT_LAYER=RunAsInvoker process-wide -- every
-# amd-smi subprocess then runs un-elevated, no per-call guard needed. setup.ps1
-# keeps per-call guards since it ALSO spawns winget installers that need elevation.
+# amd-smi auto-elevates on Windows (UAC/DiskPart prompt mid-install). This installer
+# only spawns probes and pip/uv (no elevation), so set __COMPAT_LAYER=RunAsInvoker
+# process-wide; amd-smi then runs un-elevated. setup.ps1 keeps per-call guards (it
+# also spawns winget installers that need elevation).
if IS_WINDOWS:
os.environ.setdefault("__COMPAT_LAYER", "RunAsInvoker")
# torchcodec ships wheels only for manylinux_2_28_x86_64, macosx_12_0_arm64,
@@ -74,6 +73,14 @@ _ROCM_TORCH_INDEX: dict[tuple[int, int], str] = {
(6, 0): "rocm6.0",
}
+# AMD per-arch leaves needing the torch 2.11 floor (the _grouped_mm <2.11 bug).
+# Mirrors *FloorMap in install.ps1 / setup.ps1; other arches ship <2.11 and stay bare.
+_ROCM_GFX_TORCH211_LEAVES: frozenset[str] = frozenset({"gfx120x-all", "gfx1151", "gfx1150"})
+
+# pytorch.org rocmX.Y indexes KNOWN to ship torch 2.11 (rocm7.2 only today); don't
+# floor an unknown newer rocm speculatively. Match install.sh / setup.ps1 / install.ps1.
+_ROCM_KNOWN_TORCH211_VERSIONS: frozenset[tuple[int, int]] = frozenset({(7, 2)})
+
# 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": (
@@ -81,18 +88,16 @@ _ROCM_TORCH_PKG_SPECS: dict[str, tuple[str, str, str]] = {
"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
+ # 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",
),
}
-# Windows AMD per-arch companion pins for the repo.amd.com index, mirroring the
-# install.ps1 / setup.ps1 floor maps (gfx120X and Strix Halo/Point use the rocm7.2
-# torch 2.11 trio). Pinning the companions keeps AMD's per-arch index -- which
-# publishes each independently -- from resolving an ABI-mismatched one. Unlisted
-# arches have no published floor, so stay bare. Bump with the PS maps at 2.12.x.
+# Windows AMD per-arch companion pins for the repo.amd.com index (mirrors the install.ps1 /
+# setup.ps1 floor maps): pinning stops the per-arch index (each published independently) from
+# resolving an ABI-mismatched companion. Unlisted arches have no floor, so stay bare.
_WINDOWS_ROCM_TORCH_PKG_SPECS: dict[str, tuple[str, str, str]] = {
"gfx1201": _ROCM_TORCH_PKG_SPECS["rocm7.2"],
"gfx1200": _ROCM_TORCH_PKG_SPECS["rocm7.2"],
@@ -103,19 +108,79 @@ _PYTORCH_WHL_BASE = (
os.environ.get("UNSLOTH_PYTORCH_MIRROR") or "https://download.pytorch.org/whl"
).rstrip("/")
-# CUDA torch repair specs (see _ensure_cuda_torch). torch 2.11 is allowed: its
-# torchao 0.17 cpp kernels load cleanly (0.16 crashes on cu130), and the flash-attn
-# / causal-conv1d / mamba torch2.10 wheels load and pass their upstream suites on
-# 2.11 (see wheel_utils._PREBUILT_WHEEL_TORCH_MM). torchvision/torchaudio are pinned
-# (not bare) because the install uses an exclusive --index-url (no PyPI fallback), so
-# a bare name could resolve one built against a different torch major (e.g. 0.27 for
-# torch 2.12) and fail at runtime with an ABI mismatch.
+
+def _strip_index_url_credentials(url: str) -> str:
+ """Strip userinfo (user:password@) AND query/fragment from a wheel index URL.
+
+ An authenticated pin must not leak credentials in printed output; query/fragment
+ may hold tokens and aren't part of the PEP 503 index identity. Host/path stay
+ exact. MUST match install.sh / setup.ps1 / install.ps1.
+ """
+ scheme, sep, rest = url.partition("://")
+ if not sep:
+ return url
+ rest = rest.split("?", 1)[0].split("#", 1)[0] # drop query / fragment
+ authority, slash, tail = rest.partition("/")
+ host = authority.rpartition("@")[2] # drop user:pass@ userinfo
+ return f"{scheme}://{host}{slash}{tail}"
+
+
+_URL_USERINFO_RE = re.compile(r"(https?://)[^/@\s`]+@")
+_URL_QUERY_VALUE_RE = re.compile(r"([?&][^=\s&`]+)=[^\s`]+")
+# URL-anchored so a bare "#..." (a shell comment in tool output) is never touched.
+_URL_FRAGMENT_RE = re.compile(r"(https?://[^\s`#]+)#[^\s`]+")
+
+
+def _redact_install_output(output: "bytes | str") -> str:
+ """Redact index-URL credentials (userinfo + query values + fragments) from captured
+ installer output before printing. uv/pip failure text embeds the failing --index-url
+ verbatim, which can carry a user:token@, ?token= or #token= secret. MUST match
+ install.sh / setup.ps1 / install.ps1's output sanitizers."""
+ text = output.decode(errors = "replace") if isinstance(output, bytes) else output
+ text = _URL_USERINFO_RE.sub(r"\1@", text)
+ text = _URL_QUERY_VALUE_RE.sub(r"\1=", text)
+ return _URL_FRAGMENT_RE.sub(r"\1#", text)
+
+
+def _trim_index_path_slashes(url: str) -> str:
+ """Trim trailing slashes from the URL PATH only, preserving ?query / #fragment. A
+ whole-URL rstrip("/") corrupts a token that ends in "/" (e.g. base64 ...abc/) and a
+ single-slash strip leaves .../cu128// classifying as an empty leaf. MUST match
+ install.sh / setup.ps1 / install.ps1."""
+ value = url.strip()
+ match = re.fullmatch(r"([^?#]*)([?#].*)?", value)
+ if match is None:
+ return value.rstrip("/")
+ return match.group(1).rstrip("/") + (match.group(2) or "")
+
+
+def _torch_index_leaf(url: str) -> str:
+ """Final URL path segment, lowercased, query/fragment removed first.
+
+ So a token-authenticated pin (.../cu128?token=x) classifies as cu128 (a raw leaf
+ keeps the query, never equals the +cu128 tag, and force-reinstalls every update).
+ CLASSIFICATION only; the install keeps the full URL. MUST match install.sh /
+ setup.ps1 / install.ps1.
+ """
+ path = url.split("?", 1)[0].split("#", 1)[0]
+ return path.rstrip("/").rsplit("/", 1)[-1].lower()
+
+
+# CUDA torch repair specs (see _ensure_cuda_torch). torch 2.11 is allowed (torchao
+# 0.17 cpp loads cleanly, and the flash-attn/causal-conv1d/mamba wheels pass on 2.11).
+# torchvision/torchaudio are pinned (not bare) so the exclusive --index-url can't
+# resolve one built against a different torch major -> ABI mismatch.
_CUDA_TORCH_PKG_SPEC: tuple[str, str, str] = (
"torch>=2.4,<2.12.0",
"torchvision>=0.19,<0.27.0",
"torchaudio>=2.4,<2.12.0",
)
+# CPU torch repair specs (see _ensure_cpu_torch). Same bounds/reasoning as CUDA: the
+# /cpu index also serves newer torch, so a bare trio could resolve out of range or ABI-
+# mismatched.
+_CPU_TORCH_PKG_SPEC: tuple[str, str, str] = _CUDA_TORCH_PKG_SPEC
+
# torchao's cpp extensions are pinned to ONE torch release AND CUDA major. A torch
# mismatch just skips the cpp kernels (slow Python fallback); a CUDA mismatch fails
# to import ("libcudart.so.12: cannot open shared object file"). The torch pin is a
@@ -408,9 +473,8 @@ def _detect_rocm_version() -> tuple[int, int] | None:
try:
with open(path) as fh:
parts = fh.read().strip().split("-")[0].split(".")
- # Explicit length guard so we don't rely on the broad except
- # below to swallow IndexError when the version file has a
- # single component (e.g. "6\n" on a partial install).
+ # Explicit length guard: don't rely on the broad except below to
+ # swallow IndexError on a single-component version (e.g. "6\n").
if len(parts) >= 2:
return int(parts[0]), int(parts[1])
except Exception:
@@ -455,11 +519,10 @@ def _detect_rocm_version() -> tuple[int, int] | None:
except Exception:
pass
- # Distro package-manager fallbacks. Package-managed ROCm installs can
- # expose GPUs via rocminfo/amd-smi but lack /opt/rocm/.info/version and
- # hipconfig, so probe dpkg (Debian/Ubuntu) and rpm (RHEL/Fedora/SUSE)
- # for the rocm-core version. Matches install.sh::get_torch_index_url so
- # `unsloth studio update` behaves like a fresh `curl | sh` install.
+ # Distro package-manager fallbacks: package-managed ROCm can expose GPUs via
+ # rocminfo/amd-smi but lack /opt/rocm/.info/version and hipconfig, so probe
+ # dpkg (Debian/Ubuntu) and rpm (RHEL/Fedora/SUSE) for the rocm-core version.
+ # Matches install.sh::get_torch_index_url so `studio update` == fresh install.
for cmd in (
["dpkg-query", "-W", "-f=${Version}\n", "rocm-core"],
["rpm", "-q", "--qf", "%{VERSION}\n", "rocm-core"],
@@ -561,11 +624,10 @@ def _detect_windows_gfx_arch() -> str | None:
stderr = subprocess.DEVNULL,
timeout = 10,
)
- # Accept partial output even when hipinfo crashes (e.g. exit code
- # 0xC0000005 / STATUS_ACCESS_VIOLATION on some RDNA 4 hosts): if
- # gcnArchName is present in stdout the device was enumerated before
- # the crash, so the arch is trustworthy. Ignoring it causes a
- # silent CPU PyTorch fallback (issue #6043).
+ # Accept partial output even when hipinfo crashes (e.g. 0xC0000005 /
+ # STATUS_ACCESS_VIOLATION on some RDNA 4 hosts): a gcnArchName in stdout
+ # means the device was enumerated pre-crash, so the arch is trustworthy.
+ # Ignoring it causes a silent CPU PyTorch fallback (issue #6043).
text = result.stdout.decode(errors = "replace")
# findall gets every gcnArchName line so multi-GPU hosts are
# enumerable and HIP_VISIBLE_DEVICES selects correctly.
@@ -706,9 +768,8 @@ def _detect_bnb_rocm_dll_ver() -> str | None:
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 e.g. "713" wins over "72" when both
- # variants are present. Glob order is not guaranteed, so always sort
- # rather than stopping at the first match.
+ # Highest numeric suffix wins (e.g. "713" over "72"); glob order is not
+ # guaranteed, so sort rather than take the first match.
return max(all_vers, key = lambda v: int(v)) if all_vers else None
@@ -825,17 +886,14 @@ 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 tools), the kernel exposes AMD GPUs via
- # /sys/class/kfd so `studio update` can still detect and repair.
+ # sysfs KFD topology fallback (Linux only) -- matches install.sh's runtime-only
+ # detection. On minimal package-managed installs (no rocminfo / amd-smi), the
+ # kernel exposes AMD GPUs via /sys/class/kfd so `studio update` can still repair.
#
- # Guard: reject any KFD node whose properties file reports a non-AMD
- # vendor. With the NVIDIA open kernel module (driver 560+), NVIDIA GPUs
- # can register KFD topology nodes with a non-zero gpu_id; those nodes
- # have vendor_id 4318 (0x10DE) rather than the AMD value 4098 (0x1002).
- # Without this check the fallback returns True on NVIDIA-only systems,
- # causing _ensure_rocm_torch to install ROCm wheels on NVIDIA hardware.
+ # Guard: reject any KFD node whose properties file reports a non-AMD vendor. The
+ # NVIDIA open kernel module (driver 560+) registers KFD nodes with a non-zero
+ # gpu_id and vendor_id 4318 (0x10DE), not the AMD 4098 (0x1002); without this
+ # check the fallback returns True on NVIDIA-only hosts, installing ROCm wheels.
if sys.platform != "win32":
try:
kfd_nodes = "/sys/class/kfd/kfd/topology/nodes"
@@ -849,12 +907,10 @@ def _has_rocm_gpu() -> bool:
continue
if not gpu_id or gpu_id == "0": # gpu_id 0 = CPU node
continue
- # Require AMD vendor_id 4098 (0x1002) in the properties file.
- # KFD properties files exist on every kernel that exposes
- # /sys/class/kfd, so absence of the file means we cannot
- # confirm AMD ownership -- skip the node rather than risk a
- # false positive (e.g. NVIDIA open driver KFD nodes that
- # lack a properties file on some kernel versions).
+ # Require AMD vendor_id 4098 (0x1002). KFD properties files exist
+ # on every kernel exposing /sys/class/kfd, so a missing file means
+ # AMD ownership is unconfirmed -- skip the node rather than risk a
+ # false positive (e.g. NVIDIA open-driver KFD nodes lacking it).
props_path = os.path.join(kfd_nodes, entry, "properties")
try:
with open(props_path) as fh:
@@ -981,13 +1037,10 @@ def _install_bnb_windows_rocm() -> bool:
)
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).
+ # Detect the actual ROCm DLL suffix in the wheel and set BNB_ROCM_VERSION so bnb
+ # loads the right DLL regardless of torch.version.hip (the wheel may ship "72"
+ # while torch reports 7.13). The worker subprocess inherits it; fall back to "72"
+ # if detection fails (e.g. a no-op / dry-run install).
_env_ver = os.environ.get("BNB_ROCM_VERSION")
_env_is_persisted_default = (
os.environ.get(_BNB_ROCM_VERSION_SOURCE_ENV) == _BNB_ROCM_VERSION_SOURCE_SITECUSTOMIZE
@@ -1002,13 +1055,11 @@ def _install_bnb_windows_rocm() -> bool:
_persist_detected_version = True
if _persist_detected_version:
_persist_bnb_rocm_version(_ver)
- # Make hipInfo.exe (shipped into the venv Scripts dir by the AMD torch
- # wheel) resolvable via PATH for this process and every child python the
- # installer spawns (import checks, precompile): bitsandbytes runs
- # `hipinfo.exe` at import time to detect the GPU arch and logs a scary
- # (harmless) ERROR + WARNING on every import when it is missing. The venv
- # Scripts dir is on PATH only when the venv is activated, which neither
- # Unsloth nor the installer's child processes ever do.
+ # Make hipInfo.exe (shipped into venv Scripts by the AMD torch wheel) resolvable
+ # via PATH for this process and every child python (import checks, precompile):
+ # bitsandbytes runs hipinfo.exe at import to detect the GPU arch and logs a scary
+ # (harmless) ERROR + WARNING when it is missing. Scripts is on PATH only for an
+ # activated venv, which neither Unsloth nor the installer's children ever do.
_scripts_dir = os.path.dirname(sys.executable)
if os.path.isfile(os.path.join(_scripts_dir, "hipInfo.exe")) and not shutil.which(
"hipinfo.exe"
@@ -1020,13 +1071,18 @@ def _install_bnb_windows_rocm() -> bool:
def _detect_cuda_torch_index_url() -> str:
"""Return the pytorch.org CUDA wheel index URL for the host's NVIDIA driver.
- Mirrors install.sh::get_torch_index_url's CUDA ladder so `studio update`
- repairs to the same wheel family a fresh `curl | sh` install would pick.
- Probes nvidia-smi (PATH, then /usr/bin/nvidia-smi) and parses both the
- legacy "CUDA Version:" and the newer "CUDA UMD Version:" spellings.
- Defaults to cu126 when nvidia-smi is missing or the version is unreadable
- (e.g. NVIDIA detected only via the /proc/driver/nvidia/gpus fallback).
+ Mirrors install.sh::get_torch_index_url's CUDA ladder so `studio update` repairs
+ to the same wheel family a fresh install would pick. Honours the explicit
+ overrides first (UNSLOTH_TORCH_INDEX_URL / _FAMILY) so a headless / CI install
+ never lets the host GPU decide. Otherwise probes nvidia-smi (parsing both "CUDA
+ Version:" and "CUDA UMD Version:"), defaulting to cu126 when unreadable.
"""
+ _override_url = os.environ.get("UNSLOTH_TORCH_INDEX_URL", "").strip()
+ if _override_url:
+ return _trim_index_path_slashes(_override_url)
+ _override_family = os.environ.get("UNSLOTH_TORCH_INDEX_FAMILY", "").strip()
+ if _override_family:
+ return f"{_PYTORCH_WHL_BASE}/{_override_family.strip('/')}"
exe = shutil.which("nvidia-smi")
if not exe and os.path.isfile("/usr/bin/nvidia-smi"):
exe = "/usr/bin/nvidia-smi"
@@ -1061,6 +1117,157 @@ def _detect_cuda_torch_index_url() -> str:
return f"{_PYTORCH_WHL_BASE}/{tag}"
+def _explicit_torch_index_url() -> "str | None":
+ """The wheel index URL pinned via UNSLOTH_TORCH_INDEX_URL / _FAMILY, else None.
+
+ Lets the CUDA/ROCm repair helpers honour the exact pinned family/URL instead
+ of re-probing the GPU. Mirrors install.sh::get_torch_index_url's override.
+ """
+ url = os.environ.get("UNSLOTH_TORCH_INDEX_URL", "").strip()
+ if url:
+ return _trim_index_path_slashes(url)
+ family = os.environ.get("UNSLOTH_TORCH_INDEX_FAMILY", "").strip()
+ if family:
+ return f"{_PYTORCH_WHL_BASE}/{family.strip('/')}"
+ return None
+
+
+def _is_pip_rocm_family_leaf(leaf: str) -> bool:
+ """True when a lowercased leaf names a pip --index-url ROCm family: an EXACT
+ rocm[.] leaf or a gfx leaf. A suffixed leaf (rocm-rel-7.2.1,
+ rocm7.2-private) starts with "rocm" but is a custom pin the verbatim path owns, so
+ match EXACTLY. Mirrors install.sh / setup.ps1.
+ """
+ # gfx must be followed by a digit (gfx90a, gfx1151, gfx120X-all): a gfx-prefixed
+ # custom leaf (gfx-private) is a verbatim pin, like rocm7.2-private.
+ return bool(re.fullmatch(r"rocm\d+(?:\.\d+)?", leaf)) or bool(re.match(r"gfx\d", leaf))
+
+
+def _explicit_rocm_torch_index_url() -> "str | None":
+ """The pinned wheel index URL when it names a pip ROCm family (rocm/gfx*), else None."""
+ url = _explicit_torch_index_url()
+ if url is None:
+ return None
+ return url if _is_pip_rocm_family_leaf(_torch_index_leaf(url)) else None
+
+
+def _rocm_pin_family_mismatch(pin_url: str, installed_ver: str) -> bool:
+ """True when an explicit ROCm pin names a different ROCm family than the installed
+ ROCm torch, so the pin needs a reinstall. Mirrors setup.ps1's stale-venv comparison;
+ same three pin-leaf cases as _ensure_rocm_torch. A same-family pin is NOT a mismatch.
+ """
+ leaf = _torch_index_leaf(pin_url)
+ # Pinned ROCm version. The family classifier accepts a major-only rocm leaf too,
+ # so parse the minor as optional; a major-only pin compares on the major alone.
+ _pin_rocm = re.match(r"^rocm(\d+)(?:\.(\d+))?", leaf)
+ _pin_major = int(_pin_rocm.group(1)) if _pin_rocm else None
+ _pin_ver = (
+ (int(_pin_rocm.group(1)), int(_pin_rocm.group(2)))
+ if _pin_rocm and _pin_rocm.group(2) is not None
+ else None
+ )
+ # Installed +rocmX.Y version; a THREE-part +rocmA.B.C tag is the AMD per-arch
+ # (repo.amd.com/gfx*) signature vs a two-part pytorch.org wheel.
+ _inst_rocm = re.search(r"\+rocm(\d+)\.(\d+)", installed_ver)
+ _inst_ver = (int(_inst_rocm.group(1)), int(_inst_rocm.group(2))) if _inst_rocm else None
+ _inst_is_perarch = re.search(r"\+rocm\d+\.\d+\.\d+", installed_ver) is not None
+ # A ROCm build MUST carry a +rocm tag; an untagged wheel never satisfies a ROCm pin.
+ _inst_has_rocm = re.search(r"\+rocm", installed_ver) is not None
+ # Installed torch RELEASE (before "+") is 2.11+.
+ _inst_rel = re.match(r"^(\d+)\.(\d+)", installed_ver)
+ _inst_is_211 = (
+ (int(_inst_rel.group(1)), int(_inst_rel.group(2))) >= (2, 11) if _inst_rel else False
+ )
+
+ if leaf.startswith("gfx"):
+ # 2.11-allowlist arches expect the AMD per-arch wheel (three-part +rocmA.B.C,
+ # torch 2.11+); a generic or pre-2.11 build is a mismatch.
+ if leaf in _ROCM_GFX_TORCH211_LEAVES:
+ return not (_inst_is_211 and _inst_is_perarch)
+ # Non-2.11 gfx leaf (<2.11 specs): mismatch on an untagged wheel or torch 2.11+.
+ return (not _inst_has_rocm) or _inst_is_211
+
+ # Major-only rocm pin (rocm7): compare majors only -- a +rocm6.4 wheel under a rocm7
+ # pin is a mismatch, any +rocm7.x wheel satisfies it (there is no pinned minor to
+ # compare, and the 2.11-line fallback below would invert both verdicts).
+ if _pin_major is not None and _pin_ver is None:
+ if _inst_ver is not None:
+ return _inst_ver[0] != _pin_major
+ # Untagged wheel never satisfies a ROCm pin; a +rocm tag with an unreadable
+ # version is accepted (matches the lenient unreadable fallback below).
+ return not _inst_has_rocm
+
+ # rocmX.Y pin. Only KNOWN-2.11 rocm is the 2.11 line (no speculative floor).
+ _pin_is_211 = _pin_ver in _ROCM_KNOWN_TORCH211_VERSIONS if _pin_ver is not None else False
+ if _pin_ver is not None and _inst_ver is not None:
+ # Both readable: exact (major, minor) compare (rocm7.2 pin over +rocm7.13.x ->
+ # mismatch, reinstall the pinned wheel).
+ if _pin_ver != _inst_ver:
+ return True
+ # Same family: a KNOWN-2.11 pin whose release drifted off 2.11 (2.12+rocm7.2)
+ # violates the spec -> reinstall to floor (exact compare, not >=2.11).
+ if _pin_is_211 and _inst_rel is not None:
+ if (int(_inst_rel.group(1)), int(_inst_rel.group(2))) != (2, 11):
+ return True
+ return False
+ # rocm pin, unreadable installed version: compare on the 2.11 line, but an untagged
+ # wheel never satisfies a rocmX.Y pin -> mismatch.
+ if not _inst_has_rocm:
+ return True
+ return _pin_is_211 != _inst_is_211
+
+
+def _explicit_cpu_torch_index_url() -> "str | None":
+ """The pinned wheel index URL when it names the CPU family (leaf == cpu), else None.
+
+ An explicit CPU pin (UNSLOTH_TORCH_INDEX_FAMILY=cpu or a URL ending in /cpu)
+ is authoritative -- see _ensure_cpu_torch.
+ """
+ url = _explicit_torch_index_url()
+ if url is None:
+ return None
+ return url if _torch_index_leaf(url) == "cpu" else None
+
+
+def _is_cuda_family_leaf(leaf: str) -> bool:
+ """True only for a real CUDA wheel-family leaf: "cu" + digits (cu118, cu128, ...).
+
+ A bare startswith("cu") would match "custom"/"current". The match is EXACT so
+ "cu128-private" is NOT a family leaf and routes to the verbatim path instead.
+ """
+ return re.fullmatch(r"cu[0-9]+", leaf) is not None
+
+
+def _explicit_cuda_torch_index_url() -> "str | None":
+ """The pinned wheel index URL when it names a CUDA family (leaf cuXXX), else None.
+
+ Mirrors _explicit_rocm/cpu_torch_index_url so _ensure_cuda_torch only treats a
+ *CUDA* pin as authority to override the NVIDIA-presence gate (an arbitrary mirror
+ or a ROCm/CPU pin must not force a CUDA reinstall on a non-NVIDIA host).
+ """
+ url = _explicit_torch_index_url()
+ if url is None:
+ return None
+ return url if _is_cuda_family_leaf(_torch_index_leaf(url)) else None
+
+
+def _explicit_unknown_family_torch_index_url() -> "str | None":
+ """The pinned index URL when its leaf names NO known torch family, else None.
+
+ Known = rocm* / gfx* / cpu / cuXXX. Anything else (a private mirror /simple,
+ /current) is UNKNOWN: version-tag heuristics can't judge it, so the family
+ repair helpers must leave it alone (the install applied it verbatim).
+ Matches install.sh / setup.ps1 / install.ps1.
+ """
+ url = _explicit_torch_index_url()
+ if url is None:
+ return None
+ leaf = _torch_index_leaf(url)
+ if _is_pip_rocm_family_leaf(leaf) or leaf == "cpu" or _is_cuda_family_leaf(leaf):
+ return None
+ return url
+
+
def _ensure_cuda_torch() -> None:
"""Repair a venv whose torch is a ROCm build on an NVIDIA host.
@@ -1073,44 +1280,47 @@ def _ensure_cuda_torch() -> None:
Only repairs when torch actually links against HIP/ROCm. Healthy CUDA
torch and deliberate CPU-only torch are left untouched.
"""
- # Respect an explicit backend choice from install.sh: only "" (standalone
- # `studio update`) or "cuda" should ever force CUDA wheels. "rocm"/"cpu"
- # (or any unrecognised value) are deliberate and must not be overridden.
+ # Respect install.sh's backend: only "" (standalone update) or "cuda" force CUDA
+ # wheels; "rocm"/"cpu"/unrecognised are deliberate.
if _TORCH_BACKEND not in ("", "cuda"):
return
- # No CUDA torch on macOS; Windows venv/torch lifecycle is owned by
- # install.ps1 (and the KFD poisoning bug is Linux-only), so skip both.
+ # An explicit unknown-family pin was applied VERBATIM at install time; leave it alone.
+ if _explicit_unknown_family_torch_index_url() is not None:
+ return
+ # No CUDA torch on macOS; Windows torch is owned by install.ps1 (KFD bug is Linux-only).
if IS_MACOS or IS_WINDOWS or NO_TORCH:
return
# Never undo a deliberate ROCm install (setup.ps1 sets this marker).
if os.environ.get("UNSLOTH_ROCM_TORCH_INSTALLED") == "1":
return
- # CUDA_VISIBLE_DEVICES="" / "-1" deliberately hides the NVIDIA GPU (for
- # example a mixed AMD+NVIDIA host that runs ROCm torch on the AMD card);
- # never force CUDA wheels over that choice.
+ # An explicit CUDA pin (headless / CI cross-install) commits to CUDA wheels and skips ALL
+ # GPU probing, so it clears both the CUDA_VISIBLE_DEVICES hide gate and the NVIDIA gate below.
+ _cuda_pinned = _explicit_cuda_torch_index_url() is not None
+ # CUDA_VISIBLE_DEVICES="" / "-1" deliberately hides the NVIDIA GPU; never force CUDA
+ # wheels over that unless a CUDA index is pinned.
_cvd = os.environ.get("CUDA_VISIBLE_DEVICES")
- if _cvd is not None and _cvd.strip() in ("", "-1"):
+ if not _cuda_pinned and _cvd is not None and _cvd.strip() in ("", "-1"):
return
- # Only NVIDIA hosts should carry CUDA torch. _has_usable_nvidia_gpu()
- # covers the /proc/driver/nvidia/gpus fallback when nvidia-smi is absent.
- if not _has_usable_nvidia_gpu():
+ # Only NVIDIA hosts carry CUDA torch (the CUDA pin overrides this gate too).
+ if not _cuda_pinned and not _has_usable_nvidia_gpu():
return
- # Classify the installed torch: "hip" (ROCm build -- the poisoning
- # signature), "cuda" (healthy), or "cpu" (deliberate CPU wheel). A
- # non-zero exit means torch is missing or un-importable; the base install
- # step handles that, so leave it alone.
+ # Classify the installed torch: "hip" (ROCm poisoning signature), "cuda" (healthy),
+ # or "cpu". A non-zero exit means torch is missing/un-importable: without a pin the
+ # base install owns it, but a pinned CUDA index reinstalls it below.
try:
probe = subprocess.run(
[
sys.executable,
"-c",
(
- "import torch; "
+ "import torch, re; "
"hip = getattr(torch.version, 'hip', '') or ''; "
"cuda = getattr(torch.version, 'cuda', '') or ''; "
"ver = getattr(torch, '__version__', '').lower(); "
- "print('hip' if (hip or 'rocm' in ver) else ('cuda' if cuda else 'cpu'))"
+ "m = re.search(r'\\+(cu\\d+)', ver); "
+ "marker = 'hip' if (hip or 'rocm' in ver) else ('cuda' if cuda else 'cpu'); "
+ "print(marker + '|' + (m.group(1) if m else ''))"
),
],
stdout = subprocess.PIPE,
@@ -1120,22 +1330,60 @@ def _ensure_cuda_torch() -> None:
except (OSError, subprocess.TimeoutExpired):
return
if probe.returncode != 0:
+ # torch present but can't import. Without a pin the base install owns it; but an
+ # explicit CUDA pin forces this pass (failed probe) and the base update won't
+ # reinstall an already-installed torch, so reinstall from the pin (self-resolving).
+ if not _cuda_pinned:
+ return
+ index_url = _detect_cuda_torch_index_url()
+ _torch_pkg, _vision_pkg, _audio_pkg = _CUDA_TORCH_PKG_SPEC
+ print(
+ f" torch cannot import but an explicit CUDA index is pinned -- reinstalling "
+ f"CUDA torch from {_strip_index_url_credentials(index_url)}"
+ )
+ pip_install(
+ "CUDA torch repair",
+ "--force-reinstall",
+ "--no-cache-dir",
+ _torch_pkg,
+ _vision_pkg,
+ _audio_pkg,
+ "--index-url",
+ index_url,
+ constrain = False,
+ )
return
- # Take the last non-empty stdout line: stray output from sitecustomize or
- # an import hook must not mask the marker (fail-closed either way).
+ # Last non-empty line: stray sitecustomize/import-hook output must not mask the marker.
_marker_lines = [
line.strip() for line in probe.stdout.decode(errors = "replace").splitlines() if line.strip()
]
- if not _marker_lines or _marker_lines[-1] != "hip":
- return # healthy CUDA torch, or a deliberate CPU wheel -- leave as-is
+ if not _marker_lines:
+ return
+ _marker, _, _installed_cu = _marker_lines[-1].partition("|")
+ # Reinstall CUDA torch on a ROCm build on an NVIDIA host (poisoning signature), or when a
+ # CUDA index is pinned but the venv has the wrong family (CPU or a different cuXXX). A
+ # healthy match, or a CPU wheel with no CUDA pin, is left alone.
+ _pin = _explicit_torch_index_url()
+ _pin_leaf = _torch_index_leaf(_pin) if _pin else ""
+ _pinned_cuda = _is_cuda_family_leaf(_pin_leaf)
+ if _marker == "hip":
+ _why = "torch is a ROCm build on an NVIDIA host"
+ elif _marker == "cpu" and _pinned_cuda:
+ _why = "torch is a CPU build but an explicit CUDA index is pinned"
+ elif _marker == "cuda" and _pinned_cuda and _installed_cu != _pin_leaf:
+ # Installed cuXXX differs from the pin. An untagged build (empty) counts too:
+ # the family can't be confirmed, so reinstall to enforce it (idempotent).
+ _installed_desc = _installed_cu if _installed_cu else "an untagged CUDA build"
+ _why = f"torch is {_installed_desc} but the pinned CUDA index is {_pin_leaf}"
+ else:
+ return # healthy CUDA torch matching the pin, or a deliberate CPU wheel
index_url = _detect_cuda_torch_index_url()
_torch_pkg, _vision_pkg, _audio_pkg = _CUDA_TORCH_PKG_SPEC
print(
- f" torch is a ROCm build on an NVIDIA host -- reinstalling "
- f"CUDA torch from {index_url}\n"
- f" (set UNSLOTH_TORCH_BACKEND=rocm to keep a deliberate ROCm torch "
- f"on a mixed AMD+NVIDIA host)"
+ f" {_why} -- reinstalling CUDA torch from {_strip_index_url_credentials(index_url)}\n"
+ f" (set UNSLOTH_TORCH_BACKEND=rocm or cpu to keep a deliberate "
+ f"non-CUDA torch)"
)
pip_install(
"CUDA torch repair",
@@ -1150,6 +1398,90 @@ def _ensure_cuda_torch() -> None:
)
+def _ensure_cpu_torch() -> None:
+ """Reinstall CPU torch when an explicit CPU pin is set but the venv has a GPU build.
+
+ Counterpart to _ensure_cuda/rocm_torch for the explicit-CPU case (those treat a CPU
+ backend as a skip, so a standalone `studio update` would ignore the authoritative CPU
+ pin). Only fires for an EXPLICIT pin.
+ """
+ if NO_TORCH:
+ return
+ pin = _explicit_cpu_torch_index_url()
+ if pin is None:
+ return
+
+ # Classify the installed torch family. A non-zero exit means torch is missing or
+ # un-importable: the explicit CPU pin reinstalls it below.
+ try:
+ probe = subprocess.run(
+ [
+ sys.executable,
+ "-c",
+ (
+ "import torch, re; "
+ "hip = getattr(torch.version, 'hip', '') or ''; "
+ "cuda = getattr(torch.version, 'cuda', '') or ''; "
+ "ver = getattr(torch, '__version__', '').lower(); "
+ "gpu = bool(hip) or 'rocm' in ver or bool(cuda) or bool(re.search(r'\\+cu\\d+', ver)); "
+ "print('gpu' if gpu else 'cpu')"
+ ),
+ ],
+ stdout = subprocess.PIPE,
+ stderr = subprocess.DEVNULL,
+ timeout = 90,
+ )
+ except (OSError, subprocess.TimeoutExpired):
+ return
+ if probe.returncode != 0:
+ # torch present but can't import. The explicit CPU pin forces this pass (failed
+ # probe) and the base update won't reinstall an already-installed torch, so
+ # reinstall from the pin (self-resolving, no loop).
+ _torch_pkg, _vision_pkg, _audio_pkg = _CPU_TORCH_PKG_SPEC
+ print(
+ f" torch cannot import but an explicit CPU index is pinned -- reinstalling "
+ f"CPU torch from {_strip_index_url_credentials(pin)}"
+ )
+ pip_install(
+ "CPU torch repair",
+ "--force-reinstall",
+ "--no-cache-dir",
+ _torch_pkg,
+ _vision_pkg,
+ _audio_pkg,
+ "--index-url",
+ pin,
+ constrain = False,
+ )
+ return
+ _lines = [
+ line.strip() for line in probe.stdout.decode(errors = "replace").splitlines() if line.strip()
+ ]
+ if not _lines:
+ return # unreadable -- the base install step handles a missing torch
+ if _lines[-1] != "gpu":
+ return # already a CPU build
+
+ print(
+ " torch is a GPU build but an explicit CPU index is pinned -- reinstalling "
+ f"CPU torch from {_strip_index_url_credentials(pin)}"
+ )
+ # Pin the supported torch<2.11 family (the /cpu index now serves 2.11+, so a bare
+ # trio could resolve out of range or ABI-mismatched).
+ _torch_pkg, _vision_pkg, _audio_pkg = _CPU_TORCH_PKG_SPEC
+ pip_install(
+ "CPU torch repair",
+ "--force-reinstall",
+ "--no-cache-dir",
+ _torch_pkg,
+ _vision_pkg,
+ _audio_pkg,
+ "--index-url",
+ pin,
+ constrain = False,
+ )
+
+
def _ensure_rocm_torch() -> None:
"""Reinstall torch with ROCm wheels when the venv received CPU-only torch.
@@ -1160,16 +1492,15 @@ def _ensure_rocm_torch() -> None:
Uses pip_install() to respect uv, constraints, and --python targeting.
"""
global _rocm_windows_torch_installed
- # install.sh sets UNSLOTH_TORCH_BACKEND to the resolved wheel family
- # ("cuda", "rocm", "cpu"). Skip ROCm operations entirely when install.sh
- # already selected a non-ROCm backend -- this is the authoritative signal
- # and avoids re-running GPU detection in a subprocess that may see a
- # different environment (different PATH, CUDA_VISIBLE_DEVICES, etc.).
+ # install.sh's resolved backend is authoritative: skip ROCm when it already chose a
+ # non-ROCm family (avoids re-detecting in a subprocess that may see a different env).
if _TORCH_BACKEND in ("cuda", "cpu"):
return
- # setup.ps1 sets this after installing 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.
+ # An explicit unknown-family pin was applied VERBATIM at install time; leave it alone.
+ if _explicit_unknown_family_torch_index_url() is not None:
+ return
+ # setup.ps1 sets this after installing AMD wheels; skip only when torch is actually
+ # importable as ROCm (a wiped venv leaves a stale env-var that must not suppress it).
if os.environ.get("UNSLOTH_ROCM_TORCH_INSTALLED") == "1":
_torch_ok = False
try:
@@ -1193,9 +1524,8 @@ def _ensure_rocm_torch() -> None:
pass
if _torch_ok:
_rocm_windows_torch_installed = True
- # setup.ps1 already installed ROCm torch, but we still need the AMD
- # Windows BNB wheel here -- the PyPI bitsandbytes wheel ships only
- # CUDA DLLs and fails to load on ROCm.
+ # ROCm torch is already installed, but the AMD Windows BNB wheel is still
+ # needed (the PyPI bitsandbytes ships only CUDA DLLs, fails on ROCm).
_install_bnb_windows_rocm()
return
# torch was wiped between runs; fall through to the full install path
@@ -1203,10 +1533,15 @@ def _ensure_rocm_torch() -> None:
return
if IS_WINDOWS:
- if _has_usable_nvidia_gpu():
+ # An explicit ROCm-family pin commits to ROCm wheels regardless of the visible
+ # GPU and overrides the public per-arch index (mirrors the Linux pin handling
+ # below): after a pinned setup.ps1 install fails to CPU, this repair must retry
+ # the PINNED index, not repo.amd.com.
+ _win_rocm_pin = _explicit_rocm_torch_index_url()
+ if _win_rocm_pin is None and _has_usable_nvidia_gpu():
return
gfx_arch = _detect_windows_gfx_arch()
- if not gfx_arch:
+ if not gfx_arch and _win_rocm_pin is None:
return # no AMD GPU visible via hipinfo
# Probe whether torch already links against HIP.
_torch_already_rocm = False
@@ -1231,23 +1566,24 @@ def _ensure_rocm_torch() -> None:
except (OSError, subprocess.TimeoutExpired):
pass
if not _torch_already_rocm:
- index_url = _windows_rocm_index_url(gfx_arch)
+ index_url = _win_rocm_pin or _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}")
- # Pin companions for the arches install.ps1/setup.ps1 pin (gfx120X /
- # Strix) so the per-arch index resolves an ABI-consistent trio; other
- # arches stay bare (no published floor), matching the PowerShell side.
+ print(
+ f" {gfx_arch or 'pinned ROCm index'} (Windows) -- installing torch from "
+ f"{_strip_index_url_credentials(index_url)}"
+ )
+ # Pin companions for the arches install.ps1/setup.ps1 pin (gfx120X / Strix)
+ # so the per-arch index resolves an ABI-consistent trio; other arches stay bare.
_torch_pkg, _vision_pkg, _audio_pkg = _WINDOWS_ROCM_TORCH_PKG_SPECS.get(
gfx_arch, ("torch", "torchvision", "torchaudio")
)
- # Nonfatal: a transient AMD-index failure must not abort the whole
- # install once the PowerShell side has fallen back to CPU torch.
- # --force-reinstall resolves before uninstalling, so a failed index
- # leaves the existing build intact; keep it and let the user retry.
+ # Nonfatal: a transient AMD-index failure must not abort the install.
+ # --force-reinstall resolves before uninstalling, so a failed index keeps the
+ # existing build intact; let the user retry.
if not pip_install_try(
- f"ROCm torch (Windows, {gfx_arch})",
+ f"ROCm torch (Windows, {gfx_arch or 'pinned'})",
"--force-reinstall",
"--index-url",
index_url,
@@ -1257,7 +1593,7 @@ def _ensure_rocm_torch() -> None:
constrain = False,
):
print(
- f" Warning: AMD Windows ROCm torch install failed for {gfx_arch}; "
+ f" Warning: AMD Windows ROCm torch install failed for {gfx_arch or 'the pinned index'}; "
"keeping the existing torch build. Re-run 'unsloth studio update' "
"later to retry ROCm."
)
@@ -1280,26 +1616,30 @@ def _ensure_rocm_torch() -> None:
# ── 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 a GPU is usable
- if _has_usable_nvidia_gpu():
- return
- # Use _has_rocm_gpu() (rocminfo / amd-smi GPU data rows) as the
- # authoritative "is this an AMD ROCm host?" signal. The old gate required
- # /opt/rocm or hipcc to exist, which breaks runtime-only ROCm installs
- # (minimal package-managed installs, Radeon software) that ship
- # amd-smi/rocminfo without /opt/rocm or hipcc, leaving `unsloth studio
- # update` unable to repair a CPU-only venv on those systems.
- if not _has_rocm_gpu():
- return # no AMD GPU visible
+ # An explicit ROCm pin commits to ROCm wheels regardless of the visible GPU (headless / CI).
+ # Mirror _ensure_cuda_torch: skip the NVIDIA/no-AMD/unreadable gates.
+ _rocm_pin = _explicit_rocm_torch_index_url()
+ if _rocm_pin is None:
+ # NVIDIA takes precedence on mixed hosts (only if a GPU is usable).
+ if _has_usable_nvidia_gpu():
+ return
+ # _has_rocm_gpu() (rocminfo / amd-smi rows) is the authoritative AMD-host signal;
+ # the old /opt/rocm-or-hipcc gate broke runtime-only ROCm installs.
+ if not _has_rocm_gpu():
+ return # no AMD GPU visible
ver = _detect_rocm_version()
if ver is None:
- print(" ROCm detected but version unreadable -- skipping torch reinstall")
- return
+ if _rocm_pin is None:
+ print(" ROCm detected but version unreadable -- skipping torch reinstall")
+ return
+ # Explicit pin: the pinned leaf drives the install, so an unreadable host version
+ # is fine (sentinel keeps ver comparisons defined).
+ ver = (0, 0)
- # Probe whether torch already links against HIP (ROCm already working).
- # Do NOT skip for CUDA-only builds: they are unusable on AMD-only hosts
- # (the NVIDIA check above already handled mixed AMD+NVIDIA setups).
+ # Probe whether torch links against HIP, capturing the installed ROCm tag for pin-mismatch
+ # detection. Emit ONE "|" line: marker (HIP version, "rocm" sentinel,
+ # or empty for CPU/CUDA) before "|", wheel version after.
try:
probe = subprocess.run(
[
@@ -1309,10 +1649,10 @@ def _ensure_rocm_torch() -> None:
"import torch; "
"hip=getattr(torch.version,'hip','') or ''; "
"ver=getattr(torch,'__version__','').lower(); "
- # Print the HIP version when present (back-compat), else a
- # "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 ''))"
+ # HIP version if present, else a "rocm" sentinel when only the
+ # version string flags ROCm; empty marker = CPU/CUDA torch.
+ "marker=hip if hip else ('rocm' if 'rocm' in ver else ''); "
+ "print(marker + '|' + ver)"
),
],
stdout = subprocess.PIPE,
@@ -1321,29 +1661,42 @@ def _ensure_rocm_torch() -> None:
)
except (OSError, subprocess.TimeoutExpired):
probe = None
- has_hip_torch = (
- probe is not None and probe.returncode == 0 and probe.stdout.decode().strip() != ""
+ # Last non-empty line, split on the FIRST "|" so the empty HIP field is preserved.
+ _marker_lines = (
+ [ln.strip() for ln in probe.stdout.decode(errors = "replace").splitlines() if ln.strip()]
+ if (probe is not None and probe.returncode == 0)
+ else []
+ )
+ _hip_marker, _sep, _installed_torch_ver = (
+ _marker_lines[-1].partition("|") if _marker_lines else ("", "", "")
+ )
+ # A "|"-delimited line is required; without it treat HIP as absent -> reinstall.
+ has_hip_torch = bool(_sep) and _hip_marker != ""
+
+ # An explicit ROCm pin whose family differs from the installed torch must reinstall, else a
+ # rocm7.2/gfx* pin over an older +rocm6.4/7.1 build never applies. Version-tag heuristic
+ # only: a same-tag per-arch switch (gfx1151 -> gfx120X-all, both +rocm7.13.0) isn't detectable.
+ _rocm_pin_mismatch = (
+ _rocm_pin_family_mismatch(_rocm_pin, _installed_torch_ver)
+ if (has_hip_torch and _rocm_pin is not None)
+ else False
)
- rocm_torch_ready = has_hip_torch
+ rocm_torch_ready = has_hip_torch and not _rocm_pin_mismatch
- # 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), route to the AMD per-gfx
- # index only when HIP's runtime GPU is the Strix one -- else the dGPU gets
- # an incompatible wheel. Use HIP_VISIBLE_DEVICES for the runtime target.
+ # Strix Halo / Point (gfx1151 / gfx1150) segfault under ROCm 7.1 in torch._grouped_mm;
+ # AMD's per-gfx repo ships 2.11.0+rocm7.13.0 with the fix, so route those hosts there
+ # (mirrors install.sh). On mixed hosts, reroute only when HIP's runtime GPU is the Strix one.
_strix_override_url: "str | None" = None
_strix_override_pkgs: "tuple[str, str, str] | None" = None
- if ver < (7, 2):
+ # An explicit ROCm pin is authoritative: never auto-reroute it.
+ if ver < (7, 2) and _explicit_rocm_torch_index_url() is None:
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: use the HIP_VISIBLE_DEVICES index
- # into gfx_codes, else default to the first GPU. Skip the override
- # unless the resolved GPU is Strix.
+ # Runtime-visible GPU (HIP_VISIBLE_DEVICES index into gfx_codes, else first);
+ # skip the override unless it's 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
@@ -1353,12 +1706,8 @@ def _ensure_rocm_torch() -> None:
_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 an AMD-index build
- # targeting a different torch major (e.g. 0.27 built against
- # torch 2.12), which fails at runtime with an ABI/version
- # mismatch. Matches _ROCM_TORCH_CONSTRAINT["rocm7.2"].
+ # Pin companions to the 2.11.x range: the exclusive --index-url could
+ # otherwise resolve a build for a different torch major (ABI mismatch).
"torchvision>=0.26.0,<0.27.0",
"torchaudio>=2.11.0,<2.12.0",
)
@@ -1378,14 +1727,15 @@ def _ensure_rocm_torch() -> None:
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 repairs, so skipping it leaves users on the known
- # _grouped_mm segfault.
+ # The Strix override must fire even when has_hip_torch is True: an existing
+ # torch.version.hip == "7.1" is exactly the broken combo it repairs.
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}")
+ print(
+ f" Strix ROCm 7.1 override -- installing torch from "
+ f"{_strip_index_url_credentials(index_url)}"
+ )
pip_install(
"ROCm torch (Strix arch-specific)",
"--force-reinstall",
@@ -1398,24 +1748,38 @@ def _ensure_rocm_torch() -> None:
constrain = False,
)
rocm_torch_ready = True
- elif not has_hip_torch:
- # Select best matching wheel tag (newest ROCm version <= installed)
- tag = next(
- (
- t
- for (maj, mn), t in sorted(_ROCM_TORCH_INDEX.items(), reverse = True)
- if ver >= (maj, mn)
- ),
- None,
- )
- if tag is None:
- print(f" No PyTorch wheel for ROCm {ver[0]}.{ver[1]} -- " f"skipping torch reinstall")
+ elif not has_hip_torch or _rocm_pin_mismatch:
+ # Reinstall when torch is not ROCm yet, OR a ROCm build's family differs from a pin.
+ # Honour a ROCm pin verbatim; else pick the newest wheel tag <= host.
+ _override_idx = _explicit_rocm_torch_index_url()
+ if _override_idx is not None:
+ index_url = _override_idx
+ tag = _torch_index_leaf(index_url)
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"]
+ tag = next(
+ (
+ t
+ for (maj, mn), t in sorted(_ROCM_TORCH_INDEX.items(), reverse = True)
+ if ver >= (maj, mn)
+ ),
+ None,
)
+ if tag is None:
+ print(f" No PyTorch wheel for ROCm {ver[0]}.{ver[1]} -- skipping torch reinstall")
+ else:
+ if _override_idx is None:
+ index_url = f"{_PYTORCH_WHL_BASE}/{tag}"
+ print(f" ROCm torch -- installing from {_strip_index_url_credentials(index_url)}")
+ # Only the _grouped_mm-bug gfx arches need the 2.11 spec; other gfx indexes ship
+ # <2.11 and stay on the default range (matches install.ps1 / setup.ps1).
+ if tag in _ROCM_GFX_TORCH211_LEAVES:
+ _torch_pkg, _vision_pkg, _audio_pkg = _ROCM_TORCH_PKG_SPECS["rocm7.2"]
+ elif tag.startswith("gfx"):
+ _torch_pkg, _vision_pkg, _audio_pkg = _ROCM_TORCH_PKG_SPECS["_default"]
+ else:
+ _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",
@@ -1504,11 +1868,26 @@ def _infer_no_torch() -> bool:
NO_TORCH = _infer_no_torch()
-# UNSLOTH_TORCH_BACKEND is set by install.sh after get_torch_index_url() so
-# that this script knows which torch variant was selected without re-running
-# GPU detection. Values: "cuda", "rocm", or "cpu". Empty means unknown
-# (standalone `unsloth studio update` runs, where we re-detect normally).
+# UNSLOTH_TORCH_BACKEND is set by install.sh after get_torch_index_url() ("cuda", "rocm",
+# "cpu"; empty = standalone `studio update`, where we re-detect).
_TORCH_BACKEND: str = os.environ.get("UNSLOTH_TORCH_BACKEND", "").lower()
+# Standalone update with an explicit pin: derive the backend from the override (classify on
+# the final URL/family segment, mirroring install.sh) instead of re-probing the GPU.
+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 = _torch_index_leaf(_idx_override)
+ if _idx_leaf.startswith(("rocm", "gfx")):
+ _TORCH_BACKEND = "rocm"
+ elif _idx_leaf == "cpu":
+ _TORCH_BACKEND = "cpu"
+ elif _is_cuda_family_leaf(_idx_leaf):
+ # Require a digit after "cu" so /current or /custom is NOT branded CUDA (a wrong backend
+ # makes _ensure_rocm_torch return early on AMD hosts). An unknown leaf keeps "" so the
+ # helpers probe the GPU.
+ _TORCH_BACKEND = "cuda"
def _torch_step_label(suffix: str) -> str:
@@ -1724,12 +2103,15 @@ def run(
cmd,
stdout = subprocess.PIPE if quiet else None,
stderr = subprocess.STDOUT if quiet else None,
+ env = _install_env_for_cmd(cmd),
**_windows_hidden_subprocess_kwargs(),
)
if result.returncode != 0:
_step("error", f"{label} failed (exit code {result.returncode})", _red)
if result.stdout:
- print(result.stdout.decode(errors = "replace"))
+ # Redact before printing: the failing pip command may carry a pinned --index-url
+ # with userinfo/?token= creds, so raw pip error text would leak them.
+ print(_redact_install_output(result.stdout))
sys.exit(result.returncode)
return result
@@ -1737,15 +2119,13 @@ def run(
# Packages to skip on Windows (require special build steps)
WINDOWS_SKIP_PACKAGES = {"triton_kernels"}
-# Packages to skip when torch is unavailable (Intel Mac GGUF-only mode).
-# These either *are* torch extensions or have unconditional
-# ``Requires-Dist: torch``, so installing them would pull torch back in.
-# ``librosa`` is here too despite not requiring torch: upstream ``llvmlite``
-# dropped its macOS x86_64 wheel between 0.42.0 and 0.46.0+ (see
-# https://pypi.org/project/llvmlite/0.47.0/#files -- only
-# macosx_arm64 / manylinux / win_amd64 remain), so on Intel Mac the
-# librosa -> numba -> llvmlite chain triggers a from-source build that fails
-# in CI and on hosts without LLVM 14/15 headers. Tracked in unslothai/unsloth#5046.
+# Packages to skip when torch is unavailable (Intel Mac GGUF-only mode). These
+# either *are* torch extensions or have unconditional ``Requires-Dist: torch``, so
+# installing them pulls torch back in. ``librosa`` is here despite not requiring
+# torch: upstream ``llvmlite`` dropped its macOS x86_64 wheel (0.46.0+ ships only
+# macosx_arm64 / manylinux / win_amd64), so on Intel Mac the librosa -> numba ->
+# llvmlite chain triggers a from-source build that fails without LLVM 14/15 headers.
+# Tracked in unslothai/unsloth#5046.
NO_TORCH_SKIP_PACKAGES = {
"torch-stoi",
"timm",
@@ -1767,7 +2147,8 @@ def _build_flash_attn_wheel_url(env: dict[str, str]) -> str | None:
def _print_optional_install_failure(label: str, result: subprocess.CompletedProcess[str]) -> None:
_step("warning", f"{label} failed (exit code {result.returncode})", _cyan)
if result.stdout:
- print(result.stdout.strip())
+ # Redact any pinned --index-url credentials before printing captured output.
+ print(_redact_install_output(result.stdout).strip())
def _flash_attn_install_disabled() -> bool:
@@ -1913,15 +2294,60 @@ def _build_uv_cmd(args: tuple[str, ...]) -> list[str]:
# Colab and similar).
cmd.extend(["--python", sys.executable])
cmd.extend(_translate_pip_args_for_uv(args))
- # Torch is pre-installed by install.sh/setup.ps1. Do not add
- # --torch-backend by default -- it can cause solver dead-ends on CPU-only
- # machines. Callers that need it can set UV_TORCH_BACKEND.
+ # Torch is pre-installed, so don't add --torch-backend by default (solver dead-ends on
+ # CPU-only machines); callers can set UV_TORCH_BACKEND. Never add it to a pinned-index
+ # command: uv's torch backend redirects torch to its own per-backend index, defeating the pin.
_tb = os.environ.get("UV_TORCH_BACKEND", "")
- if _tb:
+ if _tb and not _is_pinned_index_cmd(cmd):
cmd.append(f"--torch-backend={_tb}")
return cmd
+# uv resolves --index-url / --default-index at LOWEST priority, so an inherited UV_INDEX /
+# UV_EXTRA_INDEX_URL mirror wins and a pinned torch repair silently ignores the pin.
+# Neutralise these for pinned installs (as install.sh #6898 / install.ps1 / setup.ps1 do).
+# UV_TORCH_BACKEND redirects torch; PIP_* matter for the pip FALLBACK; UV_CONFIG_FILE is
+# stripped + UV_NO_CONFIG=1 (a discovered uv.toml outranks the CLI pin, uv 0.10).
+_UV_INDEX_ENV_VARS = (
+ "UV_CONFIG_FILE",
+ "UV_DEFAULT_INDEX",
+ "UV_INDEX_URL",
+ "UV_INDEX",
+ "UV_EXTRA_INDEX_URL",
+ "UV_TORCH_BACKEND",
+ "UV_FIND_LINKS",
+ "PIP_EXTRA_INDEX_URL",
+ "PIP_FIND_LINKS",
+ # PIP_NO_INDEX=1 makes the pip fallback ignore ALL indexes (defeating --index-url);
+ # PIP_INDEX_URL is dropped too so a stale mirror env can't outrank the pin.
+ "PIP_NO_INDEX",
+ "PIP_INDEX_URL",
+)
+
+
+def _is_pinned_index_cmd(cmd: "list[str] | tuple[str, ...]") -> bool:
+ """True when the command pins an index via --index-url / --default-index."""
+ return any(arg in ("--index-url", "--default-index") for arg in cmd)
+
+
+def _install_env_for_cmd(cmd: "list[str]") -> "dict[str, str] | None":
+ """Return an env with the uv index vars stripped for a pinned-index install.
+
+ None (inherit env) when the command does NOT pin an index, so ordinary installs honour
+ the user's mirror. For pinned commands, the uv index/backend vars are removed,
+ UV_NO_CONFIG=1 set (a discovered uv.toml outranks the CLI pin), and PIP_CONFIG_FILE
+ pointed at os.devnull for the pip fallback. Mirrors install.sh's gate (#6898).
+ """
+ if not _is_pinned_index_cmd(cmd):
+ return None
+ env = os.environ.copy()
+ for name in _UV_INDEX_ENV_VARS:
+ env.pop(name, None)
+ env["UV_NO_CONFIG"] = "1"
+ env["PIP_CONFIG_FILE"] = os.devnull
+ return env
+
+
def pip_install_try(
label: str,
*args: str,
@@ -1948,11 +2374,13 @@ def pip_install_try(
cmd,
stdout = subprocess.PIPE,
stderr = subprocess.STDOUT,
+ env = _install_env_for_cmd(cmd),
)
if result.returncode == 0:
return True
if VERBOSE and result.stdout:
- print(result.stdout.decode(errors = "replace"))
+ # pip/uv echo index URLs (credentials included) in failure output.
+ print(_redact_install_output(result.stdout))
return False
@@ -2000,13 +2428,14 @@ def pip_install(
uv_cmd,
stdout = subprocess.PIPE,
stderr = subprocess.STDOUT,
+ env = _install_env_for_cmd(uv_cmd),
**_windows_hidden_subprocess_kwargs(),
)
if result.returncode == 0:
return
print(_red(f" uv failed, falling back to pip..."))
if result.stdout:
- print(result.stdout.decode(errors = "replace"))
+ print(_redact_install_output(result.stdout))
pip_cmd = _build_pip_cmd(args) + constraint_args_pip + req_args_pip
run(f"{label} (pip)" if USE_UV else label, pip_cmd)
@@ -2054,10 +2483,9 @@ def install_python_stack() -> int:
global USE_UV, _STEP, _TOTAL
_STEP = 0
- # install.sh (which already installed unsloth) sets SKIP_STUDIO_BASE=1 to
- # avoid reinstalling base packages. "unsloth studio update" does NOT set it,
- # so base packages (unsloth + unsloth-zoo) are reinstalled to pick up new
- # versions.
+ # install.sh sets SKIP_STUDIO_BASE=1 to avoid reinstalling base packages;
+ # `studio update` does NOT, so unsloth + unsloth-zoo are reinstalled to pick
+ # up new versions.
skip_base = os.environ.get("SKIP_STUDIO_BASE", "0") == "1"
# --package installs a different package name (for testing).
package_name = os.environ.get("STUDIO_PACKAGE_NAME", "unsloth")
@@ -2067,9 +2495,9 @@ def install_python_stack() -> int:
if IS_MACOS:
base_total -= 1 # triton step is skipped on macOS
if not IS_MACOS and not NO_TORCH:
- base_total += 1 # ROCm torch check (line 1526) -- all non-macOS platforms
+ base_total += 1 # ROCm torch check (step 2b), non-macOS
if not IS_WINDOWS:
- base_total += 2 # flash-attn (line 1620) + ROCm torch final (line 1705) -- Linux only
+ base_total += 2 # flash-attn + torch final repair (step 13), Linux
_TOTAL = (base_total - 1) if skip_base else base_total
# 1. Try uv for faster installs (before pip upgrade -- uv venvs don't
@@ -2134,9 +2562,8 @@ def install_python_stack() -> int:
if skip_base:
pass
elif NO_TORCH:
- # No-torch update path: install unsloth + unsloth-zoo with --no-deps
- # (PyPI metadata still declares torch as a hard dep), then runtime deps
- # with --no-deps (avoids transitive torch).
+ # No-torch update path: install unsloth + unsloth-zoo, then runtime deps,
+ # both with --no-deps (PyPI metadata declares torch a hard dep; avoid it).
_progress("base packages (no torch)")
pip_install(
f"Updating {package_name} + unsloth-zoo (no-torch mode)",
@@ -2149,10 +2576,9 @@ def install_python_stack() -> int:
package_name,
"unsloth-zoo",
)
- # Resolve pydantic WITH deps so pip pins pydantic-core to the exact
- # version pydantic's metadata declares. Under --no-deps pip picks the
- # latest of each and trips pydantic's _ensure_pydantic_core_version
- # check. Transitive deps are torch-free.
+ # Resolve pydantic WITH deps so pip pins pydantic-core to the exact version
+ # its metadata declares (under --no-deps pip picks the latest of each and
+ # trips pydantic's _ensure_pydantic_core_version check). Deps are torch-free.
pip_install(
"Installing pydantic (with deps for compatible core)",
"--no-cache-dir",
@@ -2244,6 +2670,7 @@ def install_python_stack() -> int:
_progress(_torch_step_label("check"))
_ensure_cuda_torch()
_ensure_rocm_torch()
+ _ensure_cpu_torch()
# Windows + AMD GPU: warn if ROCm torch was not installed (wrong Python
# version or unknown ROCm version).
@@ -2309,11 +2736,10 @@ def install_python_stack() -> int:
req = REQ_ROOT / "extras-no-deps.txt",
)
- # 4. Overrides (torchao) -- force-reinstall. The torchao version is chosen to
- # match the torch installed in the venv so its C++ extensions load (see
- # _select_torchao_spec). Skip when torch is unavailable (e.g. Intel Mac
- # GGUF-only mode): torchao requires torch. Also skipped on Windows ROCm
- # (no working build; see below).
+ # 4. Overrides (torchao) -- force-reinstall to a version matching the venv's
+ # torch so its C++ extensions load (see _select_torchao_spec). Skipped when
+ # torch is unavailable (Intel Mac GGUF-only) and on Windows ROCm (no working
+ # build; see below).
if NO_TORCH:
_progress("dependency overrides (skipped, no torch)")
elif _rocm_windows_torch_installed or _installed_torch_is_windows_rocm():
@@ -2430,14 +2856,12 @@ def install_python_stack() -> int:
[sys.executable, str(SINGLE_ENV / "patch_metadata.py")],
)
- # 13. AMD ROCm: final torch repair. Several steps above can pull in CUDA
- # torch from PyPI (base packages, extras, overrides, studio deps, etc.).
- # Running the repair last ensures ROCm torch is in place at runtime,
- # whichever intermediate step clobbered it.
+ # 13. Final torch repair. Steps above can pull CUDA torch from PyPI, so repair last.
if not IS_WINDOWS and not IS_MACOS and not NO_TORCH:
_progress(_torch_step_label("final"))
_ensure_cuda_torch()
_ensure_rocm_torch()
+ _ensure_cpu_torch()
# 14. Final check (silent; third-party conflicts are expected)
subprocess.run(
diff --git a/studio/setup.ps1 b/studio/setup.ps1
index f7d33a1142..f523b9ff14 100644
--- a/studio/setup.ps1
+++ b/studio/setup.ps1
@@ -431,6 +431,167 @@ function Get-PytorchCudaTag {
return "cu126"
}
+# Trim trailing slashes from the URL PATH only, preserving ?query / #fragment: a whole-URL
+# TrimEnd corrupts a token ending in "/", a single strip leaves .../cu128// empty. Shared.
+function Trim-IndexPathSlashes {
+ param([string]$Url)
+ $value = $Url.Trim()
+ $idx = $value.IndexOfAny([char[]]@('?', '#'))
+ if ($idx -lt 0) {
+ return $value.TrimEnd('/')
+ }
+ return $value.Substring(0, $idx).TrimEnd('/') + $value.Substring($idx)
+}
+
+# Explicit torch-index pin (UNSLOTH_TORCH_INDEX_URL / _FAMILY), shared by the stale-venv check
+# and install selection so a pinned index wins over GPU probing (parity with the other
+# installers). URL is verbatim; _FAMILY is the leaf joined to the mirror base.
+function Get-PinnedTorchIndexUrl {
+ if (-not [string]::IsNullOrWhiteSpace($env:UNSLOTH_TORCH_INDEX_URL)) {
+ return (Trim-IndexPathSlashes $env:UNSLOTH_TORCH_INDEX_URL)
+ }
+ if (-not [string]::IsNullOrWhiteSpace($env:UNSLOTH_TORCH_INDEX_FAMILY)) {
+ $base = if ($env:UNSLOTH_PYTORCH_MIRROR) { $env:UNSLOTH_PYTORCH_MIRROR.TrimEnd('/') } else { "https://download.pytorch.org/whl" }
+ return "$base/$($env:UNSLOTH_TORCH_INDEX_FAMILY.Trim().Trim('/'))"
+ }
+ return $null
+}
+
+# Last path segment of a wheel index URL, query/fragment dropped first so a token-authenticated
+# pin (.../cu128?token=x) classifies as cu128 (else it reinstalls every update). Classification
+# only. Shared with the py / install.sh leaf extractors.
+function Get-TorchIndexLeaf {
+ param([string]$Url)
+ if ([string]::IsNullOrWhiteSpace($Url)) { return $null }
+ $path = ($Url -split '[?#]', 2)[0]
+ if ([string]::IsNullOrWhiteSpace($path)) { return $null }
+ return ($path.TrimEnd('/') -split '/')[-1].ToLowerInvariant()
+}
+
+# Redact index-URL credentials (userinfo + ?query= + #fragment) from captured installer
+# output before printing on failure; uv/pip errors echo the failing --index-url verbatim.
+# Mirrors the other installers. Verbose mode streams uncaptured, so it isn't redacted.
+function Redact-InstallOutput {
+ param([string]$Text)
+ if (-not $Text) { return $Text }
+ $Text = $Text -replace '(https?://)[^/@\s`]+@', '$1@'
+ $Text = $Text -replace '([?&][^=\s&`]+)=[^\s`]+', '$1='
+ # A #token=... fragment is as sensitive as a query; URL-anchored.
+ return $Text -replace '(https?://[^\s`#]+)#[^\s`]+', '$1#'
+}
+
+# AMD per-arch leaves needing the torch 2.11 floor (the _grouped_mm <2.11 bug). MUST match
+# the install-spec path below and the other installers; other leaves ship <2.11 and stay default.
+function Test-RocmGfx211Leaf {
+ param([string]$Leaf)
+ return @('gfx120x-all', 'gfx1151', 'gfx1150') -contains $Leaf
+}
+
+# rocmX.Y versions KNOWN to ship torch 2.11: rocm7.2 only today. Do NOT floor an unknown newer
+# rocm speculatively. MUST match _ROCM_KNOWN_TORCH211_VERSIONS and the rocm7.2 leaf elsewhere.
+function Test-RocmKnown211Version {
+ param([int]$Major, [int]$Minor)
+ return ($Major -eq 7 -and $Minor -eq 2)
+}
+
+# True only for a real CUDA family leaf: "cu" + digits (cu118, cu128, ...). A bare -like 'cu*'
+# would match "custom"/"current" and rebuild the venv every run. Mirrors _is_cuda_family_leaf.
+function Test-CudaFamilyLeaf {
+ param([string]$Leaf)
+ if ([string]::IsNullOrWhiteSpace($Leaf)) { return $false }
+ # EXACT cu+digits: cu128-private routes through the unknown-leaf path instead.
+ return $Leaf -match '^cu[0-9]+$'
+}
+
+# True only for a real pip ROCm family leaf: EXACT rocm[.] or a gfx leaf. A leaf
+# that merely STARTS with rocm (rocm-rel-7.2.1, rocm7.2-private) is a custom pin the verbatim
+# path owns, so anchor the match. Mirrors _is_pip_rocm_family_leaf / install.sh.
+function Test-PipRocmFamilyLeaf {
+ param([string]$Leaf)
+ if ([string]::IsNullOrWhiteSpace($Leaf)) { return $false }
+ # gfx must be followed by a digit (an architecture leaf); gfx-private is custom.
+ return ($Leaf -match '^gfx[0-9]') -or ($Leaf -match '^rocm[0-9]+(\.[0-9]+)?$')
+}
+
+# Stale-venv ROCm comparison for a pinned gfx*/rocm* index. Returns @{ Expected; Installed } so
+# the caller rebuilds when they differ. Mirrors _rocm_pin_family_mismatch (same rocmX.Y / gfx
+# cases). An untagged (no +rocm) wheel never satisfies a ROCm pin -> stale.
+function Get-RocmPinStaleTags {
+ param([string]$PinLeaf, [string]$TorchVersion)
+ $_pinRocm = [regex]::Match($PinLeaf, '^rocm(\d+)\.(\d+)')
+ $_pinVer = if ($_pinRocm.Success) { "$($_pinRocm.Groups[1].Value).$($_pinRocm.Groups[2].Value)" } else { $null }
+ # The family classifier accepts a major-only rocm leaf too (rocm7).
+ $_pinMajorOnly = [regex]::Match($PinLeaf, '^rocm(\d+)$')
+ # Installed rocm version and whether the wheel is a per-arch (three-part) build.
+ $_instRocm = [regex]::Match($TorchVersion, '\+rocm(\d+)\.(\d+)')
+ $_instVer = if ($_instRocm.Success) { "$($_instRocm.Groups[1].Value).$($_instRocm.Groups[2].Value)" } else { $null }
+ $_instPerArch = [regex]::IsMatch($TorchVersion, '\+rocm\d+\.\d+\.\d+')
+ # A ROCm build MUST carry a +rocm tag; an untagged wheel can't satisfy any ROCm pin.
+ $_instHasRocm = [regex]::IsMatch($TorchVersion, '\+rocm')
+ $_instRel = [regex]::Match($TorchVersion, '^(\d+)\.(\d+)')
+ $_instIs211 = $false
+ if ($_instRel.Success) {
+ $_instIs211 = ([int]$_instRel.Groups[1].Value -gt 2) -or ([int]$_instRel.Groups[1].Value -eq 2 -and [int]$_instRel.Groups[2].Value -ge 11)
+ }
+
+ if ($PinLeaf -like 'gfx*') {
+ if (Test-RocmGfx211Leaf $PinLeaf) {
+ # Expect the AMD per-arch (three-part) 2.11 wheel: satisfied only when BOTH
+ # a 2.11 release AND a three-part rocm tag are installed.
+ $installed = if ($_instIs211 -and $_instPerArch) { "rocm-perarch(torch>=2.11)" } else { "rocm-generic-or-old" }
+ return @{ Expected = "rocm-perarch(torch>=2.11)"; Installed = $installed }
+ }
+ # Non-2.11 gfx leaf (<2.11 spec): stale on an untagged wheel or a 2.11+ build.
+ $installed = if (-not $_instHasRocm) { "not-rocm" } elseif ($_instIs211) { "rocm(torch>=2.11)" } else { "rocm(torch<2.11)" }
+ return @{
+ Expected = "rocm(torch<2.11)"
+ Installed = $installed
+ }
+ }
+
+ # Major-only rocm pin (rocm7): compare majors only -- a +rocm6.4 wheel under a rocm7
+ # pin is stale, any +rocm7.x wheel satisfies it (no pinned minor to compare, and the
+ # 2.11-line fallback below would invert both verdicts). Mirrors _rocm_pin_family_mismatch.
+ if ($_pinMajorOnly.Success) {
+ $_pinMaj = [int]$_pinMajorOnly.Groups[1].Value
+ if ($_instVer) {
+ $_instMaj = [int]$_instRocm.Groups[1].Value
+ $expected = if ($_instMaj -eq $_pinMaj) { "rocm$_instVer" } else { "rocm$_pinMaj.x" }
+ return @{ Expected = $expected; Installed = "rocm$_instVer" }
+ }
+ # Untagged wheel never satisfies a ROCm pin; a +rocm tag with an unreadable
+ # version is accepted (matches the lenient unreadable fallback below).
+ $installed = if ($_instHasRocm) { "rocm" } else { "not-rocm" }
+ return @{ Expected = "rocm"; Installed = $installed }
+ }
+
+ # rocmX.Y pin.
+ if ($_pinVer -and $_instVer) {
+ # Both readable: exact compare. When they match AND the pin is KNOWN-2.11, the
+ # installed release must also be 2.11 (a +rocm7.2 wheel drifted to 2.12 shares the
+ # tag but violates the spec), so fold the release into the tag. Mirrors _rocm_pin_family_mismatch.
+ $_pinKnown211 = Test-RocmKnown211Version -Major ([int]$_pinRocm.Groups[1].Value) -Minor ([int]$_pinRocm.Groups[2].Value)
+ $_instOn211 = $_instRel.Success -and [int]$_instRel.Groups[1].Value -eq 2 -and [int]$_instRel.Groups[2].Value -eq 11
+ if ($_pinKnown211 -and -not $_instOn211) {
+ return @{ Expected = "rocm$_pinVer(torch2.11)"; Installed = "rocm$_instVer(torch-off-2.11)" }
+ }
+ return @{ Expected = "rocm$_pinVer"; Installed = "rocm$_instVer" }
+ }
+ $_pinNeeds211 = $false
+ if ($_pinRocm.Success) {
+ # Only KNOWN-2.11 rocm (rocm7.2) is on the 2.11 line (no speculative floor).
+ # Matches _ROCM_KNOWN_TORCH211_VERSIONS.
+ $_pinNeeds211 = Test-RocmKnown211Version -Major ([int]$_pinRocm.Groups[1].Value) -Minor ([int]$_pinRocm.Groups[2].Value)
+ }
+ # Fallback (installed rocm version unreadable): compare on the 2.11 line; an untagged
+ # wheel never satisfies a rocmX.Y pin -> stale.
+ $installed = if (-not $_instHasRocm) { "not-rocm" } elseif ($_instIs211) { "rocm(torch>=2.11)" } else { "rocm(torch<2.11)" }
+ return @{
+ Expected = if ($_pinNeeds211) { "rocm(torch>=2.11)" } else { "rocm(torch<2.11)" }
+ Installed = $installed
+ }
+}
+
# VS generator -> MSBuild BuildCustomizations dir; toolset tracks the VS major
# (18->v180, 17->v170), defaulting to v170 when unparseable.
function Get-VcBuildCustomizationsDir {
@@ -813,11 +974,14 @@ function Invoke-SetupCommand {
# Merge stderr into stdout so progress/warning output stays visible
# without flipping $? on successful native commands (PS 5.1 treats
# stderr records as errors that set $? = $false even on exit code 0).
- & $Command 2>&1 | Out-Host
+ # Redact per record: uv/pip echo index URLs (credentials and all) in
+ # their errors, and verbose mode must not bypass the quiet path's
+ # redaction. ForEach-Object/Out-Host leave $LASTEXITCODE untouched.
+ & $Command 2>&1 | ForEach-Object { Redact-InstallOutput "$_" } | Out-Host
} else {
$output = & $Command 2>&1 | Out-String
if ($LASTEXITCODE -ne 0) {
- Write-Host $output -ForegroundColor Red
+ Write-Host (Redact-InstallOutput $output) -ForegroundColor Red
}
}
return [int]$LASTEXITCODE
@@ -2535,6 +2699,8 @@ if ((Test-Path -LiteralPath $VenvDir -PathType Container) -and -not $NoTorchMode
$VenvPyExe = Join-Path $VenvDir "Scripts\python.exe"
$installedTorchTag = $null
$shouldRebuild = $false
+ # Set when a stale venv under a pin is repaired in place (force-reinstall) not wiped.
+ $script:PinChangedForceReinstall = $false
if (Test-Path -LiteralPath $VenvPyExe) {
try {
@@ -2551,10 +2717,14 @@ if ((Test-Path -LiteralPath $VenvDir -PathType Container) -and -not $NoTorchMode
if ($finished -and $proc.ExitCode -eq 0 -and $torchVer) {
if ($torchVer -match '\+(cu\d+)') {
$installedTorchTag = $Matches[1]
+ } elseif ($torchVer -match '\+rocm') {
+ # Any +rocm / gfx wheel -> generic "rocm" flavor (the exact version is
+ # repaired later by install_python_stack.py; here we only need the flavor).
+ $installedTorchTag = "rocm"
} elseif ($torchVer -match '\+cpu') {
$installedTorchTag = "cpu"
} else {
- # Untagged wheel (plain "2.x.y" from PyPI) -- treat as cpu
+ # Untagged wheel (plain "2.x.y" from PyPI) -> cpu.
$installedTorchTag = "cpu"
}
} else {
@@ -2570,12 +2740,71 @@ if ((Test-Path -LiteralPath $VenvDir -PathType Container) -and -not $NoTorchMode
}
if (-not $shouldRebuild) {
- $expectedTorchTag = if ($HasNvidiaSmi) { Get-PytorchCudaTag } else { "cpu" }
- if ($installedTorchTag -and $installedTorchTag -ne $expectedTorchTag) {
+ $_pinnedIdx = Get-PinnedTorchIndexUrl
+ $_expectedKnown = $true
+ if ($_pinnedIdx) {
+ $_pinLeaf = Get-TorchIndexLeaf $_pinnedIdx
+ # Digit-gated like the install selection: a custom rocm-* leaf (rocm-current /
+ # rocm-rel-7.2.1) is NOT a ROCm family and must not be stale-compared.
+ if (Test-PipRocmFamilyLeaf $_pinLeaf) {
+ # Don't collapse a pinned ROCm/gfx leaf to a generic "rocm" (would mask a family
+ # change, rocm6.4 -> gfx1151). Get-RocmPinStaleTags uses the SAME 2.11 allowlist
+ # as the install path, so a gfx110X-all/gfx90a/gfx908 pin on a <2.11 wheel is NOT stale.
+ $_rocmTags = Get-RocmPinStaleTags -PinLeaf $_pinLeaf -TorchVersion $torchVer
+ $expectedTorchTag = $_rocmTags.Expected
+ $installedTorchTag = $_rocmTags.Installed
+ } elseif ((Test-CudaFamilyLeaf $_pinLeaf) -or $_pinLeaf -eq 'cpu') {
+ # cu*/cpu leaves stay specific so a cu126-vs-cu128 mismatch rebuilds;
+ # /custom and /current fall through to the unknown-index branch below.
+ $expectedTorchTag = $_pinLeaf
+ } else {
+ # Custom index whose leaf is not a torch flavor (a /simple mirror): the
+ # flavor can't be inferred, so never treat the venv as stale over it.
+ $_expectedKnown = $false
+ $expectedTorchTag = $installedTorchTag
+ }
+ } elseif ($HasNvidiaSmi) {
+ $expectedTorchTag = Get-PytorchCudaTag
+ } elseif ($HasROCm -or $script:ROCmGfxArch) {
+ # AMD/ROCm host with no explicit pin: an existing +rocm wheel is correct (gfx arch
+ # counts even when $HasROCm is false). But only the arches the install path maps to a
+ # repo.amd.com index get ROCm torch; an unmapped arch installs CPU, so expect "cpu"
+ # for those or a correct CPU venv rebuilds every update.
+ $_rocmWheelArches = @(
+ "gfx1201", "gfx1200", # RDNA 4
+ "gfx1151", "gfx1150", # RDNA 3.5 (Strix Halo/Point)
+ "gfx1103", "gfx1102", "gfx1101", "gfx1100", # RDNA 3
+ "gfx90a", "gfx908" # MI200 / MI100
+ )
+ if ($script:ROCmGfxArch -and ($_rocmWheelArches -contains $script:ROCmGfxArch)) {
+ # A correct +rocm wheel is not stale. A CPU wheel on a supported AMD arch is
+ # NOT wiped either (the AMD Windows ROCm override below upgrades it in place);
+ # expect "cpu" for that case. A wrong CUDA wheel still rebuilds.
+ if ($installedTorchTag -eq "cpu") {
+ $expectedTorchTag = "cpu"
+ } else {
+ $expectedTorchTag = "rocm"
+ }
+ } else {
+ $expectedTorchTag = "cpu"
+ }
+ } else {
+ $expectedTorchTag = "cpu"
+ }
+ if ($_expectedKnown -and $installedTorchTag -and $installedTorchTag -ne $expectedTorchTag) {
$shouldRebuild = $true
}
}
+ # A stale venv under a pin whose torch still imports is repaired IN PLACE (the dependency
+ # pass force-reinstalls from the pin). The rebuild path wipes the venv and would strand a
+ # direct `studio update`; only a broken venv or an unpinned drift wipes.
+ if ($shouldRebuild -and $_pinnedIdx -and $installedTorchTag) {
+ substep "Torch-index pin changed ($installedTorchTag) -- reinstalling torch from the pin in place." "Cyan"
+ $script:PinChangedForceReinstall = $true
+ $shouldRebuild = $false
+ }
+
if ($shouldRebuild) {
$reason = if ($installedTorchTag) { "torch $installedTorchTag != required $expectedTorchTag" } else { "torch could not be imported" }
if ($InstallerManagedSetup) {
@@ -2653,23 +2882,41 @@ if (Get-Command uv -ErrorAction SilentlyContinue) {
# Helper: install a package, preferring uv with pip fallback
function Fast-Install {
param([Parameter(ValueFromRemainingArguments=$true)]$Args_)
- if ($UseUv) {
- $VenvPy = (Get-Command python).Source
- # An explicit --index-url must win. Inherited uv index env vars otherwise
- # override it and pull CPU torch over the CUDA/ROCm build (#6898), so drop
- # them only for index-pinned installs; mirrors still apply elsewhere.
- $saved = @{}
- if (@($Args_) -contains '--index-url') {
- foreach ($n in 'UV_DEFAULT_INDEX', 'UV_INDEX_URL', 'UV_INDEX', 'UV_EXTRA_INDEX_URL') {
- $saved[$n] = [Environment]::GetEnvironmentVariable($n)
- Remove-Item "Env:$n" -ErrorAction SilentlyContinue
- }
+ # An explicit --index-url must win: inherited uv index vars otherwise pull CPU torch over
+ # the CUDA/ROCm build (#6898), so drop them for pinned installs (scrub covers the whole
+ # function since the pip fallback honours PIP_* too). UV_TORCH_BACKEND / UV_FIND_LINKS also
+ # reroute; UV_NO_CONFIG=1 (+ dropping UV_CONFIG_FILE) stops a uv.toml index outranking the
+ # pin (uv 0.10); PIP_NO_INDEX / PIP_INDEX_URL would defeat the pinned --index-url in pip.
+ $saved = @{}
+ $pinned = @($Args_) -contains '--index-url'
+ if ($pinned) {
+ foreach ($n in 'UV_DEFAULT_INDEX', 'UV_INDEX_URL', 'UV_INDEX', 'UV_EXTRA_INDEX_URL',
+ 'UV_TORCH_BACKEND', 'UV_FIND_LINKS', 'PIP_EXTRA_INDEX_URL', 'PIP_FIND_LINKS',
+ 'PIP_NO_INDEX', 'PIP_INDEX_URL',
+ 'UV_CONFIG_FILE', 'UV_NO_CONFIG', 'PIP_CONFIG_FILE') {
+ $saved[$n] = [Environment]::GetEnvironmentVariable($n)
+ Remove-Item "Env:$n" -ErrorAction SilentlyContinue
}
- try { $result = & uv pip install --python $VenvPy @Args_ 2>&1 }
- finally { foreach ($n in $saved.Keys) { if ($null -ne $saved[$n]) { Set-Item "Env:$n" $saved[$n] } } }
- if ($LASTEXITCODE -eq 0) { return }
+ $env:UV_NO_CONFIG = '1'
+ # A `pip config` global.extra-index-url still adds indexes to the pip FALLBACK;
+ # PIP_CONFIG_FILE = 'nul' (Windows devnull) loads NO config (uv ignores pip config).
+ $env:PIP_CONFIG_FILE = 'nul'
+ }
+ try {
+ if ($UseUv) {
+ $VenvPy = (Get-Command python).Source
+ $result = & uv pip install --python $VenvPy @Args_ 2>&1
+ if ($LASTEXITCODE -eq 0) { return }
+ }
+ & python -m pip install @Args_ 2>&1
+ }
+ finally {
+ if ($pinned) {
+ Remove-Item "Env:UV_NO_CONFIG" -ErrorAction SilentlyContinue
+ Remove-Item "Env:PIP_CONFIG_FILE" -ErrorAction SilentlyContinue
+ }
+ foreach ($n in $saved.Keys) { if ($null -ne $saved[$n]) { Set-Item "Env:$n" $saved[$n] } }
}
- & python -m pip install @Args_ 2>&1
}
# ── Check if Python deps need updating ──
@@ -2752,6 +2999,10 @@ sys.exit(0 if (major, minor) >= (4, 14) else 1)
# pip install unsloth 2>&1 | Out-Null
# }
+# A torch-index pin change repairs in place: force the dependency pass so the torch install
+# below force-reinstalls from the new pin (else the fast path keeps the old wheel).
+if ($script:PinChangedForceReinstall) { $SkipPythonDeps = $false }
+
if (-not $SkipPythonDeps) {
if ($script:UnslothVerbose) {
@@ -2779,7 +3030,13 @@ $env:TORCHINDUCTOR_CACHE_DIR = $TorchCacheDir
[Environment]::SetEnvironmentVariable('TORCHINDUCTOR_CACHE_DIR', $TorchCacheDir, 'User')
substep "TORCHINDUCTOR_CACHE_DIR set to $TorchCacheDir (avoids MAX_PATH issues)"
-if ($HasNvidiaSmi) {
+# Explicit pin (URL or family) wins over GPU probing and suppresses the AMD reroute below;
+# matches install.sh / install.ps1 / install_python_stack.py.
+$PinnedTorchIndexUrl = Get-PinnedTorchIndexUrl
+$TorchIndexPinned = [bool]$PinnedTorchIndexUrl
+if ($PinnedTorchIndexUrl) {
+ $CuTag = Get-TorchIndexLeaf $PinnedTorchIndexUrl
+} elseif ($HasNvidiaSmi) {
$CuTag = Get-PytorchCudaTag
} else {
$CuTag = "cpu"
@@ -2800,7 +3057,7 @@ $ROCmIndexUrl = $null
# SDK -- which flips Unsloth out of chat-only (CHAT_ONLY) and enables Train/Export.
# Gating on $HasROCm alone left Strix Halo / Radeon 8060S on CPU torch; a failed
# ROCm install still falls back to CPU below, so this is safe.
-if (($HasROCm -or $ROCmGfxArch) -and $CuTag -eq "cpu") {
+if (-not $TorchIndexPinned -and ($HasROCm -or $ROCmGfxArch) -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
@@ -2850,8 +3107,45 @@ if (($HasROCm -or $ROCmGfxArch) -and $CuTag -eq "cpu") {
}
}
+# A pinned gfx*/rocm index skips the auto-reroute above; route it through the ROCm install path
+# with the same floor/companions the unpinned AMD path uses (mirrors install.ps1), else the CUDA
+# branch installs bare torch and resolves a known-bad wheel for gfx115x/gfx120x/rocm>=7.2.
+if ($TorchIndexPinned -and -not $ROCmIndexUrl -and $PinnedTorchIndexUrl) {
+ $_pinLeaf = Get-TorchIndexLeaf $PinnedTorchIndexUrl
+ $_pinRocm211 = $false
+ # Anchor the match ($) so a suffixed custom leaf (rocm7.2-private) falls through to the
+ # verbatim install instead of being floored by its rocm7.2 prefix.
+ if ($_pinLeaf -match '^rocm(\d+)\.(\d+)$') {
+ # Only KNOWN-2.11 rocm (rocm7.2) gets the floor (no speculative floor). Matches
+ # Test-RocmKnown211Version / _ROCM_KNOWN_TORCH211_VERSIONS.
+ $_pinRocm211 = Test-RocmKnown211Version -Major ([int]$Matches[1]) -Minor ([int]$Matches[2])
+ }
+ # Only the 2.11 gfx arches need the floor; others publish <2.11 and stay bare. Reuse
+ # Test-RocmGfx211Leaf so this allowlist and the stale-venv check never diverge.
+ $_pinGfx211 = Test-RocmGfx211Leaf $_pinLeaf
+ if ($_pinGfx211 -or $_pinRocm211) {
+ $ROCmIndexUrl = $PinnedTorchIndexUrl
+ $ROCmTorchSpec = "torch>=2.11.0,<2.12.0"
+ $ROCmVisionSpec = "torchvision>=0.26.0,<0.27.0"
+ $ROCmAudioSpec = "torchaudio>=2.11.0,<2.12.0"
+ substep "pinned ROCm index ($_pinLeaf) -- enforcing $ROCmTorchSpec" "Cyan"
+ } elseif (Test-PipRocmFamilyLeaf $_pinLeaf) {
+ # Other gfx / older rocm (<=7.1) ship torch <2.11; route via the ROCm path with
+ # bare specs. Only EXACT rocm and gfx* are --index-url families; a suffixed
+ # leaf stays on the verbatim path. Mirrors install.ps1 / _is_pip_rocm_family_leaf.
+ $ROCmIndexUrl = $PinnedTorchIndexUrl
+ $ROCmTorchSpec = "torch"
+ $ROCmVisionSpec = "torchvision"
+ $ROCmAudioSpec = "torchaudio"
+ }
+}
+
$PyTorchWhlBase = if ($env:UNSLOTH_PYTORCH_MIRROR) { $env:UNSLOTH_PYTORCH_MIRROR.TrimEnd('/') } else { "https://download.pytorch.org/whl" }
+# A full URL pin is used verbatim; a family pin already set $CuTag. A pinned ROCm install
+# goes through $ROCmIndexUrl; on failure the fallback uses the CPU index, not the ROCm pin.
+$TorchInstallIndexUrl = if ($ROCmIndexUrl) { "$PyTorchWhlBase/cpu" } elseif ($PinnedTorchIndexUrl) { $PinnedTorchIndexUrl } else { "$PyTorchWhlBase/$CuTag" }
+
$ROCmCpuFallback = $false
if ($ROCmIndexUrl) {
substep "installing PyTorch (AMD ROCm, $ROCmGfxArch)..."
@@ -2859,7 +3153,7 @@ if ($ROCmIndexUrl) {
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
+ Fast-Install $ROCmTorchSpec $ROCmVisionSpec $ROCmAudioSpec --force-reinstall --index-url $ROCmIndexUrl | ForEach-Object { Redact-InstallOutput "$_" } | Out-Host
$torchInstallExit = $LASTEXITCODE
$output = ""
} else {
@@ -2868,7 +3162,7 @@ if ($ROCmIndexUrl) {
}
if ($torchInstallExit -ne 0) {
Write-Host "[WARN] AMD ROCm PyTorch install failed -- falling back to CPU" -ForegroundColor Yellow
- Write-Host $output -ForegroundColor Yellow
+ Write-Host (Redact-InstallOutput $output) -ForegroundColor Yellow
$ROCmIndexUrl = $null
$ROCmCpuFallback = $true
} else {
@@ -2878,42 +3172,70 @@ if ($ROCmIndexUrl) {
}
}
-if (-not $ROCmIndexUrl -and $CuTag -eq "cpu") {
+if (-not $ROCmIndexUrl -and ($CuTag -eq "cpu" -or $ROCmCpuFallback)) {
substep "installing PyTorch (CPU-only)..."
- # After an AMD ROCm fallback, force-reinstall so a partially-installed ROCm torch
- # (which still satisfies the CPU torch>= range) is replaced by the CPU build. Skip
- # the forced reinstall on a genuine CPU-only host so the common path stays fast.
- # Build the array directly: an if-expression collapses @("x") to a scalar string,
- # which @splat would then enumerate char-by-char into broken single-letter args.
+ # After an AMD ROCm fallback, force-reinstall so a partial ROCm torch (which satisfies the
+ # CPU torch>= range) is replaced by the CPU build; skip on a genuine CPU host to stay fast.
+ # $ROCmCpuFallback matters when a PINNED ROCm index failed ($CuTag is still the rocm leaf).
+ # Build the array directly: an if-expression collapses @("x") to a scalar @splat would
+ # enumerate char-by-char.
$cpuForce = @()
if ($ROCmCpuFallback) { $cpuForce = @("--force-reinstall") }
+ # --force-reinstall on a pin change: a stale +cu / +rocm wheel still satisfies the CPU
+ # torch>= range, so uv would keep it and only swap companions.
+ if ($script:PinChangedForceReinstall) { $cpuForce = @("--force-reinstall") }
+ # A PINNED cpu index installs the bounded trio (parity with _CPU_TORCH_PKG_SPEC): the /cpu
+ # index serves newer torch, and _ensure_cpu_torch keeps any CPU build, so a bare trio could
+ # land an unsupported version. Unpinned CPU hosts keep the bare trio (pre-pin behavior).
+ $cpuTorchSpec = "torch"; $cpuVisionSpec = "torchvision"; $cpuAudioSpec = "torchaudio"
+ if ($TorchIndexPinned) {
+ $cpuTorchSpec = "torch>=2.4,<2.12.0"
+ $cpuVisionSpec = "torchvision>=0.19,<0.27.0"
+ $cpuAudioSpec = "torchaudio>=2.4,<2.12.0"
+ }
if ($script:UnslothVerbose) {
- Fast-Install torch torchvision torchaudio @cpuForce --index-url "$PyTorchWhlBase/cpu"
+ Fast-Install $cpuTorchSpec $cpuVisionSpec $cpuAudioSpec @cpuForce --index-url $TorchInstallIndexUrl | ForEach-Object { Redact-InstallOutput "$_" } | Out-Host
$torchInstallExit = $LASTEXITCODE
$output = ""
} else {
- $output = Fast-Install torch torchvision torchaudio @cpuForce --index-url "$PyTorchWhlBase/cpu" | Out-String
+ $output = Fast-Install $cpuTorchSpec $cpuVisionSpec $cpuAudioSpec @cpuForce --index-url $TorchInstallIndexUrl | Out-String
$torchInstallExit = $LASTEXITCODE
}
if ($torchInstallExit -ne 0) {
Write-Host "[FAILED] PyTorch install failed (exit code $torchInstallExit)" -ForegroundColor Red
- Write-Host $output -ForegroundColor Red
+ Write-Host (Redact-InstallOutput $output) -ForegroundColor Red
exit 1
}
} elseif (-not $ROCmIndexUrl) {
substep "installing PyTorch with CUDA support ($CuTag)..."
substep "(This download is ~2.8 GB -- may take a few minutes)"
+ # --force-reinstall on a pin change: an installed cuXXX wheel satisfies the bare torch
+ # requirement (PEP 440 ignores the +cuXXX tag), so without it a changed CUDA pin (cu126
+ # -> cu128) never applies.
+ $cudaForce = @()
+ if ($script:PinChangedForceReinstall) { $cudaForce = @("--force-reinstall") }
+ # An unknown-leaf custom pin (/simple, /current) routes here with $CuTag as that leaf. Bound
+ # the trio like the fresh custom-pin paths so a mirror can't pull an ABI-newer companion
+ # against the capped torch. Known cu* leaves keep bare specs.
+ $cudaTorchSpec = "torch"
+ $cudaVisionSpec = "torchvision"
+ $cudaAudioSpec = "torchaudio"
+ if ($TorchIndexPinned -and -not (Test-CudaFamilyLeaf $CuTag)) {
+ $cudaTorchSpec = "torch>=2.4,<2.11.0"
+ $cudaVisionSpec = "torchvision>=0.19,<0.26.0"
+ $cudaAudioSpec = "torchaudio>=2.4,<2.11.0"
+ }
if ($script:UnslothVerbose) {
- Fast-Install torch torchvision torchaudio --index-url "$PyTorchWhlBase/$CuTag"
+ Fast-Install $cudaTorchSpec $cudaVisionSpec $cudaAudioSpec @cudaForce --index-url $TorchInstallIndexUrl | ForEach-Object { Redact-InstallOutput "$_" } | Out-Host
$torchInstallExit = $LASTEXITCODE
$output = ""
} else {
- $output = Fast-Install torch torchvision torchaudio --index-url "$PyTorchWhlBase/$CuTag" | Out-String
+ $output = Fast-Install $cudaTorchSpec $cudaVisionSpec $cudaAudioSpec @cudaForce --index-url $TorchInstallIndexUrl | Out-String
$torchInstallExit = $LASTEXITCODE
}
if ($torchInstallExit -ne 0) {
Write-Host "[FAILED] PyTorch CUDA install failed (exit code $torchInstallExit)" -ForegroundColor Red
- Write-Host $output -ForegroundColor Red
+ Write-Host (Redact-InstallOutput $output) -ForegroundColor Red
exit 1
}
@@ -2929,7 +3251,7 @@ if (-not $ROCmIndexUrl -and $CuTag -eq "cpu") {
}
if ($tritonInstallExit -ne 0) {
substep "Triton install failed -- torch.compile may not work" "Yellow"
- Write-Host $output -ForegroundColor Yellow
+ Write-Host (Redact-InstallOutput $output) -ForegroundColor Yellow
} else {
substep "Triton for Windows installed (enables torch.compile)"
}
@@ -3026,7 +3348,7 @@ foreach ($pkg in @("transformers==5.3.0", "huggingface_hub==1.8.0", "hf_xet==1.4
}
if ($t5PkgExit -ne 0) {
Write-Host "[FAIL] Could not install $pkg into .venv_t5_530/" -ForegroundColor Red
- Write-Host $output -ForegroundColor Red
+ Write-Host (Redact-InstallOutput $output) -ForegroundColor Red
$ErrorActionPreference = $prevEAP_t5
exit 1
}
@@ -3061,7 +3383,7 @@ foreach ($pkg in @("transformers==5.5.0", "huggingface_hub==1.8.0", "hf_xet==1.4
}
if ($t5PkgExit -ne 0) {
Write-Host "[FAIL] Could not install $pkg into .venv_t5_550/" -ForegroundColor Red
- Write-Host $output -ForegroundColor Red
+ Write-Host (Redact-InstallOutput $output) -ForegroundColor Red
$ErrorActionPreference = $prevEAP_t5
exit 1
}
@@ -3096,7 +3418,7 @@ foreach ($pkg in @("transformers==5.10.2", "huggingface_hub==1.8.0", "hf_xet==1.
}
if ($t5PkgExit -ne 0) {
Write-Host "[FAIL] Could not install $pkg into .venv_t5_510/" -ForegroundColor Red
- Write-Host $output -ForegroundColor Red
+ Write-Host (Redact-InstallOutput $output) -ForegroundColor Red
$ErrorActionPreference = $prevEAP_t5
exit 1
}
diff --git a/tests/python/test_cross_platform_parity.py b/tests/python/test_cross_platform_parity.py
index 666ea7ce10..b3a9b99c55 100644
--- a/tests/python/test_cross_platform_parity.py
+++ b/tests/python/test_cross_platform_parity.py
@@ -10,6 +10,8 @@ import pytest
REPO_ROOT = Path(__file__).resolve().parents[2]
INSTALL_SH = REPO_ROOT / "install.sh"
INSTALL_PS1 = REPO_ROOT / "install.ps1"
+SETUP_PS1 = REPO_ROOT / "studio" / "setup.ps1"
+STACK_PY = REPO_ROOT / "studio" / "install_python_stack.py"
class TestNoTorchBackendAutoInInstallSh:
@@ -180,3 +182,607 @@ class TestUvBytecodeCompileTimeout:
assert (
'$env:UV_COMPILE_BYTECODE_TIMEOUT = "180"' in text
), "install.ps1 should default UV_COMPILE_BYTECODE_TIMEOUT"
+
+
+class TestTorchIndexOverrideParity:
+ """Every installer must honor UNSLOTH_TORCH_INDEX_URL / _FAMILY so a pinned wheel
+ index wins over GPU probing on all platforms (no asymmetric, per-OS coverage)."""
+
+ @pytest.mark.parametrize(
+ "path",
+ [INSTALL_SH, INSTALL_PS1, SETUP_PS1, STACK_PY],
+ ids = ["install.sh", "install.ps1", "setup.ps1", "install_python_stack.py"],
+ )
+ def test_installer_reads_override_env(self, path):
+ text = path.read_text(encoding = "utf-8")
+ for var in ("UNSLOTH_TORCH_INDEX_URL", "UNSLOTH_TORCH_INDEX_FAMILY"):
+ assert var in text, f"{path.name} does not honor {var}"
+
+ @pytest.mark.parametrize(
+ "path",
+ [INSTALL_PS1, SETUP_PS1],
+ ids = ["install.ps1", "setup.ps1"],
+ )
+ def test_amd_reroute_guarded_when_pinned(self, path):
+ # The AMD ROCm reroute must be skipped when the index is explicitly pinned,
+ # so an explicit cpu / cu* / rocm pin on an AMD host is not overwritten.
+ text = path.read_text(encoding = "utf-8")
+ assert (
+ "TorchIndexPinned" in text
+ ), f"{path.name} should gate the AMD ROCm reroute on a pinned-index flag"
+
+ def test_cuda_pin_overrides_cvd_hide_gate(self):
+ # A pinned cu* index skips ALL host-GPU probing, so the CUDA repair must clear the
+ # CUDA_VISIBLE_DEVICES hide gate too (else the GPU-less CI case bails).
+ text = STACK_PY.read_text(encoding = "utf-8")
+ m = re.search(r"def _ensure_cuda_torch\(\).*?(?=\ndef )", text, re.DOTALL)
+ assert m, "could not locate _ensure_cuda_torch"
+ body = m.group(0)
+ assert "_cuda_pinned" in body, (
+ "_ensure_cuda_torch should compute a CUDA-pin flag so the pin can "
+ "override the CVD hide gate"
+ )
+ assert re.search(
+ r"if not _cuda_pinned and _cvd is not None", body
+ ), "the CVD hide gate must be bypassed when a CUDA index is pinned"
+
+ def test_cpu_repair_pins_supported_torch_range(self):
+ # The explicit-CPU repair must use the bounded CPU/CUDA spec, not a bare trio (the
+ # /cpu index serves torch 2.11+, so a bare install could resolve out of range).
+ text = STACK_PY.read_text(encoding = "utf-8")
+ m = re.search(r"def _ensure_cpu_torch\(\).*?(?=\ndef )", text, re.DOTALL)
+ assert m, "could not locate _ensure_cpu_torch"
+ body = m.group(0)
+ assert "_CPU_TORCH_PKG_SPEC" in body, (
+ "_ensure_cpu_torch should install the bounded _CPU_TORCH_PKG_SPEC, "
+ "not a bare torch/torchvision/torchaudio trio"
+ )
+
+ def test_setup_ps1_stale_check_gates_rocm_on_supported_arch(self):
+ # The stale check must expect ROCm torch only for arches the install path maps to a
+ # repo.amd.com index; expecting "rocm" for an unmapped arch marks a good CPU venv stale.
+ text = SETUP_PS1.read_text(encoding = "utf-8")
+ assert "_rocmWheelArches" in text, (
+ "setup.ps1 stale check should restrict the ROCm expected-tag to the "
+ "supported gfx wheel arches"
+ )
+
+
+class TestGfx211AllowlistParity:
+ """The gfx per-arch 2.11-floor leaves (gfx120X-all / gfx1151 / gfx1150) must be the
+ SAME set in every installer and its stale/mismatch check. When they diverged, a
+ pinned gfx110X-all / gfx90a / gfx908 wheel (<2.11) was force-reinstalled every update."""
+
+ EXPECTED = {"gfx120x-all", "gfx1151", "gfx1150"}
+
+ def test_install_sh_allowlist(self):
+ text = INSTALL_SH.read_text(encoding = "utf-8").lower()
+ # install.sh: the TORCH_CONSTRAINT case (rocm7.2|gfx120x-all|gfx1151|gfx1150).
+ m = re.search(r"rocm7\.2\|gfx120x-all\|gfx1151\|gfx1150", text)
+ assert m, "install.sh gfx-2.11 allowlist case not found / changed"
+
+ def test_install_ps1_allowlist(self):
+ text = INSTALL_PS1.read_text(encoding = "utf-8").lower()
+ m = re.search(r"@\('gfx120x-all',\s*'gfx1151',\s*'gfx1150'\)", text)
+ assert m, "install.ps1 $_pinGfx211 allowlist not found / changed"
+
+ def test_setup_ps1_defines_single_allowlist_helper(self):
+ # setup.ps1 must define the allowlist once (Test-RocmGfx211Leaf) and reuse it, so
+ # the stale check and install spec can't disagree.
+ text = SETUP_PS1.read_text(encoding = "utf-8")
+ assert (
+ "function Test-RocmGfx211Leaf" in text
+ ), "setup.ps1 should define a single Test-RocmGfx211Leaf allowlist helper"
+ assert re.search(
+ r"@\('gfx120x-all',\s*'gfx1151',\s*'gfx1150'\)", text.lower()
+ ), "Test-RocmGfx211Leaf should hold the gfx-2.11 allowlist"
+ assert "$_pinGfx211 = Test-RocmGfx211Leaf" in text, (
+ "setup.ps1 install-spec path should reuse Test-RocmGfx211Leaf, not "
+ "re-hardcode the allowlist (they must not diverge)"
+ )
+
+ def test_stack_py_allowlist(self):
+ text = STACK_PY.read_text(encoding = "utf-8").lower()
+ assert (
+ '"gfx120x-all", "gfx1151", "gfx1150"' in text
+ ), "install_python_stack.py _ROCM_GFX_TORCH211_LEAVES not found / changed"
+
+
+class TestCudaLeafDigitParity:
+ """A wheel-family leaf is CUDA only when it is "cu" + digits (cu118/cu128/...).
+ A bare cu* glob wrongly catches mirror leaves like /custom or /current; when
+ that happened the venv was marked stale and rebuilt on every run. Every
+ installer must require a digit after "cu" in its family/CUDA classification."""
+
+ def test_stack_py_requires_cu_digit(self):
+ text = STACK_PY.read_text(encoding = "utf-8")
+ # EXACT cu+digits: a custom leaf like cu128-private must route to the
+ # verbatim/unknown path, not be compared against the installed +cu128 tag.
+ assert re.search(
+ r'r"cu\[0-9\]\+"', text
+ ), "install_python_stack.py _is_cuda_family_leaf must fullmatch cu[0-9]+"
+
+ def test_setup_ps1_requires_cu_digit(self):
+ text = SETUP_PS1.read_text(encoding = "utf-8")
+ # EXACT cu+digits: cu128-private must not classify as CUDA (it would become
+ # the expected tag and rebuild the venv on every update).
+ assert re.search(
+ r"'\^cu\[0-9\]\+\$'", text
+ ), "setup.ps1 Test-CudaFamilyLeaf must match ^cu[0-9]+$, not a cu* prefix"
+ # The stale-venv branch must go through the digit-guarded helper.
+ assert (
+ "Test-CudaFamilyLeaf $_pinLeaf" in text
+ ), "setup.ps1 stale check should classify CUDA via Test-CudaFamilyLeaf"
+
+ def test_install_ps1_requires_cu_digit_in_gpu_branch(self):
+ text = INSTALL_PS1.read_text(encoding = "utf-8")
+ assert re.search(
+ r"'\^cu\[0-9\]'", text
+ ), "install.ps1 Get-TauriGpuBranch must require a digit after cu"
+
+ def test_install_sh_requires_cu_digit_in_gpu_branch(self):
+ text = INSTALL_SH.read_text(encoding = "utf-8")
+ # The _tauri_gpu_branch cuda case must be cu[0-9]*, not a bare cu*.
+ assert re.search(
+ r"cu\[0-9\]\*\)\s*echo \"cuda\"", text
+ ), "install.sh _tauri_gpu_branch cuda case must be cu[0-9]*, not cu*"
+
+ def test_install_sh_backend_export_requires_cu_digit(self):
+ text = INSTALL_SH.read_text(encoding = "utf-8")
+ # Brand CUDA only on cu[0-9]*; a bare catch-all *) -> cuda would mis-brand
+ # /current, /custom pins and skip ROCm repair on AMD hosts.
+ assert re.search(
+ r'cu\[0-9\]\*\)\s*export UNSLOTH_TORCH_BACKEND="cuda"', text
+ ), "install.sh backend export must brand cuda only on cu[0-9]*"
+ # An unknown leaf must NOT commit a cuda backend (it unsets instead).
+ assert re.search(
+ r"\*\)\s*unset UNSLOTH_TORCH_BACKEND", text
+ ), "install.sh backend export must unset (not force cuda) on an unknown leaf"
+
+ def test_install_sh_lowercases_backend_leaf(self):
+ text = INSTALL_SH.read_text(encoding = "utf-8")
+ # The leaf feeding both the backend case and the 2.11 floor case must be
+ # lowercased so the canonical gfx120X-all (capital X) matches.
+ assert re.search(
+ r"_torch_index_leaf=\$\(printf '%s' \"\$_torch_index_leaf\" \| tr '\[:upper:\]' '\[:lower:\]'\)",
+ text,
+ ), "install.sh must lowercase _torch_index_leaf before the gfx/rocm/cu case matches"
+
+
+class TestKnown211SetParity:
+ """The KNOWN-2.11 rocm/gfx set must be identical across all four installers:
+ exactly {rocm7.2} plus the gfx allowlist {gfx120x-all, gfx1151, gfx1150}.
+ rocm7.3 / torch 2.12 do not exist, so no side may floor them speculatively."""
+
+ def test_install_sh_known_211_leaf_is_rocm72_and_gfx_allowlist(self):
+ text = INSTALL_SH.read_text(encoding = "utf-8")
+ # The 2.11 floor case matches exactly rocm7.2 + the three gfx leaves.
+ assert re.search(
+ r"rocm7\.2\|gfx120x-all\|gfx1151\|gfx1150\)", text
+ ), "install.sh 2.11 floor must be exactly rocm7.2|gfx120x-all|gfx1151|gfx1150"
+ # No speculative rocm7.3 anywhere.
+ assert "rocm7.3" not in text, "install.sh must not reference a non-existent rocm7.3"
+
+ def test_python_known_211_versions_is_only_rocm72(self):
+ text = STACK_PY.read_text(encoding = "utf-8")
+ assert "_ROCM_KNOWN_TORCH211_VERSIONS" in text
+ # The frozenset literal is exactly {(7, 2)}.
+ m = re.search(r"_ROCM_KNOWN_TORCH211_VERSIONS[^=]*=\s*frozenset\(\{([^}]*)\}\)", text)
+ assert m is not None, "install_python_stack.py must define _ROCM_KNOWN_TORCH211_VERSIONS"
+ assert "(7, 2)" in m.group(1)
+ assert "7, 3" not in m.group(1) and "7, 1" not in m.group(1)
+
+ def test_setup_ps1_known_211_helper_is_only_rocm72(self):
+ text = SETUP_PS1.read_text(encoding = "utf-8")
+ assert "Test-RocmKnown211Version" in text
+ # The predicate is Major -eq 7 -and Minor -eq 2 (only rocm7.2).
+ assert re.search(
+ r"Test-RocmKnown211Version[\s\S]{0,400}\$Major -eq 7 -and \$Minor -eq 2", text
+ ), "setup.ps1 Test-RocmKnown211Version must accept only rocm7.2"
+
+ def test_install_ps1_pin_floor_is_only_rocm72(self):
+ text = INSTALL_PS1.read_text(encoding = "utf-8")
+ # The pinned-ROCm install-spec floor must be Major -eq 7 -and Minor -eq 2,
+ # not the speculative >= 2 that would floor a non-existent rocm7.3.
+ assert re.search(
+ r"\$_pinRocm211 = \(\[int\]\$Matches\[1\] -eq 7 -and \[int\]\$Matches\[2\] -eq 2\)",
+ text,
+ ), "install.ps1 pinned-ROCm floor must be rocm7.2 only (no speculative >= 2)"
+
+ def test_ps1_pin_floor_gate_is_anchored(self):
+ """The floor-selection gate that reads $_pinRocm211 from the raw leaf must anchor
+ the rocm match ($), or a suffixed custom leaf (rocm7.2-private) matches the rocm7.2
+ prefix, takes the 2.11-floor branch, and is force-routed through the ROCm path
+ before the exact-match elseif can send it to the verbatim install (Codex P2)."""
+ for path, label in ((INSTALL_PS1, "install.ps1"), (SETUP_PS1, "setup.ps1")):
+ text = path.read_text(encoding = "utf-8")
+ assert "-match '^rocm(\\d+)\\.(\\d+)$'" in text, (
+ f"{label} floor gate must anchor the rocm match (^rocm(\\d+)\\.(\\d+)$) so a "
+ "suffixed custom leaf is not floored/routed as rocm7.2"
+ )
+ assert (
+ "-match '^rocm(\\d+)\\.(\\d+)'\n" not in text
+ ), f"{label} floor gate must not use the unanchored ^rocm(\\d+)\\.(\\d+) prefix"
+
+ def test_install_ps1_bounds_unknown_leaf_pinned_torch(self):
+ """install.ps1's pinned-torch install must bound BOTH companions on EVERY
+ index, cu families included: torchaudio 2.11 dropped its exact torch
+ pin from the wheel metadata, so a bare companion beside torch<2.11 can
+ resolve a mismatched 2.11.0 build (Codex P2, then unconditional per the
+ torchaudio 2.11 unpinning)."""
+ text = INSTALL_PS1.read_text(encoding = "utf-8")
+ assert (
+ '$_pinVisionSpec = "torchvision>=0.19,<0.26.0"' in text
+ ), "install.ps1 custom-pin install must bound torchvision (>=0.19,<0.26.0)"
+ assert (
+ '$_pinAudioSpec = "torchaudio>=2.4,<2.11.0"' in text
+ ), "install.ps1 custom-pin install must bound torchaudio (>=2.4,<2.11.0)"
+ # No cu-family exemption: the bounds apply unconditionally.
+ assert (
+ "$_pinCuLeaf" not in text
+ ), "install.ps1 must bound companions on every index (no cu-family exemption)"
+ # The bounded companions must actually be passed to the install command.
+ assert re.search(
+ r'"torch>=2\.4,<2\.11\.0" \$_pinVisionSpec \$_pinAudioSpec --default-index \$TorchIndexUrl',
+ text,
+ ), "install.ps1 custom-pin install must pass the bounded companion specs to uv"
+
+ def test_gfx_allowlist_matches_across_installers(self):
+ # The gfx 2.11 allowlist {gfx120x-all, gfx1151, gfx1150} must appear in each.
+ gfx = ("gfx120x-all", "gfx1151", "gfx1150")
+ for path, label in (
+ (INSTALL_SH, "install.sh"),
+ (INSTALL_PS1, "install.ps1"),
+ (SETUP_PS1, "setup.ps1"),
+ (STACK_PY, "install_python_stack.py"),
+ ):
+ low = path.read_text(encoding = "utf-8").lower()
+ for g in gfx:
+ assert g in low, f"{label} missing gfx 2.11 allowlist member {g}"
+
+
+class TestPinnedRocmLeafDigitParity:
+ """A pinned index is a pip ROCm --default-index family only when its leaf is an
+ EXACT rocm+digits (rocm7 / rocm7.2) or gfx*. A ^rocm[0-9] PREFIX (or a bare rocm*
+ glob) wrongly catches a custom mirror / find-links leaf (rocm-current /
+ rocm-rel-7.2.1) AND a suffixed private-mirror leaf (rocm7.2-private / rocm7-current),
+ routing it through the ROCm install path (which silently falls back to CPU on
+ failure) or skipping the custom-index companion bounds, instead of the verbatim
+ --default-index install. All installers must match the family EXACTLY: Python and
+ install.sh via a shared _is_pip_rocm_family_leaf, setup.ps1 via Test-PipRocmFamilyLeaf,
+ install.ps1 via an anchored ^rocm[0-9]+(\\.[0-9]+)?$ reroute."""
+
+ def test_install_ps1_pinned_reroute_requires_rocm_digit(self):
+ text = INSTALL_PS1.read_text(encoding = "utf-8")
+ # The pinned gfx*/rocm reroute must match rocm EXACTLY (anchored), so a suffixed
+ # rocm7.2-private / rocm-current falls through to the verbatim --default-index path.
+ assert "-match '^rocm[0-9]+(\\.[0-9]+)?$'" in text, (
+ "install.ps1 pinned-index reroute must anchor the rocm match "
+ "(^rocm[0-9]+(\\.[0-9]+)?$), not a bare -like 'rocm*' or an unanchored ^rocm\\d"
+ )
+ # Neither the broad glob nor the unanchored prefix may drive that reroute.
+ assert (
+ "-like 'rocm*'" not in text
+ ), "install.ps1 must not route a pinned index on a bare -like 'rocm*' glob"
+ assert (
+ "-match '^rocm\\d'" not in text
+ ), "install.ps1 must not route a pinned index on an unanchored -match '^rocm\\d'"
+
+ def test_setup_ps1_pinned_reroute_requires_rocm_digit(self):
+ text = SETUP_PS1.read_text(encoding = "utf-8")
+ # setup.ps1 routes every family decision through Test-PipRocmFamilyLeaf, which
+ # anchors the rocm match so a suffixed custom leaf stays on the verbatim path.
+ assert (
+ "function Test-PipRocmFamilyLeaf" in text
+ ), "setup.ps1 must define Test-PipRocmFamilyLeaf (the exact rocm/gfx family gate)"
+ assert "'^rocm[0-9]+(\\.[0-9]+)?$'" in text, (
+ "setup.ps1 Test-PipRocmFamilyLeaf must anchor the rocm match "
+ "(^rocm[0-9]+(\\.[0-9]+)?$) so rocm7.2-private / rocm-current stay verbatim"
+ )
+ pinned_block = text[text.find("$_pinGfx211 = Test-RocmGfx211Leaf") :][:2000]
+ assert (
+ "-like 'rocm*'" not in pinned_block
+ ), "setup.ps1 pinned reroute must not route on a bare -like 'rocm*' glob"
+
+ def test_install_sh_repairable_requires_rocm_digit(self):
+ text = INSTALL_SH.read_text(encoding = "utf-8")
+ # _torch_index_repairable routes rocm/gfx through the exact-match helper.
+ assert (
+ "_is_pip_rocm_family_leaf" in text
+ ), "install.sh must define/use _is_pip_rocm_family_leaf for the exact rocm gate"
+ # gfx needs a following digit: gfx-private / gfxfoo are custom verbatim pins.
+ assert re.search(
+ r'case "\$1" in\n\s*gfx\[0-9\]\*\) return 0', text
+ ), "install.sh _is_pip_rocm_family_leaf must treat only gfx* as a family"
+ assert not re.search(
+ r'case "\$1" in\n\s*gfx\*\) return 0', text
+ ), "install.sh _is_pip_rocm_family_leaf must not family-match a bare gfx* glob"
+
+ def test_stack_py_pip_rocm_family_requires_digit(self):
+ text = STACK_PY.read_text(encoding = "utf-8")
+ assert re.search(
+ r'fullmatch\(r"rocm\\d\+\(\?:\\\.\\d\+\)\?", leaf\)', text
+ ), "install_python_stack.py _is_pip_rocm_family_leaf must fullmatch rocm\\d+(?:\\.\\d+)?"
+ # The unanchored prefix must be gone from the family/flavor gates.
+ assert (
+ 're.match(r"^rocm\\d"' not in text
+ ), "install_python_stack.py must not gate a family on an unanchored re.match(^rocm\\d)"
+
+ def test_install_sh_rocm_side_effects_digit_gated(self):
+ """The AMD bitsandbytes + 'repair ROCm torch' side effects must fire only on
+ an EXACT ROCm family (rocm7.2/gfx*), not a bare */rocm* whole-URL glob nor a
+ ^rocm[0-9] prefix that catches a custom CPU/CUDA index like /rocm-current or a
+ suffixed /rocm7.2-private and force-repairs it from the wrong --default-index."""
+ text = INSTALL_SH.read_text(encoding = "utf-8")
+ assert (
+ 'if _is_pip_rocm_family_leaf "$_torch_index_leaf"; then\n _torch_index_is_rocm_family=true'
+ in text
+ ), "install.sh must set _torch_index_is_rocm_family from the exact-match helper"
+ assert (
+ '[ "$_torch_index_is_rocm_family" = true ]' in text
+ ), "install.sh ROCm bnb/repair hooks must gate on _torch_index_is_rocm_family"
+ assert (
+ "*/rocm*|*/gfx*)\n _install_bnb_rocm" not in text
+ ), "install.sh must not gate _install_bnb_rocm on a bare */rocm* whole-URL glob"
+
+
+class TestPinnedIndexClearsUvEnvParity:
+ """Every installer must neutralise the uv index env vars for a pinned torch
+ install (#6898). uv treats the default index (--index-url / --default-index) as
+ lowest priority, so an inherited UV_INDEX / UV_EXTRA_INDEX_URL mirror would win
+ under uv's first-index strategy and pull torch from the wrong index -- after
+ which the pinned wheel index is silently never used."""
+
+ UV_VARS = ("UV_DEFAULT_INDEX", "UV_INDEX_URL", "UV_INDEX", "UV_EXTRA_INDEX_URL")
+
+ def test_install_sh_clears_uv_index_vars(self):
+ text = INSTALL_SH.read_text(encoding = "utf-8")
+ assert (
+ "env -u UV_DEFAULT_INDEX -u UV_INDEX_URL -u UV_INDEX -u UV_EXTRA_INDEX_URL" in text
+ ), "install.sh run_install_cmd must clear the uv index vars for --default-index installs"
+
+ def test_install_ps1_clears_uv_index_vars(self):
+ text = INSTALL_PS1.read_text(encoding = "utf-8")
+ for var in self.UV_VARS:
+ assert var in text, f"install.ps1 must clear {var} for pinned installs"
+
+ def test_setup_ps1_clears_uv_index_vars(self):
+ text = SETUP_PS1.read_text(encoding = "utf-8")
+ for var in self.UV_VARS:
+ assert var in text, f"setup.ps1 must clear {var} for pinned installs"
+
+ def test_stack_py_clears_uv_index_vars(self):
+ text = STACK_PY.read_text(encoding = "utf-8")
+ assert "_install_env_for_cmd" in text, (
+ "install_python_stack.py must scrub inherited uv index vars for pinned "
+ "installs via _install_env_for_cmd (parity with install.sh #6898)"
+ )
+ for var in self.UV_VARS:
+ assert var in text, f"install_python_stack.py must clear {var} for pinned installs"
+
+ def test_all_installers_clear_uv_torch_backend(self):
+ """uv's torch backend redirects torch resolution to its own per-backend
+ index even against an explicit pin, so every installer's pinned-install
+ scrub must clear UV_TORCH_BACKEND too."""
+ sh = INSTALL_SH.read_text(encoding = "utf-8")
+ assert "-u UV_TORCH_BACKEND" in sh, "install.sh pinned scrub must clear UV_TORCH_BACKEND"
+ for path in (INSTALL_PS1, SETUP_PS1):
+ text = path.read_text(encoding = "utf-8")
+ assert (
+ "'UV_TORCH_BACKEND'" in text
+ ), f"{path.name} pinned scrub must clear UV_TORCH_BACKEND"
+ stack = STACK_PY.read_text(encoding = "utf-8")
+ assert (
+ '"UV_TORCH_BACKEND",' in stack
+ ), "install_python_stack.py strip tuple must include UV_TORCH_BACKEND"
+
+ def test_stack_py_strips_pip_extra_index_for_pip_fallback(self):
+ """The pip fallback honours PIP_EXTRA_INDEX_URL (pip adds it IN ADDITION
+ to --index-url), so the pinned-command scrub must strip it."""
+ stack = STACK_PY.read_text(encoding = "utf-8")
+ assert (
+ '"PIP_EXTRA_INDEX_URL",' in stack
+ ), "install_python_stack.py strip tuple must include PIP_EXTRA_INDEX_URL"
+
+ def test_all_installers_scrub_find_links(self):
+ """uv's --find-links (env UV_FIND_LINKS) adds candidate locations that can
+ satisfy torch off a pinned index; every pinned-install scrub must clear it."""
+ sh = INSTALL_SH.read_text(encoding = "utf-8")
+ assert "-u UV_FIND_LINKS" in sh
+ for path in (INSTALL_PS1, SETUP_PS1):
+ assert "'UV_FIND_LINKS'" in path.read_text(encoding = "utf-8"), path.name
+ stack = STACK_PY.read_text(encoding = "utf-8")
+ assert '"UV_FIND_LINKS",' in stack and '"PIP_FIND_LINKS",' in stack
+
+ def test_setup_ps1_scrub_covers_pip_fallback(self):
+ """setup.ps1's Fast-Install must keep the scrub active through the pip
+ fallback (pip honours PIP_EXTRA_INDEX_URL / PIP_FIND_LINKS in addition to
+ --index-url); restoring the vars before the fallback reopens the hole."""
+ text = SETUP_PS1.read_text(encoding = "utf-8")
+ fi = text[text.find("function Fast-Install") :][:2500]
+ assert "'PIP_EXTRA_INDEX_URL'" in fi and "'PIP_FIND_LINKS'" in fi
+ # the pip fallback must sit INSIDE the try whose finally restores the vars
+ assert fi.find("python -m pip install") < fi.find(
+ "finally"
+ ), "pip fallback must run before the scrub is restored"
+
+ def test_all_installers_disable_uv_config_for_pinned_installs(self):
+ """A DISCOVERED uv.toml / pyproject [tool.uv] outranks the CLI pin
+ (verified with uv 0.10: [pip] torch-backend = "cpu" and a non-default
+ [[index]] both resolve torch+cpu against an explicit --index-url /
+ --default-index cu126 pin; UV_NO_CONFIG=1 restores the pin). Every
+ installer's pinned scrub must set UV_NO_CONFIG=1 and drop UV_CONFIG_FILE."""
+ sh = INSTALL_SH.read_text(encoding = "utf-8")
+ assert "-u UV_CONFIG_FILE UV_NO_CONFIG=1" in sh, (
+ "install.sh run_install_cmd must set UV_NO_CONFIG=1 and drop "
+ "UV_CONFIG_FILE for --default-index installs"
+ )
+ for path in (INSTALL_PS1, SETUP_PS1):
+ text = path.read_text(encoding = "utf-8")
+ assert "'UV_CONFIG_FILE'" in text, f"{path.name} must drop UV_CONFIG_FILE"
+ assert (
+ "$env:UV_NO_CONFIG = '1'" in text
+ ), f"{path.name} must set UV_NO_CONFIG=1 for pinned installs"
+ stack = STACK_PY.read_text(encoding = "utf-8")
+ assert (
+ '"UV_CONFIG_FILE",' in stack
+ ), "install_python_stack.py strip tuple must include UV_CONFIG_FILE"
+ assert (
+ 'env["UV_NO_CONFIG"] = "1"' in stack
+ ), "_install_env_for_cmd must set UV_NO_CONFIG=1 for pinned installs"
+
+ def test_pip_fallbacks_disable_pip_config_files(self):
+ """The pip FALLBACK (uv missing/failed) honours user/site pip config files
+ even with the PIP_* env vars stripped: `pip config set
+ global.extra-index-url` still adds indexes to a pinned install. pip loads
+ NO configuration files when PIP_CONFIG_FILE is the platform devnull, so
+ the two installers that HAVE a pip fallback (install_python_stack.py and
+ setup.ps1's Fast-Install) must set it in their pinned scrub. install.sh
+ and install.ps1 are uv-only (no python -m pip fallback) and need no
+ equivalent."""
+ stack = STACK_PY.read_text(encoding = "utf-8")
+ assert 'env["PIP_CONFIG_FILE"] = os.devnull' in stack, (
+ "_install_env_for_cmd must point PIP_CONFIG_FILE at os.devnull for "
+ "pinned installs (pip fallback isolation)"
+ )
+ setup = SETUP_PS1.read_text(encoding = "utf-8")
+ assert "$env:PIP_CONFIG_FILE = 'nul'" in setup, (
+ "setup.ps1 Fast-Install pinned scrub must point PIP_CONFIG_FILE at nul "
+ "(Windows devnull) so the pip fallback ignores user/site pip config"
+ )
+ assert (
+ "'PIP_CONFIG_FILE'" in setup
+ ), "setup.ps1 must save/restore PIP_CONFIG_FILE around the pinned scrub"
+
+ def test_setup_ps1_bounds_unknown_leaf_pinned_torch(self):
+ """A first-time/changed unknown-leaf custom pin routes through setup.ps1's
+ CUDA branch; install.ps1's fresh pinned install, install.sh, and the Python
+ verbatim path bound the WHOLE trio, so the Windows update path must too -- a
+ private mirror serving newer torch OR newer companions must not lift the venv
+ above the supported range under the pin."""
+ text = SETUP_PS1.read_text(encoding = "utf-8")
+ # The custom-leaf branch bounds torch AND both companions (parity with the
+ # other installers' custom-pin trio bounds), gated on a non-cu-family leaf.
+ for spec in (
+ '$cudaTorchSpec = "torch>=2.4,<2.11.0"',
+ '$cudaVisionSpec = "torchvision>=0.19,<0.26.0"',
+ '$cudaAudioSpec = "torchaudio>=2.4,<2.11.0"',
+ ):
+ assert spec in text, f"setup.ps1 must bound the custom-leaf trio: {spec}"
+ assert (
+ "if ($TorchIndexPinned -and -not (Test-CudaFamilyLeaf $CuTag)) {" in text
+ ), "the custom-leaf trio bounds must be gated on a pinned non-cu-family leaf"
+ assert (
+ "Fast-Install $cudaTorchSpec $cudaVisionSpec $cudaAudioSpec" in text
+ ), "setup.ps1's CUDA branch must install via the bounded spec variables"
+
+ def test_setup_ps1_bounds_pinned_cpu_torch(self):
+ """setup.ps1's CPU branch must bound the trio under an explicit pin (parity with
+ _CPU_TORCH_PKG_SPEC): the /cpu index serves newer torch, and _ensure_cpu_torch
+ keeps any CPU build, so a bare pinned trio could land an unsupported version.
+ An unpinned CPU host keeps the bare trio (pre-pin behavior unchanged)."""
+ text = SETUP_PS1.read_text(encoding = "utf-8")
+ for spec in (
+ '$cpuTorchSpec = "torch>=2.4,<2.12.0"',
+ '$cpuVisionSpec = "torchvision>=0.19,<0.27.0"',
+ '$cpuAudioSpec = "torchaudio>=2.4,<2.12.0"',
+ ):
+ assert spec in text, f"setup.ps1 must bound the pinned CPU trio: {spec}"
+ assert (
+ "if ($TorchIndexPinned) {" in text
+ ), "the CPU trio bounds must be gated on an explicit pin"
+ assert (
+ "Fast-Install $cpuTorchSpec $cpuVisionSpec $cpuAudioSpec @cpuForce" in text
+ ), "setup.ps1's CPU branch must install via the spec variables"
+ # The ceilings mirror the Python repair spec exactly.
+ stack = STACK_PY.read_text(encoding = "utf-8")
+ spec_block = re.search(r"_CUDA_TORCH_PKG_SPEC[^(]*\(\s*(.*?)\)", stack, re.DOTALL)
+ assert spec_block and '"torch>=2.4,<2.12.0"' in spec_block.group(1), (
+ "_CPU_TORCH_PKG_SPEC (via _CUDA_TORCH_PKG_SPEC) must keep the torch<2.12 "
+ "ceiling the setup.ps1 pinned CPU branch mirrors"
+ )
+
+ def test_setup_ps1_stale_check_requires_rocm_digit(self):
+ """The stale-venv check must use the same EXACT rocm/gfx gate as the install
+ selection (Test-PipRocmFamilyLeaf), or a custom rocm-* / suffixed rocm7.2-private
+ leaf is stale-compared as a family and force-reinstalls on every studio update."""
+ text = SETUP_PS1.read_text(encoding = "utf-8")
+ anchor = text.find("$_pinLeaf = Get-TorchIndexLeaf $_pinnedIdx")
+ assert anchor >= 0, "setup.ps1 stale check must classify the pinned leaf"
+ stale = text[anchor:][:2500]
+ assert (
+ "Test-PipRocmFamilyLeaf" in stale
+ ), "setup.ps1 stale check must gate rocm leaves via the exact Test-PipRocmFamilyLeaf"
+ assert (
+ stale.count("-like 'rocm*'") == 0
+ ), "setup.ps1 stale check must not use a bare -like 'rocm*' glob"
+ assert (
+ "-match '^rocm\\d'" not in stale
+ ), "setup.ps1 stale check must not use an unanchored -match '^rocm\\d'"
+
+
+class TestIndexPathSlashTrimParity:
+ """Every installer must trim trailing PATH slashes only on the verbatim
+ UNSLOTH_TORCH_INDEX_URL override, preserving a ?query/#fragment token: a whole-URL
+ strip corrupts a base64 token ending in "/", a single strip leaves a double-slash leaf
+ empty. The helper must be DEFINED and WIRED into the override return in all four."""
+
+ def test_helper_defined_in_all_installers(self):
+ assert "def _trim_index_path_slashes(" in STACK_PY.read_text(encoding = "utf-8")
+ assert "_trim_index_path_slashes()" in INSTALL_SH.read_text(encoding = "utf-8")
+ assert "function Trim-IndexPathSlashes" in INSTALL_PS1.read_text(encoding = "utf-8")
+ assert "function Trim-IndexPathSlashes" in SETUP_PS1.read_text(encoding = "utf-8")
+
+ def test_helper_wired_into_override_in_all_installers(self):
+ assert "_trim_index_path_slashes(url)" in STACK_PY.read_text(encoding = "utf-8")
+ assert '_url=$(_trim_index_path_slashes "$_url")' in INSTALL_SH.read_text(encoding = "utf-8")
+ assert "Trim-IndexPathSlashes $env:UNSLOTH_TORCH_INDEX_URL" in INSTALL_PS1.read_text(
+ encoding = "utf-8"
+ )
+ assert "Trim-IndexPathSlashes $env:UNSLOTH_TORCH_INDEX_URL" in SETUP_PS1.read_text(
+ encoding = "utf-8"
+ )
+
+
+class TestInstallOutputRedactionParity:
+ """uv/pip failure text embeds the failing --index-url verbatim, so a captured install
+ log dumped on error can leak a user:token@ or ?token= secret. Every installer must
+ DEFINE a redaction helper and WIRE it into the captured-output print path."""
+
+ def test_helper_defined_in_all_installers(self):
+ assert "def _redact_install_output(" in STACK_PY.read_text(encoding = "utf-8")
+ assert "_redact_install_output()" in INSTALL_SH.read_text(encoding = "utf-8")
+ assert "function Redact-InstallOutput" in INSTALL_PS1.read_text(encoding = "utf-8")
+ assert "function Redact-InstallOutput" in SETUP_PS1.read_text(encoding = "utf-8")
+
+ def test_helper_wired_into_failure_print(self):
+ # install.sh dumps the captured log through the redactor on failure.
+ assert '_redact_install_output "$_log"' in INSTALL_SH.read_text(encoding = "utf-8")
+ # Both ps1 installers redact the captured $output before Write-Host on non-zero exit.
+ assert (
+ "Write-Host (Redact-InstallOutput $output) -ForegroundColor Red"
+ in INSTALL_PS1.read_text(encoding = "utf-8")
+ )
+ assert (
+ "Write-Host (Redact-InstallOutput $output) -ForegroundColor Red"
+ in SETUP_PS1.read_text(encoding = "utf-8")
+ )
+ # Python redacts the captured stdout before printing.
+ assert "_redact_install_output(" in STACK_PY.read_text(encoding = "utf-8")
+
+
+class TestPipNoIndexScrubParity:
+ """The plain-pip fallback honours PIP_*: PIP_NO_INDEX=1 makes it ignore ALL indexes
+ (defeating the pinned --index-url) and PIP_INDEX_URL replaces the pin. The two installers
+ that HAVE a plain-pip fallback (Python + setup.ps1) must scrub both for a pinned install.
+ install.sh / install.ps1 are uv-only (--default-index), which ignores pip config/env."""
+
+ def test_python_scrubs_pip_no_index_and_pip_index_url(self):
+ text = STACK_PY.read_text(encoding = "utf-8")
+ assert '"PIP_NO_INDEX"' in text
+ assert '"PIP_INDEX_URL"' in text
+
+ def test_setup_ps1_scrubs_pip_no_index_and_pip_index_url(self):
+ text = SETUP_PS1.read_text(encoding = "utf-8")
+ assert "'PIP_NO_INDEX'" in text
+ assert "'PIP_INDEX_URL'" in text
diff --git a/tests/python/test_install_python_stack.py b/tests/python/test_install_python_stack.py
index 9015ff8c9d..3a12e53f95 100644
--- a/tests/python/test_install_python_stack.py
+++ b/tests/python/test_install_python_stack.py
@@ -54,6 +54,24 @@ class TestBuildUvCmdTorchBackend:
a.startswith("--torch-backend") for a in cmd
), f"Empty UV_TORCH_BACKEND should not add flag, got: {cmd}"
+ def test_uv_torch_backend_skipped_for_pinned_index(self):
+ """A pinned-index command must NOT get --torch-backend: uv's torch backend
+ redirects torch resolution to its own per-backend index even when
+ --index-url is given (verified: cu128 pin + backend cpu installs
+ torch+cpu), defeating the pin."""
+ for pin_flag in ("--index-url", "--default-index"):
+ with mock.patch.dict(os.environ, {"UV_TORCH_BACKEND": "cpu"}):
+ cmd = self._call(("torch", pin_flag, "https://download.pytorch.org/whl/cu128"))
+ assert not any(
+ a.startswith("--torch-backend") for a in cmd
+ ), f"{pin_flag} command must not carry --torch-backend, got: {cmd}"
+
+ def test_uv_torch_backend_kept_for_unpinned(self):
+ """Non-pinned commands still honour UV_TORCH_BACKEND."""
+ with mock.patch.dict(os.environ, {"UV_TORCH_BACKEND": "cpu"}):
+ cmd = self._call(("somepackage",))
+ assert "--torch-backend=cpu" in cmd
+
class TestUvSafePath:
"""_uv_safe_path hands uv a space-free `-c`/`-r` path (issue #6503)."""
@@ -148,3 +166,119 @@ class TestUvSafePathHardening:
assert " " not in value
assert Path(value).read_text() == "transformers>=4.57.6\n"
+
+
+class TestPinnedIndexClearsUvEnv:
+ """A pinned torch install (--index-url / --default-index) must neutralise an
+ inherited UV_INDEX / UV_EXTRA_INDEX_URL so the pinned wheel index wins.
+
+ uv treats the default index (--index-url / --default-index) as LOWEST priority,
+ so an inherited UV_INDEX / UV_EXTRA_INDEX_URL (a corporate/CPU mirror) would be
+ searched first and, under uv's default first-index strategy, resolve torch from
+ the wrong mirror -- after which the marker records a wheel index that was never
+ used. install.sh (#6898), install.ps1 and setup.ps1 already clear these for
+ pinned installs; install_python_stack must match (parity across all installers).
+ """
+
+ UV_VARS = ("UV_DEFAULT_INDEX", "UV_INDEX_URL", "UV_INDEX", "UV_EXTRA_INDEX_URL")
+
+ def test_pinned_index_url_strips_uv_index_vars(self):
+ cmd = [
+ "uv",
+ "pip",
+ "install",
+ "--force-reinstall",
+ "torch",
+ "torchvision",
+ "torchaudio",
+ "--index-url",
+ "https://download.pytorch.org/whl/cu128",
+ ]
+ with mock.patch.dict(
+ os.environ,
+ {
+ "UV_INDEX": "https://mirror.corp/simple",
+ "UV_EXTRA_INDEX_URL": "https://mirror.corp/extra",
+ "UV_INDEX_URL": "https://mirror.corp/root",
+ "UV_DEFAULT_INDEX": "https://mirror.corp/default",
+ },
+ ):
+ env = ips._install_env_for_cmd(cmd)
+ assert env is not None, "a --index-url install must run with a scrubbed env"
+ for var in self.UV_VARS:
+ assert var not in env, f"{var} must be cleared for a pinned-index install"
+
+ def test_pinned_default_index_strips_uv_index_vars(self):
+ # --default-index must be gated too (matches install.sh / install.ps1).
+ cmd = ["uv", "pip", "install", "torch", "--default-index", "https://x/cu126"]
+ with mock.patch.dict(os.environ, {"UV_INDEX": "https://mirror.corp/simple"}):
+ env = ips._install_env_for_cmd(cmd)
+ assert env is not None
+ assert "UV_INDEX" not in env
+
+ def test_non_pinned_install_keeps_user_mirror(self):
+ # A plain install (no --index-url) must NOT scrub the env, so a user's mirror
+ # still applies to base packages.
+ cmd = ["uv", "pip", "install", "unsloth", "unsloth-zoo"]
+ with mock.patch.dict(os.environ, {"UV_INDEX": "https://mirror.corp/simple"}):
+ env = ips._install_env_for_cmd(cmd)
+ assert env is None, "non-pinned installs must inherit the caller env unchanged"
+
+ def test_scrubbed_env_preserves_other_vars(self):
+ cmd = ["uv", "pip", "install", "torch", "--index-url", "https://x/cu128"]
+ with mock.patch.dict(
+ os.environ,
+ {"UV_INDEX": "https://mirror.corp/simple", "PATH_SENTINEL_XYZ": "keepme"},
+ ):
+ env = ips._install_env_for_cmd(cmd)
+ assert env is not None
+ assert env.get("PATH_SENTINEL_XYZ") == "keepme", "only uv index vars are removed"
+
+ def test_pinned_cmd_strips_pip_extra_index_url(self):
+ """PIP_EXTRA_INDEX_URL is stripped for pinned commands so the pip
+ fallback cannot satisfy torch from an inherited extra index."""
+ with mock.patch.dict(os.environ, {"PIP_EXTRA_INDEX_URL": "https://mirror/simple"}):
+ env = ips._install_env_for_cmd(
+ ["pip", "install", "torch", "--index-url", "https://x/cu128"]
+ )
+ assert env is not None and "PIP_EXTRA_INDEX_URL" not in env
+
+ def test_pinned_cmd_strips_uv_torch_backend(self):
+ """UV_TORCH_BACKEND is stripped for pinned commands so uv cannot read it
+ from the environment and reroute torch off the pinned index."""
+ with mock.patch.dict(os.environ, {"UV_TORCH_BACKEND": "cpu"}):
+ env = ips._install_env_for_cmd(
+ ["uv", "pip", "install", "torch", "--index-url", "https://x/cu128"]
+ )
+ assert env is not None and "UV_TORCH_BACKEND" not in env
+
+ def test_pinned_cmd_disables_uv_config_discovery(self):
+ """A DISCOVERED uv.toml / pyproject [tool.uv] outranks the CLI pin too
+ (verified with uv 0.10: [pip] torch-backend = "cpu" and a non-default
+ [[index]] both resolve torch+cpu against an explicit --index-url /
+ --default-index cu126 pin). Pinned commands must run with UV_NO_CONFIG=1
+ and without an inherited UV_CONFIG_FILE."""
+ with mock.patch.dict(os.environ, {"UV_CONFIG_FILE": "/etc/uv/uv.toml"}):
+ env = ips._install_env_for_cmd(
+ ["uv", "pip", "install", "torch", "--index-url", "https://x/cu128"]
+ )
+ assert env is not None
+ assert env.get("UV_NO_CONFIG") == "1"
+ assert "UV_CONFIG_FILE" not in env
+
+ def test_pinned_cmd_disables_pip_config_files(self):
+ """The pip FALLBACK honours user/site pip config files (pip config set
+ global.extra-index-url) even with the PIP_* env vars stripped; pip loads
+ NO configuration files when PIP_CONFIG_FILE is os.devnull. Harmless for
+ uv, decisive for the fallback."""
+ env = ips._install_env_for_cmd(
+ ["uv", "pip", "install", "torch", "--index-url", "https://x/cu128"]
+ )
+ assert env is not None
+ assert env.get("PIP_CONFIG_FILE") == os.devnull
+
+ def test_non_pinned_cmd_keeps_uv_config_discovery(self):
+ """Non-pinned installs inherit the caller env unchanged, so a user's uv
+ configuration still applies to base packages."""
+ env = ips._install_env_for_cmd(["uv", "pip", "install", "unsloth"])
+ assert env is None
diff --git a/tests/run_all.sh b/tests/run_all.sh
index d03f4c4d4f..a31103a85b 100755
--- a/tests/run_all.sh
+++ b/tests/run_all.sh
@@ -15,6 +15,7 @@ sh "$TESTS_DIR/sh/test_resolve_cuda_archs.sh"
sh "$TESTS_DIR/sh/test_strixhalo_wsl_reroute.sh"
sh "$TESTS_DIR/sh/test_uninstall_shared_icon.sh"
sh "$TESTS_DIR/sh/test_torch_flavor.sh"
+sh "$TESTS_DIR/sh/test_redact_install_output.sh"
sh "$TESTS_DIR/sh/test_install_uv_override_space.sh"
echo ""
diff --git a/tests/sh/test_get_torch_index_url.sh b/tests/sh/test_get_torch_index_url.sh
index 6656142625..23902097ef 100755
--- a/tests/sh/test_get_torch_index_url.sh
+++ b/tests/sh/test_get_torch_index_url.sh
@@ -23,6 +23,8 @@ _FAKE_SMI_DIR=$(mktemp -d)
echo ""
sed -n '/^_has_usable_nvidia_gpu()/,/^}/p' "$INSTALL_SH"
echo ""
+ sed -n '/^_trim_index_path_slashes()/,/^}/p' "$INSTALL_SH"
+ echo ""
sed -n '/^get_torch_index_url()/,/^}/p' "$INSTALL_SH"
} | sed "s|/usr/bin/nvidia-smi|$_FAKE_SMI_DIR/nvidia-smi-absent|g" \
> "$_FUNC_FILE"
@@ -379,6 +381,61 @@ _result=$(run_func "$_dir" " -1 ")
assert_eq "CVD=' -1 ' hides NVIDIA -> cpu" "https://download.pytorch.org/whl/cpu" "$_result"
rm -rf "$_dir"
+# --- explicit overrides (headless / container / CI; no GPU probing) ----------
+# 39) UNSLOTH_TORCH_INDEX_FAMILY pins the family with no GPU present (not the cpu fallback).
+_result=$(UNSLOTH_TORCH_INDEX_FAMILY="cu128" run_func "none")
+assert_eq "family override (no GPU) -> cu128" "https://download.pytorch.org/whl/cu128" "$_result"
+
+# 40) Family override beats real detection: an nvidia-smi 12.6 host still gets cu128
+# (the Docker-build case -- builder sees the host driver but publishes a cu128 image).
+_dir=$(make_mock_smi "12.6")
+_result=$(UNSLOTH_TORCH_INDEX_FAMILY="cu128" run_func "$_dir")
+assert_eq "family override beats detected 12.6 -> cu128" "https://download.pytorch.org/whl/cu128" "$_result"
+rm -rf "$_dir"
+
+# 41) UNSLOTH_TORCH_INDEX_URL is used verbatim and wins over detection.
+_dir=$(make_mock_smi "12.6")
+_result=$(UNSLOTH_TORCH_INDEX_URL="https://mirror.example.com/whl/cu999" run_func "$_dir")
+assert_eq "url override beats detection -> verbatim" "https://mirror.example.com/whl/cu999" "$_result"
+rm -rf "$_dir"
+
+# 42) Family override is appended to UNSLOTH_PYTORCH_MIRROR (mirror still honoured).
+_result=$(UNSLOTH_PYTORCH_MIRROR="https://mirror.example.com/whl" UNSLOTH_TORCH_INDEX_FAMILY="cu128" run_func "none")
+assert_eq "mirror + family override -> mirror/cu128" "https://mirror.example.com/whl/cu128" "$_result"
+
+# 43) Trailing slash in UNSLOTH_TORCH_INDEX_URL is stripped.
+_result=$(UNSLOTH_TORCH_INDEX_URL="https://mirror.example.com/whl/cu128/" run_func "none")
+assert_eq "url override trailing slash stripped" "https://mirror.example.com/whl/cu128" "$_result"
+
+# 44) URL override takes precedence over family override.
+_result=$(UNSLOTH_TORCH_INDEX_URL="https://mirror.example.com/whl/cu130" UNSLOTH_TORCH_INDEX_FAMILY="cu128" run_func "none")
+assert_eq "url override beats family override -> url" "https://mirror.example.com/whl/cu130" "$_result"
+
+# 45) An empty override is ignored (falls through to normal detection).
+_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"
+
+# 48) A ?query token that ends in "/" is PRESERVED: only PATH slashes are trimmed, so a
+# base64 token ending in "/" is not corrupted (path-only trim, not whole-URL rstrip).
+_result=$(UNSLOTH_TORCH_INDEX_URL="https://mirror.example.com/whl/cu128?token=ab12cd/" run_func "none")
+assert_eq "url override preserves query token slash" "https://mirror.example.com/whl/cu128?token=ab12cd/" "$_result"
+
+# 49) Double PATH slash before a query is collapsed while the query survives intact.
+_result=$(UNSLOTH_TORCH_INDEX_URL="https://mirror.example.com/whl/cu128//?token=ab12cd/" run_func "none")
+assert_eq "url override path slash trimmed, query kept" "https://mirror.example.com/whl/cu128?token=ab12cd/" "$_result"
+
+# 50) A #fragment ending in "/" is likewise preserved.
+_result=$(UNSLOTH_TORCH_INDEX_URL="https://mirror.example.com/whl/cu128#anchor/" run_func "none")
+assert_eq "url override preserves fragment slash" "https://mirror.example.com/whl/cu128#anchor/" "$_result"
+
rm -f "$_FUNC_FILE"
rm -rf "$_FAKE_SMI_DIR"
rm -rf "$_TOOLS_DIR"
diff --git a/tests/sh/test_redact_install_output.sh b/tests/sh/test_redact_install_output.sh
new file mode 100755
index 0000000000..0f10122aea
--- /dev/null
+++ b/tests/sh/test_redact_install_output.sh
@@ -0,0 +1,89 @@
+#!/bin/bash
+# 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 install.sh's _redact_install_output helper. uv/pip failure text embeds the
+# failing --index-url verbatim, so a captured install log dumped on error can leak a
+# user:token@ or ?token= secret. The helper redacts both before printing. Mirrors
+# _redact_install_output (install_python_stack.py) / Redact-InstallOutput (install.ps1 /
+# setup.ps1).
+set -e
+
+SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
+INSTALL_SH="$SCRIPT_DIR/../../install.sh"
+PASS=0
+FAIL=0
+
+_FUNC_FILE=$(mktemp)
+sed -n '/^_redact_install_output()/,/^}/p' "$INSTALL_SH" > "$_FUNC_FILE"
+# shellcheck disable=SC1090
+. "$_FUNC_FILE"
+rm -f "$_FUNC_FILE"
+
+assert_eq() {
+ _label="$1"; _expected="$2"; _actual="$3"
+ if [ "$_actual" = "$_expected" ]; then
+ echo " PASS: $_label"; PASS=$((PASS + 1))
+ else
+ echo " FAIL: $_label (expected '$_expected', got '$_actual')"; FAIL=$((FAIL + 1))
+ fi
+}
+
+# Redact from a file (the actual call site passes a captured-log tempfile).
+redact_str() {
+ _rs_tmp=$(mktemp)
+ printf '%s\n' "$1" > "$_rs_tmp"
+ _rs_out=$(_redact_install_output "$_rs_tmp")
+ rm -f "$_rs_tmp"
+ printf '%s' "$_rs_out"
+}
+
+echo "=== _redact_install_output ==="
+assert_eq "userinfo user:token@ redacted" \
+ "ERROR: failed https://@download.pytorch.org/whl/cu128" \
+ "$(redact_str 'ERROR: failed https://alice:s3cr3t@download.pytorch.org/whl/cu128')"
+
+assert_eq "bare-token@ userinfo redacted" \
+ "fetch https://@host/whl/cu128 failed" \
+ "$(redact_str 'fetch https://ghp_deadbeef@host/whl/cu128 failed')"
+
+assert_eq "single ?token= query redacted" \
+ "url https://host/whl/cu128?token= unreachable" \
+ "$(redact_str 'url https://host/whl/cu128?token=abcd1234 unreachable')"
+
+assert_eq "multiple query values redacted" \
+ "https://host/whl/cu128?token=&channel=" \
+ "$(redact_str 'https://host/whl/cu128?token=abcd1234&channel=beta')"
+
+assert_eq "http (not https) userinfo redacted" \
+ "http://@host/simple" \
+ "$(redact_str 'http://u:p@host/simple')"
+
+assert_eq "fragment token redacted" \
+ "ERROR: could not fetch https://mirror.local/whl/cu128# (403)" \
+ "$(redact_str 'ERROR: could not fetch https://mirror.local/whl/cu128#token=SECRET123 (403)')"
+
+assert_eq "query and fragment both redacted" \
+ "https://host/whl/cu128?token=# done" \
+ "$(redact_str 'https://host/whl/cu128?token=abc#sig=xyz done')"
+
+# Non-secret text is untouched (no false positives on ordinary log lines).
+assert_eq "plain line untouched" \
+ "Resolved 42 packages in 1.2s" \
+ "$(redact_str 'Resolved 42 packages in 1.2s')"
+assert_eq "plain url without creds untouched" \
+ "downloading https://download.pytorch.org/whl/cu128/torch-2.8.0.whl" \
+ "$(redact_str 'downloading https://download.pytorch.org/whl/cu128/torch-2.8.0.whl')"
+assert_eq "bare hash comment untouched" \
+ "# retrying with --no-cache-dir" \
+ "$(redact_str '# retrying with --no-cache-dir')"
+
+# Regression guard: no secret substring survives.
+_leak=$(redact_str 'https://alice:s3cr3t@host/whl/cu128?token=SUPERSECRET#frag=ALSOSECRET')
+case "$_leak" in
+ *s3cr3t*|*SUPERSECRET*|*ALSOSECRET*) assert_eq "no secret leak" "clean" "leaked:$_leak" ;;
+ *) assert_eq "no secret leak" "clean" "clean" ;;
+esac
+
+echo ""
+echo "Results: $PASS passed, $FAIL failed"
+[ "$FAIL" -eq 0 ]
diff --git a/tests/sh/test_torch_constraint.sh b/tests/sh/test_torch_constraint.sh
index d60dfc9f90..bfafbd161b 100644
--- a/tests/sh/test_torch_constraint.sh
+++ b/tests/sh/test_torch_constraint.sh
@@ -108,6 +108,25 @@ assert_eq "\$TORCH_CONSTRAINT used in pip install" "yes" "$_has_var"
_hardcoded=$(grep -c '"torch>=2.4,<2.11.0"' "$INSTALL_SH" || true)
assert_eq "hardcoded torch>=2.4 appears exactly once" "1" "$_hardcoded"
+# Companions must be bounded to torch's window everywhere: the <2.11 bound appears
+# twice (default assignments + the pinned custom-leaf block), never bare. torchaudio
+# 2.11 dropped its exact torch pin, so a bare companion next to a <2.11-capped torch
+# resolves a mismatched 2.11 build.
+_count=$(grep -c 'TORCHVISION_CONSTRAINT="torchvision>=0.19,<0.26.0"' "$INSTALL_SH" || true)
+assert_eq "torchvision bounded (<0.26) at default + custom-leaf" "2" "$_count"
+_count=$(grep -c 'TORCHAUDIO_CONSTRAINT="torchaudio>=2.4,<2.11.0"' "$INSTALL_SH" || true)
+assert_eq "torchaudio bounded (<2.11) at default + custom-leaf" "2" "$_count"
+_count=$(grep -c 'TORCHVISION_CONSTRAINT="torchvision"$' "$INSTALL_SH" || true)
+assert_eq "no bare torchvision companion remains" "0" "$_count"
+_count=$(grep -c 'TORCHAUDIO_CONSTRAINT="torchaudio"$' "$INSTALL_SH" || true)
+assert_eq "no bare torchaudio companion remains" "0" "$_count"
+# The cu* widen must carry the companions with it (torch <2.12 with torchaudio <2.11
+# would cap a mismatched pair the other way).
+assert_eq "cu widen pairs torchaudio (<2.12)" "1" "$(grep -c 'TORCHAUDIO_CONSTRAINT="torchaudio>=2.4,<2.12.0"' "$INSTALL_SH" || true)"
+_gated=$(grep -c '_expected_torch_flavor_tag "$TORCH_INDEX_URL"' "$INSTALL_SH" || true)
+_has_gate=$([ "$_gated" -ge 1 ] && echo "yes" || echo "no")
+assert_eq "custom-companion bound gated on empty flavor tag" "yes" "$_has_gate"
+
# A fresh CUDA install widens the ceiling to <2.12.0 so cu12x/cu13x land torch
# 2.11.x (matches the base image and _CUDA_TORCH_PKG_SPEC).
_cuda_widen=$(grep -c 'TORCH_CONSTRAINT="torch>=2.4,<2.12.0"' "$INSTALL_SH" || true)
@@ -285,6 +304,61 @@ bash -c "
_uv_got2=$(cat "$_UV_LOG2" 2>/dev/null || echo "")
assert_contains "mock uv arm64+py312 receives torch>=2.4" "$_uv_got2" "torch>=2.4,<2.11.0"
+# ======================================================================
+# ROCm 2.11 floor: leaf is lowercased before the gfx*/rocm* allowlist match
+# ======================================================================
+echo ""
+echo "=== ROCm 2.11 floor case (leaf normalization) ==="
+
+# Structural: install.sh lowercases _torch_index_leaf before the floor case, so the
+# canonical gfx120X-all (capital X) matches gfx120x-all.
+_has_lc=$(grep -c '_torch_index_leaf=$(printf .* | tr .\[:upper:\]. .\[:lower:\].)' "$INSTALL_SH" || true)
+_has_lc_ok=$([ "$_has_lc" -ge 1 ] && echo "yes" || echo "no")
+assert_eq "install.sh lowercases _torch_index_leaf" "yes" "$_has_lc_ok"
+
+# Runtime: replicate install.sh's normalization + floor case and assert both gfx120X-all
+# and gfx120x-all get the floor, while non-2.11 leaves keep the default.
+run_floor_case() {
+ _url="$1"
+ bash -c '
+ TORCH_CONSTRAINT="torch>=2.4,<2.11.0"
+ TORCHVISION_CONSTRAINT="torchvision"
+ TORCHAUDIO_CONSTRAINT="torchaudio"
+ _torch_index_leaf="${1%/}"
+ _torch_index_leaf="${_torch_index_leaf##*/}"
+ _torch_index_leaf=$(printf "%s" "$_torch_index_leaf" | tr "[:upper:]" "[:lower:]")
+ case "$_torch_index_leaf" in
+ rocm7.2|gfx120x-all|gfx1151|gfx1150)
+ TORCH_CONSTRAINT="torch>=2.11.0,<2.12.0"
+ TORCHVISION_CONSTRAINT="torchvision>=0.26.0,<0.27.0"
+ TORCHAUDIO_CONSTRAINT="torchaudio>=2.11.0,<2.12.0"
+ ;;
+ esac
+ echo "$TORCH_CONSTRAINT"
+ ' _ "$_url"
+}
+
+assert_eq "gfx120X-all (capital) -> 2.11 floor" "torch>=2.11.0,<2.12.0" \
+ "$(run_floor_case 'https://repo.amd.com/rocm/whl/gfx120X-all')"
+assert_eq "gfx120X-all trailing slash -> 2.11 floor" "torch>=2.11.0,<2.12.0" \
+ "$(run_floor_case 'https://repo.amd.com/rocm/whl/gfx120X-all/')"
+assert_eq "gfx120x-all (lowercase) -> 2.11 floor" "torch>=2.11.0,<2.12.0" \
+ "$(run_floor_case 'https://repo.amd.com/rocm/whl/gfx120x-all')"
+assert_eq "gfx1151 -> 2.11 floor" "torch>=2.11.0,<2.12.0" \
+ "$(run_floor_case 'https://repo.amd.com/rocm/whl/gfx1151')"
+assert_eq "gfx1150 -> 2.11 floor" "torch>=2.11.0,<2.12.0" \
+ "$(run_floor_case 'https://repo.amd.com/rocm/whl/gfx1150')"
+assert_eq "rocm7.2 -> 2.11 floor" "torch>=2.11.0,<2.12.0" \
+ "$(run_floor_case 'https://download.pytorch.org/whl/rocm7.2')"
+assert_eq "gfx110X-all -> default (no floor)" "torch>=2.4,<2.11.0" \
+ "$(run_floor_case 'https://repo.amd.com/rocm/whl/gfx110X-all')"
+assert_eq "rocm6.4 -> default (no floor)" "torch>=2.4,<2.11.0" \
+ "$(run_floor_case 'https://download.pytorch.org/whl/rocm6.4')"
+assert_eq "cu128 -> default (no floor)" "torch>=2.4,<2.11.0" \
+ "$(run_floor_case 'https://download.pytorch.org/whl/cu128')"
+assert_eq "cpu -> default (no floor)" "torch>=2.4,<2.11.0" \
+ "$(run_floor_case 'https://download.pytorch.org/whl/cpu')"
+
# ======================================================================
# Summary
# ======================================================================
diff --git a/tests/sh/test_torch_flavor.sh b/tests/sh/test_torch_flavor.sh
index ead2c4164f..55da2c0f07 100755
--- a/tests/sh/test_torch_flavor.sh
+++ b/tests/sh/test_torch_flavor.sh
@@ -11,14 +11,21 @@ INSTALL_SH="$SCRIPT_DIR/../../install.sh"
PASS=0
FAIL=0
-# Extract the three helper functions from install.sh and source them.
+# Extract the helper functions from install.sh and source them
+# (_torch_index_url_leaf is the shared leaf extractor the classifiers call).
_FUNC_FILE=$(mktemp)
{
sed -n '/^_torch_flavor_tag()/,/^}/p' "$INSTALL_SH"
echo ""
+ sed -n '/^_torch_index_url_leaf()/,/^}/p' "$INSTALL_SH"
+ echo ""
+ sed -n '/^_is_pip_rocm_family_leaf()/,/^}/p' "$INSTALL_SH"
+ echo ""
sed -n '/^_expected_torch_flavor_tag()/,/^}/p' "$INSTALL_SH"
echo ""
sed -n '/^_torch_index_repairable()/,/^}/p' "$INSTALL_SH"
+ echo ""
+ sed -n '/^_tauri_torch_index_family()/,/^}/p' "$INSTALL_SH"
} > "$_FUNC_FILE"
# shellcheck disable=SC1090
. "$_FUNC_FILE"
@@ -56,6 +63,26 @@ assert_eq "amd gfx index" "rocm" "$(_expected_torch_flavor_tag 'https://re
assert_eq "mirror cu130 leaf" "cu130" "$(_expected_torch_flavor_tag 'https://my.mirror/pytorch/whl/cu130')"
assert_eq "unrecognized leaf" "" "$(_expected_torch_flavor_tag 'https://my.mirror/whl/simple')"
assert_eq "empty url" "" "$(_expected_torch_flavor_tag '')"
+# Query/fragment dropped before classification: .../cu128?token=x classifies as cu128,
+# not an opaque leaf that reinstalls every run.
+assert_eq "query-bearing cu128" "cu128" "$(_expected_torch_flavor_tag 'https://m/whl/cu128?token=x')"
+assert_eq "fragment-bearing cpu" "cpu" "$(_expected_torch_flavor_tag 'https://m/whl/cpu#frag')"
+# A cu-suffixed CUSTOM leaf (cu128-private, cu128x) is NOT the cu128 family (exact
+# cu+digits only). Mirrors Python re.fullmatch(cu[0-9]+) / PowerShell.
+assert_eq "cu-suffix custom leaf" "" "$(_expected_torch_flavor_tag 'https://m/whl/cu128-private')"
+assert_eq "cu-alnum custom leaf" "" "$(_expected_torch_flavor_tag 'https://m/whl/cu128x')"
+assert_eq "bare cu digits stays" "cu126" "$(_expected_torch_flavor_tag 'https://m/whl/cu126')"
+# A custom leaf merely STARTING with rocm (rocm-current, rocm-rel-7.2.1) is NOT a pip
+# rocm family -> "" (custom); real families (rocm7.2) and gfx indexes stay "rocm".
+assert_eq "custom rocm-current" "" "$(_expected_torch_flavor_tag 'https://mirror/whl/rocm-current')"
+assert_eq "radeon rocm-rel leaf" "" "$(_expected_torch_flavor_tag 'https://repo.radeon.com/rocm/manylinux/rocm-rel-7.2.1')"
+assert_eq "real rocm7.2 stays" "rocm" "$(_expected_torch_flavor_tag 'https://download.pytorch.org/whl/rocm7.2')"
+# A rocm-SUFFIX private mirror (rocm7.2-private, rocm7-current) is a custom pin ->
+# "" (custom); match the family exactly, not the prefix.
+assert_eq "suffixed rocm7.2-private" "" "$(_expected_torch_flavor_tag 'https://co.internal/whl/rocm7.2-private')"
+assert_eq "suffixed rocm7-current" "" "$(_expected_torch_flavor_tag 'https://co.internal/whl/rocm7-current')"
+assert_eq "two-dot rocm7.2.1" "" "$(_expected_torch_flavor_tag 'https://co.internal/whl/rocm7.2.1')"
+assert_eq "bare rocm7 stays" "rocm" "$(_expected_torch_flavor_tag 'https://download.pytorch.org/whl/rocm7')"
echo "=== _torch_index_repairable ==="
assert_eq "cu130 repairable" "yes" "$(_torch_index_repairable 'https://download.pytorch.org/whl/cu130')"
@@ -64,6 +91,66 @@ assert_eq "gfx repairable" "yes" "$(_torch_index_repairable 'https://repo.
assert_eq "gfx1151 repairable" "yes" "$(_torch_index_repairable 'https://repo.amd.com/rocm/whl/gfx1151/')"
assert_eq "cpu NOT repairable" "no" "$(_torch_index_repairable 'https://download.pytorch.org/whl/cpu')"
assert_eq "unknown NOT repair" "no" "$(_torch_index_repairable 'https://my.mirror/whl/simple')"
+# A suffixed rocm leaf is a verbatim pin, not a --default-index repairable family.
+assert_eq "rocm-private NOT repair" "no" "$(_torch_index_repairable 'https://co.internal/whl/rocm7.2-private')"
+
+echo "=== _is_pip_rocm_family_leaf ==="
+assert_family() {
+ _label="$1"; _expected="$2"; _leaf="$3"
+ if _is_pip_rocm_family_leaf "$_leaf"; then _actual="yes"; else _actual="no"; fi
+ assert_eq "$_label" "$_expected" "$_actual"
+}
+assert_family "rocm7.2 family" "yes" "rocm7.2"
+assert_family "rocm6.4 family" "yes" "rocm6.4"
+assert_family "bare rocm7 family" "yes" "rocm7"
+assert_family "gfx120x-all family" "yes" "gfx120x-all"
+assert_family "gfx1151 family" "yes" "gfx1151"
+assert_family "rocm7.2-private custom" "no" "rocm7.2-private"
+assert_family "rocm7-current custom" "no" "rocm7-current"
+assert_family "rocm-current custom" "no" "rocm-current"
+assert_family "rocm-rel-7.2.1 custom" "no" "rocm-rel-7.2.1"
+assert_family "rocm7.2.1 custom" "no" "rocm7.2.1"
+# A trailing dot (rocm7.) or leading/double dot is NOT a family: both major and minor must
+# be non-empty all-digits, matching Python re.fullmatch(rocm\d+(?:\.\d+)?). Bash previously
+# accepted rocm7. via a bare %/-style trim while Python rejected it (validator asymmetry).
+assert_family "rocm7. trailing-dot custom" "no" "rocm7."
+assert_family "rocm.7 leading-dot custom" "no" "rocm.7"
+assert_family "rocm7..2 double-dot custom" "no" "rocm7..2"
+assert_family "cpu not rocm" "no" "cpu"
+assert_family "cu128 not rocm" "no" "cu128"
+assert_family "simple not rocm" "no" "simple"
+
+echo "=== _torch_index_url_leaf (ALL trailing slashes stripped -> non-empty leaf) ==="
+# A double (or triple) trailing slash must yield the real leaf, not an empty string that
+# fails every classifier arm. Python .rstrip("/") drops them all; bash must match (a bare
+# %/ left .../cu128// classifying as "").
+assert_eq "double slash cu128 leaf" "cu128" "$(_torch_index_url_leaf 'https://m/whl/cu128//')"
+assert_eq "triple slash rocm7.2 leaf" "rocm7.2" "$(_torch_index_url_leaf 'https://m/whl/rocm7.2///')"
+assert_eq "double slash + token leaf" "cu128" "$(_torch_index_url_leaf 'https://m/whl/cu128//?token=x')"
+assert_eq "single slash cu128 leaf" "cu128" "$(_torch_index_url_leaf 'https://m/whl/cu128/')"
+# The classifier that consumes the leaf must therefore still tag a double-slash index.
+assert_eq "double-slash cu128 tag" "cu128" "$(_expected_torch_flavor_tag 'https://m/whl/cu128//')"
+assert_eq "double-slash rocm7.2 tag" "rocm" "$(_expected_torch_flavor_tag 'https://m/whl/rocm7.2//')"
+
+echo "=== _tauri_torch_index_family (credential redaction) ==="
+# A token/fragment must be stripped BEFORE classification so it never reaches the
+# [TAURI:DIAG] line (the family is the last path segment, which else carries the query).
+SKIP_TORCH=false
+assert_eq "token stripped from rocm" "rocm7.2" "$(_tauri_torch_index_family 'https://mirror/whl/rocm7.2?token=SECRET')"
+assert_eq "token-bearing cu classifies" "cu128" "$(_tauri_torch_index_family 'https://m/whl/cu128?token=x')"
+assert_eq "fragment stripped cpu" "cpu" "$(_tauri_torch_index_family 'https://m/whl/cpu#frag')"
+assert_eq "plain rocm7.2 unchanged" "rocm7.2" "$(_tauri_torch_index_family 'https://download.pytorch.org/whl/rocm7.2')"
+# A trailing slash must be stripped too, or the */cu128 and */cpu arms miss .../cu128/
+# and it falls through to "auto".
+assert_eq "trailing slash cu128" "cu128" "$(_tauri_torch_index_family 'https://download.pytorch.org/whl/cu128/')"
+assert_eq "slash + token cu128" "cu128" "$(_tauri_torch_index_family 'https://m/whl/cu128/?token=x')"
+assert_eq "trailing slash cpu" "cpu" "$(_tauri_torch_index_family 'https://m/whl/cpu/')"
+# Regression guard: no secret token substring may survive in any classification.
+_leak=$(_tauri_torch_index_family 'https://mirror/whl/rocm7.2?token=SECRET')
+case "$_leak" in
+ *SECRET*|*token*) assert_eq "no token leak in family" "clean" "leaked:$_leak" ;;
+ *) assert_eq "no token leak in family" "clean" "clean" ;;
+esac
echo ""
echo "Results: $PASS passed, $FAIL failed"
diff --git a/tests/studio/install/test_cuda_repair.py b/tests/studio/install/test_cuda_repair.py
index cea4383268..c6d2b95316 100644
--- a/tests/studio/install/test_cuda_repair.py
+++ b/tests/studio/install/test_cuda_repair.py
@@ -64,15 +64,23 @@ def _run_cuda_repair(
rocm_marker = False,
smi_path = "/usr/bin/nvidia-smi",
cvd = None,
+ index_family = None,
+ index_url = None,
):
"""Invoke _ensure_cuda_torch under a fully mocked host; return the pip mock.
- cvd controls CUDA_VISIBLE_DEVICES: None removes it from the env, any string sets it."""
+ cvd controls CUDA_VISIBLE_DEVICES: None removes it from the env, any string sets it.
+ index_family sets UNSLOTH_TORCH_INDEX_FAMILY (the explicit wheel-index pin).
+ index_url sets UNSLOTH_TORCH_INDEX_URL (the full-URL pin form)."""
env = {}
if rocm_marker:
env["UNSLOTH_ROCM_TORCH_INSTALLED"] = "1"
if cvd is not None:
env["CUDA_VISIBLE_DEVICES"] = cvd
+ if index_family is not None:
+ env["UNSLOTH_TORCH_INDEX_FAMILY"] = index_family
+ if index_url is not None:
+ env["UNSLOTH_TORCH_INDEX_URL"] = index_url
def _which(name, *a, **k):
if name == "nvidia-smi":
@@ -99,6 +107,10 @@ def _run_cuda_repair(
stack_mod.os.environ.pop("UNSLOTH_ROCM_TORCH_INSTALLED", None)
if cvd is None:
stack_mod.os.environ.pop("CUDA_VISIBLE_DEVICES", None)
+ if index_family is None:
+ stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
+ if index_url is None:
+ stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_URL", None)
_ensure_cuda_torch()
return mock_pip
@@ -123,11 +135,74 @@ class TestCudaRepairFires:
assert mock_pip.call_args.kwargs["constrain"] is False
def test_rocm_in_version_string_triggers_repair(self):
- # AMD SDK / Radeon wheels may encode rocm in __version__ without
- # torch.version.hip; the probe prints "hip" for both.
+ # AMD SDK / Radeon wheels may encode rocm in __version__ without torch.version.hip;
+ # the probe prints "hip" for both.
mock_pip = _run_cuda_repair(torch_state = "hip")
assert mock_pip.call_count == 1
+ def test_no_gpu_but_explicit_cuda_pin_repairs(self):
+ # Headless / CI cross-install: an explicit cu* pin commits to CUDA wheels with no
+ # NVIDIA GPU visible, so a ROCm-poisoned venv is still repaired to the pinned family.
+ mock_pip = _run_cuda_repair(
+ nvidia = False,
+ backend = "cuda",
+ index_family = "cu128",
+ torch_state = "hip",
+ )
+ assert mock_pip.call_count == 1
+ assert "cu128" in _index_url(mock_pip)
+
+ def test_cvd_hidden_but_explicit_cuda_pin_repairs(self):
+ # CVD=-1/"" hides the GPU, but an explicit cu* pin skips ALL host-GPU probing, so the
+ # CVD hide gate must not suppress the repair (GPU-less CI: CVD=-1, FAMILY=cu128).
+ for _cvd in ("-1", ""):
+ mock_pip = _run_cuda_repair(
+ nvidia = False,
+ backend = "cuda",
+ cvd = _cvd,
+ index_family = "cu128",
+ torch_state = "hip",
+ )
+ assert mock_pip.call_count == 1
+ assert "cu128" in _index_url(mock_pip)
+
+ def test_tagged_cuda_mismatch_repairs(self):
+ # A healthy CUDA torch whose +cuXXX differs from the pin is repaired.
+ mock_pip = _run_cuda_repair(
+ index_family = "cu128",
+ torch_state = "cuda|cu126",
+ cuda_version = "12.8",
+ )
+ assert mock_pip.call_count == 1
+ assert "cu128" in _index_url(mock_pip)
+
+ def test_untagged_cuda_build_under_pin_repairs(self):
+ # An untagged CUDA build (no +cuXXX tag -> empty installed cu) can't be confirmed
+ # to match the pin, so the pin is enforced with a reinstall.
+ mock_pip = _run_cuda_repair(
+ index_family = "cu128",
+ torch_state = "cuda", # marker cuda, empty installed cu
+ cuda_version = "12.8",
+ )
+ assert mock_pip.call_count == 1
+ assert "cu128" in _index_url(mock_pip)
+
+ def test_broken_probe_with_cuda_pin_repairs(self):
+ # torch present but unimportable under a CUDA pin: the base update won't repair a
+ # broken already-installed torch, so reinstall from the pin instead of stranding it.
+ mock_pip = _run_cuda_repair(torch_state = "hip", torch_rc = 1, index_family = "cu128")
+ assert mock_pip.call_count == 1
+ assert "cu128" in _index_url(mock_pip)
+
+ def test_broken_probe_with_cuda_url_pin_repairs(self):
+ mock_pip = _run_cuda_repair(
+ torch_state = "cpu",
+ torch_rc = 1,
+ index_url = "https://mirror.local/cu128",
+ )
+ assert mock_pip.call_count == 1
+ assert "https://mirror.local/cu128" in _index_url(mock_pip)
+
# No-op cases.
@@ -157,8 +232,9 @@ class TestCudaRepairSkips:
mock_pip = _run_cuda_repair(nvidia = False, torch_state = "hip")
mock_pip.assert_not_called()
- def test_torch_missing_skips(self):
- # Non-zero probe exit = torch missing / un-importable.
+ def test_torch_missing_no_pin_skips(self):
+ # Non-zero probe exit = torch missing/un-importable. With NO CUDA pin the base
+ # install owns it, so leave it alone (a pinned build reinstalls).
mock_pip = _run_cuda_repair(torch_state = "hip", torch_rc = 1)
mock_pip.assert_not_called()
@@ -191,6 +267,109 @@ class TestCudaRepairSkips:
mock_pip = _run_cuda_repair(cvd = "0", torch_state = "hip")
assert mock_pip.call_count == 1
+ def test_matching_tagged_cuda_pin_no_repair(self):
+ # Healthy CUDA torch whose +cuXXX already matches the pin: no reinstall.
+ mock_pip = _run_cuda_repair(
+ index_family = "cu128",
+ torch_state = "cuda|cu128",
+ cuda_version = "12.8",
+ )
+ mock_pip.assert_not_called()
+
+ def test_custom_mirror_leaf_not_treated_as_cuda_pin(self):
+ # A mirror leaf starting with "cu" but not cuXXX (.../custom, .../current) must
+ # NOT be treated as a CUDA pin, so it can't bypass the NVIDIA gate.
+ for _leaf in ("custom", "current"):
+ mock_pip = _run_cuda_repair(
+ nvidia = False,
+ backend = "cuda",
+ index_url = f"https://mymirror.example/{_leaf}",
+ torch_state = "hip",
+ )
+ mock_pip.assert_not_called()
+
+ def test_explicit_cuda_family_leaf_helper(self):
+ # _explicit_cuda_torch_index_url matches cuXXX narrowly, not any cu* leaf.
+ import contextlib
+
+ def _with(url):
+ with patch.dict(stack_mod.os.environ, {"UNSLOTH_TORCH_INDEX_URL": url}, clear = False):
+ stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
+ return stack_mod._explicit_cuda_torch_index_url()
+
+ assert _with("https://download.pytorch.org/whl/cu128") is not None
+ assert _with("https://download.pytorch.org/whl/cu126") is not None
+ assert _with("https://mymirror.example/custom") is None
+ assert _with("https://mymirror.example/current") is None
+ assert _with("https://download.pytorch.org/whl/cpu") is None
+ with contextlib.suppress(Exception):
+ stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_URL", None)
+
+
+class TestTorchBackendDerivationFromPin:
+ """The module-level _TORCH_BACKEND derivation (standalone `studio update`
+ with no install.sh-set UNSLOTH_TORCH_BACKEND) must classify the pinned index
+ leaf via _is_cuda_family_leaf (^cu[0-9]), NOT a bare startswith("cu"). A
+ full-override URL ending in /current or /custom must fall through to backend
+ "" (probe the GPU) so _ensure_rocm_torch() still repairs a wrong/CPU torch on
+ AMD hosts, instead of being wrongly branded "cuda" and returning early."""
+
+ @staticmethod
+ def _derive(env):
+ # Re-run the module's import-time derivation, using its own _is_cuda_family_leaf
+ # so this stays in lockstep.
+ idx_override = (
+ env.get("UNSLOTH_TORCH_INDEX_URL", "").strip()
+ or env.get("UNSLOTH_TORCH_INDEX_FAMILY", "").strip()
+ )
+ backend = env.get("UNSLOTH_TORCH_BACKEND", "").lower()
+ if not backend:
+ leaf = idx_override.rstrip("/").rsplit("/", 1)[-1].lower()
+ if leaf.startswith(("rocm", "gfx")):
+ backend = "rocm"
+ elif leaf == "cpu":
+ backend = "cpu"
+ elif stack_mod._is_cuda_family_leaf(leaf):
+ backend = "cuda"
+ return backend
+
+ def test_cu128_pin_is_cuda(self):
+ assert (
+ self._derive({"UNSLOTH_TORCH_INDEX_URL": "https://download.pytorch.org/whl/cu128"})
+ == "cuda"
+ )
+
+ def test_cu128_family_is_cuda(self):
+ assert self._derive({"UNSLOTH_TORCH_INDEX_FAMILY": "cu128"}) == "cuda"
+
+ def test_current_leaf_not_cuda(self):
+ # ^cu[0-9] rejects /current -> backend stays "" (probe GPU), so an AMD host still
+ # repairs a CPU/wrong torch instead of short-circuiting.
+ assert self._derive({"UNSLOTH_TORCH_INDEX_URL": "https://mymirror.example/current"}) == ""
+
+ def test_custom_leaf_not_cuda(self):
+ assert self._derive({"UNSLOTH_TORCH_INDEX_URL": "https://mymirror.example/custom"}) == ""
+
+ def test_rocm_and_gfx_pins_are_rocm(self):
+ assert self._derive({"UNSLOTH_TORCH_INDEX_FAMILY": "rocm7.2"}) == "rocm"
+ assert (
+ self._derive({"UNSLOTH_TORCH_INDEX_URL": "https://repo.amd.com/rocm/whl/gfx120X-all"})
+ == "rocm"
+ )
+
+ def test_cpu_pin_is_cpu(self):
+ assert self._derive({"UNSLOTH_TORCH_INDEX_FAMILY": "cpu"}) == "cpu"
+
+ def test_source_uses_helper_not_bare_startswith(self):
+ # Guard against a regression back to elif _idx_leaf.startswith("cu").
+ src = _STACK_PATH.read_text(encoding = "utf-8")
+ assert (
+ "elif _is_cuda_family_leaf(_idx_leaf):" in src
+ ), "_TORCH_BACKEND derivation must classify CUDA via _is_cuda_family_leaf"
+ assert (
+ 'elif _idx_leaf.startswith("cu"):' not in src
+ ), "_TORCH_BACKEND derivation must not use a bare startswith('cu')"
+
# CUDA index ladder.
diff --git a/tests/studio/install/test_gpu_detection_followups.py b/tests/studio/install/test_gpu_detection_followups.py
index f5ad9566d7..d2fd7ae8db 100644
--- a/tests/studio/install/test_gpu_detection_followups.py
+++ b/tests/studio/install/test_gpu_detection_followups.py
@@ -284,7 +284,9 @@ class TestBackendExportLeafClassification:
def test_export_block_uses_leaf(self, install_src):
anchor = install_src.find("_torch_index_leaf=")
assert anchor >= 0, "backend export must classify on the final path segment"
- window = install_src[anchor : anchor + 500]
+ # Window spans the leaf-normalization prelude (query/frag drop + all-slash trim loop)
+ # through the export case arms.
+ window = install_src[anchor : anchor + 900]
assert 'export UNSLOTH_TORCH_BACKEND="rocm"' in window
assert 'export UNSLOTH_TORCH_BACKEND="cpu"' in window
assert 'export UNSLOTH_TORCH_BACKEND="cuda"' in window
@@ -459,3 +461,130 @@ class TestHiddenCvdNotUsable:
cvd,
)
assert out == expected
+
+
+class TestRedactInstallOutput:
+ """_redact_install_output scrubs index-URL credentials from a captured install log
+ before it is printed on failure (uv/pip embeds the failing --index-url verbatim)."""
+
+ def test_userinfo_redacted(self):
+ out = stack_mod._redact_install_output(
+ "ERROR: failed https://alice:s3cr3t@download.pytorch.org/whl/cu128"
+ )
+ assert out == "ERROR: failed https://@download.pytorch.org/whl/cu128"
+
+ def test_bytes_input_decoded_and_redacted(self):
+ out = stack_mod._redact_install_output(b"fetch https://ghp_deadbeef@host/whl/cu128 failed")
+ assert out == "fetch https://@host/whl/cu128 failed"
+
+ def test_query_values_redacted(self):
+ out = stack_mod._redact_install_output(
+ "url https://host/whl/cu128?token=abcd1234&channel=beta unreachable"
+ )
+ assert out == "url https://host/whl/cu128?token=&channel= unreachable"
+
+ def test_fragment_redacted(self):
+ out = stack_mod._redact_install_output(
+ "ERROR: could not fetch https://mirror.local/whl/cu128#token=SECRET123 (403)"
+ )
+ assert out == "ERROR: could not fetch https://mirror.local/whl/cu128# (403)"
+
+ def test_query_and_fragment_both_redacted(self):
+ out = stack_mod._redact_install_output("https://host/whl/cu128?token=abc#sig=xyz done")
+ assert out == "https://host/whl/cu128?token=# done"
+
+ def test_bare_hash_comment_untouched(self):
+ # The fragment redaction is URL-anchored: a shell comment in tool output survives.
+ assert (
+ stack_mod._redact_install_output("# retrying with --no-cache-dir")
+ == "# retrying with --no-cache-dir"
+ )
+
+ def test_plain_line_untouched(self):
+ assert (
+ stack_mod._redact_install_output("Resolved 42 packages in 1.2s")
+ == "Resolved 42 packages in 1.2s"
+ )
+
+ def test_no_secret_substring_survives(self):
+ out = stack_mod._redact_install_output(
+ "https://alice:s3cr3t@host/whl/cu128?token=SUPERSECRET#frag=ALSOSECRET"
+ )
+ assert "s3cr3t" not in out and "SUPERSECRET" not in out and "ALSOSECRET" not in out
+
+
+class TestTrimIndexPathSlashes:
+ """_trim_index_path_slashes strips trailing PATH slashes only; a ?query/#fragment token
+ ending in "/" must survive (a whole-URL rstrip would corrupt a base64 token)."""
+
+ def test_double_path_slash_collapsed(self):
+ assert stack_mod._trim_index_path_slashes("https://h/whl/cu128//") == "https://h/whl/cu128"
+
+ def test_query_token_slash_preserved(self):
+ assert (
+ stack_mod._trim_index_path_slashes("https://h/whl/cu128?token=ab12cd/")
+ == "https://h/whl/cu128?token=ab12cd/"
+ )
+
+ def test_path_slash_trimmed_query_kept(self):
+ assert (
+ stack_mod._trim_index_path_slashes("https://h/whl/cu128//?token=ab12cd/")
+ == "https://h/whl/cu128?token=ab12cd/"
+ )
+
+ def test_fragment_slash_preserved(self):
+ assert (
+ stack_mod._trim_index_path_slashes("https://h/whl/cu128#anchor/")
+ == "https://h/whl/cu128#anchor/"
+ )
+
+
+class TestRocmFamilyLeafParity:
+ """_is_pip_rocm_family_leaf must match re.fullmatch(rocm\\d+(?:\\.\\d+)?): a trailing dot
+ (rocm7.) is a CUSTOM pin, not a family (the historical bash/py validator asymmetry)."""
+
+ @pytest.mark.parametrize(
+ "leaf, expected",
+ [
+ ("rocm7", True),
+ ("rocm7.2", True),
+ ("gfx1151", True),
+ ("rocm7.", False),
+ ("rocm.7", False),
+ ("rocm7..2", False),
+ ("rocm7.2.1", False),
+ ("rocm7.2-private", False),
+ ("cpu", False),
+ ("cu128", False),
+ ],
+ )
+ def test_family_classification(self, leaf, expected):
+ assert stack_mod._is_pip_rocm_family_leaf(leaf) is expected
+
+
+class TestTorchIndexLeafAllSlashes:
+ """_torch_index_leaf drops query/fragment then strips ALL trailing slashes, so a
+ double-slash index still yields the real leaf (not an empty string)."""
+
+ @pytest.mark.parametrize(
+ "url, expected",
+ [
+ ("https://m/whl/cu128//", "cu128"),
+ ("https://m/whl/rocm7.2///", "rocm7.2"),
+ ("https://m/whl/cu128//?token=x", "cu128"),
+ ("https://m/whl/cu128/", "cu128"),
+ ],
+ )
+ def test_leaf_never_empty_on_double_slash(self, url, expected):
+ assert stack_mod._torch_index_leaf(url) == expected
+
+
+class TestUvIndexEnvVarsScrub:
+ """The pinned-install env scrub must drop PIP_NO_INDEX (which makes the pip fallback
+ ignore ALL indexes, defeating the pin) and PIP_INDEX_URL (replaces the pinned index)."""
+
+ def test_pip_no_index_scrubbed(self):
+ assert "PIP_NO_INDEX" in stack_mod._UV_INDEX_ENV_VARS
+
+ def test_pip_index_url_scrubbed(self):
+ assert "PIP_INDEX_URL" in stack_mod._UV_INDEX_ENV_VARS
diff --git a/tests/studio/install/test_pr5940_followups.py b/tests/studio/install/test_pr5940_followups.py
index d6dc8b2f8e..dd2a7ec487 100644
--- a/tests/studio/install/test_pr5940_followups.py
+++ b/tests/studio/install/test_pr5940_followups.py
@@ -784,7 +784,7 @@ def test_install_python_stack_windows_rocm_repair_pins_and_is_nonfatal():
assert re.search(
r'"' + gfx + r'":\s*_ROCM_TORCH_PKG_SPECS\["rocm7\.2"\]', text
), f"{gfx} must pin to the rocm7.2 trio like install.ps1/setup.ps1"
- i = text.find('f"ROCm torch (Windows, {gfx_arch})"')
+ i = text.find("f\"ROCm torch (Windows, {gfx_arch or 'pinned'})\"")
assert i != -1, "Windows ROCm repair pip call not found"
# The nearest preceding call must be the nonfatal pip_install_try, not pip_install.
j = text.rfind("pip_install_try(", 0, i)
diff --git a/tests/studio/install/test_rocm_support.py b/tests/studio/install/test_rocm_support.py
index e7ac0ec82d..b94578369d 100644
--- a/tests/studio/install/test_rocm_support.py
+++ b/tests/studio/install/test_rocm_support.py
@@ -569,19 +569,27 @@ 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)
@patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
@patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 1))
- def test_torch_already_has_cuda_skips(self, mock_ver, mock_gpu, mock_nvidia, mock_pip):
- """If torch already has CUDA, should skip ROCm reinstall."""
+ def test_cuda_torch_on_amd_host_reinstalls(
+ self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
+ ):
+ """A CUDA-only torch build is unusable on an AMD-only host, so it must be
+ reinstalled to ROCm (has_hip_torch is driven by the empty HIP marker, not
+ by treating the CUDA version string as a HIP marker)."""
mock_probe = MagicMock()
mock_probe.returncode = 0
- mock_probe.stdout = b"12.6\n" # CUDA version
+ # Single-line probe: empty HIP marker before "|" for a CUDA build.
+ mock_probe.stdout = b"|2.10.0+cu126\n"
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
- mock_pip.assert_not_called()
+ assert mock_pip.call_count == 1
+ assert "rocm7.1" in str(mock_pip.call_args_list[0])
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@@ -591,12 +599,31 @@ class TestEnsureRocmTorch:
"""If torch already has HIP, should skip ROCm reinstall."""
mock_probe = MagicMock()
mock_probe.returncode = 0
- mock_probe.stdout = b"7.1.12345\n" # HIP version
+ mock_probe.stdout = b"7.1.12345|2.10.0+rocm7.1\n" # HIP marker + version
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_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)
+ @patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 1))
+ def test_cpu_torch_probe_line_not_read_as_hip(self, mock_ver, mock_gpu, mock_nvidia, mock_pip):
+ """A CPU build's probe line ("|2.10.0+cpu") must not read as HIP: the version
+ after the "|" separator is data, not a HIP marker, so has_hip_torch stays False
+ and the reinstall fires."""
+ mock_probe = MagicMock()
+ mock_probe.returncode = 0
+ mock_probe.stdout = b"|2.10.0+cpu\n"
+ with patch("os.path.isdir", return_value = True):
+ with patch("subprocess.run", return_value = mock_probe):
+ with patch.object(stack_mod, "pip_install_try", return_value = True):
+ _ensure_rocm_torch()
+ assert mock_pip.call_count == 1
+ assert "rocm7.1" in str(mock_pip.call_args_list[0])
+
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@@ -680,6 +707,295 @@ class TestEnsureRocmTorch:
torch_call = mock_pip.call_args_list[0]
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)
+ @patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
+ @patch.object(stack_mod, "_detect_rocm_version", return_value = (6, 4))
+ def test_explicit_gfx_index_honored_and_skips_strix_reroute(
+ self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
+ ):
+ """An explicit gfx wheel-index pin is authoritative: install from it verbatim
+ with torch 2.11, and never re-probe gfx codes to second-guess it (host ROCm 6.4
+ would otherwise pick the rocm6.4 wheel / trigger the Strix re-route)."""
+ mock_probe = MagicMock()
+ mock_probe.returncode = 0
+ mock_probe.stdout = b"\n" # cpu torch -> reinstall
+ env = {"UNSLOTH_TORCH_INDEX_URL": "https://repo.amd.com/rocm/whl/gfx1151"}
+ with patch.dict(stack_mod.os.environ, env, clear = False):
+ stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
+ with patch("os.path.isdir", return_value = True):
+ with patch("subprocess.run", return_value = mock_probe):
+ # Would raise if the Strix block ran (it is skipped on an explicit pin).
+ with patch.object(
+ stack_mod, "_detect_amd_gfx_codes", side_effect = AssertionError
+ ):
+ _ensure_rocm_torch()
+ assert mock_pip.call_count == 1
+ torch_call = str(mock_pip.call_args_list[0])
+ assert "gfx1151" in torch_call
+ assert "torch>=2.11.0,<2.12.0" in torch_call
+
+ def test_rocm_pin_family_mismatch_helper(self):
+ """_rocm_pin_family_mismatch: exact rocm compare, else the 2.11 line."""
+ f = stack_mod._rocm_pin_family_mismatch
+ base = "https://download.pytorch.org/whl"
+ amd = "https://repo.amd.com/rocm/whl"
+ # Exact rocm version comparison.
+ assert f(f"{base}/rocm7.2", "2.11.0+rocm7.2") is False
+ assert f(f"{base}/rocm7.2", "2.10.0+rocm6.4") is True
+ assert f(f"{base}/rocm6.4", "2.10.0+rocm6.4") is False
+ # rocm7.2 is KNOWN-2.11. A +rocm7.2 wheel whose RELEASE drifted off 2.11 shares the
+ # tag but violates the spec -> mismatch (a plain version compare would accept it).
+ assert f(f"{base}/rocm7.2", "2.12.0+rocm7.2") is True
+ assert f(f"{base}/rocm7.2", "2.13.0+rocm7.2") is True
+ assert f(f"{base}/rocm7.2", "2.11.5+rocm7.2") is False # patch on 2.11 is in-spec
+ # An UNKNOWN newer rocm (not on the 2.11 allowlist) is not floored to 2.11, so a
+ # matching rocm version at any release line is NOT a mismatch on this branch.
+ assert f(f"{base}/rocm8.0", "2.12.0+rocm8.0") is False
+ # gfx pin (2.11 line) vs installed release line.
+ assert f(f"{amd}/gfx1151", "2.10.0+rocm6.4") is True
+ assert f(f"{amd}/gfx1151", "2.11.0+rocm7.13.0") is False
+ # rocm7.2 pin vs an untagged (no +rocm) wheel: a CPU/CUDA build never
+ # satisfies a ROCm pin, regardless of its release line -> always a mismatch.
+ assert f(f"{base}/rocm7.2", "2.10.0") is True
+ assert f(f"{base}/rocm7.2", "2.11.0") is True
+ assert f(f"{base}/rocm6.4", "2.10.0") is True
+ # A 2.11-allowlist gfx pin over a GENERIC (two-part +rocm7.2) 2.11 wheel mismatches:
+ # the user wants AMD's per-arch (three-part) wheel, not the generic one.
+ assert f(f"{amd}/gfx1151", "2.11.0+rocm7.2") is True
+ assert f(f"{amd}/gfx120X-all", "2.11.0+rocm7.2") is True
+ # ...but an already-installed per-arch (three-part) wheel is NOT re-flagged
+ # (no reinstall loop once the correct gfx wheel is present).
+ assert f(f"{amd}/gfx120X-all", "2.11.0+rocm7.13.0") is False
+ assert f(f"{amd}/gfx1150", "2.11.0+rocm7.13.0") is False
+ # A NON-2.11 gfx pin (gfx110X-all/gfx90a/gfx908) tracks the default <2.11 spec: a
+ # correct 2.10+rocm wheel is NOT a mismatch, a 2.11 build is.
+ assert f(f"{amd}/gfx110X-all", "2.10.0+rocm6.4") is False
+ assert f(f"{amd}/gfx90a", "2.10.0+rocm6.3") is False
+ assert f(f"{amd}/gfx908", "2.10.0+rocm7.0") is False
+ assert f(f"{amd}/gfx110X-all", "2.11.0+rocm7.2") is True
+ # A non-2.11 gfx pin over an untagged (no +rocm) wheel is a mismatch even
+ # when torch is already <2.11: a CPU/CUDA build never satisfies the ROCm pin.
+ assert f(f"{amd}/gfx110X-all", "2.10.0") is True
+ assert f(f"{amd}/gfx90a", "2.10.0") is True
+ # A major-only rocm pin (rocm7) compares on the major alone: rocm6.x mismatches,
+ # any rocm7.x satisfies it, an untagged wheel never does, a bare +rocm is lenient.
+ assert f(f"{base}/rocm7", "2.10.0+rocm6.4") is True
+ assert f(f"{base}/rocm7", "2.11.0+rocm7.2") is False
+ assert f(f"{base}/rocm7", "2.11.0+rocm7.13.0") is False
+ assert f(f"{base}/rocm7", "2.10.0") is True
+ assert f(f"{base}/rocm7", "2.10.0+rocm") is False
+
+ @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)
+ @patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
+ @patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
+ def test_rocm_pin_mismatch_over_installed_rocm_reinstalls(
+ self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
+ ):
+ """A rocm7.2 pin over an already-installed OLDER +rocm6.4 build must reinstall,
+ even though has_hip_torch is True (the ROCm analogue of the CUDA cuXXX mismatch)."""
+ mock_probe = MagicMock()
+ mock_probe.returncode = 0
+ # HIP marker present (has_hip_torch=True) + installed +rocm6.4 wheel.
+ mock_probe.stdout = b"6.4.12345|2.10.0+rocm6.4\n"
+ env = {"UNSLOTH_TORCH_INDEX_FAMILY": "rocm7.2"}
+ with patch.dict(stack_mod.os.environ, env, clear = False):
+ stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_URL", None)
+ with patch("os.path.isdir", return_value = True):
+ with patch("subprocess.run", return_value = mock_probe):
+ _ensure_rocm_torch()
+ torch_call = str(mock_pip.call_args_list[0])
+ assert "rocm7.2" in torch_call
+ assert "torch>=2.11.0,<2.12.0" in 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)
+ @patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
+ @patch.object(stack_mod, "_detect_rocm_version", return_value = (6, 4))
+ def test_gfx_pin_over_installed_pre211_rocm_reinstalls(
+ self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
+ ):
+ """A gfx* pin (2.11 line) over an installed pre-2.11 +rocm6.4 build reinstalls."""
+ mock_probe = MagicMock()
+ mock_probe.returncode = 0
+ mock_probe.stdout = b"6.4.12345|2.10.0+rocm6.4\n"
+ env = {"UNSLOTH_TORCH_INDEX_URL": "https://repo.amd.com/rocm/whl/gfx1151"}
+ with patch.dict(stack_mod.os.environ, env, clear = False):
+ stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
+ with patch("os.path.isdir", return_value = True):
+ with patch("subprocess.run", return_value = mock_probe):
+ with patch.object(
+ stack_mod, "_detect_amd_gfx_codes", side_effect = AssertionError
+ ):
+ _ensure_rocm_torch()
+ torch_call = str(mock_pip.call_args_list[0])
+ assert "gfx1151" in torch_call
+ assert "torch>=2.11.0,<2.12.0" in 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)
+ @patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
+ @patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
+ def test_rocm_pin_matches_installed_no_torch_reinstall(
+ self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
+ ):
+ """A rocm7.2 pin over an already-matching +rocm7.2 build must NOT reinstall torch
+ (no false reinstall of a correct ROCm venv)."""
+ mock_probe = MagicMock()
+ mock_probe.returncode = 0
+ mock_probe.stdout = b"7.2.12345|2.11.0+rocm7.2\n"
+ env = {"UNSLOTH_TORCH_INDEX_FAMILY": "rocm7.2"}
+ with patch.dict(stack_mod.os.environ, env, clear = False):
+ stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_URL", None)
+ with patch("os.path.isdir", return_value = True):
+ with patch("subprocess.run", return_value = mock_probe):
+ _ensure_rocm_torch()
+ # No torch reinstall: any pip_install call must not target a torch index.
+ for _call in mock_pip.call_args_list:
+ _args = [str(a) for a in _call.args]
+ if "--index-url" in _args:
+ _url = _args[_args.index("--index-url") + 1]
+ assert "rocm7.2" not in _url or "torch" not in " ".join(
+ _args
+ ), "torch must not be reinstalled when the pin already matches"
+ # A torch reinstall would pass torch>=... as a positional; assert none did.
+ assert not any(
+ any(str(a).startswith("torch") for a in _c.args) for _c in mock_pip.call_args_list
+ )
+
+ @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)
+ @patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
+ @patch.object(stack_mod, "_detect_rocm_version", return_value = (6, 4))
+ def test_non211_gfx_pin_over_210_rocm_no_reinstall(
+ self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
+ ):
+ """A gfx110X-all pin (NOT in the 2.11 allowlist) over a correct 2.10+rocm
+ wheel must NOT be flagged stale -- the install path uses the default <2.11
+ specs for that arch, so re-flagging would reinstall-loop on every update."""
+ mock_probe = MagicMock()
+ mock_probe.returncode = 0
+ mock_probe.stdout = b"6.4.12345|2.10.0+rocm6.4\n"
+ env = {"UNSLOTH_TORCH_INDEX_URL": "https://repo.amd.com/rocm/whl/gfx110X-all"}
+ with patch.dict(stack_mod.os.environ, env, clear = False):
+ stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
+ with patch("os.path.isdir", return_value = True):
+ with patch("subprocess.run", return_value = mock_probe):
+ _ensure_rocm_torch()
+ # has_hip_torch True + no mismatch -> torch must NOT be reinstalled.
+ assert not any(
+ any(str(a).startswith("torch") for a in _c.args) for _c in mock_pip.call_args_list
+ )
+
+ @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)
+ @patch.object(stack_mod, "_has_rocm_gpu", return_value = True)
+ @patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 2))
+ def test_gfx_pin_over_generic_rocm211_reinstalls(
+ self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
+ ):
+ """A gfx1151 pin over a GENERIC (two-part +rocm7.2) 2.11 wheel must reinstall
+ the AMD per-arch wheel -- even though both are torch 2.11, the generic wheel
+ is not the per-arch build the user pinned (Strix stays off the generic wheel)."""
+ mock_probe = MagicMock()
+ mock_probe.returncode = 0
+ mock_probe.stdout = b"7.2.12345|2.11.0+rocm7.2\n"
+ env = {"UNSLOTH_TORCH_INDEX_URL": "https://repo.amd.com/rocm/whl/gfx1151"}
+ with patch.dict(stack_mod.os.environ, env, clear = False):
+ stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
+ with patch("os.path.isdir", return_value = True):
+ with patch("subprocess.run", return_value = mock_probe):
+ with patch.object(
+ stack_mod, "_detect_amd_gfx_codes", side_effect = AssertionError
+ ):
+ _ensure_rocm_torch()
+ torch_call = str(mock_pip.call_args_list[0])
+ assert "gfx1151" in torch_call
+ assert "torch>=2.11.0,<2.12.0" in torch_call
+
+ def test_radeon_url_not_classified_as_pip_rocm_family(self):
+ """A repo.radeon.com find-links dir (leaf rocm-rel-7.2.1) starts with "rocm" but is
+ NOT a pip --index-url ROCm family: it must route to the verbatim path, not a
+ --index-url reinstall that fails against a find-links listing."""
+ leaf_f = stack_mod._is_pip_rocm_family_leaf
+ # Real pip ROCm families (download.pytorch.org/whl/rocmX.Y, repo.amd.com gfx).
+ assert leaf_f("rocm7.2") is True
+ assert leaf_f("rocm6.4") is True
+ assert leaf_f("gfx120x-all") is True
+ assert leaf_f("gfx1151") is True
+ # A bare rocm (no minor) is still an exact family.
+ assert leaf_f("rocm7") is True
+ # A Radeon find-links dir leaf, a custom mirror, cpu and cuda are NOT pip rocm.
+ assert leaf_f("rocm-rel-7.2.1") is False
+ assert leaf_f("simple") is False
+ assert leaf_f("current") is False
+ assert leaf_f("cpu") is False
+ assert leaf_f("cu128") is False
+ # A rocm-SUFFIX private mirror shares the family prefix but is a custom pin
+ # the verbatim path owns: a ^rocm\d PREFIX match would wrongly treat it as a
+ # --index-url family. Match EXACTLY.
+ assert leaf_f("rocm7.2-private") is False
+ assert leaf_f("rocm7-current") is False
+ assert leaf_f("rocm7.2.1") is False # two-part local suffix -> custom, not rocm7.2
+
+ radeon = "https://repo.radeon.com/rocm/manylinux/rocm-rel-7.2.1"
+ pip_rocm = "https://download.pytorch.org/whl/rocm7.2"
+ amd_gfx = "https://repo.amd.com/rocm/whl/gfx120X-all"
+
+ def _classify(url, fn):
+ with patch.dict(stack_mod.os.environ, {"UNSLOTH_TORCH_INDEX_URL": url}, clear = False):
+ stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
+ return fn()
+
+ rocm_fn = stack_mod._explicit_rocm_torch_index_url
+ unk_fn = stack_mod._explicit_unknown_family_torch_index_url
+ # Real pip rocm/gfx pins ARE a ROCm family (reinstallable via --index-url) and
+ # are NOT "unknown".
+ assert _classify(pip_rocm, rocm_fn) == pip_rocm
+ assert _classify(amd_gfx, rocm_fn) == amd_gfx
+ assert _classify(pip_rocm, unk_fn) is None
+ assert _classify(amd_gfx, unk_fn) is None
+ # The Radeon find-links URL is NOT a pip ROCm family (so _ensure_rocm_torch skips
+ # it) and IS unknown, so the family repair helpers leave it alone.
+ assert _classify(radeon, rocm_fn) is None
+ assert _classify(radeon, unk_fn) == radeon
+
+ # A rocm-suffix private mirror routes the same way: NOT a pip rocm family,
+ # IS an unknown-family (verbatim) pin.
+ suffixed = "https://co.internal/whl/rocm7.2-private"
+ assert _classify(suffixed, rocm_fn) is None
+ assert _classify(suffixed, unk_fn) == suffixed
+
+ @patch.object(stack_mod, "pip_install")
+ def test_ensure_cpu_torch_broken_probe_reinstalls(self, mock_pip):
+ """_ensure_cpu_torch: torch present but unimportable (probe exit != 0) under an
+ explicit CPU pin must reinstall from the pin, not return -- the base update does
+ not repair a broken installed torch, so returning would strand it (Codex P2)."""
+ mock_probe = MagicMock()
+ mock_probe.returncode = 1 # torch present but cannot import
+ mock_probe.stdout = b""
+ env = {"UNSLOTH_TORCH_INDEX_URL": "https://mirror.local/cpu"}
+ with patch.dict(stack_mod.os.environ, env, clear = False):
+ stack_mod.os.environ.pop("UNSLOTH_TORCH_INDEX_FAMILY", None)
+ with patch("subprocess.run", return_value = mock_probe):
+ with patch.object(stack_mod, "NO_TORCH", False):
+ stack_mod._ensure_cpu_torch()
+ assert mock_pip.call_count == 1
+ assert "https://mirror.local/cpu" in str(mock_pip.call_args)
+
@patch.object(stack_mod, "IS_WINDOWS", False)
@patch.object(stack_mod, "pip_install_try", return_value = True)
@patch.object(stack_mod, "pip_install")
@@ -732,7 +1048,7 @@ class TestEnsureRocmTorch:
mock_pip.assert_not_called()
-# TEST: install_python_stack.py -- _has_rocm_gpu KFD sysfs vendor_id guard
+# TEST: install_python_stack.py -- torch-index MARKER mechanism (PR #6692)
class TestHasRocmGpuKfdVendorGuard:
@@ -1716,9 +2032,8 @@ class TestDetectWindowsGfxArch:
assert result == "gfx1200"
def test_returns_arch_on_crash_with_gcnarchname_in_output(self):
- # Regression #6043: hipinfo may crash (0xC0000005 on RDNA 4) after
- # printing gcnArchName. Accept the arch whenever gcnArchName is in
- # stdout, regardless of exit code (previously a CPU fallback).
+ # Regression #6043: hipinfo may crash (0xC0000005 on RDNA 4) after printing
+ # gcnArchName. Accept the arch whenever gcnArchName is in stdout, any exit code.
mock_result = MagicMock()
mock_result.returncode = -1073741819 # 0xC0000005 STATUS_ACCESS_VIOLATION
mock_result.stdout = b"gcnArchName : gfx1200\nsome other line\n"
@@ -2360,6 +2675,49 @@ class TestWindowsRocmTorchaoGuard:
assert not any("torchao" in arg for arg in installed_specs)
+class TestProgressStepCountMatchesTotal:
+ """The progress bar must reach exactly _TOTAL: every _progress() step is counted in
+ base_total. Regression for a repair step added without incrementing base_total,
+ which pushed _STEP past _TOTAL (Codex P2)."""
+
+ def _run_stack(self, tmp_path, *, is_windows, is_macos, is_mac_arm):
+ unstructured_plugin = tmp_path / "unstructured"
+ github_plugin = tmp_path / "github"
+ unstructured_plugin.mkdir()
+ github_plugin.mkdir()
+ sub = MagicMock()
+ sub.returncode = 0
+ sub.stdout = ""
+ with (
+ patch.dict(os.environ, {"SKIP_STUDIO_BASE": "1"}),
+ patch.object(stack_mod, "IS_WINDOWS", is_windows),
+ patch.object(stack_mod, "IS_MACOS", is_macos),
+ patch.object(stack_mod, "IS_MAC_ARM", is_mac_arm),
+ patch.object(stack_mod, "NO_TORCH", False),
+ patch.object(stack_mod, "_rocm_windows_torch_installed", False),
+ patch.object(stack_mod, "_bootstrap_uv", return_value = False),
+ patch.object(stack_mod, "_installed_torch_is_windows_rocm", return_value = False),
+ patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = True),
+ patch.object(stack_mod, "_repair_bad_anyio"),
+ patch.object(stack_mod, "_ensure_cuda_torch"),
+ patch.object(stack_mod, "_ensure_rocm_torch"),
+ patch.object(stack_mod, "_ensure_cpu_torch"),
+ patch.object(stack_mod, "LOCAL_DD_UNSTRUCTURED_PLUGIN", unstructured_plugin),
+ patch.object(stack_mod, "LOCAL_DD_GITHUB_PLUGIN", github_plugin),
+ patch.object(stack_mod.subprocess, "run", return_value = sub),
+ ):
+ assert stack_mod.install_python_stack() == 0
+ return stack_mod._STEP, stack_mod._TOTAL
+
+ def test_windows_progress_reaches_total(self, tmp_path):
+ step, total = self._run_stack(tmp_path, is_windows = True, is_macos = False, is_mac_arm = False)
+ assert step == total, f"Windows progress {step} != total {total} (final step uncounted)"
+
+ def test_linux_progress_reaches_total(self, tmp_path):
+ step, total = self._run_stack(tmp_path, is_windows = False, is_macos = False, is_mac_arm = False)
+ assert step == total, f"Linux progress {step} != total {total}"
+
+
# TEST: worker.py -- Windows ROCm patches (source-level checks)
@@ -2846,6 +3204,24 @@ class TestStrixRocm71Override:
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
assert "TORCH_CONSTRAINT" in source and "2.11" in source
+ def test_torch_constraint_211_matches_leaf_not_whole_url(self):
+ """The 2.11 constraint case must match the index LEAF, not the whole URL.
+
+ A custom UNSLOTH_PYTORCH_MIRROR whose base path contains a gfx/rocm7.2
+ segment (e.g. https://mirror.local/gfx-cache) with a cu*/cpu family must
+ not be pushed to the torch 2.11 line -- same leaf-only reasoning the
+ UNSLOTH_TORCH_BACKEND classification uses.
+ """
+ source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
+ # The 2.11 constraint block must switch on $_torch_index_leaf, not the full
+ # $TORCH_INDEX_URL (a */gfx* match false-positives on a mirror base path). Only the
+ # _grouped_mm-bug gfx families (gfx120X-all / gfx1151 / gfx1150) are pushed to 2.11;
+ # a bare gfx* would also floor gfx110X-all/gfx90a/gfx908, left bare on purpose.
+ assert 'case "$_torch_index_leaf" in\n rocm7.2|gfx120x-all|gfx1151|gfx1150)' in source, (
+ "the torch>=2.11 constraint must match the specific gfx leaves that need "
+ "it (rocm7.2|gfx120x-all|gfx1151|gfx1150), not a bare gfx* or the whole URL"
+ )
+
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")
@@ -2934,9 +3310,9 @@ class TestServerStartupRocmFixes:
assert '"BNB_ROCM_VERSION" not in os.environ' in source
# ── hipInfo.exe PATH prepend (bitsandbytes arch-probe fix) ────────────────
- # bnb's get_rocm_gpu_arch() runs hipinfo.exe via PATH at import; the AMD
- # wheel ships it in venv Scripts (on PATH only for activated venvs), so
- # without the prepend bnb logs "[WinError 2]" when launched directly.
+ # bnb's get_rocm_gpu_arch() runs hipinfo.exe via PATH at import; the AMD wheel ships it
+ # in venv Scripts (on PATH only for activated venvs), so without the prepend bnb logs
+ # "[WinError 2]" when launched directly.
def test_main_py_prepends_hipinfo_dir_to_path(self):
"""main.py must make hipInfo.exe resolvable before bnb imports."""
@@ -3178,11 +3554,10 @@ class TestRocmGfxForwarding:
assert '$HelperReleaseRepo = "unslothai/llama.cpp"' in source
assert "$HelperReleaseRepo = if (" not in source
- # The text pins above guard the literal. The tests below *execute* the real
- # routing line from setup.sh / setup.ps1 and assert the resolved release repo,
- # so a refactor that reintroduces a conditional (or a ggml-org branch) is still
- # caught. Inputs are varied -- CPU-only, inferred/forwarded gfx, usable NVIDIA --
- # to prove no host slips back onto ggml-org. No GPU, no tooling, no network.
+ # The text pins above guard the literal. The tests below execute the real routing line
+ # from setup.sh / setup.ps1 and assert the resolved release repo, so a refactor that
+ # reintroduces a conditional (or a ggml-org branch) is still caught. Inputs vary
+ # (CPU-only, inferred/forwarded gfx, usable NVIDIA) to prove no host hits ggml-org.
@staticmethod
def _resolve_setup_sh_repo(
@@ -3277,8 +3652,8 @@ class TestRocmGfxForwarding:
# TEST: _pick_rocm_gfx_target -- visible-device selection from rocminfo output.
-# Honours CUDA/HIP_VISIBLE_DEVICES so a mixed-arch host installs the prebuilt
-# for the selected GPU, not GPU 0.
+# Honours CUDA/HIP_VISIBLE_DEVICES so a mixed-arch host installs the prebuilt for the
+# selected GPU, not GPU 0.
_pick_rocm_gfx_target = prebuilt_mod._pick_rocm_gfx_target
@@ -3454,7 +3829,11 @@ class TestWslRerouteNvidiaGuard:
source = _INSTALL_SH_PATH.read_text(encoding = "utf-8")
start = source.find("_maybe_reroute_strixhalo_to_2404()")
assert start != -1
- body = source[start : start + 1200]
+ # Slice the WHOLE function body (to its closing brace at column 0), not a
+ # fixed-length window: preamble growth must not push the signals out of view.
+ end = source.find("\n}", start)
+ assert end != -1
+ body = source[start:end]
nv = body.find("_has_usable_nvidia_gpu")
wmi = body.find("_wsl_amd_gpu_name")
assert nv != -1, "reroute must consult _has_usable_nvidia_gpu before deciding to reroute"
diff --git a/tests/studio/test_setup_pin_stale.ps1 b/tests/studio/test_setup_pin_stale.ps1
new file mode 100644
index 0000000000..2c92ae317f
--- /dev/null
+++ b/tests/studio/test_setup_pin_stale.ps1
@@ -0,0 +1,114 @@
+#!/usr/bin/env pwsh
+# 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 test for studio/setup.ps1's pinned-torch-index stale-venv helpers
+# (Test-RocmGfx211Leaf, Test-CudaFamilyLeaf, Get-RocmPinStaleTags). Pure helpers,
+# AST-extracted and run in-process. Mirrors the Python _rocm_pin_family_mismatch /
+# _is_cuda_family_leaf tests.
+# Run: pwsh -NoProfile -File tests/studio/test_setup_pin_stale.ps1
+
+$ErrorActionPreference = "Stop"
+$setupPath = [System.IO.Path]::Combine($PSScriptRoot, "..", "..", "studio", "setup.ps1")
+$setupPath = (Resolve-Path $setupPath).Path
+
+# --- Parse setup.ps1 (also serves as a syntax gate) and extract the helpers ---
+$tokens = $null; $errors = $null
+$ast = [System.Management.Automation.Language.Parser]::ParseFile($setupPath, [ref]$tokens, [ref]$errors)
+if ($errors) { $errors | ForEach-Object { $_.ToString() }; throw "setup.ps1 has parse errors" }
+
+foreach ($name in @("Test-RocmGfx211Leaf", "Test-RocmKnown211Version", "Test-CudaFamilyLeaf", "Get-RocmPinStaleTags")) {
+ $fn = $ast.FindAll({ param($n)
+ $n -is [System.Management.Automation.Language.FunctionDefinitionAst] -and $n.Name -eq $name
+ }, $true)
+ if ($fn.Count -ne 1) { throw "expected exactly one $name in setup.ps1, found $($fn.Count)" }
+ # Pure helpers (no exit / external calls) -- safe to define in this scope.
+ Invoke-Expression $fn[0].Extent.Text
+}
+
+$failures = 0
+function Check($name, $cond) {
+ if ($cond) { Write-Host " PASS $name" }
+ else { Write-Host " FAIL $name" -ForegroundColor Red; $script:failures++ }
+}
+
+# A pinned gfx/rocm index is stale when Expected != Installed.
+function IsStale($leaf, $ver) {
+ $t = Get-RocmPinStaleTags -PinLeaf $leaf -TorchVersion $ver
+ return $t.Expected -ne $t.Installed
+}
+
+Write-Host "Test-RocmGfx211Leaf (the 2.11 gfx allowlist)"
+Check "gfx1151 -> true" (Test-RocmGfx211Leaf "gfx1151")
+Check "gfx1150 -> true" (Test-RocmGfx211Leaf "gfx1150")
+Check "gfx120x-all -> true" (Test-RocmGfx211Leaf "gfx120x-all")
+Check "gfx110x-all -> false" (-not (Test-RocmGfx211Leaf "gfx110x-all"))
+Check "gfx90a -> false" (-not (Test-RocmGfx211Leaf "gfx90a"))
+Check "gfx908 -> false" (-not (Test-RocmGfx211Leaf "gfx908"))
+
+Write-Host "Test-CudaFamilyLeaf (^cu[0-9])"
+Check "cu118 -> true" (Test-CudaFamilyLeaf "cu118")
+Check "cu128 -> true" (Test-CudaFamilyLeaf "cu128")
+Check "cu130 -> true" (Test-CudaFamilyLeaf "cu130")
+Check "custom -> false" (-not (Test-CudaFamilyLeaf "custom"))
+Check "current -> false" (-not (Test-CudaFamilyLeaf "current"))
+Check "cpu -> false" (-not (Test-CudaFamilyLeaf "cpu"))
+Check "empty -> false" (-not (Test-CudaFamilyLeaf ""))
+
+Write-Host "Get-RocmPinStaleTags (mirror of _rocm_pin_family_mismatch)"
+# Exact rocm version comparison.
+Check "rocm7.2 pin + 2.11.0+rocm7.2 -> not stale" (-not (IsStale "rocm7.2" "2.11.0+rocm7.2"))
+Check "rocm7.2 pin + 2.10.0+rocm6.4 -> stale" (IsStale "rocm7.2" "2.10.0+rocm6.4")
+Check "rocm6.4 pin + 2.10.0+rocm6.4 -> not stale" (-not (IsStale "rocm6.4" "2.10.0+rocm6.4"))
+# rocm7.2 is a KNOWN-2.11 index. A +rocm7.2 wheel whose RELEASE drifted off 2.11 shares
+# the tag but violates the spec -> stale (mirror of _rocm_pin_family_mismatch).
+Check "rocm7.2 pin + 2.12.0+rocm7.2 -> stale" (IsStale "rocm7.2" "2.12.0+rocm7.2")
+Check "rocm7.2 pin + 2.13.0+rocm7.2 -> stale" (IsStale "rocm7.2" "2.13.0+rocm7.2")
+Check "rocm7.2 pin + 2.11.5+rocm7.2 -> not stale" (-not (IsStale "rocm7.2" "2.11.5+rocm7.2"))
+# An UNKNOWN newer rocm (off the 2.11 allowlist) isn't floored, so a matching version at
+# any release line is NOT stale on this exact-compare branch.
+Check "rocm8.0 pin + 2.12.0+rocm8.0 -> not stale" (-not (IsStale "rocm8.0" "2.12.0+rocm8.0"))
+# An untagged (no +rocm) wheel never satisfies a ROCm pin -> always stale.
+Check "rocm7.2 pin + 2.10.0 (untagged) -> stale" (IsStale "rocm7.2" "2.10.0")
+Check "rocm7.2 pin + 2.11.0 (untagged) -> stale" (IsStale "rocm7.2" "2.11.0")
+Check "rocm6.4 pin + 2.10.0 (untagged) -> stale" (IsStale "rocm6.4" "2.10.0")
+# 2.11-allowlist gfx pin: per-arch (three-part) wheel is satisfied, generic is stale.
+Check "gfx1151 pin + 2.11.0+rocm7.13.0 -> not stale" (-not (IsStale "gfx1151" "2.11.0+rocm7.13.0"))
+Check "gfx1150 pin + 2.11.0+rocm7.13.0 -> not stale" (-not (IsStale "gfx1150" "2.11.0+rocm7.13.0"))
+Check "gfx120x-all pin + 2.11.0+rocm7.13.0 -> not stale" (-not (IsStale "gfx120x-all" "2.11.0+rocm7.13.0"))
+Check "gfx1151 pin + 2.11.0+rocm7.2 (generic) -> stale" (IsStale "gfx1151" "2.11.0+rocm7.2")
+Check "gfx1151 pin + 2.10.0+rocm6.4 -> stale" (IsStale "gfx1151" "2.10.0+rocm6.4")
+# Non-2.11 gfx pin (gfx110X-all/gfx90a/gfx908): a valid <2.11 wheel is NOT stale.
+Check "gfx110x-all pin + 2.10.0+rocm6.4 -> not stale" (-not (IsStale "gfx110x-all" "2.10.0+rocm6.4"))
+Check "gfx90a pin + 2.10.0+rocm6.3 -> not stale" (-not (IsStale "gfx90a" "2.10.0+rocm6.3"))
+Check "gfx908 pin + 2.10.0+rocm7.0 -> not stale" (-not (IsStale "gfx908" "2.10.0+rocm7.0"))
+Check "gfx110x-all pin + 2.11.0+rocm7.2 -> stale" (IsStale "gfx110x-all" "2.11.0+rocm7.2")
+# Non-2.11 gfx pin over an untagged wheel: never satisfies the pin -> stale, so the
+# explicit ROCm index is applied even when torch is already <2.11.
+Check "gfx110x-all pin + 2.10.0 (untagged) -> stale" (IsStale "gfx110x-all" "2.10.0")
+Check "gfx90a pin + 2.10.0 (untagged) -> stale" (IsStale "gfx90a" "2.10.0")
+# Capital gfx120X-all is lowercased by Get-TorchIndexLeaf before this helper, so the
+# 2.11-allowlist branch fires and a generic/untagged wheel is stale.
+Check "gfx120x-all pin + 2.11.0+rocm7.2 (generic) -> stale" (IsStale "gfx120x-all" "2.11.0+rocm7.2")
+Check "gfx120x-all pin + 2.10.0 (untagged) -> stale" (IsStale "gfx120x-all" "2.10.0")
+
+# Major-only rocm pin (rocm7): majors compared alone; mirrors _rocm_pin_family_mismatch.
+Check "rocm7 pin + 2.10.0+rocm6.4 -> stale" (IsStale "rocm7" "2.10.0+rocm6.4")
+Check "rocm7 pin + 2.11.0+rocm7.2 -> not stale" (-not (IsStale "rocm7" "2.11.0+rocm7.2"))
+Check "rocm7 pin + 2.11.0+rocm7.13.0 -> not stale" (-not (IsStale "rocm7" "2.11.0+rocm7.13.0"))
+Check "rocm7 pin + 2.10.0 (untagged) -> stale" (IsStale "rocm7" "2.10.0")
+Check "rocm7 pin + 2.10.0+rocm (unreadable) -> not stale" (-not (IsStale "rocm7" "2.10.0+rocm"))
+
+Write-Host "Test-RocmKnown211Version + KNOWN-2.11 fallback (rocm7.2 only; no speculative rocm7.3)"
+Check "rocm7.2 -> known 2.11" (Test-RocmKnown211Version -Major 7 -Minor 2)
+Check "rocm7.1 -> not known" (-not (Test-RocmKnown211Version -Major 7 -Minor 1))
+Check "rocm7.3 -> not known" (-not (Test-RocmKnown211Version -Major 7 -Minor 3))
+Check "rocm8.0 -> not known" (-not (Test-RocmKnown211Version -Major 8 -Minor 0))
+# Unreadable-installed fallback: a rocm7.3 pin (unknown -> <2.11 line) over a <2.11 +rocm
+# wheel with an unreadable version is NOT stale; rocm7.2 (KNOWN-2.11) over the same wheel
+# IS stale (#2534 alignment).
+Check "rocm7.3 pin + 2.10.0+rocm (unreadable ver) -> not stale" (-not (IsStale "rocm7.3" "2.10.0+rocm"))
+Check "rocm7.2 pin + 2.10.0+rocm (unreadable ver) -> stale" (IsStale "rocm7.2" "2.10.0+rocm")
+
+Write-Host ""
+if ($failures -gt 0) { Write-Host "$failures check(s) FAILED" -ForegroundColor Red; exit 1 }
+Write-Host "All checks passed" -ForegroundColor Green
diff --git a/tests/studio/test_torch_flavor.ps1 b/tests/studio/test_torch_flavor.ps1
index 50be4814b0..f2cc55c21c 100644
--- a/tests/studio/test_torch_flavor.ps1
+++ b/tests/studio/test_torch_flavor.ps1
@@ -15,7 +15,7 @@ $tokens = $null; $errors = $null
$ast = [System.Management.Automation.Language.Parser]::ParseFile($installPath, [ref]$tokens, [ref]$errors)
if ($errors) { $errors | ForEach-Object { $_.ToString() }; throw "install.ps1 has parse errors" }
-foreach ($name in @("ConvertTo-TorchFlavorTag", "Get-ExpectedTorchFlavorTag")) {
+foreach ($name in @("ConvertTo-TorchFlavorTag", "Get-ExpectedTorchFlavorTag", "Trim-IndexPathSlashes", "Redact-InstallOutput")) {
$fn = $ast.FindAll({ param($n)
$n -is [System.Management.Automation.Language.FunctionDefinitionAst] -and $n.Name -eq $name
}, $true)
@@ -49,6 +49,19 @@ Check "mirror cu130 leaf -> cu130" ((Get-ExpectedTorchFlavorTag -TorchIndexUrl
Check "unrecognized leaf -> null" ($null -eq (Get-ExpectedTorchFlavorTag -TorchIndexUrl "https://my.mirror/whl/simple"))
Check "empty url -> null" ($null -eq (Get-ExpectedTorchFlavorTag -TorchIndexUrl ""))
+Write-Host "Trim-IndexPathSlashes (install.ps1 parity: path-only, token-preserving)"
+Check "double path slash collapsed" ((Trim-IndexPathSlashes "https://h/whl/cu128//") -eq "https://h/whl/cu128")
+Check "single trailing slash trimmed" ((Trim-IndexPathSlashes "https://h/whl/cu128/") -eq "https://h/whl/cu128")
+Check "query token slash preserved" ((Trim-IndexPathSlashes "https://h/whl/cu128?token=ab12cd/") -eq "https://h/whl/cu128?token=ab12cd/")
+Check "path slash trimmed, query kept" ((Trim-IndexPathSlashes "https://h/whl/cu128//?token=ab12cd/") -eq "https://h/whl/cu128?token=ab12cd/")
+
+Write-Host "Redact-InstallOutput (install.ps1 parity: credential redaction)"
+Check "userinfo redacted" ((Redact-InstallOutput "ERROR https://alice:s3cr3t@download.pytorch.org/whl/cu128") -eq "ERROR https://@download.pytorch.org/whl/cu128")
+Check "query value redacted" ((Redact-InstallOutput "https://host/whl/cu128?token=abcd1234&channel=beta") -eq "https://host/whl/cu128?token=&channel=")
+Check "fragment token redacted" ((Redact-InstallOutput "ERROR https://mirror.local/whl/cu128#token=SECRET123 (403)") -eq "ERROR https://mirror.local/whl/cu128# (403)")
+Check "bare hash comment untouched" ((Redact-InstallOutput "# retrying with --no-cache-dir") -eq "# retrying with --no-cache-dir")
+Check "plain line untouched" ((Redact-InstallOutput "Resolved 42 packages in 1.2s") -eq "Resolved 42 packages in 1.2s")
+
Write-Host ""
if ($failures -gt 0) { Write-Host "$failures check(s) FAILED" -ForegroundColor Red; exit 1 }
Write-Host "All checks passed" -ForegroundColor Green
diff --git a/tests/studio/test_torch_index_pin_hardening.ps1 b/tests/studio/test_torch_index_pin_hardening.ps1
new file mode 100644
index 0000000000..b8edf3da25
--- /dev/null
+++ b/tests/studio/test_torch_index_pin_hardening.ps1
@@ -0,0 +1,78 @@
+#!/usr/bin/env pwsh
+# 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 setup.ps1's torch-index pin-hardening helpers: Trim-IndexPathSlashes
+# (path-only slash trim, token-preserving), Redact-InstallOutput (credential redaction of
+# captured install logs), Get-TorchIndexLeaf (ALL trailing slashes stripped) and
+# Test-PipRocmFamilyLeaf (rocm7. is a custom pin, not a family). Pure helpers, AST-extracted
+# and run in-process. Run: pwsh -NoProfile -File tests/studio/test_torch_index_pin_hardening.ps1
+
+$ErrorActionPreference = "Stop"
+$setupPath = [System.IO.Path]::Combine($PSScriptRoot, "..", "..", "studio", "setup.ps1")
+$setupPath = (Resolve-Path $setupPath).Path
+$setupText = Get-Content -Raw $setupPath
+
+# --- Parse setup.ps1 (also a syntax gate) and extract the pure helpers ---
+$tokens = $null; $errors = $null
+$ast = [System.Management.Automation.Language.Parser]::ParseFile($setupPath, [ref]$tokens, [ref]$errors)
+if ($errors) { $errors | ForEach-Object { $_.ToString() }; throw "setup.ps1 has parse errors" }
+
+foreach ($name in @("Trim-IndexPathSlashes", "Redact-InstallOutput", "Get-TorchIndexLeaf", "Test-PipRocmFamilyLeaf")) {
+ $fn = $ast.FindAll({ param($n)
+ $n -is [System.Management.Automation.Language.FunctionDefinitionAst] -and $n.Name -eq $name
+ }, $true)
+ if ($fn.Count -ne 1) { throw "expected exactly one $name in setup.ps1, found $($fn.Count)" }
+ Invoke-Expression $fn[0].Extent.Text
+}
+
+$failures = 0
+function Check($name, $cond) {
+ if ($cond) { Write-Host " PASS $name" }
+ else { Write-Host " FAIL $name" -ForegroundColor Red; $script:failures++ }
+}
+
+Write-Host "Trim-IndexPathSlashes (path-only, token-preserving)"
+Check "double path slash collapsed" ((Trim-IndexPathSlashes "https://h/whl/cu128//") -eq "https://h/whl/cu128")
+Check "single trailing slash trimmed" ((Trim-IndexPathSlashes "https://h/whl/cu128/") -eq "https://h/whl/cu128")
+Check "no slash unchanged" ((Trim-IndexPathSlashes "https://h/whl/cu128") -eq "https://h/whl/cu128")
+Check "query token slash preserved" ((Trim-IndexPathSlashes "https://h/whl/cu128?token=ab12cd/") -eq "https://h/whl/cu128?token=ab12cd/")
+Check "path slash trimmed, query kept" ((Trim-IndexPathSlashes "https://h/whl/cu128//?token=ab12cd/") -eq "https://h/whl/cu128?token=ab12cd/")
+Check "fragment slash preserved" ((Trim-IndexPathSlashes "https://h/whl/cu128#anchor/") -eq "https://h/whl/cu128#anchor/")
+
+Write-Host "Redact-InstallOutput (credential redaction)"
+Check "userinfo redacted" ((Redact-InstallOutput "ERROR https://alice:s3cr3t@download.pytorch.org/whl/cu128") -eq "ERROR https://@download.pytorch.org/whl/cu128")
+Check "bare-token@ redacted" ((Redact-InstallOutput "fetch https://ghp_deadbeef@host/whl/cu128 failed") -eq "fetch https://@host/whl/cu128 failed")
+Check "single query value redacted" ((Redact-InstallOutput "url https://host/whl/cu128?token=abcd1234 unreachable") -eq "url https://host/whl/cu128?token= unreachable")
+Check "multiple query values redacted" ((Redact-InstallOutput "https://host/whl/cu128?token=abcd1234&channel=beta") -eq "https://host/whl/cu128?token=&channel=")
+Check "fragment token redacted" ((Redact-InstallOutput "ERROR https://mirror.local/whl/cu128#token=SECRET123 (403)") -eq "ERROR https://mirror.local/whl/cu128# (403)")
+Check "query and fragment both redacted" ((Redact-InstallOutput "https://host/whl/cu128?token=abc#sig=xyz done") -eq "https://host/whl/cu128?token=# done")
+Check "bare hash comment untouched" ((Redact-InstallOutput "# retrying with --no-cache-dir") -eq "# retrying with --no-cache-dir")
+Check "plain line untouched" ((Redact-InstallOutput "Resolved 42 packages in 1.2s") -eq "Resolved 42 packages in 1.2s")
+$leak = Redact-InstallOutput "https://alice:s3cr3t@host/whl/cu128?token=SUPERSECRET#frag=ALSOSECRET"
+Check "no secret substring survives" (($leak -notmatch "s3cr3t") -and ($leak -notmatch "SUPERSECRET") -and ($leak -notmatch "ALSOSECRET"))
+
+Write-Host "Get-TorchIndexLeaf (ALL trailing slashes stripped)"
+Check "double slash cu128 -> cu128" ((Get-TorchIndexLeaf "https://m/whl/cu128//") -eq "cu128")
+Check "triple slash rocm7.2 -> rocm7.2" ((Get-TorchIndexLeaf "https://m/whl/rocm7.2///") -eq "rocm7.2")
+Check "double slash + token -> cu128" ((Get-TorchIndexLeaf "https://m/whl/cu128//?token=x") -eq "cu128")
+Check "single slash cu128 -> cu128" ((Get-TorchIndexLeaf "https://m/whl/cu128/") -eq "cu128")
+
+Write-Host "Test-PipRocmFamilyLeaf (rocm7. is a custom pin, not a family)"
+Check "rocm7 family" (Test-PipRocmFamilyLeaf "rocm7")
+Check "rocm7.2 family" (Test-PipRocmFamilyLeaf "rocm7.2")
+Check "gfx1151 family" (Test-PipRocmFamilyLeaf "gfx1151")
+Check "rocm7. trailing-dot NOT family" (-not (Test-PipRocmFamilyLeaf "rocm7."))
+Check "rocm.7 leading-dot NOT family" (-not (Test-PipRocmFamilyLeaf "rocm.7"))
+Check "rocm7.2.1 two-dot NOT family" (-not (Test-PipRocmFamilyLeaf "rocm7.2.1"))
+Check "rocm7.2-private NOT family" (-not (Test-PipRocmFamilyLeaf "rocm7.2-private"))
+Check "cu128 NOT family" (-not (Test-PipRocmFamilyLeaf "cu128"))
+
+Write-Host "Fast-Install pinned-install env scrub (source assertion)"
+# The pip fallback honours PIP_*; PIP_NO_INDEX=1 would make it ignore the pinned --index-url
+# and PIP_INDEX_URL would replace it, so both must be scrubbed for a pinned install.
+Check "PIP_NO_INDEX scrubbed" ($setupText -match "'PIP_NO_INDEX'")
+Check "PIP_INDEX_URL scrubbed" ($setupText -match "'PIP_INDEX_URL'")
+
+Write-Host ""
+if ($failures -gt 0) { Write-Host "$failures check(s) FAILED" -ForegroundColor Red; exit 1 }
+Write-Host "All checks passed" -ForegroundColor Green
From 65587c2be7b6167e369bc5d95155d315c2717ac8 Mon Sep 17 00:00:00 2001
From: Michael Han <107991372+shimmyshimmer@users.noreply.github.com>
Date: Mon, 20 Jul 2026 04:57:44 -0700
Subject: [PATCH 15/41] Studio: Data settings tab, uploaded files manager,
quant pinning, and chat image preview fix (#7029)
* Studio: Data settings tab, uploaded files manager, quant pinning, image preview fix
Settings
- New Data tab in the settings sidebar, under Connections. Chat data
management (archived chats, confirm before deleting, exports, import,
clear all) moved there from the Chat tab.
- New Archive all chats action with confirmation. Archives every chat in
Recents and Projects; compare pairs count as one chat.
- New Uploaded files manager listing RAG documents (chats, projects,
knowledge bases) and chat message attachments with location, size and
date. Files can be opened in a new tab or deleted. Deleting a chat
attachment keeps the message text.
Backend
- GET /api/rag/documents lists all uploaded RAG documents with file size
plus KB and project names.
- GET /api/chat/attachments lists chat message attachments; per
attachment file and delete endpoints included.
Model selector
- Downloaded GGUF quants can be pinned from the quant row (next to the
settings and delete actions). Pinned quants show at the top of On
Device under a Pinned heading as model name plus a grey quant chip and
load directly with one click. Non GGUF cached repos pin as a whole.
- Toned down the green of the downloaded label.
Fix
- Clicking an image attachment in chat now opens the preview overlay.
The tooltip trigger wrapper called preventDefault before composed
handlers ran, which made Radix DialogTrigger skip opening.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Studio: image previews and file type chips in uploaded files list
Image attachments now show a small thumbnail (lazy loaded from the
stored bytes, object URL revoked on unmount) and every row shows a grey
uppercase type chip derived from the extension or content type. Non
image rows keep a file icon. Name cell floors its width and clips
overflow so narrow dialogs stay aligned.
* Harden attachment serving, add tests, and polish pinned rows and previews
- Strict base64 decoding for attachment files: corrupt payloads now return
422 instead of silently serving empty or garbled bytes; whitespace,
missing padding, the URL-safe alphabet, and RFC 2397 percent-encoded
data URLs are all handled
- New backend test suite covering attachment listing, size accounting,
malformed rows, deletion semantics, and every file-serving edge case
- Pinned quant rows show a Loaded tag when that exact quant is active,
and reveal unpin, settings, and delete actions on hover
- Uploaded files dialog is wider and chat locations link straight to the
thread the attachment belongs to
- Chat image preview is now a chrome-free lightbox: dimmed backdrop,
rounded image, corner close button, click outside to dismiss
- File opens go through a synchronous window.open so Safari and Firefox
popup blockers do not eat them
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Uploaded files: click a file to jump to its chat, square thumbs, new Data icon
- Clicking a file row (thumbnail or name) now goes straight to the chat it
belongs to; files without a chat open directly as before
- File thumbnails pin a small 7px radius: the theme scales rounded-md up
to a near circle at this size
- Settings Data tab now uses the database-setting icon
* Uploaded files is now a Data tab subpage instead of a popup
- Manage swaps the tab body for an inline Uploaded files page with a back
header, matching the rest of settings navigation
- Size column header and values are left aligned like the other columns
- Column widths tightened so the table fits the settings panel
* Lightbox polish and Data tab row order
- Image preview close button is transparent until hovered
- Preview image no longer rounds its corners
- Import chats now sits below Clear all chats in the Data tab
* Data tab: export chats as fine-tuning data and open them in Recipes
- New Fine-tuning section in Settings > Data converts every chat into a
JSONL dataset in the OpenAI messages format, one conversation per line
with string-only system/user/assistant turns
- The Train tab detects this file as chatml natively: no column mapping
and no standardization pass, and it works with train on completions
since every assistant turn sits behind the chat template response marker
- Consecutive same-role turns merge, trailing turns without an assistant
reply drop, and reasoning, tool calls, and images are excluded so chat
templates format the data cleanly
- Open in Recipes stages the JSONL as a local seed upload, creates a new
Data Recipe with the seed block preconfigured, and jumps to the editor
* Data tab: load chats straight into the Train tab, row moved to the top
- New Load in Train tab button uploads the fine-tuning JSONL through the
training dataset endpoint, selects it in the training config store, and
opens the Train tab with the dataset loaded and format-checked
- Use chats as training data now sits at the very top of the Data tab
- The Chats subheading is gone; chat rows flow directly under it
* Address review findings on the uploads manager and quant pins
- Deleting the last attachment stores '[]' instead of NULL: a NULL reads
back as a missing field and triggers the legacy IndexedDB backfill,
which resurrected the deleted attachment on the next chat load
- The attachment file endpoint now serves audio: adapter parts store
{data, format} raw base64 and compare chats store a bare base64 string;
media type comes from the attachment contentType or the format
- Compare-chat uploads live in message content parts, not attachments;
the uploads list now includes those blobs via synthetic content-part
ids that the same get and delete routes resolve
- Deleting a quant from the expanded repo row also unpins it so a pinned
row cannot try to load a file that no longer exists
- Thumbnails in the uploads list fetch their blob only once the row is
visible, so a long screenshot history does not download everything
- Nine new backend tests cover audio serving, content-part listing,
serving, deletion, and the empty-list delete behavior
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Data tab: single action dropdown with format choices for chat training data
- The three fine-tune buttons collapse into one dropdown plus a run
button; pick Load in Train tab, Open in Recipes, or Export JSONL,
then click the arrow to run it
- The dropdown's Format section adds ShareGPT and Alpaca alongside the
default OpenAI messages format, ticked like a checklist; all three
shapes are auto-detected by the Train tab's format check
- Alpaca is single-turn, so each user to assistant pair becomes its own
record with the system prompt and earlier turns carried in the input
column
- Shorter description on the training data row
- Uploaded files rows show the size under the file name instead of a
separate column, matching the tighter layout
* Polish the training data action control
- Run button is a true circle (icon-sm plus rounded-full) with a
heavier arrow stroke
- Dropdown trigger uses the shared standard chevron and a fixed width
so switching actions no longer resizes the control
* Shorten the training data row description
* Use the standard chevron for the run button and enlarge the ticks
- Run button uses the shared standard right chevron so it matches the
dropdown chevron instead of the hugeicons arrow
- Dropdown ticks bumped up a size for legibility
* Reword the training data row description
* Shorten Data Recipes to Recipes in the training data description
* List Export JSONL first and rename the default format to Chat Completions
* Handle legacy string content in fine-tune exports and gate Train on chat-only hosts
- messageToPlainText now accepts plain-string message content, the shape
legacy and imported histories store, so those conversations export
instead of being skipped as having no exchange
- The Load in Train tab action is disabled on chat-only hosts the same
way the sidebar gates Train; the default action falls back to Export
JSONL there so the run button never uploads a dataset that /studio
would immediately redirect away from
* Narrow the training data action dropdown slightly
* Drop the format picker from the training data dropdown
Chat Completions (OpenAI messages) is the only export format we ship, so
the ShareGPT and Alpaca options and the Format section are removed. The
export always uses the OpenAI messages shape.
* Address the second round of review findings
Security
- Chat attachment data URLs no longer echo their embedded media type:
anything that is not a plain raster image serves as octet-stream, so
imported text/html or SVG payloads cannot render under the app origin
- Uploaded .html/.htm RAG documents serve as text/plain for the same
reason; the preview sheet only uses the file URL for PDFs
Uploads manager
- Remote image URLs in imported chats are no longer listed as stored
uploads (nothing to serve, and delete would strip the chat reference);
the delete guard mirrors the same data:-only rule
- Deleting a content-part upload refetches the list since the remaining
parts re-index, keeping sibling row ids current
- Deleting a project document from the Data tab invalidates the project
sources cache like the sources panel does
- Data-tab deletions now patch the loaded thread's in-memory copy via a
small event, so a later repo sync cannot write the attachment back
Fine-tune export
- Branch siblings from retries stay out of the exported conversation;
only the selected chain converts (full exports still keep everything)
- Assistant turns before the first user turn drop, preserving leading
system prompts, so no unconditioned assistant targets are emitted
Four new backend tests cover the media type clamp and remote-URL rows;
two existing tests updated for the clamped types
* Fix uploaded file lifecycle and model state
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Make archived chats a Data settings subpage
* Studio: fix attachment route tests and pinned quant edge cases
- test_chat_attachments: drop asyncio.run around the synchronous
/attachments routes (list/get/delete are plain def, so asyncio.run
raised 'a coroutine was expected' and failed the Repo tests CI job).
- test_chat_attachments: align compare-chat content-part assertions with
the stable content-hash id scheme (content-part-sha256-...) instead of
the removed array-index ids; resolve ids from the listing.
- pickers: pass disabled={deleteDisabled} to the pinned-quant delete
action so a quant cannot be deleted mid model-load, matching the
expanded variant rows.
- pickers: build the pinned-quant existence set from the query-unfiltered
cached GGUF repos (format filter still applied) so a pinned quant stays
findable when the search term matches only its quant name.
* Fix Studio review regressions
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Guard fine-tune export content blocks
* Add Export button for archived chats
Adds an Export action to the Archived chats view in Settings > Data that
downloads only the archived chats as a JSON backup (their threads, messages
and projects). The button sits in the archived header row and appears only
when archived chats exist.
* Refactor archived export into pure, testable units
Split the archived-chats export into a dependency-free filter
(archived-chat-export.ts) and a shared JSON download helper
(download-json.ts). Skip the download when nothing is archived so a
stray call never drops an empty file. No behavior change to the button.
---------
Co-authored-by: shimmyshimmer
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Etherll <61019402+Etherll@users.noreply.github.com>
Co-authored-by: Unsloth
Co-authored-by: Lee Jackson <130007945+Imagineer99@users.noreply.github.com>
---
studio/backend/core/rag/store.py | 10 +
studio/backend/routes/chat_history.py | 134 ++-
studio/backend/routes/rag.py | 39 +-
studio/backend/storage/studio_db.py | 813 +++++++++++++++++-
studio/backend/tests/test_chat_attachments.py | 634 ++++++++++++++
.../frontend/src/components/app-sidebar.tsx | 4 +-
.../components/assistant-ui/attachment.tsx | 27 +-
.../assistant-ui/model-selector/pickers.tsx | 519 ++++++++++-
.../model-selector/pinned-models.ts | 72 ++
studio/frontend/src/components/ui/tooltip.tsx | 7 +-
.../src/features/chat/api/chat-api.ts | 74 +-
.../chat/hooks/use-chat-sidebar-items.ts | 34 +
studio/frontend/src/features/chat/index.ts | 19 +-
.../prompt-storage/prompt-storage-dialog.tsx | 223 ++++-
.../src/features/chat/runtime-provider.tsx | 167 ++++
.../chat/utils/archived-chat-export.ts | 62 ++
.../chat/utils/chat-attachment-events.ts | 123 +++
.../src/features/chat/utils/download-json.ts | 18 +
.../chat/utils/export-chat-history.ts | 35 +-
.../frontend/src/features/rag/api/rag-api.ts | 27 +-
.../rag/components/use-rag-documents.ts | 136 +--
studio/frontend/src/features/rag/index.ts | 7 +-
studio/frontend/src/features/rag/types/rag.ts | 7 +
.../components/archived-chats-dialog.tsx | 135 ++-
.../settings/components/finetune-recipe.ts | 98 +++
.../components/uploaded-files-dialog.tsx | 644 ++++++++++++++
.../src/features/settings/settings-dialog.tsx | 11 +
.../src/features/settings/settings-search.ts | 13 +-
.../settings/stores/settings-dialog-store.ts | 6 +-
.../src/features/settings/tabs/chat-tab.tsx | 329 +------
.../src/features/settings/tabs/data-tab.tsx | 727 ++++++++++++++++
studio/frontend/src/i18n/locales/en.ts | 50 +-
32 files changed, 4671 insertions(+), 533 deletions(-)
create mode 100644 studio/backend/tests/test_chat_attachments.py
create mode 100644 studio/frontend/src/components/assistant-ui/model-selector/pinned-models.ts
create mode 100644 studio/frontend/src/features/chat/utils/archived-chat-export.ts
create mode 100644 studio/frontend/src/features/chat/utils/chat-attachment-events.ts
create mode 100644 studio/frontend/src/features/chat/utils/download-json.ts
create mode 100644 studio/frontend/src/features/settings/components/finetune-recipe.ts
create mode 100644 studio/frontend/src/features/settings/components/uploaded-files-dialog.tsx
create mode 100644 studio/frontend/src/features/settings/tabs/data-tab.tsx
diff --git a/studio/backend/core/rag/store.py b/studio/backend/core/rag/store.py
index f9128d1715..1165b6bb0e 100644
--- a/studio/backend/core/rag/store.py
+++ b/studio/backend/core/rag/store.py
@@ -158,6 +158,16 @@ def list_documents(conn: sqlite3.Connection, scope: str) -> list[dict]:
return [dict(r) for r in rows]
+def list_all_documents(conn: sqlite3.Connection) -> list[dict]:
+ """Every uploaded document across all scopes (KBs, threads, projects)."""
+ rows = conn.execute(
+ "SELECT id, scope, kb_id, thread_id, project_id, filename, sha256, status, error, "
+ "num_chunks, stored_path, created_at "
+ "FROM documents ORDER BY created_at DESC"
+ ).fetchall()
+ return [dict(r) for r in rows]
+
+
def get_document(conn: sqlite3.Connection, document_id: str) -> dict | None:
row = conn.execute("SELECT * FROM documents WHERE id=?", (document_id,)).fetchone()
return dict(row) if row else None
diff --git a/studio/backend/routes/chat_history.py b/studio/backend/routes/chat_history.py
index 7a27a58a52..24b6dfb36d 100644
--- a/studio/backend/routes/chat_history.py
+++ b/studio/backend/routes/chat_history.py
@@ -5,7 +5,7 @@
Chat history API routes backed by studio.db.
"""
-from typing import Any, Literal, Optional
+from typing import Annotated, Any, Literal, Optional
from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel, ConfigDict, Field, ValidationError
@@ -19,13 +19,16 @@ from storage.studio_db import (
clear_chat_history,
count_chat_threads,
count_forks_for_message,
+ delete_chat_attachment,
delete_chat_threads,
delete_chat_project,
ensure_chat_project_workspace,
fork_chat_thread,
+ get_chat_attachment,
get_chat_project,
get_chat_thread,
get_chat_message,
+ list_chat_attachments_page,
list_chat_projects,
list_chat_legacy_imports,
list_chat_settings,
@@ -279,6 +282,131 @@ async def delete_threads(
return {"status": "deleted"}
+@router.get("/attachments")
+def list_attachments(
+ limit: Annotated[int, Query(ge = 1, le = 100)] = 50,
+ offset: Annotated[int, Query(ge = 0)] = 0,
+ current_subject: str = Depends(get_current_subject),
+) -> dict:
+ """One bounded page of chat uploads for the settings Data tab."""
+ attachments, next_offset = list_chat_attachments_page(limit = limit, offset = offset)
+ return {"attachments": attachments, "nextOffset": next_offset}
+
+
+def _decode_attachment_base64(payload: str) -> bytes:
+ """Strict base64 decode of a stored payload.
+
+ Normalizes first: strips whitespace, fixes padding, accepts the URL-safe
+ alphabet. validate=False would silently drop bad characters and serve
+ corrupted bytes instead of failing, so raise 422 on anything else.
+ """
+ import base64
+
+ normalized = "".join(payload.split())
+ altchars = b"-_" if ("-" in normalized or "_" in normalized) else None
+ normalized += "=" * (-len(normalized) % 4)
+ try:
+ return base64.b64decode(normalized, altchars = altchars, validate = True)
+ except Exception as exc: # noqa: BLE001 - corrupt stored payload
+ raise HTTPException(status_code = 422, detail = "Attachment data is corrupt") from exc
+
+
+_AUDIO_FORMAT_MEDIA_TYPES = {
+ "mp3": "audio/mpeg",
+ "wav": "audio/wav",
+ "ogg": "audio/ogg",
+ "flac": "audio/flac",
+}
+
+
+def _safe_image_media_type(media_type: str) -> str:
+ """Clamp a data-URL media type to something inert to render.
+
+ Imported chats store image parts verbatim, so the embedded type can be
+ text/html or image/svg+xml; echoing those would execute markup with the
+ app origin when opened. Anything not a plain raster type downloads as
+ bytes instead.
+ """
+ lowered = media_type.strip().lower()
+ if lowered.startswith("image/") and lowered != "image/svg+xml":
+ return lowered
+ return "application/octet-stream"
+
+
+@router.get("/attachments/{message_id}/{attachment_id}/file")
+def get_attachment_file(
+ message_id: str,
+ attachment_id: str,
+ current_subject: str = Depends(get_current_subject),
+):
+ """Serve one attachment's stored content: image or audio bytes, or
+ extracted text."""
+ import urllib.parse
+
+ from fastapi.responses import Response
+
+ attachment = get_chat_attachment(message_id, attachment_id)
+ if attachment is None:
+ raise HTTPException(status_code = 404, detail = "Attachment not found")
+
+ attachment_content_type = attachment.get("contentType")
+ texts: list[str] = []
+ for part in attachment.get("content") or []:
+ if not isinstance(part, dict):
+ continue
+ image = part.get("image")
+ if isinstance(image, str) and image[:5].lower() == "data:":
+ header, _, payload = image.partition(",")
+ media_type = _safe_image_media_type(
+ header[5:].split(";", 1)[0] or "application/octet-stream"
+ )
+ if "base64" not in header.lower():
+ # RFC 2397 non-base64 form stores percent-encoded bytes.
+ data = urllib.parse.unquote_to_bytes(payload)
+ return Response(content = data, media_type = media_type)
+ data = _decode_attachment_base64(payload)
+ return Response(content = data, media_type = media_type)
+ # Audio parts: the attachment adapter stores {data, format} with raw
+ # base64; compare chats store a bare base64 string.
+ audio = part.get("audio")
+ if isinstance(audio, dict) or (isinstance(audio, str) and audio):
+ if isinstance(audio, dict):
+ payload = audio.get("data")
+ audio_format = audio.get("format")
+ else:
+ payload = audio.rsplit(",", 1)[-1]
+ audio_format = None
+ if isinstance(payload, str) and payload:
+ data = _decode_attachment_base64(payload)
+ media_type = (
+ attachment_content_type
+ if isinstance(attachment_content_type, str)
+ and attachment_content_type.startswith("audio/")
+ else _AUDIO_FORMAT_MEDIA_TYPES.get(
+ str(audio_format or "").lower(), "application/octet-stream"
+ )
+ )
+ return Response(content = data, media_type = media_type)
+ text = part.get("text")
+ if isinstance(text, str) and text:
+ texts.append(text)
+ if texts:
+ return Response(content = "\n".join(texts), media_type = "text/plain; charset=utf-8")
+ raise HTTPException(status_code = 404, detail = "Attachment has no stored content")
+
+
+@router.delete("/attachments/{message_id}/{attachment_id}")
+def delete_attachment(
+ message_id: str,
+ attachment_id: str,
+ current_subject: str = Depends(get_current_subject),
+) -> dict:
+ """Remove one attachment from its chat message."""
+ if not delete_chat_attachment(message_id, attachment_id):
+ raise HTTPException(status_code = 404, detail = "Attachment not found")
+ return {"ok": True}
+
+
@router.get("/projects", response_model = ChatProjectListResponse)
async def list_projects(
include_archived: bool = Query(False), current_subject: str = Depends(get_current_subject)
@@ -409,7 +537,7 @@ async def get_thread_message(
@router.put("/threads/{thread_id}/messages/{message_id}", response_model = ChatMessage)
-async def save_thread_message(
+def save_thread_message(
thread_id: str,
message_id: str,
payload: ChatMessage,
@@ -432,7 +560,7 @@ async def save_thread_message(
@router.put("/threads/{thread_id}/messages", response_model = ChatMessageListResponse)
-async def replace_thread_messages(
+def replace_thread_messages(
thread_id: str,
payload: ChatMessageSyncRequest,
current_subject: str = Depends(get_current_subject),
diff --git a/studio/backend/routes/rag.py b/studio/backend/routes/rag.py
index e20fea74a3..392a4e0d02 100644
--- a/studio/backend/routes/rag.py
+++ b/studio/backend/routes/rag.py
@@ -318,6 +318,39 @@ def list_project_documents(project_id: str, subject: str = Depends(get_current_s
conn.close()
+@router.get("/documents")
+def list_all_uploaded_documents(subject: str = Depends(get_current_subject)) -> dict:
+ """Every uploaded file across chats, projects, and knowledge bases (settings
+ Data tab)."""
+ _require_rag()
+ conn = rag_db.get_connection()
+ try:
+ docs = store.list_all_documents(conn)
+ kb_names = {kb["id"]: kb["name"] for kb in store.list_kbs(conn)}
+ finally:
+ conn.close()
+
+ from storage.studio_db import list_chat_projects
+
+ project_names = {p["id"]: p["name"] for p in list_chat_projects(include_archived = True)}
+
+ out = []
+ for doc in docs:
+ view = _doc_view(doc)
+ stored_path = doc.get("stored_path")
+ size = None
+ if stored_path:
+ try:
+ size = os.path.getsize(stored_path)
+ except OSError:
+ size = None
+ view["sizeBytes"] = size
+ view["kbName"] = kb_names.get(doc.get("kb_id"))
+ view["projectName"] = project_names.get(doc.get("project_id"))
+ out.append(view)
+ return {"documents": out}
+
+
@router.delete("/documents/{document_id}")
def delete_document(document_id: str, subject: str = Depends(get_current_subject)) -> dict:
_require_rag()
@@ -424,8 +457,10 @@ _CONTENT_TYPES = {
".txt": "text/plain; charset=utf-8",
".md": "text/markdown; charset=utf-8",
".markdown": "text/markdown; charset=utf-8",
- ".html": "text/html; charset=utf-8",
- ".htm": "text/html; charset=utf-8",
+ # Served as plain text, never text/html: an uploaded HTML document rendered
+ # same-origin would execute its scripts with access to the app's storage.
+ ".html": "text/plain; charset=utf-8",
+ ".htm": "text/plain; charset=utf-8",
".docx": "application/vnd.openxmlformats-officedocument.wordprocessingml.document",
}
diff --git a/studio/backend/storage/studio_db.py b/studio/backend/storage/studio_db.py
index 4e0c711b69..d889894d04 100644
--- a/studio/backend/storage/studio_db.py
+++ b/studio/backend/storage/studio_db.py
@@ -7,6 +7,7 @@ Like auth/storage.py (module-level functions, raw sqlite3, per-function
connections) plus WAL mode and PRAGMA foreign_keys = ON for CASCADE deletes.
"""
+import hashlib
import json
import logging
import os
@@ -100,6 +101,7 @@ _schema_lock = threading.Lock()
_schema_ready = False
_SQLITE_IN_CHUNK_SIZE = 900
_PROJECT_WORKSPACE_SUBDIRS = ("sandbox",)
+_CHAT_ATTACHMENT_INVENTORY_VERSION = 1
def _project_slug(name: str) -> str:
@@ -313,6 +315,141 @@ def _ensure_schema(conn: sqlite3.Connection) -> None:
)
"""
)
+ tombstone_schema = """
+ CREATE TABLE chat_attachment_tombstones (
+ thread_id TEXT NOT NULL REFERENCES chat_threads(id) ON DELETE CASCADE,
+ message_id TEXT NOT NULL,
+ attachment_id TEXT NOT NULL,
+ deleted_at INTEGER NOT NULL,
+ PRIMARY KEY(thread_id, message_id, attachment_id)
+ ) WITHOUT ROWID
+ """
+ tombstone_table = conn.execute(
+ """
+ SELECT 1 FROM sqlite_master
+ WHERE type = 'table' AND name = 'chat_attachment_tombstones'
+ """
+ ).fetchone()
+ if tombstone_table is None:
+ conn.execute(tombstone_schema)
+ else:
+ tombstone_columns = {
+ row[1] for row in conn.execute("PRAGMA table_info(chat_attachment_tombstones)")
+ }
+ tombstone_fk_targets = {
+ row[2] for row in conn.execute("PRAGMA foreign_key_list(chat_attachment_tombstones)")
+ }
+ if "thread_id" not in tombstone_columns or "chat_threads" not in tombstone_fk_targets:
+ # The first implementation cascaded through chat_messages, which
+ # erased deletion knowledge during pruneMissing. Rebuild once,
+ # retaining every tombstone whose owning thread still exists.
+ conn.execute("SAVEPOINT migrate_chat_attachment_tombstones")
+ try:
+ conn.execute(
+ "ALTER TABLE chat_attachment_tombstones "
+ "RENAME TO chat_attachment_tombstones_legacy"
+ )
+ conn.execute(tombstone_schema)
+ if "thread_id" in tombstone_columns:
+ conn.execute(
+ """
+ INSERT OR IGNORE INTO chat_attachment_tombstones
+ (thread_id, message_id, attachment_id, deleted_at)
+ SELECT legacy.thread_id, legacy.message_id,
+ legacy.attachment_id, legacy.deleted_at
+ FROM chat_attachment_tombstones_legacy legacy
+ JOIN chat_threads thread ON thread.id = legacy.thread_id
+ """
+ )
+ else:
+ conn.execute(
+ """
+ INSERT OR IGNORE INTO chat_attachment_tombstones
+ (thread_id, message_id, attachment_id, deleted_at)
+ SELECT message.thread_id, legacy.message_id,
+ legacy.attachment_id, legacy.deleted_at
+ FROM chat_attachment_tombstones_legacy legacy
+ JOIN chat_messages message ON message.id = legacy.message_id
+ """
+ )
+ conn.execute("DROP TABLE chat_attachment_tombstones_legacy")
+ conn.execute("RELEASE SAVEPOINT migrate_chat_attachment_tombstones")
+ except Exception:
+ conn.execute("ROLLBACK TO SAVEPOINT migrate_chat_attachment_tombstones")
+ conn.execute("RELEASE SAVEPOINT migrate_chat_attachment_tombstones")
+ raise
+ conn.execute(
+ """
+ CREATE TABLE IF NOT EXISTS chat_attachment_inventory (
+ message_id TEXT NOT NULL REFERENCES chat_messages(id) ON DELETE CASCADE,
+ attachment_id TEXT NOT NULL,
+ name TEXT NOT NULL,
+ type TEXT,
+ content_type TEXT,
+ size_bytes INTEGER,
+ PRIMARY KEY(message_id, attachment_id)
+ ) WITHOUT ROWID
+ """
+ )
+ conn.execute(
+ """
+ CREATE TABLE IF NOT EXISTS chat_attachment_inventory_state (
+ singleton INTEGER NOT NULL PRIMARY KEY CHECK(singleton = 1),
+ inventory_version INTEGER NOT NULL DEFAULT 0,
+ dirty INTEGER NOT NULL DEFAULT 1,
+ backfilled_at INTEGER NOT NULL
+ )
+ """
+ )
+ inventory_state_columns = {
+ row[1] for row in conn.execute("PRAGMA table_info(chat_attachment_inventory_state)")
+ }
+ if "inventory_version" not in inventory_state_columns:
+ conn.execute(
+ "ALTER TABLE chat_attachment_inventory_state "
+ "ADD COLUMN inventory_version INTEGER NOT NULL DEFAULT 0"
+ )
+ if "dirty" not in inventory_state_columns:
+ conn.execute(
+ "ALTER TABLE chat_attachment_inventory_state "
+ "ADD COLUMN dirty INTEGER NOT NULL DEFAULT 1"
+ )
+ conn.execute(
+ """
+ CREATE TRIGGER IF NOT EXISTS chat_attachment_inventory_dirty_insert
+ AFTER INSERT ON chat_messages
+ BEGIN
+ INSERT INTO chat_attachment_inventory_state
+ (singleton, inventory_version, dirty, backfilled_at)
+ VALUES (1, 0, 1, 0)
+ ON CONFLICT(singleton) DO UPDATE SET dirty = 1;
+ END
+ """
+ )
+ conn.execute(
+ """
+ CREATE TRIGGER IF NOT EXISTS chat_attachment_inventory_dirty_update
+ AFTER UPDATE ON chat_messages
+ BEGIN
+ INSERT INTO chat_attachment_inventory_state
+ (singleton, inventory_version, dirty, backfilled_at)
+ VALUES (1, 0, 1, 0)
+ ON CONFLICT(singleton) DO UPDATE SET dirty = 1;
+ END
+ """
+ )
+ conn.execute(
+ """
+ CREATE TRIGGER IF NOT EXISTS chat_attachment_inventory_dirty_delete
+ AFTER DELETE ON chat_messages
+ BEGIN
+ INSERT INTO chat_attachment_inventory_state
+ (singleton, inventory_version, dirty, backfilled_at)
+ VALUES (1, 0, 1, 0)
+ ON CONFLICT(singleton) DO UPDATE SET dirty = 1;
+ END
+ """
+ )
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_chat_threads_model_type_created_at ON chat_threads(model_type, created_at)"
)
@@ -391,6 +528,21 @@ def _ensure_schema(conn: sqlite3.Connection) -> None:
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_prompt_lists_created_at ON prompt_lists(created_at)"
)
+ inventory_state = conn.execute(
+ """
+ SELECT inventory_version, dirty
+ FROM chat_attachment_inventory_state
+ WHERE singleton = 1
+ """
+ ).fetchone()
+ if (
+ inventory_state is None
+ or inventory_state["inventory_version"] != _CHAT_ATTACHMENT_INVENTORY_VERSION
+ or inventory_state["dirty"]
+ ):
+ _rebuild_chat_attachment_inventory(conn)
+ _mark_chat_attachment_inventory_clean(conn)
+ conn.commit()
def _prompt_entry_from_row(row: sqlite3.Row) -> dict:
@@ -1219,7 +1371,14 @@ def delete_chat_threads(ids: list[str]) -> None:
return
conn = get_connection()
try:
+ conn.execute("BEGIN IMMEDIATE")
+ _ensure_chat_attachment_inventory_current(conn)
+ conn.executemany(
+ "DELETE FROM chat_attachment_tombstones WHERE thread_id = ?",
+ [(id,) for id in ids],
+ )
conn.executemany("DELETE FROM chat_threads WHERE id = ?", [(id,) for id in ids])
+ _mark_chat_attachment_inventory_clean(conn)
conn.commit()
finally:
conn.close()
@@ -1228,7 +1387,11 @@ def delete_chat_threads(ids: list[str]) -> None:
def clear_chat_history() -> None:
conn = get_connection()
try:
+ conn.execute("BEGIN IMMEDIATE")
+ _ensure_chat_attachment_inventory_current(conn)
+ conn.execute("DELETE FROM chat_attachment_tombstones")
conn.execute("DELETE FROM chat_threads")
+ _mark_chat_attachment_inventory_clean(conn)
conn.commit()
finally:
conn.close()
@@ -1354,6 +1517,7 @@ def delete_chat_project(id: str, delete_files: bool = False) -> Optional[dict]:
conn = get_connection()
try:
conn.execute("BEGIN IMMEDIATE")
+ _ensure_chat_attachment_inventory_current(conn)
row = conn.execute("SELECT * FROM chat_projects WHERE id = ?", (id,)).fetchone()
if row is None:
conn.rollback()
@@ -1361,6 +1525,7 @@ def delete_chat_project(id: str, delete_files: bool = False) -> Optional[dict]:
project = _chat_project_from_row(row)
conn.execute("DELETE FROM chat_threads WHERE project_id = ?", (id,))
conn.execute("DELETE FROM chat_projects WHERE id = ?", (id,))
+ _mark_chat_attachment_inventory_clean(conn)
conn.commit()
if delete_files:
_delete_project_workspace(project)
@@ -1483,15 +1648,285 @@ def _recompute_chat_thread_updated_at(conn: sqlite3.Connection, thread_id: str)
)
+_CONTENT_PART_ID_PREFIX = "content-part-sha256-"
+_URI_SCHEME_RE = re.compile(r"^[A-Za-z][A-Za-z0-9+.-]*:")
+
+
+def _is_locally_stored_blob(value: str) -> bool:
+ """True for data URIs or bare base64, never external/blob URI references."""
+ candidate = value.lstrip()
+ if not candidate:
+ return False
+ if candidate[:5].lower() == "data:":
+ return True
+ if candidate.startswith(("//", "\\\\")):
+ return False
+ return _URI_SCHEME_RE.match(candidate) is None
+
+
+def _managed_content_part_payload(part: dict) -> Optional[tuple[str, Any]]:
+ """Return the locally stored blob payload used to identify a content part."""
+ image = part.get("image")
+ if isinstance(image, str) and image[:5].lower() == "data:":
+ return "image", image
+
+ audio = part.get("audio")
+ if isinstance(audio, str) and _is_locally_stored_blob(audio):
+ return "audio", audio
+ if isinstance(audio, dict):
+ data = audio.get("data")
+ if isinstance(data, str) and _is_locally_stored_blob(data):
+ return "audio", audio
+ return None
+
+
+def _content_part_id(part: dict) -> Optional[str]:
+ """Stable managed id derived from blob data, without mutating inference content."""
+ payload = _managed_content_part_payload(part)
+ if payload is None:
+ return None
+ canonical = json.dumps(
+ payload,
+ ensure_ascii = False,
+ separators = (",", ":"),
+ sort_keys = True,
+ ).encode("utf-8")
+ return f"{_CONTENT_PART_ID_PREFIX}{hashlib.sha256(canonical).hexdigest()}"
+
+
+def _chat_attachment_tombstones_for_messages(
+ conn: sqlite3.Connection, thread_id: str, message_ids: list[str]
+) -> dict[str, set[str]]:
+ tombstones = {message_id: set() for message_id in message_ids}
+ unique_ids = list(dict.fromkeys(message_ids))
+ for start in range(0, len(unique_ids), _SQLITE_IN_CHUNK_SIZE):
+ chunk = unique_ids[start : start + _SQLITE_IN_CHUNK_SIZE]
+ placeholders = ",".join("?" for _ in chunk)
+ rows = conn.execute(
+ f"""
+ SELECT message_id, attachment_id
+ FROM chat_attachment_tombstones
+ WHERE thread_id = ? AND message_id IN ({placeholders})
+ """,
+ (thread_id, *chunk),
+ ).fetchall()
+ for row in rows:
+ tombstones[row["message_id"]].add(row["attachment_id"])
+ return tombstones
+
+
+def _reconcile_chat_message_uploads(message: dict, tombstones: set[str]) -> dict:
+ """Strip uploads previously deleted through the Data tab from a stale write."""
+ if not tombstones:
+ return message
+
+ reconciled = dict(message)
+ attachments = message.get("attachments")
+ if isinstance(attachments, list):
+ reconciled["attachments"] = [
+ attachment
+ for attachment in attachments
+ if not (isinstance(attachment, dict) and str(attachment.get("id") or "") in tombstones)
+ ]
+
+ content = message.get("content")
+ if isinstance(content, list):
+ reconciled["content"] = [
+ part
+ for part in content
+ if not (isinstance(part, dict) and (_content_part_id(part) or "") in tombstones)
+ ]
+ return reconciled
+
+
+def _chat_attachment_metadata_text(value, fallback: Optional[str] = None) -> Optional[str]:
+ """Keep untyped legacy/import metadata safe for SQLite binding."""
+ if value is None:
+ return fallback
+ if isinstance(value, str):
+ return value or fallback
+ if isinstance(value, (bool, int, float)):
+ return str(value)
+ # Objects and arrays are not useful display metadata and sqlite3 rejects
+ # binding them directly.
+ return fallback
+
+
+def _chat_attachment_inventory_entries(
+ attachments_json: Optional[str],
+ content_json: Optional[str],
+ tombstones: Optional[set[str]] = None,
+) -> list[dict]:
+ tombstones = tombstones or set()
+ attachments = _json_loads(attachments_json, None)
+ if not isinstance(attachments, list):
+ attachments = []
+ attachments = [
+ attachment
+ for attachment in attachments
+ if isinstance(attachment, dict) and attachment.get("id")
+ ]
+ attachments.extend(_content_part_attachments(content_json))
+
+ entries: list[dict] = []
+ seen: set[str] = set()
+ for attachment in attachments:
+ attachment_id = str(attachment["id"])
+ if attachment_id in seen or attachment_id in tombstones:
+ continue
+ seen.add(attachment_id)
+ entries.append(
+ {
+ "id": attachment_id,
+ "name": _chat_attachment_metadata_text(attachment.get("name"), "attachment"),
+ "type": _chat_attachment_metadata_text(attachment.get("type")),
+ "contentType": _chat_attachment_metadata_text(attachment.get("contentType")),
+ "sizeBytes": _chat_attachment_size_bytes(attachment),
+ }
+ )
+ return entries
+
+
+def _replace_chat_attachment_inventory(
+ conn: sqlite3.Connection,
+ message_id: str,
+ attachments_json: Optional[str],
+ content_json: Optional[str],
+ tombstones: Optional[set[str]] = None,
+) -> None:
+ conn.execute("DELETE FROM chat_attachment_inventory WHERE message_id = ?", (message_id,))
+ entries = _chat_attachment_inventory_entries(
+ attachments_json,
+ content_json,
+ tombstones,
+ )
+ conn.executemany(
+ """
+ INSERT INTO chat_attachment_inventory
+ (message_id, attachment_id, name, type, content_type, size_bytes)
+ VALUES (?, ?, ?, ?, ?, ?)
+ """,
+ [
+ (
+ message_id,
+ entry["id"],
+ entry["name"],
+ entry["type"],
+ entry["contentType"],
+ entry["sizeBytes"],
+ )
+ for entry in entries
+ ],
+ )
+
+
+def _mark_chat_attachment_inventory_clean(conn: sqlite3.Connection) -> None:
+ conn.execute(
+ """
+ INSERT INTO chat_attachment_inventory_state
+ (singleton, inventory_version, dirty, backfilled_at)
+ VALUES (1, ?, 0, ?)
+ ON CONFLICT(singleton) DO UPDATE SET
+ inventory_version = excluded.inventory_version,
+ dirty = 0,
+ backfilled_at = excluded.backfilled_at
+ """,
+ (
+ _CHAT_ATTACHMENT_INVENTORY_VERSION,
+ int(datetime.now(timezone.utc).timestamp() * 1000),
+ ),
+ )
+
+
+def _rebuild_chat_attachment_inventory(conn: sqlite3.Connection) -> None:
+ """Rebuild after schema upgrade or a write from an older Studio build."""
+ conn.execute("DELETE FROM chat_attachment_inventory")
+ tombstones: dict[tuple[str, str], set[str]] = {}
+ for row in conn.execute(
+ "SELECT thread_id, message_id, attachment_id FROM chat_attachment_tombstones"
+ ).fetchall():
+ tombstones.setdefault((row["thread_id"], row["message_id"]), set()).add(
+ row["attachment_id"]
+ )
+ rows = conn.execute(
+ "SELECT id, thread_id, attachments_json, content_json FROM chat_messages"
+ ).fetchall()
+ for row in rows:
+ _replace_chat_attachment_inventory(
+ conn,
+ row["id"],
+ row["attachments_json"],
+ row["content_json"],
+ tombstones.get((row["thread_id"], row["id"]), set()),
+ )
+
+
+def _ensure_chat_attachment_inventory_current(conn: sqlite3.Connection) -> None:
+ state = conn.execute(
+ """
+ SELECT inventory_version, dirty
+ FROM chat_attachment_inventory_state
+ WHERE singleton = 1
+ """
+ ).fetchone()
+ if (
+ state is not None
+ and state["inventory_version"] == _CHAT_ATTACHMENT_INVENTORY_VERSION
+ and not state["dirty"]
+ ):
+ return
+
+ owns_transaction = not conn.in_transaction
+ if owns_transaction:
+ conn.execute("BEGIN IMMEDIATE")
+ try:
+ state = conn.execute(
+ """
+ SELECT inventory_version, dirty
+ FROM chat_attachment_inventory_state
+ WHERE singleton = 1
+ """
+ ).fetchone()
+ if (
+ state is None
+ or state["inventory_version"] != _CHAT_ATTACHMENT_INVENTORY_VERSION
+ or state["dirty"]
+ ):
+ _rebuild_chat_attachment_inventory(conn)
+ _mark_chat_attachment_inventory_clean(conn)
+ if owns_transaction:
+ conn.commit()
+ except Exception:
+ if owns_transaction:
+ conn.rollback()
+ raise
+
+
def upsert_chat_message(message: dict) -> dict:
conn = get_connection()
try:
conn.execute("BEGIN IMMEDIATE")
+ _ensure_chat_attachment_inventory_current(conn)
_raise_if_chat_message_thread_conflicts(
conn,
message["threadId"],
[message["id"]],
)
+ tombstones = _chat_attachment_tombstones_for_messages(
+ conn,
+ message["threadId"],
+ [message["id"]],
+ )
+ reconciled = _reconcile_chat_message_uploads(
+ message,
+ tombstones.get(message["id"], set()),
+ )
+ content_json = json.dumps(reconciled.get("content", []))
+ attachments_json = (
+ json.dumps(reconciled.get("attachments"))
+ if reconciled.get("attachments") is not None
+ else None
+ )
conn.execute(
"""
INSERT INTO chat_messages
@@ -1507,23 +1942,32 @@ def upsert_chat_message(message: dict) -> dict:
WHERE excluded.thread_id = chat_messages.thread_id
""",
(
- message["id"],
- message["threadId"],
- message.get("parentId"),
- message["role"],
- json.dumps(message.get("content", [])),
- json.dumps(message.get("attachments"))
- if message.get("attachments") is not None
+ reconciled["id"],
+ reconciled["threadId"],
+ reconciled.get("parentId"),
+ reconciled["role"],
+ content_json,
+ attachments_json,
+ json.dumps(reconciled.get("metadata"))
+ if reconciled.get("metadata") is not None
else None,
- json.dumps(message.get("metadata"))
- if message.get("metadata") is not None
- else None,
- int(message["createdAt"]),
+ int(reconciled["createdAt"]),
),
)
- _bump_chat_thread_updated_at(conn, message["threadId"], int(message["createdAt"]))
+ _replace_chat_attachment_inventory(
+ conn,
+ reconciled["id"],
+ attachments_json,
+ content_json,
+ )
+ _bump_chat_thread_updated_at(
+ conn,
+ reconciled["threadId"],
+ int(reconciled["createdAt"]),
+ )
+ _mark_chat_attachment_inventory_clean(conn)
conn.commit()
- return message
+ return reconciled
except Exception:
conn.rollback()
raise
@@ -1539,13 +1983,28 @@ def sync_chat_messages(
conn = get_connection()
try:
conn.execute("BEGIN IMMEDIATE")
+ _ensure_chat_attachment_inventory_current(conn)
_raise_if_chat_message_thread_conflicts(
conn,
thread_id,
[m["id"] for m in messages],
)
- if prune_missing:
- conn.execute("DELETE FROM chat_messages WHERE thread_id = ?", (thread_id,))
+ tombstones = _chat_attachment_tombstones_for_messages(
+ conn,
+ thread_id,
+ [m["id"] for m in messages],
+ )
+ reconciled_messages = [
+ _reconcile_chat_message_uploads(m, tombstones.get(m["id"], set())) for m in messages
+ ]
+ serialized_messages = [
+ (
+ m,
+ json.dumps(m.get("content", [])),
+ json.dumps(m.get("attachments")) if m.get("attachments") is not None else None,
+ )
+ for m in reconciled_messages
+ ]
conn.executemany(
"""
INSERT INTO chat_messages
@@ -1566,20 +2025,46 @@ def sync_chat_messages(
thread_id,
m.get("parentId"),
m["role"],
- json.dumps(m.get("content", [])),
- json.dumps(m.get("attachments")) if m.get("attachments") is not None else None,
+ content_json,
+ attachments_json,
json.dumps(m.get("metadata")) if m.get("metadata") is not None else None,
int(m["createdAt"]),
)
- for m in messages
+ for m, content_json, attachments_json in serialized_messages
],
)
- if prune_missing:
- _recompute_chat_thread_updated_at(conn, thread_id)
- elif messages:
- _bump_chat_thread_updated_at(
- conn, thread_id, max(int(m["createdAt"]) for m in messages)
+ for m, content_json, attachments_json in serialized_messages:
+ _replace_chat_attachment_inventory(
+ conn,
+ m["id"],
+ attachments_json,
+ content_json,
)
+ if prune_missing:
+ retained_ids = {m["id"] for m in reconciled_messages}
+ existing_ids = {
+ row["id"]
+ for row in conn.execute(
+ "SELECT id FROM chat_messages WHERE thread_id = ?",
+ (thread_id,),
+ ).fetchall()
+ }
+ missing_ids = sorted(existing_ids - retained_ids)
+ for start in range(0, len(missing_ids), _SQLITE_IN_CHUNK_SIZE):
+ chunk = missing_ids[start : start + _SQLITE_IN_CHUNK_SIZE]
+ placeholders = ",".join("?" for _ in chunk)
+ conn.execute(
+ f"DELETE FROM chat_messages WHERE thread_id = ? AND id IN ({placeholders})",
+ (thread_id, *chunk),
+ )
+ _recompute_chat_thread_updated_at(conn, thread_id)
+ elif reconciled_messages:
+ _bump_chat_thread_updated_at(
+ conn,
+ thread_id,
+ max(int(m["createdAt"]) for m in reconciled_messages),
+ )
+ _mark_chat_attachment_inventory_clean(conn)
conn.commit()
return list_chat_messages(thread_id)
except ChatMessageConflictError:
@@ -1613,6 +2098,7 @@ def fork_chat_thread(
conn = get_connection()
try:
conn.execute("BEGIN IMMEDIATE")
+ _ensure_chat_attachment_inventory_current(conn)
src = conn.execute(
"SELECT * FROM chat_threads WHERE id = ?", (source_thread_id,)
).fetchone()
@@ -1686,6 +2172,14 @@ def fork_chat_thread(
for row in ancestry
],
)
+ for row in ancestry:
+ _replace_chat_attachment_inventory(
+ conn,
+ id_map[row["id"]],
+ row["attachments_json"],
+ row["content_json"],
+ )
+ _mark_chat_attachment_inventory_clean(conn)
conn.commit()
thread_row = conn.execute(
"SELECT * FROM chat_threads WHERE id = ?", (new_thread_id,)
@@ -1744,6 +2238,279 @@ def get_chat_message(thread_id: str, message_id: str) -> Optional[dict]:
conn.close()
+def _blob_part_base64_len(part: dict) -> int:
+ """Base64 payload length of an image or audio content part, or 0."""
+ image = part.get("image")
+ if isinstance(image, str) and image[:5].lower() == "data:":
+ return len(image.rsplit(",", 1)[-1])
+ audio = part.get("audio")
+ if isinstance(audio, str) and _is_locally_stored_blob(audio):
+ return len(audio.rsplit(",", 1)[-1])
+ if isinstance(audio, dict):
+ data = audio.get("data")
+ if isinstance(data, str) and _is_locally_stored_blob(data):
+ return len(data)
+ return 0
+
+
+def _chat_attachment_size_bytes(attachment: dict) -> Optional[int]:
+ """Approximate stored size of one attachment's content parts.
+
+ Image and audio parts hold base64 payloads (decoded bytes ~= 3/4 of the
+ encoded length); text parts count their character length. None when there
+ is no sizable content (e.g. a stripped/legacy attachment).
+ """
+ total = 0
+ found = False
+ for part in attachment.get("content") or []:
+ if not isinstance(part, dict):
+ continue
+ blob_len = _blob_part_base64_len(part)
+ if blob_len > 0:
+ total += (blob_len * 3) // 4
+ found = True
+ continue
+ text = part.get("text")
+ if isinstance(text, str) and text:
+ total += len(text.encode("utf-8", errors = "ignore"))
+ found = True
+ return total if found else None
+
+
+def _content_part_attachments(content_json: Optional[str]) -> list[dict]:
+ """Managed local blobs stored in content_json, with stable payload ids.
+
+ Exact duplicate blobs intentionally share one inventory id. Deleting that
+ id removes every identical copy, avoiding ambiguous index-based addressing.
+ """
+ content = _json_loads(content_json, None)
+ if not isinstance(content, list):
+ return []
+ out: list[dict] = []
+ seen: set[str] = set()
+ for part in content:
+ if not isinstance(part, dict):
+ continue
+ attachment_id = _content_part_id(part)
+ payload = _managed_content_part_payload(part)
+ if attachment_id is None or payload is None or attachment_id in seen:
+ continue
+ seen.add(attachment_id)
+ kind, value = payload
+ content_type = None
+ if kind == "image" and isinstance(value, str):
+ content_type = value[5:].split(";", 1)[0].split(",", 1)[0] or None
+ out.append(
+ {
+ "id": attachment_id,
+ "type": kind,
+ "name": "Chat image" if kind == "image" else "Chat audio",
+ "contentType": content_type,
+ "content": [part],
+ }
+ )
+ return out
+
+
+def list_chat_attachments_page(
+ limit: int = 50, offset: int = 0
+) -> tuple[list[dict], Optional[int]]:
+ """One bounded page from the normalized attachment inventory."""
+ if not 1 <= limit <= 100:
+ raise ValueError("limit must be between 1 and 100")
+ if offset < 0:
+ raise ValueError("offset must be non-negative")
+
+ conn = get_connection()
+ try:
+ _ensure_chat_attachment_inventory_current(conn)
+ rows = conn.execute(
+ """
+ SELECT i.attachment_id, i.name, i.type, i.content_type,
+ i.size_bytes, m.id AS message_id, m.thread_id,
+ m.created_at, t.title AS thread_title, t.pair_id
+ FROM chat_attachment_inventory i
+ JOIN chat_messages m ON m.id = i.message_id
+ LEFT JOIN chat_threads t ON t.id = m.thread_id
+ ORDER BY m.created_at DESC, m.id ASC, i.attachment_id ASC
+ LIMIT ? OFFSET ?
+ """,
+ (limit + 1, offset),
+ ).fetchall()
+ finally:
+ conn.close()
+
+ has_more = len(rows) > limit
+ page_rows = rows[:limit]
+ attachments = [
+ {
+ "id": row["attachment_id"],
+ "messageId": row["message_id"],
+ "threadId": row["thread_id"],
+ "pairId": row["pair_id"],
+ "threadTitle": row["thread_title"],
+ "name": row["name"],
+ "type": row["type"],
+ "contentType": row["content_type"],
+ "sizeBytes": row["size_bytes"],
+ "createdAt": row["created_at"],
+ }
+ for row in page_rows
+ ]
+ return attachments, offset + limit if has_more else None
+
+
+def list_chat_attachments() -> list[dict]:
+ """Compatibility helper returning the full normalized inventory."""
+ attachments: list[dict] = []
+ offset = 0
+ while True:
+ page, next_offset = list_chat_attachments_page(limit = 100, offset = offset)
+ attachments.extend(page)
+ if next_offset is None:
+ return attachments
+ offset = next_offset
+
+
+def get_chat_attachment(message_id: str, attachment_id: str) -> Optional[dict]:
+ """One attachment record (full content) from a message, or None."""
+ conn = get_connection()
+ try:
+ row = conn.execute(
+ """
+ SELECT message.attachments_json, message.content_json,
+ EXISTS(
+ SELECT 1 FROM chat_attachment_tombstones tombstone
+ WHERE tombstone.thread_id = message.thread_id
+ AND tombstone.message_id = message.id
+ AND tombstone.attachment_id = ?
+ ) AS tombstoned
+ FROM chat_messages message
+ WHERE message.id = ?
+ """,
+ (attachment_id, message_id),
+ ).fetchone()
+ finally:
+ conn.close()
+ if row is None or row["tombstoned"]:
+ return None
+ attachments = _json_loads(row["attachments_json"], None)
+ if isinstance(attachments, list):
+ for attachment in attachments:
+ if isinstance(attachment, dict) and str(attachment.get("id") or "") == attachment_id:
+ return attachment
+ if attachment_id.startswith(_CONTENT_PART_ID_PREFIX):
+ for attachment in _content_part_attachments(row["content_json"]):
+ if attachment["id"] == attachment_id:
+ return attachment
+ return None
+
+
+def _record_chat_attachment_tombstone(
+ conn: sqlite3.Connection, thread_id: str, message_id: str, attachment_id: str
+) -> None:
+ conn.execute(
+ """
+ INSERT INTO chat_attachment_tombstones
+ (thread_id, message_id, attachment_id, deleted_at)
+ VALUES (?, ?, ?, ?)
+ ON CONFLICT(thread_id, message_id, attachment_id) DO UPDATE SET
+ deleted_at = excluded.deleted_at
+ """,
+ (
+ thread_id,
+ message_id,
+ attachment_id,
+ int(datetime.now(timezone.utc).timestamp() * 1000),
+ ),
+ )
+
+
+def delete_chat_attachment(message_id: str, attachment_id: str) -> bool:
+ """Remove one stored upload from a message.
+
+ The tombstone is retained while the thread exists, so pruning and later
+ recreating the same message id cannot restore the deleted upload. If an
+ ordinary attachment id collides with a content-blob id, both are deleted as
+ one managed item.
+ """
+ conn = get_connection()
+ try:
+ conn.execute("BEGIN IMMEDIATE")
+ _ensure_chat_attachment_inventory_current(conn)
+ row = conn.execute(
+ """
+ SELECT thread_id, attachments_json, content_json
+ FROM chat_messages WHERE id = ?
+ """,
+ (message_id,),
+ ).fetchone()
+ if row is None:
+ conn.rollback()
+ return False
+
+ attachments = _json_loads(row["attachments_json"], None)
+ updated_attachments_json = row["attachments_json"]
+ deleted_attachment = False
+ if isinstance(attachments, list):
+ remaining_attachments = [
+ attachment
+ for attachment in attachments
+ if not (
+ isinstance(attachment, dict)
+ and str(attachment.get("id") or "") == attachment_id
+ )
+ ]
+ deleted_attachment = len(remaining_attachments) != len(attachments)
+ if deleted_attachment:
+ updated_attachments_json = json.dumps(remaining_attachments)
+
+ content = _json_loads(row["content_json"], None)
+ updated_content_json = row["content_json"]
+ deleted_content = False
+ if attachment_id.startswith(_CONTENT_PART_ID_PREFIX) and isinstance(content, list):
+ remaining_content = [
+ part
+ for part in content
+ if not (isinstance(part, dict) and _content_part_id(part) == attachment_id)
+ ]
+ deleted_content = len(remaining_content) != len(content)
+ if deleted_content:
+ updated_content_json = json.dumps(remaining_content)
+
+ if not deleted_attachment and not deleted_content:
+ conn.rollback()
+ return False
+ conn.execute(
+ """
+ UPDATE chat_messages
+ SET attachments_json = ?, content_json = ?
+ WHERE id = ?
+ """,
+ (updated_attachments_json, updated_content_json, message_id),
+ )
+ _record_chat_attachment_tombstone(
+ conn,
+ row["thread_id"],
+ message_id,
+ attachment_id,
+ )
+ _replace_chat_attachment_inventory(
+ conn,
+ message_id,
+ updated_attachments_json,
+ updated_content_json,
+ )
+ _mark_chat_attachment_inventory_clean(conn)
+ conn.commit()
+ return True
+ except Exception:
+ conn.rollback()
+ raise
+ finally:
+ conn.close()
+
+
def list_chat_messages_for_threads(thread_ids: list[str]) -> list[dict]:
if not thread_ids:
return []
diff --git a/studio/backend/tests/test_chat_attachments.py b/studio/backend/tests/test_chat_attachments.py
new file mode 100644
index 0000000000..459587ca9e
--- /dev/null
+++ b/studio/backend/tests/test_chat_attachments.py
@@ -0,0 +1,634 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+import base64
+import json
+import os
+import sqlite3
+import sys
+
+import pytest
+from fastapi import HTTPException
+
+_backend = os.path.join(os.path.dirname(__file__), "..")
+sys.path.insert(0, _backend)
+
+from routes import chat_history
+from storage import studio_db
+from utils.paths import studio_db_path
+
+PNG_BYTES = base64.b64decode(
+ "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8BQDwAEhQGAhKmMIQAAAABJRU5ErkJggg=="
+)
+PNG_DATA_URL = "data:image/png;base64," + base64.b64encode(PNG_BYTES).decode("ascii")
+
+
+def _reset_studio_db(tmp_path, monkeypatch):
+ monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(tmp_path))
+ monkeypatch.setenv("UNSLOTH_STUDIO_PROJECTS_HOME", str(tmp_path / "Projects"))
+ monkeypatch.setattr(studio_db, "_schema_ready", False)
+
+
+def _thread(
+ thread_id: str = "thread-1",
+ title: str = "Test Chat",
+ pair_id: str | None = None,
+) -> dict:
+ return {
+ "id": thread_id,
+ "title": title,
+ "modelType": "base",
+ "modelId": "test-model",
+ "pairId": pair_id,
+ "archived": False,
+ "createdAt": 1_700_000_000_000,
+ }
+
+
+def _message(
+ message_id: str,
+ created_at: int = 1_700_000_000_000,
+ attachments = None,
+ thread_id: str = "thread-1",
+) -> dict:
+ message = {
+ "id": message_id,
+ "threadId": thread_id,
+ "parentId": None,
+ "role": "user",
+ "content": [{"type": "text", "text": "hello"}],
+ "createdAt": created_at,
+ }
+ if attachments is not None:
+ message["attachments"] = attachments
+ return message
+
+
+def _image_attachment(attachment_id: str = "att-1", name: str = "photo.png") -> dict:
+ return {
+ "id": attachment_id,
+ "type": "image",
+ "name": name,
+ "contentType": "image/png",
+ "content": [{"type": "image", "image": PNG_DATA_URL}],
+ "status": {"type": "complete"},
+ }
+
+
+def _seed(
+ tmp_path,
+ monkeypatch,
+ attachments,
+ message_id: str = "msg-1",
+):
+ _reset_studio_db(tmp_path, monkeypatch)
+ studio_db.upsert_chat_thread(_thread())
+ studio_db.upsert_chat_message(_message(message_id, attachments = attachments))
+
+
+def _set_raw_attachments_json(message_id: str, raw: str) -> None:
+ conn = sqlite3.connect(studio_db_path())
+ try:
+ conn.execute(
+ "UPDATE chat_messages SET attachments_json = ? WHERE id = ?",
+ (raw, message_id),
+ )
+ conn.commit()
+ finally:
+ conn.close()
+
+
+def _raw_attachments_json(message_id: str):
+ conn = sqlite3.connect(studio_db_path())
+ try:
+ row = conn.execute(
+ "SELECT attachments_json FROM chat_messages WHERE id = ?",
+ (message_id,),
+ ).fetchone()
+ return row[0] if row is not None else None
+ finally:
+ conn.close()
+
+
+# ---------------------------------------------------------------------------
+# Storage: list_chat_attachments
+# ---------------------------------------------------------------------------
+
+
+def test_list_chat_attachments_empty_db(tmp_path, monkeypatch):
+ _reset_studio_db(tmp_path, monkeypatch)
+ assert studio_db.list_chat_attachments() == []
+
+
+def test_list_chat_attachments_round_trip(tmp_path, monkeypatch):
+ _seed(tmp_path, monkeypatch, [_image_attachment()])
+ records = studio_db.list_chat_attachments()
+ assert len(records) == 1
+ record = records[0]
+ assert record["id"] == "att-1"
+ assert record["messageId"] == "msg-1"
+ assert record["threadId"] == "thread-1"
+ assert record["threadTitle"] == "Test Chat"
+ assert record["name"] == "photo.png"
+ assert record["type"] == "image"
+ assert record["contentType"] == "image/png"
+ assert record["createdAt"] == 1_700_000_000_000
+ # Base64 length estimate is within padding error of the decoded size.
+ assert abs(record["sizeBytes"] - len(PNG_BYTES)) <= 2
+
+
+def test_list_chat_attachments_counts_text_utf8(tmp_path, monkeypatch):
+ text = "héllo wörld é世界"
+ attachment = {
+ "id": "att-txt",
+ "type": "document",
+ "name": "notes.txt",
+ "content": [{"type": "text", "text": text}],
+ }
+ _seed(tmp_path, monkeypatch, [attachment])
+ records = studio_db.list_chat_attachments()
+ assert records[0]["sizeBytes"] == len(text.encode("utf-8"))
+
+
+def test_list_chat_attachments_no_content_size_is_none(tmp_path, monkeypatch):
+ attachment = {"id": "att-empty", "name": "ghost.bin", "content": []}
+ _seed(tmp_path, monkeypatch, [attachment])
+ records = studio_db.list_chat_attachments()
+ assert records[0]["sizeBytes"] is None
+ assert records[0]["name"] == "ghost.bin"
+
+
+def test_list_chat_attachments_defaults_missing_name(tmp_path, monkeypatch):
+ attachment = {"id": "att-noname", "content": []}
+ _seed(tmp_path, monkeypatch, [attachment])
+ assert studio_db.list_chat_attachments()[0]["name"] == "attachment"
+
+
+def test_list_chat_attachments_sanitizes_structured_metadata(tmp_path, monkeypatch):
+ attachment = {
+ "id": "att-weird",
+ "name": {"nested": "name"},
+ "type": ["image"],
+ "contentType": {"mime": "image/png"},
+ "content": [],
+ }
+ _seed(tmp_path, monkeypatch, [attachment])
+ record = studio_db.list_chat_attachments()[0]
+ assert record["name"] == "attachment"
+ assert record["type"] is None
+ assert record["contentType"] is None
+
+
+def test_list_chat_attachments_skips_malformed_rows(tmp_path, monkeypatch):
+ _reset_studio_db(tmp_path, monkeypatch)
+ studio_db.upsert_chat_thread(_thread())
+ for i, raw in enumerate(
+ [
+ "not json at all",
+ '{"id": "att-obj"}',
+ "null",
+ "[]",
+ '[{"noid": true}, "just a string", 42]',
+ '[{"id": ""}]',
+ ]
+ ):
+ message_id = f"msg-bad-{i}"
+ studio_db.upsert_chat_message(_message(message_id))
+ _set_raw_attachments_json(message_id, raw)
+ studio_db.upsert_chat_message(_message("msg-good", attachments = [_image_attachment("att-ok")]))
+ records = studio_db.list_chat_attachments()
+ assert [r["id"] for r in records] == ["att-ok"]
+
+
+def test_list_chat_attachments_orders_newest_first(tmp_path, monkeypatch):
+ _reset_studio_db(tmp_path, monkeypatch)
+ studio_db.upsert_chat_thread(_thread())
+ studio_db.upsert_chat_message(
+ _message("msg-old", 1_700_000_000_000, [_image_attachment("att-old")])
+ )
+ studio_db.upsert_chat_message(
+ _message("msg-new", 1_700_000_100_000, [_image_attachment("att-new")])
+ )
+ assert [r["id"] for r in studio_db.list_chat_attachments()] == ["att-new", "att-old"]
+
+
+def test_list_chat_attachments_survives_missing_thread_row(tmp_path, monkeypatch):
+ _reset_studio_db(tmp_path, monkeypatch)
+ studio_db.upsert_chat_thread(_thread())
+ studio_db.upsert_chat_message(_message("msg-1", attachments = [_image_attachment()]))
+ conn = sqlite3.connect(studio_db_path())
+ try:
+ conn.execute("DELETE FROM chat_threads WHERE id = 'thread-1'")
+ conn.commit()
+ finally:
+ conn.close()
+ records = studio_db.list_chat_attachments()
+ assert len(records) == 1
+ assert records[0]["threadTitle"] is None
+
+
+def test_list_chat_attachments_includes_compare_pair_id(tmp_path, monkeypatch):
+ _reset_studio_db(tmp_path, monkeypatch)
+ studio_db.upsert_chat_thread(_thread(pair_id = "pair-1"))
+ studio_db.upsert_chat_message(_message("msg-compare", attachments = [_image_attachment()]))
+ record = studio_db.list_chat_attachments()[0]
+ assert record["threadId"] == "thread-1"
+ assert record["pairId"] == "pair-1"
+
+
+def test_list_chat_attachments_gone_after_thread_delete(tmp_path, monkeypatch):
+ _seed(tmp_path, monkeypatch, [_image_attachment()])
+ studio_db.delete_chat_threads(["thread-1"])
+ assert studio_db.list_chat_attachments() == []
+
+
+# ---------------------------------------------------------------------------
+# Storage: get_chat_attachment / delete_chat_attachment
+# ---------------------------------------------------------------------------
+
+
+def test_get_chat_attachment_found_and_missing(tmp_path, monkeypatch):
+ _seed(tmp_path, monkeypatch, [_image_attachment()])
+ attachment = studio_db.get_chat_attachment("msg-1", "att-1")
+ assert attachment is not None
+ assert attachment["content"][0]["image"] == PNG_DATA_URL
+ assert studio_db.get_chat_attachment("msg-1", "att-missing") is None
+ assert studio_db.get_chat_attachment("msg-missing", "att-1") is None
+
+
+def test_delete_chat_attachment_keeps_others(tmp_path, monkeypatch):
+ _seed(
+ tmp_path,
+ monkeypatch,
+ [_image_attachment("att-1"), _image_attachment("att-2", "other.png")],
+ )
+ assert studio_db.delete_chat_attachment("msg-1", "att-1") is True
+ assert studio_db.get_chat_attachment("msg-1", "att-1") is None
+ assert studio_db.get_chat_attachment("msg-1", "att-2") is not None
+ assert [r["id"] for r in studio_db.list_chat_attachments()] == ["att-2"]
+
+
+def test_delete_last_chat_attachment_stores_empty_list(tmp_path, monkeypatch):
+ _seed(tmp_path, monkeypatch, [_image_attachment()])
+ assert studio_db.delete_chat_attachment("msg-1", "att-1") is True
+ # '[]' rather than NULL: a NULL attachments field reads back as missing
+ # and triggers the legacy IndexedDB backfill, resurrecting the deleted
+ # attachment on the next chat load.
+ assert _raw_attachments_json("msg-1") == "[]"
+ assert studio_db.list_chat_attachments() == []
+ # The message itself must survive with its content intact.
+ message = studio_db.get_chat_message("thread-1", "msg-1")
+ assert message is not None
+ assert message["content"] == [{"type": "text", "text": "hello"}]
+ assert message["attachments"] == []
+
+
+def test_delete_chat_attachment_missing_targets(tmp_path, monkeypatch):
+ _seed(tmp_path, monkeypatch, [_image_attachment()])
+ assert studio_db.delete_chat_attachment("msg-missing", "att-1") is False
+ assert studio_db.delete_chat_attachment("msg-1", "att-missing") is False
+ _set_raw_attachments_json("msg-1", "not json")
+ assert studio_db.delete_chat_attachment("msg-1", "att-1") is False
+
+
+# ---------------------------------------------------------------------------
+# Routes: /attachments endpoints (real storage, direct calls)
+# ---------------------------------------------------------------------------
+
+
+def test_list_attachments_route(tmp_path, monkeypatch):
+ _seed(tmp_path, monkeypatch, [_image_attachment()])
+ result = chat_history.list_attachments(current_subject = "unsloth")
+ assert [a["id"] for a in result["attachments"]] == ["att-1"]
+
+
+def test_attachment_file_serves_image_bytes(tmp_path, monkeypatch):
+ _seed(tmp_path, monkeypatch, [_image_attachment()])
+ response = chat_history.get_attachment_file("msg-1", "att-1", current_subject = "unsloth")
+ assert response.body == PNG_BYTES
+ assert response.media_type == "image/png"
+
+
+def test_attachment_file_tolerates_whitespace_in_base64(tmp_path, monkeypatch):
+ encoded = base64.b64encode(PNG_BYTES).decode("ascii")
+ wrapped = "\n".join(encoded[i : i + 8] for i in range(0, len(encoded), 8))
+ attachment = _image_attachment()
+ attachment["content"] = [{"type": "image", "image": "data:image/png;base64," + wrapped}]
+ _seed(tmp_path, monkeypatch, [attachment])
+ response = chat_history.get_attachment_file("msg-1", "att-1", current_subject = "unsloth")
+ assert response.body == PNG_BYTES
+
+
+def test_attachment_file_corrupt_base64_is_422(tmp_path, monkeypatch):
+ attachment = _image_attachment()
+ attachment["content"] = [{"type": "image", "image": "data:image/png;base64,%%%"}]
+ _seed(tmp_path, monkeypatch, [attachment])
+ with pytest.raises(HTTPException) as excinfo:
+ chat_history.get_attachment_file("msg-1", "att-1", current_subject = "unsloth")
+ assert excinfo.value.status_code == 422
+
+
+def test_attachment_file_accepts_urlsafe_base64(tmp_path, monkeypatch):
+ data = bytes(range(251, 256)) * 3 # encodes to characters remapped by urlsafe
+ payload = base64.urlsafe_b64encode(data).decode("ascii")
+ assert "-" in payload or "_" in payload
+ attachment = _image_attachment()
+ attachment["content"] = [{"type": "image", "image": "data:image/png;base64," + payload}]
+ _seed(tmp_path, monkeypatch, [attachment])
+ response = chat_history.get_attachment_file("msg-1", "att-1", current_subject = "unsloth")
+ assert response.body == data
+
+
+def test_attachment_file_accepts_missing_padding(tmp_path, monkeypatch):
+ payload = base64.b64encode(PNG_BYTES).decode("ascii").rstrip("=")
+ attachment = _image_attachment()
+ attachment["content"] = [{"type": "image", "image": "data:image/png;base64," + payload}]
+ _seed(tmp_path, monkeypatch, [attachment])
+ response = chat_history.get_attachment_file("msg-1", "att-1", current_subject = "unsloth")
+ assert response.body == PNG_BYTES
+
+
+def test_attachment_file_serves_percent_encoded_data_url(tmp_path, monkeypatch):
+ attachment = _image_attachment()
+ attachment["content"] = [{"type": "image", "image": "data:text/plain,hello%20world"}]
+ _seed(tmp_path, monkeypatch, [attachment])
+ response = chat_history.get_attachment_file("msg-1", "att-1", current_subject = "unsloth")
+ assert response.body == b"hello world"
+ # Non-image data URL types are clamped so markup never renders same-origin.
+ assert response.media_type == "application/octet-stream"
+
+
+def test_attachment_file_serves_text_parts(tmp_path, monkeypatch):
+ attachment = {
+ "id": "att-txt",
+ "type": "document",
+ "name": "notes.txt",
+ "content": [
+ {"type": "text", "text": "first"},
+ {"type": "text", "text": "second"},
+ ],
+ }
+ _seed(tmp_path, monkeypatch, [attachment])
+ response = chat_history.get_attachment_file("msg-1", "att-txt", current_subject = "unsloth")
+ assert response.body.decode("utf-8") == "first\nsecond"
+ assert response.media_type.startswith("text/plain")
+
+
+def test_attachment_file_no_content_is_404(tmp_path, monkeypatch):
+ _seed(tmp_path, monkeypatch, [{"id": "att-empty", "name": "ghost", "content": []}])
+ with pytest.raises(HTTPException) as excinfo:
+ chat_history.get_attachment_file("msg-1", "att-empty", current_subject = "unsloth")
+ assert excinfo.value.status_code == 404
+
+
+def test_attachment_file_missing_message_is_404(tmp_path, monkeypatch):
+ _reset_studio_db(tmp_path, monkeypatch)
+ with pytest.raises(HTTPException) as excinfo:
+ chat_history.get_attachment_file("nope", "att-1", current_subject = "unsloth")
+ assert excinfo.value.status_code == 404
+
+
+def test_attachment_file_non_data_url_image_is_404(tmp_path, monkeypatch):
+ attachment = _image_attachment()
+ attachment["content"] = [{"type": "image", "image": "https://example.com/a.png"}]
+ _seed(tmp_path, monkeypatch, [attachment])
+ with pytest.raises(HTTPException) as excinfo:
+ chat_history.get_attachment_file("msg-1", "att-1", current_subject = "unsloth")
+ assert excinfo.value.status_code == 404
+
+
+def test_attachment_file_defaults_media_type(tmp_path, monkeypatch):
+ payload = base64.b64encode(b"raw-bytes").decode("ascii")
+ attachment = _image_attachment()
+ attachment["content"] = [{"type": "image", "image": "data:;base64," + payload}]
+ _seed(tmp_path, monkeypatch, [attachment])
+ response = chat_history.get_attachment_file("msg-1", "att-1", current_subject = "unsloth")
+ assert response.body == b"raw-bytes"
+ assert response.media_type == "application/octet-stream"
+
+
+def test_attachment_file_svg_media_type(tmp_path, monkeypatch):
+ svg = b""
+ payload = base64.b64encode(svg).decode("ascii")
+ attachment = _image_attachment()
+ attachment["content"] = [{"type": "image", "image": "data:image/svg+xml;base64," + payload}]
+ _seed(tmp_path, monkeypatch, [attachment])
+ response = chat_history.get_attachment_file("msg-1", "att-1", current_subject = "unsloth")
+ assert response.body == svg
+ # SVG can carry scripts, so it downloads as bytes instead of rendering.
+ assert response.media_type == "application/octet-stream"
+
+
+def test_delete_attachment_route_then_404(tmp_path, monkeypatch):
+ _seed(tmp_path, monkeypatch, [_image_attachment()])
+ result = chat_history.delete_attachment("msg-1", "att-1", current_subject = "unsloth")
+ assert result == {"ok": True}
+ with pytest.raises(HTTPException) as excinfo:
+ chat_history.delete_attachment("msg-1", "att-1", current_subject = "unsloth")
+ assert excinfo.value.status_code == 404
+
+
+# ---------------------------------------------------------------------------
+# Audio attachments (adapter {data, format} and compare-chat bare base64)
+# ---------------------------------------------------------------------------
+
+WAV_BYTES = b"RIFF$\x00\x00\x00WAVEfmt \x10\x00\x00\x00\x01\x00\x01\x00"
+WAV_B64 = base64.b64encode(WAV_BYTES).decode("ascii")
+
+
+def _audio_attachment(attachment_id: str = "att-audio") -> dict:
+ return {
+ "id": attachment_id,
+ "type": "file",
+ "name": "clip.wav",
+ "contentType": "audio/wav",
+ "content": [{"type": "audio", "audio": {"data": WAV_B64, "format": "wav"}}],
+ "status": {"type": "complete"},
+ }
+
+
+def test_audio_attachment_lists_with_size(tmp_path, monkeypatch):
+ _seed(tmp_path, monkeypatch, [_audio_attachment()])
+ records = studio_db.list_chat_attachments()
+ assert len(records) == 1
+ assert records[0]["id"] == "att-audio"
+ assert abs(records[0]["sizeBytes"] - len(WAV_BYTES)) <= 2
+
+
+def test_audio_attachment_file_serves_bytes(tmp_path, monkeypatch):
+ _seed(tmp_path, monkeypatch, [_audio_attachment()])
+ response = chat_history.get_attachment_file("msg-1", "att-audio", current_subject = "unsloth")
+ assert response.body == WAV_BYTES
+ assert response.media_type == "audio/wav"
+
+
+def test_audio_attachment_media_type_from_format(tmp_path, monkeypatch):
+ attachment = _audio_attachment()
+ attachment["contentType"] = None
+ attachment["content"] = [{"type": "audio", "audio": {"data": WAV_B64, "format": "mp3"}}]
+ _seed(tmp_path, monkeypatch, [attachment])
+ response = chat_history.get_attachment_file("msg-1", "att-audio", current_subject = "unsloth")
+ assert response.media_type == "audio/mpeg"
+
+
+def test_audio_attachment_corrupt_payload_is_422(tmp_path, monkeypatch):
+ attachment = _audio_attachment()
+ attachment["content"] = [{"type": "audio", "audio": {"data": "%%%", "format": "wav"}}]
+ _seed(tmp_path, monkeypatch, [attachment])
+ with pytest.raises(HTTPException) as excinfo:
+ chat_history.get_attachment_file("msg-1", "att-audio", current_subject = "unsloth")
+ assert excinfo.value.status_code == 422
+
+
+# ---------------------------------------------------------------------------
+# Compare-chat uploads stored as message content parts
+# ---------------------------------------------------------------------------
+
+
+def _compare_message(message_id: str = "msg-cmp") -> dict:
+ return {
+ "id": message_id,
+ "threadId": "thread-1",
+ "parentId": None,
+ "role": "user",
+ "content": [
+ {"type": "image", "image": PNG_DATA_URL},
+ {"type": "audio", "audio": WAV_B64},
+ {"type": "text", "text": "compare these"},
+ ],
+ "createdAt": 1_700_000_000_000,
+ }
+
+
+def _seed_compare(tmp_path, monkeypatch):
+ _reset_studio_db(tmp_path, monkeypatch)
+ studio_db.upsert_chat_thread(_thread())
+ studio_db.upsert_chat_message(_compare_message())
+
+
+_CONTENT_PART_PREFIX = "content-part-sha256-"
+
+
+def _content_part_id_for(message_id: str, kind: str) -> str:
+ """Resolve the stable content-hash id for a message's stored blob.
+
+ Content-part ids are SHA-256 hashes of the blob payload, not array
+ indices, so tests look them up from the listing instead of hardcoding an
+ index that would shift when an earlier part is deleted.
+ """
+ for record in studio_db.list_chat_attachments():
+ if record["messageId"] == message_id and record["type"] == kind:
+ return record["id"]
+ raise AssertionError(f"no {kind} content-part upload for {message_id}")
+
+
+def test_content_part_uploads_are_listed(tmp_path, monkeypatch):
+ _seed_compare(tmp_path, monkeypatch)
+ records = studio_db.list_chat_attachments()
+ # Ids are stable content hashes, not array indices.
+ assert all(r["id"].startswith(_CONTENT_PART_PREFIX) for r in records)
+ assert {r["type"] for r in records} == {"image", "audio"}
+ image = next(r for r in records if r["type"] == "image")
+ assert image["contentType"] == "image/png"
+ assert abs(image["sizeBytes"] - len(PNG_BYTES)) <= 2
+ audio = next(r for r in records if r["type"] == "audio")
+ assert audio["type"] == "audio"
+
+
+def test_content_part_file_serves_image_bytes(tmp_path, monkeypatch):
+ _seed_compare(tmp_path, monkeypatch)
+ image_id = _content_part_id_for("msg-cmp", "image")
+ response = chat_history.get_attachment_file("msg-cmp", image_id, current_subject = "unsloth")
+ assert response.body == PNG_BYTES
+ assert response.media_type == "image/png"
+
+
+def test_content_part_delete_keeps_text(tmp_path, monkeypatch):
+ _seed_compare(tmp_path, monkeypatch)
+ image_id = _content_part_id_for("msg-cmp", "image")
+ assert studio_db.delete_chat_attachment("msg-cmp", image_id) is True
+ message = studio_db.get_chat_message("thread-1", "msg-cmp")
+ types = [p["type"] for p in message["content"]]
+ assert types == ["audio", "text"]
+ # The surviving audio blob keeps its own stable hash id after the delete.
+ remaining = studio_db.list_chat_attachments()
+ assert [r["type"] for r in remaining] == ["audio"]
+ assert remaining[0]["id"].startswith(_CONTENT_PART_PREFIX)
+ assert remaining[0]["id"] != image_id
+
+
+def test_content_part_delete_rejects_non_blob(tmp_path, monkeypatch):
+ _seed_compare(tmp_path, monkeypatch)
+ # The text part is not a stored upload, so it never gets an id: only the
+ # image and audio blobs are addressable.
+ assert len(studio_db.list_chat_attachments()) == 2
+ # A well-formed but unknown content-hash id, and malformed ids, all no-op.
+ assert studio_db.delete_chat_attachment("msg-cmp", _CONTENT_PART_PREFIX + "0" * 64) is False
+ assert studio_db.delete_chat_attachment("msg-cmp", "content-part-99") is False
+ assert studio_db.delete_chat_attachment("msg-cmp", "content-part-x") is False
+
+
+def test_text_only_messages_not_listed_as_uploads(tmp_path, monkeypatch):
+ _reset_studio_db(tmp_path, monkeypatch)
+ studio_db.upsert_chat_thread(_thread())
+ # The word "image" inside text must not create phantom upload rows.
+ message = _message("msg-txt")
+ message["content"] = [{"type": "text", "text": 'discussing an "image" and "audio" here'}]
+ studio_db.upsert_chat_message(message)
+ assert studio_db.list_chat_attachments() == []
+
+
+def test_remote_image_urls_are_not_listed_as_uploads(tmp_path, monkeypatch):
+ _reset_studio_db(tmp_path, monkeypatch)
+ studio_db.upsert_chat_thread(_thread())
+ message = _message("msg-remote")
+ message["content"] = [
+ {"type": "image", "image": "https://example.com/cat.png"},
+ {"type": "text", "text": "look at this"},
+ ]
+ studio_db.upsert_chat_message(message)
+ # No stored bytes: nothing to list, open, or delete.
+ assert studio_db.list_chat_attachments() == []
+ assert studio_db.get_chat_attachment("msg-remote", "content-part-0") is None
+ assert studio_db.delete_chat_attachment("msg-remote", "content-part-0") is False
+ stored = studio_db.get_chat_message("thread-1", "msg-remote")
+ assert [p["type"] for p in stored["content"]] == ["image", "text"]
+
+
+def test_html_data_url_serves_as_octet_stream(tmp_path, monkeypatch):
+ _reset_studio_db(tmp_path, monkeypatch)
+ studio_db.upsert_chat_thread(_thread())
+ html_b64 = base64.b64encode(b"").decode()
+ message = _message("msg-html")
+ message["content"] = [
+ {"type": "image", "image": f"data:text/html;base64,{html_b64}"},
+ ]
+ studio_db.upsert_chat_message(message)
+ attachment_id = _content_part_id_for("msg-html", "image")
+ response = chat_history.get_attachment_file(
+ "msg-html", attachment_id, current_subject = "unsloth"
+ )
+ # Never echo a script-capable media type back under the app origin.
+ assert response.media_type == "application/octet-stream"
+ assert response.body == b""
+
+
+def test_svg_data_url_serves_as_octet_stream(tmp_path, monkeypatch):
+ _reset_studio_db(tmp_path, monkeypatch)
+ studio_db.upsert_chat_thread(_thread())
+ svg_b64 = base64.b64encode(b"").decode()
+ message = _message("msg-svg")
+ message["content"] = [
+ {"type": "image", "image": f"data:image/svg+xml;base64,{svg_b64}"},
+ ]
+ studio_db.upsert_chat_message(message)
+ attachment_id = _content_part_id_for("msg-svg", "image")
+ response = chat_history.get_attachment_file("msg-svg", attachment_id, current_subject = "unsloth")
+ assert response.media_type == "application/octet-stream"
+
+
+def test_png_data_url_keeps_its_media_type(tmp_path, monkeypatch):
+ _seed_compare(tmp_path, monkeypatch)
+ image_id = _content_part_id_for("msg-cmp", "image")
+ response = chat_history.get_attachment_file("msg-cmp", image_id, current_subject = "unsloth")
+ assert response.media_type == "image/png"
diff --git a/studio/frontend/src/components/app-sidebar.tsx b/studio/frontend/src/components/app-sidebar.tsx
index b6e218793a..b8601b00f6 100644
--- a/studio/frontend/src/components/app-sidebar.tsx
+++ b/studio/frontend/src/components/app-sidebar.tsx
@@ -969,10 +969,10 @@ export function AppSidebar() {
))}
- {/* Bulk export and import live in Settings -> Chat -> Data. */}
+ {/* Bulk export and import live in Settings -> Data. */}
- useSettingsDialogStore.getState().openDialog("chat")
+ useSettingsDialogStore.getState().openDialog("data")
}
>
Export all chats…
diff --git a/studio/frontend/src/components/assistant-ui/attachment.tsx b/studio/frontend/src/components/assistant-ui/attachment.tsx
index 98ebe5ab5f..b26840cc02 100644
--- a/studio/frontend/src/components/assistant-ui/attachment.tsx
+++ b/studio/frontend/src/components/assistant-ui/attachment.tsx
@@ -7,6 +7,7 @@
import { TooltipIconButton } from "@/components/assistant-ui/tooltip-icon-button";
import {
Dialog,
+ DialogClose,
DialogContent,
DialogTitle,
DialogTrigger,
@@ -27,12 +28,7 @@ import {
import { AudioWave01Icon, File02Icon } from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
import { PlusIcon, XIcon } from "lucide-react";
-import {
- type FC,
- type PropsWithChildren,
- useEffect,
- useState,
-} from "react";
+import { type FC, type PropsWithChildren, useEffect, useState } from "react";
import { useShallow } from "zustand/shallow";
const useFileSrc = (file: File | undefined): string | undefined => {
@@ -83,7 +79,7 @@ const AttachmentPreview: FC = ({ src }) => {
src={src}
alt="Preview"
className={cn(
- "block h-auto max-h-[80vh] w-auto max-w-full object-contain",
+ "block h-auto max-h-[90dvh] w-auto max-w-[92vw] object-contain",
isLoaded
? "aui-attachment-preview-image-loaded"
: "aui-attachment-preview-image-loading invisible",
@@ -108,12 +104,23 @@ const AttachmentPreviewDialog: FC = ({ children }) => {
>
{children}
-
+ {/* Chrome-free lightbox: the image floats on the dimmed backdrop with
+ no dialog panel, and the close button sits in the screen corner. */}
+
Image Attachment Preview
-
-
+ {/* Clicking the backdrop (anywhere off the image) closes the preview. */}
+
+
+
+
+ );
+}
diff --git a/studio/frontend/src/i18n/locales/en.ts b/studio/frontend/src/i18n/locales/en.ts
index cbddc9f0c2..5ee2805a33 100644
--- a/studio/frontend/src/i18n/locales/en.ts
+++ b/studio/frontend/src/i18n/locales/en.ts
@@ -99,6 +99,7 @@ export const en = {
chat: "Chat",
voice: "Voice",
connections: "Connections",
+ data: "Data",
apiKeys: "API",
about: "About",
},
@@ -511,7 +512,7 @@ export const en = {
},
chat: {
title: "Chat",
- description: "Manage chat history stored on this device.",
+ description: "Customize how chat behaves on this device.",
modelDisclaimer: "Show model disclaimer",
modelDisclaimerDescription:
'Show "LLMs can make mistakes" under the chat box.',
@@ -580,6 +581,53 @@ export const en = {
"A storage clear failed; {count} chats may remain. Please retry.",
failedToClearChats: "Failed to clear chats",
},
+ data: {
+ title: "Data",
+ description:
+ "Manage chat history and uploaded files stored on this device.",
+ archivedChats: "Archived chats",
+ archivedChatsDescription: "View and manage chats you have archived.",
+ manageAction: "Manage",
+ exportArchivedChats: "Export",
+ exportingArchivedChats: "Exporting...",
+ exportedOneArchivedChat: "Exported 1 archived chat",
+ exportedArchivedChatCount: "Exported {count} archived chats",
+ noArchivedChatsToExport: "No archived chats to export.",
+ failedToExportArchivedChats: "Failed to export archived chats",
+ archiveAllChats: "Archive all chats",
+ archiveAllChatsDescription:
+ "Move every chat in Recents and Projects to the archive.",
+ noChatsToArchive: "No chats to archive.",
+ archiveAllAction: "Archive all",
+ archivingAction: "Archiving...",
+ archiveAllChatsTitle: "Archive all chats?",
+ archiveAllChatsConfirmDescription:
+ "Moves every chat on this device to the archive. Archived chats stay available and can be unarchived at any time.",
+ archivedAllChats: "Archived all chats",
+ archivedOneChat: "Archived 1 chat",
+ archivedChatCount: "Archived {count} chats",
+ failedToArchiveChats: "Failed to archive chats",
+ confirmBeforeDeleting: "Confirm before deleting",
+ confirmBeforeDeletingDescription:
+ "Ask for confirmation before a chat is deleted. Turn off to delete instantly.",
+ filesSection: "Files",
+ uploadedFiles: "Uploaded files",
+ uploadedFilesDescription:
+ "View and manage files uploaded to chats, projects, and knowledge bases.",
+ fineTuneExport: "Use chats as training data",
+ fineTuneExportDescription:
+ "Create a fine-tuning JSONL dataset from your chats. Load it in Train, refine in Recipes, or export it.",
+ fineTuneExportAction: "Export JSONL",
+ fineTuneRunAction: "Run",
+ fineTuneExportingAction: "Exporting...",
+ fineTuneOpenRecipesAction: "Open in Recipes",
+ fineTuneOpeningRecipesAction: "Opening...",
+ fineTuneTrainAction: "Load in Train tab",
+ fineTuneTrainingAction: "Loading...",
+ fineTuneExportFailed: "Failed to export training data",
+ fineTuneRecipeFailed: "Failed to open chats in Recipes",
+ fineTuneTrainFailed: "Failed to load dataset in the Train tab",
+ },
connections: {
title: "Connections",
description: "Manage providers and external connections.",
From 66808ab25dcb7655dab05a0837a0ec0bb6cf080b Mon Sep 17 00:00:00 2001
From: Daniel Han
Date: Mon, 20 Jul 2026 05:27:53 -0700
Subject: [PATCH 16/41] Studio: fix per-GPU VRAM reporting on Windows ROCm
(#7238)
* Studio: fix per-GPU VRAM reporting on Windows ROCm
On Windows ROCm without a HIP SDK, amd-smi is disabled and the System tab fell
back to torch mem_get_info, which reports free==total there (ROCm/ROCm#1909), so
used VRAM showed as 0. The perf-counter fallback also summed every adapter into a
single device with only GPU 0's total, hiding the second GPU.
Read per-adapter Dedicated Usage (LUID-instanced) for used and take each GPU's
total from torch properties, and treat the free==total case as unknown rather
than 0, so every GPU shows real usage. NVIDIA, Linux ROCm, Apple and CPU paths
are unchanged. Final validation needs a real Windows AMD box.
Fixes #7072
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Studio: report unknown VRAM instead of fabricating or zeroing it
Two gaps in the Windows ROCm VRAM path. When more adapters are actively using
VRAM than are visible to the process (a GPU outside the visibility mask), the
per-adapter attribution paired usage by size and fabricated a per-GPU value;
report unknown for every device in that case rather than mis-assign. And the
System API turned an unknown (None) used value into 0 with ``or 0``, then
reported the full card as free, re-hiding the exact case this change surfaces;
keep None so the UI shows unknown.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* Render unknown VRAM as Unknown instead of zero in the System tab
The backend reports null usage when it is unknown (e.g. the Windows ROCm
perf counter is unavailable or localized), but the System tab coerced
null to 0 and derived free from it, fabricating a 0-used/full-free total.
Preserve null and render the translated Unknown for per-device used, free
and utilization, and mark the aggregate VRAM tile unknown when any device
is unknown.
* Render unknown VRAM as Unknown in the floating monitor and the util tile
The floating VRAM monitor and the aggregate utilization ring both still
coerced a null usage to 0, showing a fabricated 0.00 GiB / full free / 0%
on the same Windows ROCm no-counter case the resources tab already
handles. Guard both on whether every device reports a finite usage and
render Unknown (value and percent) instead of a concrete 0.
* Attribute per-adapter VRAM usage only when capacity forces the mapping
On Windows/ROCm there is no shared key between LUID performance-counter
instances and torch ordinals, so usage was paired to devices purely by capacity
ranking. That pairing is only trustworthy when capacity forces it (a usage
larger than every smaller device can sit on one card). When a smaller-capacity
device could equally hold a strictly larger usage (for example an 8 GiB card
near full beside a lightly used 48 GiB card), the two values are swappable
without violating any capacity, so the ranking is a guess with no key to break
the tie. A wrong guess both mislabels the System tab and feeds
routes/training_vram.py a wrong per-index free value, driving a wrong
keep-resident decision.
Report unknown for every device when the assignment is ambiguous, keeping the
attribution only for the capacity-forced case. Returning None is the
conservative direction: training_vram treats a missing index as zero free, so it
never keeps a chat model into an OOM. Add regression tests for the
not-capacity-ordered, same-capacity, single-fits-both, and capacity-forced
cases.
* Report unknown VRAM usage when a hidden adapter survives the noise filter
When HIP_VISIBLE_DEVICES exposes a subset of the physical adapters, the LUID
usage counters cover cards outside the visibility mask too. The sub-64 MiB noise
filter could drop a genuinely-idle visible card's real usage while keeping a
hidden larger card's high usage, which was then clamped onto the smaller visible
device and reported as fully used (for example a hidden 48 GiB card at 40 GiB
shown as a visible 8 GiB card fully used, with its true 10 MiB usage filtered
out). That fabricated reading also feeds routes/training_vram.py a wrong
per-index free value.
Flag extra adapters on the raw counter count (before the noise filter, since an
idle visible card can itself fall below the floor) and, when a kept usage exceeds
its ranked visible capacity, report unknown rather than clamp a hidden card's
usage onto a visible device. The genuinely-idle-noise and capacity-forced
single-model cases are unchanged. Add a regression test for the hidden
high-use-adapter case in both counter orders.
* Report unknown when only a placeholder adapter counter survives the noise filter
When more raw counters than visible devices are present but every counter sits
below the 64 MiB noise floor (an idle real GPU alongside a Windows Basic Render
Driver placeholder), the non_trivial-or-raw fallback resurrected the raw
magnitude-sorted counters and could attribute the placeholder to a real GPU while
dropping a real card's reading. With a single visible device the swap-ambiguity
check cannot catch it (it needs at least two ranks), so the fabricated value
reached the System tab and automatic GPU selection.
Return unknown for every device in that case instead of falling back to raw
counters. With the earlier guards this completes the invariant: a concrete
per-GPU usage is emitted only when the assignment is capacity-forced, and every
ambiguous, extra-adapter, placeholder-fallback, or count-mismatch path reports
unknown. Add a regression test for the placeholder fallback in both counter
orders and the two-idle-GPU case.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* studio: attribute Windows/ROCm VRAM only when capacity forces a clean bijection
With more raw adapter counters than visible devices, a survivor that merely
fits a visible card was pinned to it by magnitude ranking, fabricating a hidden
GPU's usage onto an idle visible card whose true reading was dropped by the
sub-threshold noise filter (two visible 48/8 GiB cards using 40 GiB / 10 MiB
beside a hidden 6 GiB adapter returned [40, 6]). Emit a concrete per-device
value only when the supra-threshold counters number exactly the visible devices
(every visible card has one real reading, the extras were sub-threshold
placeholders) AND the ranked usage strictly exceeds every smaller visible card's
capacity. When a visible card is idle (fewer supra-threshold counters than
devices) a survivor could be the hidden GPU's usage, so every device reports
unknown; more active counters than visible cards, the smallest card, and any
merely-fitting usage stay unknown too. The reporter's loaded-card display is
preserved (40 GiB / 0.5 GiB across 48/8 GiB -> [40, None]). Adds a regression
test for the reported case plus an exhaustive capacity-forced/bijection matrix.
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
* studio: keep the unified-memory total when Windows-ROCm used is unknown
_apply_unified_memory_correction gated both the total and the used update on
torch_used_gb being known, so on a unified-memory APU (Strix Halo) where torch
reports used=None (the Windows-ROCm free==total sentinel) but an authoritative
full-GTT total, the device kept amd-smi's small dedicated carve-out and
underreported its capacity on the System tab. Adopt torch's larger total
independently of used; overwrite used only when torch's is known (otherwise keep
amd-smi's dedicated-usage figure) and recompute utilization against the
corrected total. Adds regression tests.
* Tighten comments in the ROCm/Windows VRAM reporting path
* Tighten comments further in the ROCm/Windows VRAM reporting path
---------
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: danielhanchen
---
studio/backend/main.py | 8 +-
.../tests/test_rocm_windows_vram_7072.py | 361 ++++++++++++++++++
studio/backend/utils/hardware/hardware.py | 308 ++++++++++++---
.../src/components/floating-monitor.tsx | 26 +-
.../features/settings/tabs/resources-tab.tsx | 99 +++--
5 files changed, 716 insertions(+), 86 deletions(-)
create mode 100644 studio/backend/tests/test_rocm_windows_vram_7072.py
diff --git a/studio/backend/main.py b/studio/backend/main.py
index a1ff4d60da..0ffc4489a1 100644
--- a/studio/backend/main.py
+++ b/studio/backend/main.py
@@ -1149,11 +1149,15 @@ def _get_cached_system_gpu_info(logger) -> dict[str, Any]:
util = util_devices.get(idx, {})
total_vram = util.get("vram_total_gb") or dev.get("memory_total_gb") or 0
- used_vram = util.get("vram_used_gb") or 0
+ # Keep None (usage unknown, e.g. Windows ROCm perf counter) so the UI
+ # shows unknown, not a fabricated 0 used / full free.
+ used_vram = util.get("vram_used_gb")
enriched_dev = dict(dev)
enriched_dev["vram_used_gb"] = used_vram
- enriched_dev["vram_free_gb"] = round(total_vram - used_vram, 2) if total_vram else 0
+ enriched_dev["vram_free_gb"] = (
+ round(total_vram - used_vram, 2) if total_vram and used_vram is not None else None
+ )
enriched_dev["vram_utilization_pct"] = util.get("vram_utilization_pct")
enriched_devices.append(enriched_dev)
diff --git a/studio/backend/tests/test_rocm_windows_vram_7072.py b/studio/backend/tests/test_rocm_windows_vram_7072.py
new file mode 100644
index 0000000000..b4079831b7
--- /dev/null
+++ b/studio/backend/tests/test_rocm_windows_vram_7072.py
@@ -0,0 +1,361 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""Regression tests for issue #7072 -- "VRAM Usage in System Tab is wrong".
+
+Reporter: dual AMD (Radeon PRO W7900 ~48GB + W7500 8GB), Windows 10, ROCm 7.13,
+torch 2.11.0+rocm7.13. On Windows without a HIP SDK, amd-smi is permanently
+disabled (avoids a UAC/DiskPart prompt) and hipMemGetInfo returns free==total
+(used 0). Two symptoms followed:
+
+ * System tab (/api/system -> get_visible_gpu_utilization) showed ~0 VRAM used
+ on every GPU (torch mem_get_info free==total quirk; ROCm/ROCm#1909).
+ * get_gpu_utilization()'s Windows fallback SUMMED "GPU Adapter Memory\\Dedicated
+ Usage" across all adapters into ONE fake device with only GPU 0's total, so
+ the second GPU never appeared.
+
+The fix reads the per-adapter (LUID-instanced) Dedicated Usage performance
+counter -- Task Manager's source -- for per-GPU used, takes per-GPU total from
+torch device properties, and guards the free==total mem_get_info quirk. CI has no
+AMD GPU/Windows, so torch, the performance counter, and platform are all mocked.
+"""
+
+from __future__ import annotations
+
+import subprocess
+import sys
+import types
+
+import pytest
+
+from utils.hardware import hardware as hw
+
+GB = 1024**3
+MiB = 1024**2
+
+
+# ----------------------------------------------------------------------------- #
+# Fakes
+# ----------------------------------------------------------------------------- #
+def _fake_torch(
+ devices,
+ *,
+ free_equals_total = False,
+ used_per_device = None,
+):
+ """Build a fake `torch` module. devices: list of (name, total_bytes)."""
+ dev = list(devices)
+
+ class _Props:
+ def __init__(self, name, total):
+ self.name = name
+ self.total_memory = total
+
+ def get_device_properties(i):
+ name, total = dev[i]
+ return _Props(name, total)
+
+ def mem_get_info(i):
+ _, total = dev[i]
+ if free_equals_total:
+ return (total, total)
+ used = used_per_device[i] if used_per_device is not None else 0
+ return (total - used, total)
+
+ t = types.ModuleType("torch")
+ t.__version__ = "2.11.0+rocm7.13"
+ t.version = types.SimpleNamespace(hip = "7.13", cuda = None)
+ t.cuda = types.SimpleNamespace(
+ is_available = lambda: len(dev) > 0,
+ device_count = lambda: len(dev),
+ current_device = lambda: 0,
+ get_device_properties = get_device_properties,
+ mem_get_info = mem_get_info,
+ memory_allocated = lambda i: 0,
+ memory_reserved = lambda i: 0,
+ )
+ return t
+
+
+def _adapter_output(adapters):
+ if not adapters:
+ return "__NONE__\n"
+ return "".join(f"{name}|{int(used)}\n" for name, used in adapters)
+
+
+def _subprocess_run(*, adapter_output = "__NONE__\n", util_output = "12.0\n"):
+ def fake_run(cmd, *a, **k):
+ joined = " ".join(cmd) if isinstance(cmd, list) else str(cmd)
+ if "GPU Adapter Memory" in joined and "InstanceName" in joined:
+ out = adapter_output
+ elif "engtype_3D" in joined or "GPU Engine" in joined:
+ out = util_output
+ else:
+ out = "-1\n"
+ return subprocess.CompletedProcess(args = cmd, returncode = 0, stdout = out, stderr = "")
+
+ return fake_run
+
+
+@pytest.fixture
+def win_rocm(monkeypatch):
+ """Configure the hardware module as a Windows ROCm host with 2 visible GPUs."""
+ monkeypatch.setattr(hw, "get_device", lambda: hw.DeviceType.CUDA)
+ monkeypatch.setattr(hw, "IS_ROCM", True)
+ monkeypatch.setattr(hw.platform, "system", lambda: "Windows")
+ monkeypatch.setattr(hw.sys, "platform", "win32")
+ monkeypatch.setattr(hw, "_smi_query", lambda *a, **k: None) # amd-smi disabled
+ # Visible set via HIP mask so we don't shell out to amd-smi for the count.
+ monkeypatch.setenv("HIP_VISIBLE_DEVICES", "0,1")
+ monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising = False)
+ monkeypatch.delenv("ROCR_VISIBLE_DEVICES", raising = False)
+ return monkeypatch
+
+
+REPORTER_ADAPTERS = [
+ ("luid_0x00000000_0x0000d1e2_phys_0", 40.0 * GB), # W7900, model loaded
+ ("luid_0x00000000_0x0000e34a_phys_0", 0.5 * GB), # W7500, idle
+ ("luid_0x00000000_0x0000f001_phys_0", 3 * MiB), # Basic Render Driver
+]
+DEVICES = [("AMD Radeon PRO W7900", 48 * GB), ("AMD Radeon PRO W7500", 8 * GB)]
+
+
+# ----------------------------------------------------------------------------- #
+# System tab (get_visible_gpu_utilization) -- the reporter's screenshot
+# ----------------------------------------------------------------------------- #
+def test_system_tab_shows_per_gpu_used(win_rocm, monkeypatch):
+ monkeypatch.setitem(sys.modules, "torch", _fake_torch(DEVICES, free_equals_total = True))
+ monkeypatch.setattr(
+ hw.subprocess, "run", _subprocess_run(adapter_output = _adapter_output(REPORTER_ADAPTERS))
+ )
+
+ devices = hw.get_visible_gpu_utilization()["devices"]
+ by_idx = {d["index"]: d for d in devices}
+ assert len(devices) == 2
+ assert by_idx[0]["vram_total_gb"] == 48.0
+ assert by_idx[0]["vram_used_gb"] == pytest.approx(40.0, abs = 0.01) # not 0
+ assert by_idx[1]["vram_total_gb"] == 8.0 # own total
+ # The 3 MiB Basic Render Driver counter makes this a hidden-adapter case: only
+ # the 40 GiB is forced onto the 48 GiB card; the idle card reads Unknown.
+ assert by_idx[1]["vram_used_gb"] is None
+ assert by_idx[1]["vram_utilization_pct"] is None
+ assert all(
+ d["vram_used_gb"] <= d["vram_total_gb"] for d in devices if d["vram_used_gb"] is not None
+ )
+
+
+def test_gpu_utilization_does_not_collapse(win_rocm, monkeypatch):
+ monkeypatch.setitem(sys.modules, "torch", _fake_torch(DEVICES, free_equals_total = True))
+ monkeypatch.setattr(
+ hw.subprocess, "run", _subprocess_run(adapter_output = _adapter_output(REPORTER_ADAPTERS))
+ )
+
+ result = hw.get_gpu_utilization()
+ devices = result["devices"]
+ assert sorted(d["index"] for d in devices) == [0, 1] # both GPUs, no collapse
+ assert {d["vram_total_gb"] for d in devices} == {48.0, 8.0}
+ assert result["vram_total_gb"] == 48.0 # legacy primary mirror preserved
+
+
+def test_localized_counter_reports_unknown_not_zero(win_rocm, monkeypatch):
+ monkeypatch.setitem(sys.modules, "torch", _fake_torch(DEVICES, free_equals_total = True))
+ monkeypatch.setattr(hw.subprocess, "run", _subprocess_run(adapter_output = "__NONE__\n"))
+
+ devices = hw.get_visible_gpu_utilization()["devices"]
+ assert len(devices) == 2 # both still shown with correct totals
+ assert {d["vram_total_gb"] for d in devices} == {48.0, 8.0}
+ assert all(d["vram_used_gb"] is None for d in devices) # unknown, not fake 0
+ assert all(d["vram_utilization_pct"] is None for d in devices)
+
+
+# ----------------------------------------------------------------------------- #
+# mem_get_info free==total guard scoping
+# ----------------------------------------------------------------------------- #
+def test_mem_get_info_guard_scopes_to_windows_rocm(monkeypatch):
+ torch_mod = _fake_torch(DEVICES, free_equals_total = True)
+ monkeypatch.setattr(hw, "get_device", lambda: hw.DeviceType.CUDA)
+ monkeypatch.setitem(sys.modules, "torch", torch_mod)
+
+ # Windows ROCm -> used unknown (None), total kept.
+ monkeypatch.setattr(hw, "IS_ROCM", True)
+ monkeypatch.setattr(hw.sys, "platform", "win32")
+ win = hw._torch_get_per_device_info([0, 1])
+ assert [d["used_gb"] for d in win] == [None, None]
+ assert [d["total_gb"] for d in win] == [48.0, 8.0]
+
+ # Linux ROCm -> unchanged numeric used.
+ monkeypatch.setattr(hw.sys, "platform", "linux")
+ assert [d["used_gb"] for d in hw._torch_get_per_device_info([0, 1])] == [0.0, 0.0]
+
+ # Windows NVIDIA -> guard must not fire.
+ monkeypatch.setattr(hw, "IS_ROCM", False)
+ monkeypatch.setattr(hw.sys, "platform", "win32")
+ assert [d["used_gb"] for d in hw._torch_get_per_device_info([0, 1])] == [0.0, 0.0]
+
+
+# ----------------------------------------------------------------------------- #
+# Per-adapter attribution helpers (pure unit)
+# ----------------------------------------------------------------------------- #
+def test_match_adapter_pairs_and_clamps():
+ assert hw._match_adapter_used_to_devices([40 * GB, 0.5 * GB], [48 * GB, 8 * GB]) == [
+ 40 * GB,
+ 0.5 * GB,
+ ]
+ assert hw._match_adapter_used_to_devices([100 * GB], [48 * GB]) == [48 * GB] # clamp
+ assert hw._match_adapter_used_to_devices([40 * GB], [48 * GB, 8 * GB]) == [40 * GB, None]
+
+
+def test_match_adapter_reports_unknown_when_more_active_than_visible():
+ # More adapters actively using VRAM than are visible (a GPU outside the mask):
+ # attribution would fabricate a value, so report unknown for every device.
+ assert hw._match_adapter_used_to_devices([40 * GB, 0.5 * GB], [8 * GB]) == [None]
+
+
+def test_match_adapter_reports_unknown_when_hidden_high_use_adapter_survives_filter():
+ # Idle 8 GiB card (10 MiB noise) beside a hidden 48 GiB card at 40 GiB: the
+ # 40 GiB can't fit the 8 GiB device, so clamping there would fabricate. Unknown.
+ assert hw._match_adapter_used_to_devices([40 * GB, 10 * MiB], [8 * GB]) == [None]
+ # Order of the counters must not matter.
+ assert hw._match_adapter_used_to_devices([10 * MiB, 40 * GB], [8 * GB]) == [None]
+
+
+def test_match_adapter_reports_unknown_for_placeholder_fallback():
+ # Every counter below the 64 MiB floor plus a placeholder: no LUID-to-ordinal
+ # mapping tells placeholder from idle GPU, so report unknown, not fabricate.
+ # Single visible 8 GiB card idle (10 MiB) beside a 50 MiB placeholder counter.
+ assert hw._match_adapter_used_to_devices([50 * MiB, 10 * MiB], [8 * GB]) == [None]
+ # Order of the counters must not matter.
+ assert hw._match_adapter_used_to_devices([10 * MiB, 50 * MiB], [8 * GB]) == [None]
+ # Two idle visible GPUs plus a placeholder: all three counters below the floor.
+ assert hw._match_adapter_used_to_devices([50 * MiB, 10 * MiB, 5 * MiB], [48 * GB, 8 * GB]) == [
+ None,
+ None,
+ ]
+
+
+def test_match_adapter_reports_unknown_when_usage_not_capacity_ordered():
+ # 8 GiB card at 7 GiB beside a 48 GiB card at 5 GiB: the bigger usage still fits
+ # the smaller card, so both pairings are feasible -> unknown.
+ assert hw._match_adapter_used_to_devices([7 * GB, 5 * GB], [8 * GB, 48 * GB]) == [None, None]
+ # Device order must not matter (same physical situation, ordinals flipped).
+ assert hw._match_adapter_used_to_devices([7 * GB, 5 * GB], [48 * GB, 8 * GB]) == [None, None]
+ # Same-capacity cards with unequal usage are equally unattributable.
+ assert hw._match_adapter_used_to_devices([12 * GB, 8 * GB], [24 * GB, 24 * GB]) == [None, None]
+ # A single usage that fits both cards can sit on either -> unknown.
+ assert hw._match_adapter_used_to_devices([5 * GB], [48 * GB, 8 * GB]) == [None, None]
+ # But a capacity-forced assignment (usage exceeds the smaller card) is kept:
+ # 40 GiB can only be the 48 GiB card, so it is not fabrication.
+ assert hw._match_adapter_used_to_devices([40 * GB], [48 * GB, 8 * GB]) == [40 * GB, None]
+
+
+def test_match_adapter_reports_unknown_when_hidden_usage_fits_visible_card():
+ # A survivor that merely *fits* a visible card must not be pinned onto it. Two
+ # cards (48/8 GiB) at 40 GiB / 10 MiB beside a hidden 6 GiB adapter: the 6 GiB
+ # fits the idle 8 GiB card but isn't forced -> Unknown; only 40 GiB is forced.
+ assert hw._match_adapter_used_to_devices([40 * GB, 10 * MiB, 6 * GB], [48 * GB, 8 * GB]) == [
+ 40 * GB,
+ None,
+ ]
+ # Counter order must not matter.
+ assert hw._match_adapter_used_to_devices([6 * GB, 40 * GB, 10 * MiB], [48 * GB, 8 * GB]) == [
+ 40 * GB,
+ None,
+ ]
+ # A single visible card with a hidden adapter is never attributable: a fitting
+ # survivor could be the hidden GPU's while the visible card is idle.
+ assert hw._match_adapter_used_to_devices([6 * GB, 10 * MiB], [8 * GB]) == [None]
+
+
+def test_match_adapter_capacity_forced_matrix():
+ """Exhaustive hidden-adapter matrix for the capacity-forced rule.
+
+ A value is emitted only when the supra-threshold counters number exactly the
+ visible devices AND a device's ranked usage strictly exceeds every smaller
+ card's capacity. Otherwise (a visible card idle, a merely-fitting usage, or the
+ smallest card) every device reports unknown.
+ """
+ m = hw._match_adapter_used_to_devices
+ # -- exactly-n supra-threshold counters, capacity-forced survivors are kept - #
+ # Both visible cards have a real reading (the 3 MiB is a placeholder): 40 GiB
+ # forced onto the 48 GiB card, 0.5 GiB not forced -> None.
+ assert m([40 * GB, 0.5 * GB, 3 * MiB], [48 * GB, 8 * GB]) == [40 * GB, None]
+ # Three visible cards all active (supra-threshold) + placeholder: 40 > 24 and
+ # 20 > 8, both forced; the 8 GiB card is not forced -> None.
+ assert m([40 * GB, 20 * GB, 5 * GB, 3 * MiB], [48 * GB, 24 * GB, 8 * GB]) == [
+ 40 * GB,
+ 20 * GB,
+ None,
+ ]
+ # -- fewer supra-threshold counters than visible cards -> all unknown ------ #
+ # A visible card is idle, so even a "forced" 40 could be the hidden GPU's.
+ assert m([40 * GB, 3 * MiB, 3 * MiB], [48 * GB, 8 * GB]) == [None, None]
+ assert m([40 * GB, 10 * MiB, 10 * MiB], [48 * GB, 8 * GB]) == [None, None]
+ assert m([40 * GB, 20 * GB, 3 * MiB, 3 * MiB], [48 * GB, 24 * GB, 8 * GB]) == [
+ None,
+ None,
+ None,
+ ]
+ # Middle usage (6 GiB) fits both the 24 and 8 GiB cards, and only two cards are
+ # active for three visible -> not a bijection -> all unknown.
+ assert m([40 * GB, 6 * GB, 3 * MiB, 3 * MiB], [48 * GB, 24 * GB, 8 * GB]) == [
+ None,
+ None,
+ None,
+ ]
+ # -- hidden larger than every visible card -> all unknown ----------------- #
+ assert m([40 * GB, 10 * MiB], [8 * GB]) == [None]
+ assert m([48 * GB, 3 * MiB, 3 * MiB], [24 * GB, 8 * GB]) == [None, None]
+ # -- more active adapters than visible cards -> all unknown --------------- #
+ assert m([40 * GB, 7 * GB, 6 * GB, 3 * MiB], [48 * GB, 8 * GB]) == [None, None]
+ assert m([40 * GB, 7 * GB, 6 * GB, 3 * MiB, 3 * MiB], [48 * GB, 8 * GB]) == [None, None]
+ # -- every counter below the noise floor (placeholder fallback) -> unknown - #
+ assert m([50 * MiB, 10 * MiB], [8 * GB]) == [None]
+ assert m([50 * MiB, 10 * MiB, 5 * MiB], [48 * GB, 8 * GB]) == [None, None]
+ # -- equal-capacity cards with a hidden adapter: nothing is forced -------- #
+ assert m([40 * GB, 40 * GB, 3 * MiB], [48 * GB, 48 * GB]) == [None, None]
+ assert m([40 * GB, 30 * GB, 3 * MiB], [48 * GB, 48 * GB]) == [None, None]
+
+
+def test_perf_counter_parser_and_sentinel(monkeypatch):
+ monkeypatch.setattr(hw.platform, "system", lambda: "Windows")
+ monkeypatch.setattr(
+ hw.subprocess, "run", _subprocess_run(adapter_output = _adapter_output(REPORTER_ADAPTERS))
+ )
+ parsed = hw._rocm_windows_perf_counter_vram_by_adapter()
+ assert parsed is not None and len(parsed) == 3
+ assert parsed[0][0].startswith("luid_")
+ monkeypatch.setattr(hw.subprocess, "run", _subprocess_run(adapter_output = "__NONE__\n"))
+ assert hw._rocm_windows_perf_counter_vram_by_adapter() is None
+
+
+# ----------------------------------------------------------------------------- #
+# Unified-memory (Strix Halo APU) total reconciliation (Codex #7238)
+# ----------------------------------------------------------------------------- #
+def test_unified_memory_adopts_torch_total_even_when_used_unknown():
+ """Windows ROCm unified-memory APU: torch's used is None but its total (the full
+ GTT pool) is authoritative. The correction must still adopt the larger total;
+ used stays at amd-smi's figure when torch's is unknown."""
+ metrics = {"vram_total_gb": 8.0, "vram_used_gb": 2.0, "vram_utilization_pct": 25.0}
+ hw._apply_unified_memory_correction(metrics, {"total_gb": 124.0, "used_gb": None, "index": 0})
+ assert metrics["vram_total_gb"] == 124.0 # full unified pool, not the 8 GB carve-out
+ assert metrics["vram_used_gb"] == 2.0 # amd-smi used preserved (torch's was None)
+ assert metrics["vram_utilization_pct"] == pytest.approx(round(2.0 / 124.0 * 100, 1))
+
+
+def test_unified_memory_overwrites_used_when_torch_used_known():
+ """When torch reports both a larger total and a known used, both are adopted
+ and utilization is recomputed against the corrected total (unchanged path)."""
+ metrics = {"vram_total_gb": 8.0, "vram_used_gb": 2.0, "vram_utilization_pct": 25.0}
+ hw._apply_unified_memory_correction(metrics, {"total_gb": 124.0, "used_gb": 40.0, "index": 0})
+ assert metrics["vram_total_gb"] == 124.0
+ assert metrics["vram_used_gb"] == 40.0
+ assert metrics["vram_utilization_pct"] == pytest.approx(round(40.0 / 124.0 * 100, 1))
+
+
+def test_unified_memory_no_op_when_torch_total_not_larger():
+ """A discrete GPU where torch total does not exceed amd-smi's is left untouched."""
+ metrics = {"vram_total_gb": 48.0, "vram_used_gb": 10.0, "vram_utilization_pct": 20.8}
+ hw._apply_unified_memory_correction(metrics, {"total_gb": 48.0, "used_gb": None, "index": 0})
+ assert metrics["vram_total_gb"] == 48.0
+ assert metrics["vram_used_gb"] == 10.0
+ assert metrics["vram_utilization_pct"] == 20.8
diff --git a/studio/backend/utils/hardware/hardware.py b/studio/backend/utils/hardware/hardware.py
index adc9a54aab..9fef53e65e 100644
--- a/studio/backend/utils/hardware/hardware.py
+++ b/studio/backend/utils/hardware/hardware.py
@@ -538,21 +538,31 @@ def _torch_get_physical_gpu_count() -> Optional[int]:
def _torch_get_per_device_info(device_indices: list[int]) -> list[Dict[str, Any]]:
- """Query torch for per-GPU name, total VRAM, and used VRAM."""
+ """Query torch for per-GPU name, total VRAM, and used VRAM.
+
+ ``used_gb`` is ``None`` on Windows ROCm when ``hipMemGetInfo`` reports
+ ``free == total`` (ROCm/ROCm#1909): that 0 means unknown, not empty.
+ """
mod, _ = _torch_get_device_module()
if mod is None:
return []
+ # free==total is a Windows-ROCm-only quirk.
+ _win_rocm = sys.platform == "win32" and IS_ROCM
devices = []
for ordinal, phys_idx in enumerate(device_indices):
try:
# torch ordinals are 0-based relative to CUDA_VISIBLE_DEVICES.
props = mod.get_device_properties(ordinal)
total_bytes = props.total_memory
+ used_bytes: Optional[int]
# Prefer mem_get_info (system-wide) so auto-select sees other consumers.
if hasattr(mod, "mem_get_info"):
free_bytes, total_bytes = mod.mem_get_info(ordinal)
used_bytes = total_bytes - free_bytes
+ # free==total is the broken-API sentinel, not an idle GPU.
+ if _win_rocm and free_bytes == total_bytes:
+ used_bytes = None
else:
used_bytes = mod.memory_allocated(ordinal)
devices.append(
@@ -561,7 +571,7 @@ def _torch_get_per_device_info(device_indices: list[int]) -> list[Dict[str, Any]
"visible_ordinal": ordinal,
"name": props.name,
"total_gb": round(total_bytes / (1024**3), 2),
- "used_gb": round(used_bytes / (1024**3), 2),
+ "used_gb": round(used_bytes / (1024**3), 2) if used_bytes is not None else None,
}
)
except Exception as e:
@@ -724,20 +734,30 @@ def _rocm_linux_sysfs_vram_gb() -> tuple[Optional[float], Optional[float]]:
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.
+# ── Windows AMD/ROCm per-adapter VRAM (issue #7072) ──────────────────────────
+# amd-smi is disabled and hipMemGetInfo reports free==total, so read used from the
+# per-LUID "GPU Adapter Memory" perf counters and take each total from torch, so
+# every GPU shows instead of one fake device with GPU 0's total.
+# Placeholder adapters (Basic Render Driver / idle iGPU) drop only when they would
+# outnumber the real torch devices.
+_ROCM_WIN_ADAPTER_MIN_BYTES = 64 * 1024 * 1024 # 64 MiB
- Same data source as Task Manager, so cross-process usage is accurate.
- Works for any GPU vendor without amd-smi or nvidia-smi.
- Returns (used_gb, total_gb) or (None, None) on failure.
+
+def _rocm_windows_perf_counter_vram_by_adapter() -> Optional[list[tuple[str, float]]]:
+ """Per-adapter dedicated VRAM usage on Windows via Performance Counters.
+
+ Returns ``[(instance_name, used_bytes)]`` (one per LUID-named adapter), or
+ ``None`` when the counter is unavailable/localized/empty so callers fall back.
"""
if platform.system() != "Windows":
- return None, None
+ return None
try:
+ # Emit "|" per sample, or a __NONE__ sentinel.
ps = (
"$s=(Get-Counter '\\GPU Adapter Memory(*)\\Dedicated Usage'"
" -ErrorAction SilentlyContinue).CounterSamples;"
- "if($s){($s|Measure-Object CookedValue -Sum).Sum}else{-1}"
+ "if($s){$s|ForEach-Object{'{0}|{1}' -f $_.InstanceName,[int64]$_.CookedValue}}"
+ "else{'__NONE__'}"
)
r = subprocess.run(
["powershell", "-NoProfile", "-NonInteractive", "-Command", ps],
@@ -746,16 +766,167 @@ def _rocm_windows_perf_counter_vram_gb() -> tuple[Optional[float], Optional[floa
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)
+ return None
+ adapters: list[tuple[str, float]] = []
+ for line in r.stdout.splitlines():
+ line = line.strip()
+ if not line or line == "__NONE__" or "|" not in line:
+ continue
+ instance, _, raw = line.rpartition("|")
+ try:
+ used = float(raw.strip())
+ except (ValueError, TypeError):
+ continue
+ if used < 0:
+ continue
+ adapters.append((instance.strip(), used))
+ return adapters or None
except Exception:
- return None, None
+ return None
+
+
+def _match_adapter_used_to_devices(
+ adapter_useds: list[float], device_totals: list[float]
+) -> list[Optional[float]]:
+ """Attribute per-adapter used bytes to torch devices by capacity ranking.
+
+ Windows shares no key between LUID counters and torch ordinals, so usages are
+ ranked against device totals and each is trusted only when capacity *forces* it
+ (it exceeds every smaller device); an ambiguous ranking reports unknown
+ (``None``) rather than fabricate a per-index free.
+
+ Extra counters mean a hidden/display adapter, and the noise filter may have
+ dropped a real reading, so values are emitted only when the supra-threshold
+ counters number EXACTLY the visible devices AND capacity forces the mapping;
+ otherwise every device is unknown. Best-effort but correct for the common
+ loaded-card case (#7072). Returns a list aligned to ``device_totals``.
+ """
+ n = len(device_totals)
+ if n == 0:
+ return []
+ useds = sorted(adapter_useds, reverse = True)
+ ranked_positions = sorted(range(n), key = lambda i: -device_totals[i])
+ ranked_totals = [device_totals[pos] for pos in ranked_positions]
+ assigned: list[Optional[float]]
+ # More counters than devices -> a hidden/display adapter (check before noise filter).
+ if len(useds) > n:
+ non_trivial = [u for u in useds if u >= _ROCM_WIN_ADAPTER_MIN_BYTES]
+ if len(non_trivial) != n:
+ # Not a clean bijection (a masked GPU is busy or a visible card idle):
+ # no counter maps to a specific card, so report unknown.
+ return [None] * n
+ # Exactly n supra-threshold counters: extras were placeholders, so a
+ # capacity-ranked bijection is plausible.
+ useds = non_trivial
+ ranked_useds = [useds[rank] for rank in range(n)]
+ # A usage above its ranked capacity is a hidden larger GPU; clamping onto the
+ # smaller card would fabricate a fully-used reading.
+ for rank in range(n):
+ if ranked_useds[rank] > ranked_totals[rank]:
+ return [None] * n
+ # Capacity forces the mapping only when the usage exceeds the next-smaller
+ # capacity; the smallest card and merely-fitting usages stay unknown.
+ # Keeps 40 GiB over 48/8 GiB -> [40, None].
+ assigned = [None] * n
+ for rank, pos in enumerate(ranked_positions):
+ if rank + 1 < n and ranked_useds[rank] > ranked_totals[rank + 1]:
+ assigned[pos] = min(ranked_useds[rank], device_totals[pos])
+ return assigned
+ # No hidden adapters: every counter is a visible card, so ranking is a permutation.
+ ranked_useds = [useds[rank] if rank < len(useds) else 0.0 for rank in range(n)]
+ # Ambiguous if a strictly larger usage also fits the next smaller card: the two
+ # could be swapped without breaking capacity, so ranking can't tell them apart.
+ for rank in range(n - 1):
+ upper, lower = ranked_useds[rank], ranked_useds[rank + 1]
+ if upper > lower and upper <= ranked_totals[rank + 1]:
+ return [None] * n
+ assigned = [None] * n
+ for rank, pos in enumerate(ranked_positions):
+ if rank < len(useds):
+ assigned[pos] = min(useds[rank], device_totals[pos])
+ return assigned
+
+
+def _rocm_windows_per_device_vram(device_indices: list[int]) -> list[Dict[str, Any]]:
+ """Per-GPU VRAM on Windows AMD/ROCm: total from torch properties (reliable),
+ used from the per-adapter Dedicated Usage counter.
+
+ Returns ``{index, visible_ordinal, name, used_gb, total_gb}`` per visible GPU
+ (``used_gb`` may be ``None`` when the counter is unavailable), or ``[]`` when
+ torch can't enumerate devices so callers fall through to the torch last resort.
+ """
+ if platform.system() != "Windows":
+ return []
+ mod, _ = _torch_get_device_module()
+ if mod is None:
+ return []
+ # Totals/names from torch properties (mem_get_info's free==total quirk zeroes used).
+ dev_meta: list[Dict[str, Any]] = []
+ for ordinal, phys_idx in enumerate(device_indices):
+ try:
+ props = mod.get_device_properties(ordinal)
+ dev_meta.append(
+ {
+ "index": phys_idx,
+ "visible_ordinal": ordinal,
+ "name": props.name,
+ "total_bytes": int(props.total_memory),
+ }
+ )
+ except Exception as e:
+ logger.debug("torch property probe failed for ordinal %d: %s", ordinal, e)
+ if not dev_meta:
+ return []
+
+ adapters = _rocm_windows_perf_counter_vram_by_adapter()
+ if adapters:
+ assigned = _match_adapter_used_to_devices(
+ [used for _, used in adapters],
+ [d["total_bytes"] for d in dev_meta],
+ )
+ else:
+ # Counter unavailable: show every GPU with a correct total, used unknown.
+ assigned = [None] * len(dev_meta)
+
+ devices: list[Dict[str, Any]] = []
+ for meta, used_bytes in zip(dev_meta, assigned):
+ total_gb = round(meta["total_bytes"] / (1024**3), 2)
+ used_gb = round(used_bytes / (1024**3), 2) if used_bytes is not None else None
+ devices.append(
+ {
+ "index": meta["index"],
+ "visible_ordinal": meta["visible_ordinal"],
+ "name": meta["name"],
+ "used_gb": used_gb,
+ "total_gb": total_gb,
+ }
+ )
+ return devices
+
+
+def _rocm_windows_device_payload_entry(
+ device: DeviceType, dev: Dict[str, Any], gpu_util_pct: Optional[float]
+) -> Dict[str, Any]:
+ """Build a ``get_gpu_utilization`` device entry from a per-device VRAM dict."""
+ total_gb = dev["total_gb"]
+ used_gb = dev["used_gb"]
+ return {
+ "available": True,
+ "backend": _backend_label(device),
+ "index": dev["index"],
+ "visible_ordinal": dev["visible_ordinal"],
+ "name": dev.get("name", "Unknown"),
+ "gpu_utilization_pct": gpu_util_pct,
+ "temperature_c": None,
+ "vram_used_gb": used_gb,
+ "vram_total_gb": total_gb,
+ "vram_utilization_pct": round((used_gb / total_gb) * 100, 1)
+ if total_gb and total_gb > 0 and used_gb is not None
+ else None,
+ "power_draw_w": None,
+ "power_limit_w": None,
+ "power_utilization_pct": None,
+ }
def _gpu_utilization_payload(
@@ -821,30 +992,24 @@ def get_gpu_utilization() -> Dict[str, Any]:
index_kind = result.get("index_kind"),
)
- # Fallback Windows ROCm
+ # Fallback Windows ROCm: per-adapter VRAM attribution (issue #7072), so
+ # every visible GPU is shown instead of a sum collapsed onto one device.
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()
+ _win_ids = _get_parent_visible_gpu_spec().get("numeric_ids")
+ if not _win_ids:
+ _win_ids = list(range(_torch_get_physical_gpu_count() or 0))
+ _win_devices = _rocm_windows_per_device_vram(_win_ids)
+ if _win_devices:
+ # A single visible GPU can own the aggregate 3D-engine utilization;
+ # across several GPUs the sum isn't per-device, so leave it unset.
+ _win_util = (
+ _rocm_windows_perf_counter_gpu_util_pct() if len(_win_devices) == 1 else None
+ )
return _gpu_utilization_payload(
device,
[
- {
- "available": True,
- "backend": _backend_label(device),
- "index": 0,
- "visible_ordinal": 0,
- "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,
- }
+ _rocm_windows_device_payload_entry(device, _wd, _win_util)
+ for _wd in _win_devices
],
)
@@ -901,7 +1066,7 @@ def get_gpu_utilization() -> Dict[str, Any]:
"vram_used_gb": _used,
"vram_total_gb": _total,
"vram_utilization_pct": round((_used / _total) * 100, 1)
- if _total > 0
+ if _total > 0 and _used is not None
else None,
"power_draw_w": None,
"power_limit_w": None,
@@ -995,19 +1160,27 @@ def _apply_unified_memory_correction(
endpoints stay in sync on AMD iGPUs with unified memory.
"""
torch_total_gb = torch_info["total_gb"]
+ torch_used_gb = torch_info.get("used_gb")
smi_total_gb = device_metrics.get("vram_total_gb") or 0.0
+ # torch sees the full unified (GTT) pool; amd-smi only the dedicated carve-out.
+ # Adopt torch's larger total regardless of used: on Windows ROCm torch_used is
+ # None (free==total sentinel) but its total stays authoritative. Overwrite used
+ # only when torch's is known, then recompute utilization against whatever remains.
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
+ if torch_used_gb is not None:
+ device_metrics["vram_used_gb"] = torch_used_gb
+ _used_for_pct = device_metrics.get("vram_used_gb")
device_metrics["vram_utilization_pct"] = (
- round((torch_used_gb / torch_total_gb) * 100, 1) if torch_total_gb > 0 else None
+ round((_used_for_pct / torch_total_gb) * 100, 1)
+ if torch_total_gb > 0 and _used_for_pct is not None
+ 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,
+ "ROCm unified memory: adopted torch mem_get_info total (%.2f GB) over "
+ "amd-smi (%.2f GB) for device %s",
torch_total_gb,
+ smi_total_gb,
torch_info.get("index"),
)
@@ -1067,6 +1240,49 @@ def get_visible_gpu_utilization() -> Dict[str, Any]:
_reconcile_rocm_unified_memory(result, numeric_ids)
return result
+ # Windows AMD/ROCm (issue #7072): the System tab's VRAM source. The torch
+ # fallback below would report used==0 (free==total), so read per-adapter
+ # Dedicated Usage instead; total from torch properties.
+ if IS_ROCM and platform.system() == "Windows":
+ win_numeric_ids = parent_visible_spec.get("numeric_ids")
+ if win_numeric_ids:
+ win_ids = win_numeric_ids
+ win_index_kind = "physical"
+ else:
+ win_ids = list(range(_torch_get_physical_gpu_count() or 0))
+ win_index_kind = "relative"
+ win_devices = _rocm_windows_per_device_vram(win_ids)
+ if win_devices:
+ devices = []
+ for wd in win_devices:
+ total = wd["total_gb"]
+ used = wd["used_gb"]
+ devices.append(
+ {
+ "index": wd["index"],
+ "index_kind": win_index_kind,
+ "visible_ordinal": wd["visible_ordinal"],
+ "name": wd.get("name"),
+ "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 and total > 0 and used is not None
+ else None,
+ "power_draw_w": None,
+ "power_limit_w": None,
+ "power_utilization_pct": None,
+ }
+ )
+ return {
+ "available": True,
+ "backend": _backend_label(device),
+ "parent_visible_gpu_ids": win_numeric_ids or [],
+ "devices": devices,
+ "index_kind": win_index_kind,
+ }
+
# Torch-based fallback for CUDA (nvidia-smi unavailable, AMD ROCm) and XPU (Intel)
if device in (DeviceType.CUDA, DeviceType.XPU):
parent_ids = get_parent_visible_gpu_ids()
@@ -1094,7 +1310,7 @@ def get_visible_gpu_utilization() -> Dict[str, Any]:
"vram_used_gb": used,
"vram_total_gb": total,
"vram_utilization_pct": round((used / total) * 100, 1)
- if total > 0
+ if total > 0 and used is not None
else None,
"power_draw_w": None,
"power_limit_w": None,
diff --git a/studio/frontend/src/components/floating-monitor.tsx b/studio/frontend/src/components/floating-monitor.tsx
index 1272a577e9..e38e2e5882 100644
--- a/studio/frontend/src/components/floating-monitor.tsx
+++ b/studio/frontend/src/components/floating-monitor.tsx
@@ -70,13 +70,18 @@ export function FloatingMonitor() {
(sum, device) => sum + (device.memory_total_gb ?? 0),
0,
);
- const vramUsed = devices.reduce(
- (sum, device) => sum + (device.vram_used_gb ?? 0),
- 0,
- );
+ // null usage = unknown (e.g. Windows ROCm perf counter): treating it as 0
+ // fabricates a 0-used readout, so the aggregate is unknown if any device is.
+ const vramUsageKnown =
+ devices.length > 0 &&
+ devices.every((device) => Number.isFinite(device.vram_used_gb));
+ const vramUsed = vramUsageKnown
+ ? devices.reduce((sum, device) => sum + (device.vram_used_gb ?? 0), 0)
+ : 0;
const vramPercent = clampPercent(
- vramTotal > 0 ? (vramUsed / vramTotal) * 100 : 0,
+ vramUsageKnown && vramTotal > 0 ? (vramUsed / vramTotal) * 100 : 0,
);
+ const unknownLabel = t("settings.resources.environment.unknown");
const hasGpu = (systemInfo.gpu?.available ?? false) && devices.length > 0;
@@ -164,17 +169,20 @@ export function FloatingMonitor() {
- {Math.round(vramPercent)}%
+ {vramUsageKnown ? `${Math.round(vramPercent)}%` : "--"}