Merge branch 'unslothai:main' into fix/rocm-strix-halo-unified-memory

This commit is contained in:
Leo Borcherding 2026-05-06 13:01:36 -07:00 committed by GitHub
commit 4e7a083271
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
78 changed files with 7543 additions and 793 deletions

200
.github/workflows/studio-backend-ci.yml vendored Normal file
View file

@ -0,0 +1,200 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
# Runs the existing studio/backend/tests/ suite (~860 tests, all CPU-friendly)
# on every PR that touches the backend or unsloth library. Until this lands,
# none of those tests run automatically. Verified locally on Python 3.13 with
# the surgical exclusions below: 861 pass, 4 skipped.
#
# Exclusions:
# - tests/test_studio_api.py: end-to-end against a live model + GGUF download,
# too heavy for free runners. Run separately when GPU CI is available.
# - -k 'not llama_cpp_load_progress_live': spawns a real llama.cpp process,
# not appropriate for CPU-only runners.
#
# ruff is non-blocking initially; remove `|| true` once the backend lints clean.
name: Backend CI
on:
pull_request:
paths:
- 'studio/**'
- 'unsloth/**'
- 'unsloth_cli/**'
- 'tests/**'
- 'pyproject.toml'
- '.github/workflows/studio-backend-ci.yml'
push:
branches: [main, pip]
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
cancel-in-progress: true
jobs:
pytest:
name: (Python ${{ matrix.python }})
runs-on: ubuntu-latest
timeout-minutes: 15
strategy:
fail-fast: false
matrix:
python: ['3.10', '3.11', '3.12', '3.13']
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: '${{ matrix.python }}'
cache: 'pip'
- name: Install backend test dependencies (CPU only)
run: |
python -m pip install --upgrade pip
# Studio's declared backend deps:
pip install -r studio/backend/requirements/studio.txt
# Extras that studio.txt does not list but the import chain needs
# (python-multipart for FastAPI form/file uploads, sqlalchemy/cryptography
# for the auth DB, yaml/jinja2 for utils.models.model_config, etc.):
pip install \
python-multipart aiofiles sqlalchemy cryptography \
pyyaml jinja2 mammoth unpdf requests \
'numpy<3' pytest pytest-asyncio httpx
# Torch CPU + transformers are required by a chunk of the backend test
# suite (gpu_selection, kv_cache_estimation, utils). CPU-only torch
# keeps the install ~250 MB / ~1 min on a clean runner.
pip install --index-url https://download.pytorch.org/whl/cpu 'torch>=2.4,<2.11'
pip install 'transformers>=4.51,<5.5'
- name: Backend tests
working-directory: studio/backend
# Locally validated against this dep set: 831 passed, 5 skipped, 35 deselected.
# Deselections (all environment-specific, would never pass on a GPU-less
# `ubuntu-latest` runner regardless of code correctness):
# - llama_cpp_load_progress_live: spawns a real llama.cpp process
# - TestGpuAutoSelection / TestPreSpawnGpuResolution / TestPerGpuFitGuardAllCounts:
# require live transformers config introspection on real GPUs
# - TestTransformersIntrospection: same
# - test_returns_cuda_when_cuda_available / test_calls_cuda_cache_when_cuda:
# assume CUDA-capable GPU
run: |
python -m pytest tests/ -q --tb=short \
--ignore=tests/test_studio_api.py \
-k 'not llama_cpp_load_progress_live and not TestGpuAutoSelection and not TestPreSpawnGpuResolution and not TestPerGpuFitGuardAllCounts and not TestTransformersIntrospection and not test_returns_cuda_when_cuda_available and not test_calls_cuda_cache_when_cuda'
repo-cpu-tests:
# Auto-discover everything under tests/ that is not GPU-bound by
# design. New tests added in covered directories are picked up
# without a workflow edit. Locally validated: 779 passed, 11
# skipped, 23 deselected. tests/conftest.py (mirroring unsloth-zoo
# PR #624) pre-loads unsloth_zoo.device_type and unsloth.device_type
# under a mocked torch.cuda.is_available so the unsloth import
# chain succeeds on CPU.
name: Repo tests (CPU)
runs-on: ubuntu-latest
timeout-minutes: 10
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: '3.12'
cache: 'pip'
- name: Install deps (shared shape with backend pytest job)
run: |
python -m pip install --upgrade pip
pip install -r studio/backend/requirements/studio.txt
pip install \
python-multipart aiofiles sqlalchemy cryptography \
pyyaml jinja2 mammoth unpdf requests typer \
'numpy<3' pytest pytest-asyncio httpx
# torchvision is needed because unsloth_zoo.vision_utils imports
# it at module scope and is reached via unsloth.models._utils.
pip install --index-url https://download.pytorch.org/whl/cpu \
'torch>=2.4,<2.11' 'torchvision<0.26'
pip install 'transformers>=4.51,<5.5'
# bitsandbytes is a hard import in unsloth/models/_utils.py.
# Recent versions ship a CPU build so it installs on a free
# Linux runner; the kernels still raise on use, but import
# succeeds and the package collects.
pip install 'bitsandbytes>=0.45'
# unsloth.device_type imports unsloth_zoo.utils.Version at module
# scope, so the conftest harness needs unsloth_zoo on the path
# even though it is an optional dep of unsloth.
pip install 'unsloth_zoo>=2026.5.1'
pip install -e . --no-deps
- name: Repo tests (CPU, auto-discovered)
env:
# tests/python/* import install_python_stack from studio/.
PYTHONPATH: ${{ github.workspace }}/studio
# Skip lazy compilation work the unsloth import chain wants to
# do at import time on a real GPU.
UNSLOTH_COMPILE_DISABLE: '1'
# --ignore: GPU-bound directories (qlora and saving need real
# weights / GPU; tests/sh is a shell suite the next step
# handles; tests/utils is a helpers folder, not tests).
# State-sensitive hardware-spoofing files are pulled out and run
# in isolation in the next step because they mutate
# hardware.py module globals (IS_ROCM / DEVICE) and pollute
# downstream tests.
# -m: honour markers already declared in tests/python/conftest.py
# (`server` = needs studio venv, `e2e` = needs network).
# --deselect: two registry tests that hit huggingface_hub for
# live model existence checks; they belong on a network job.
run: |
python -m pytest tests/ -q --tb=short \
--ignore=tests/qlora \
--ignore=tests/saving \
--ignore=tests/utils \
--ignore=tests/sh \
--ignore=tests/studio/test_hardware_dispatch_matrix.py \
--ignore=tests/studio/test_is_mlx_dispatch_gate.py \
-m 'not server and not e2e' \
--deselect tests/test_model_registry.py::test_model_registration \
--deselect tests/test_model_registry.py::test_all_model_registration
- name: Hardware-spoof tests (state-sensitive, run in isolation)
env:
PYTHONPATH: ${{ github.workspace }}/studio
UNSLOTH_COMPILE_DISABLE: '1'
# These two files mutate hardware.py module globals at runtime
# via the spoof fixtures, which leaks state into any other test
# that imports hardware. Run them in their own pytest invocation
# so the leak does not cross file boundaries.
run: |
python -m pytest -q --tb=short \
tests/studio/test_hardware_dispatch_matrix.py \
tests/studio/test_is_mlx_dispatch_gate.py
- name: Shell installer tests
# Subset that does not depend on a writable / pristine install.sh
# tree; test_install_host_defaults.sh checks install.ps1 layout
# which has drifted (separate followup).
run: |
set -e
for s in \
tests/sh/test_get_torch_index_url.sh \
tests/sh/test_mac_intel_compat.sh \
tests/sh/test_tauri_install_exit_order.sh \
tests/sh/test_torch_constraint.sh; do
echo "::group::$s"
bash "$s"
echo "::endgroup::"
done
ruff:
name: Backend ruff lint (non-blocking)
runs-on: ubuntu-latest
timeout-minutes: 5
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: '3.12'
cache: 'pip'
- run: pip install ruff
- name: ruff check (non-blocking until accumulated drift is cleared)
run: ruff check studio/backend || true

108
.github/workflows/studio-frontend-ci.yml vendored Normal file
View file

@ -0,0 +1,108 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
# Frontend PR gate: lockfile freshness, typecheck, build, and a bundle grep
# that catches the 2026.5.1 chat-history regression at the JS level.
#
# biome runs as non-blocking for now: the codebase currently has accumulated
# ~470 errors and ~1650 warnings against the existing biome config. Surfacing
# the count in CI lets us drive it down without forcing a fleet-wide cleanup
# in the same PR. Drop `continue-on-error` once that number is zero.
name: Frontend CI
on:
pull_request:
paths:
- 'studio/frontend/**'
- '.github/workflows/studio-frontend-ci.yml'
push:
branches: [main, pip]
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
cancel-in-progress: true
jobs:
build:
name: Frontend build + bundle sanity
runs-on: ubuntu-latest
timeout-minutes: 10
defaults:
run:
working-directory: studio/frontend
steps:
- uses: actions/checkout@v4
# FIXME: drop this step once @assistant-ui/* and assistant-stream
# leave 0.x -- on 1.x, caret ranges are conventional. Until then,
# every 0.minor on this surface is a SemVer-major (this is exactly
# how 2026.5.1 shipped a broken chat runtime: ^0.12.19 quietly
# resolved to 0.12.28).
- name: '@assistant-ui must be pinned exactly (no caret/tilde)'
working-directory: ${{ github.workspace }}
run: |
set -e
if grep -nE '"(@assistant-ui/[a-z-]+|assistant-stream)":[[:space:]]*"[\^~]' studio/frontend/package.json; then
echo "::error file=studio/frontend/package.json::These packages must be pinned to exact versions until they leave 0.x. Drop the leading ^ or ~."
exit 1
fi
echo "All assistant-ui packages are pinned exactly."
- uses: actions/setup-node@v4
with:
node-version: '22'
cache: 'npm'
cache-dependency-path: studio/frontend/package-lock.json
- name: Lockfile must agree with package.json (npm ci is strict)
run: npm ci --no-fund --no-audit
- name: npm ci must not have modified the working tree
working-directory: ${{ github.workspace }}
run: |
if ! git diff --quiet -- studio/frontend; then
echo "::error::npm ci modified files; commit the updated lockfile"
git status -- studio/frontend
exit 1
fi
- name: Typecheck
run: npm run typecheck
- name: Build
run: npm run build
- name: Built bundle must not contain Studio's unstable_Provider call site
run: |
set -e
JS=$(ls dist/assets/index-*.js | head -1)
HITS=$(grep -c 'unstable_Provider:' "$JS" || echo 0)
echo "main bundle: $JS"
echo "unstable_Provider: hits=$HITS (assistant-ui internals contribute up to 3)"
if [ "$HITS" -gt 3 ]; then
echo "::error file=studio/frontend/src/features/chat/runtime-provider.tsx::Studio bundle still passes unstable_Provider through useRemoteThreadListRuntime; this is the 2026.5.1 chat-history regression. Pass adapters directly into useLocalRuntime instead."
exit 1
fi
- name: Bundle size budget (75 MB)
run: |
SIZE=$(du -sb dist | cut -f1)
BUDGET=$((75 * 1024 * 1024))
echo "dist size: $SIZE bytes ($((SIZE/1024/1024)) MB), budget: $BUDGET bytes (75 MB)"
if [ "$SIZE" -gt "$BUDGET" ]; then
echo "::error::studio/frontend/dist/ exceeded the 75 MB budget. Drop dead deps (e.g. the unused next dep) or split chunks."
exit 1
fi
- name: Biome (non-blocking until accumulated drift is cleared)
continue-on-error: true
run: npm run biome:check
- name: Upload built dist on failure
if: failure()
uses: actions/upload-artifact@v4
with:
name: studio-frontend-dist
path: studio/frontend/dist
retention-days: 3

View file

@ -0,0 +1,185 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
# End-to-end smoke: install Studio via install.sh --local --no-torch, download
# a tiny GGUF, boot Studio, log in, change password, load the model, send a
# chat completion, assert a non-empty response. Only workflow that tests "the
# app actually works".
#
# Model: Qwen3.5-2B UD-IQ3_XXS (~890 MiB) -- small enough that the cache miss
# is cheap and inference fits in the 25 min CPU-runner budget. GGUF is cached
# across runs via actions/cache.
name: Studio GGUF CI
on:
pull_request:
paths:
- 'studio/**'
- 'unsloth/**'
- 'unsloth_cli/**'
- 'install.sh'
- 'pyproject.toml'
- '.github/workflows/studio-inference-smoke.yml'
push:
branches: [main, pip]
# Manual trigger for pre-warming the GGUF cache on main, or re-running
# against an arbitrary branch without pushing a no-op commit.
workflow_dispatch:
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
cancel-in-progress: true
env:
GGUF_REPO: unsloth/Qwen3.5-2B-GGUF
GGUF_FILE: Qwen3.5-2B-UD-IQ3_XXS.gguf
STUDIO_PORT: '18888'
jobs:
inference:
name: Studio boots, loads a GGUF, answers a chat completion
runs-on: ubuntu-latest
timeout-minutes: 25
steps:
- uses: actions/checkout@v4
- name: Linux dependencies for llama.cpp prebuilt
run: |
sudo apt-get update
sudo apt-get install -y --no-install-recommends \
libcurl4-openssl-dev libssl-dev jq
- uses: actions/setup-node@v4
with:
node-version: '22'
cache: 'npm'
cache-dependency-path: studio/frontend/package-lock.json
- uses: actions/setup-python@v5
with:
python-version: '3.12'
cache: 'pip'
- name: Cache GGUF model file
id: cache-gguf
uses: actions/cache@v4
with:
path: gguf-cache
key: ${{ runner.os }}-gguf-${{ env.GGUF_REPO }}-${{ env.GGUF_FILE }}-v1
- name: Download GGUF if cache miss
if: steps.cache-gguf.outputs.cache-hit != 'true'
run: |
# huggingface-cli was deprecated in huggingface_hub 1.13; the new CLI is `hf`.
python -m pip install --upgrade huggingface_hub hf_transfer
mkdir -p gguf-cache
HF_HUB_ENABLE_HF_TRANSFER=1 \
hf download "$GGUF_REPO" "$GGUF_FILE" --local-dir gguf-cache
- name: Install Studio (--local, --no-torch keeps the install lean)
run: |
mkdir -p logs
set -o pipefail
bash install.sh --local --no-torch 2>&1 | tee logs/install.log
- name: Assert llama.cpp prebuilt was installed (no source-build fallback)
# ubuntu-latest is CPU-only x86_64, so studio/setup.sh should route
# to ggml-org/llama.cpp and grab bin-ubuntu-x64.tar.gz. A source
# build here means the routing regressed.
run: |
if grep -q "falling back to source build" logs/install.log; then
echo "::error::llama.cpp prebuilt path failed on ubuntu-latest. studio/setup.sh routing regressed; CPU-only Linux x86_64 should hit ggml-org/llama.cpp's bin-ubuntu-x64.tar.gz."
grep -E "llama-prebuilt|llama.cpp" logs/install.log | tail -60
exit 1
fi
if ! grep -qE "prebuilt installed and validated|prebuilt up to date and validated" logs/install.log; then
echo "::error::install.log does not contain the success marker for the llama.cpp prebuilt path. Did setup.sh skip the prebuilt install?"
grep -E "llama-prebuilt|llama.cpp" logs/install.log | tail -60
exit 1
fi
echo "llama.cpp prebuilt path used successfully"
- name: Reset auth + start Studio in the background
run: |
unsloth studio reset-password
mkdir -p logs
UNSLOTH_API_ONLY=1 unsloth studio -H 127.0.0.1 -p "$STUDIO_PORT" \
> logs/studio.log 2>&1 &
echo "STUDIO_PID=$!" >> "$GITHUB_ENV"
- name: Wait for /api/health
run: |
for i in $(seq 1 60); do
if curl -fs "http://127.0.0.1:${STUDIO_PORT}/api/health" > /tmp/health.json; then
echo "ready after ${i}s"
cat /tmp/health.json
jq -e '.status == "healthy"' /tmp/health.json
exit 0
fi
sleep 1
done
echo "Studio did not become healthy in 60s"
tail -200 logs/studio.log
exit 1
- name: Login + change bootstrap password
run: |
PW=$(cat ~/.unsloth/studio/auth/.bootstrap_password)
NEW="CIPasswordSmoke12345!"
TOKEN=$(curl -fs -X POST "http://127.0.0.1:${STUDIO_PORT}/api/auth/login" \
-H 'content-type: application/json' \
-d "{\"username\":\"unsloth\",\"password\":\"$PW\"}" | jq -r .access_token)
curl -fs -X POST "http://127.0.0.1:${STUDIO_PORT}/api/auth/change-password" \
-H "Authorization: Bearer $TOKEN" -H 'content-type: application/json' \
-d "{\"current_password\":\"$PW\",\"new_password\":\"$NEW\"}" > /dev/null
# Re-login to clear must_change_password flag.
NEW_TOKEN=$(curl -fs -X POST "http://127.0.0.1:${STUDIO_PORT}/api/auth/login" \
-H 'content-type: application/json' \
-d "{\"username\":\"unsloth\",\"password\":\"$NEW\"}" | jq -r .access_token)
echo "TOKEN=$NEW_TOKEN" >> "$GITHUB_ENV"
- name: Load the GGUF into Studio
run: |
GGUF_PATH="$GITHUB_WORKSPACE/gguf-cache/${GGUF_FILE}"
ls -lh "$GGUF_PATH"
curl -fs -X POST "http://127.0.0.1:${STUDIO_PORT}/api/inference/load" \
-H "Authorization: Bearer $TOKEN" -H 'content-type: application/json' \
--max-time 600 \
-d "{\"model_path\":\"$GGUF_PATH\",\"is_lora\":false,\"max_seq_length\":2048}" \
| jq '{status, display_name, is_gguf, context_length}'
- name: Send a chat completion + assert non-empty response
run: |
RESP=$(curl -fs -X POST "http://127.0.0.1:${STUDIO_PORT}/api/inference/chat/completions" \
-H "Authorization: Bearer $TOKEN" -H 'content-type: application/json' \
--max-time 900 \
-d '{
"messages":[{"role":"user","content":"Say hello in one short sentence."}],
"max_tokens":40,
"stream":false
}')
echo "raw response: $RESP"
CONTENT=$(echo "$RESP" | jq -r '.choices[0].message.content // empty')
echo "model response: $CONTENT"
if [ -z "$CONTENT" ]; then
echo "::error::Empty assistant response from Studio"
exit 1
fi
- name: Stop Studio
if: always()
run: |
kill "${STUDIO_PID}" || true
sleep 2
ss -tln | grep ":${STUDIO_PORT}" || true
- name: Upload Studio + install logs on failure
if: failure()
uses: actions/upload-artifact@v4
with:
name: studio-inference-log
path: |
logs/studio.log
logs/install.log
retention-days: 7

105
.github/workflows/studio-tauri-smoke.yml vendored Normal file
View file

@ -0,0 +1,105 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
# PR-time smoke for the Tauri desktop wrapper. Builds the frontend and the
# Tauri Linux debug binary, with no codesigning. Catches:
# - tauri.conf.json drift
# - src-tauri Cargo.toml or rust source breakage
# - Tauri CLI version drift (we pin 2.10.1, matching release-desktop.yml)
# - frontend output not picked up by Tauri's distDir
#
# Linux-only on a free `ubuntu-latest` runner. Mac and Windows desktop builds
# stay in release-desktop.yml (manual `workflow_dispatch`) because they need
# code-signing secrets and ~30 min of runner time each.
name: Studio Tauri CI
on:
pull_request:
paths:
- 'studio/frontend/**'
- 'studio/src-tauri/**'
- '.github/workflows/studio-tauri-smoke.yml'
push:
branches: [main, pip]
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
cancel-in-progress: true
jobs:
linux-debug-build:
name: Tauri Linux debug build (no codesign)
runs-on: ubuntu-22.04
timeout-minutes: 25
steps:
- uses: actions/checkout@v4
- name: Linux native deps for Tauri / WebKit2GTK
run: |
sudo apt-get update
sudo apt-get install -y \
libwebkit2gtk-4.1-dev libayatana-appindicator3-dev \
librsvg2-dev libxdo-dev libssl-dev patchelf
- uses: actions/setup-node@v4
with:
node-version: '24'
cache: 'npm'
cache-dependency-path: studio/frontend/package-lock.json
- uses: dtolnay/rust-toolchain@stable
- uses: swatinem/rust-cache@v2
with:
workspaces: studio/src-tauri -> target
- name: Install pinned Tauri CLI (matches release-desktop.yml)
run: npm install --save-dev --prefix studio @tauri-apps/cli@2.10.1
- name: Verify pinned Tauri CLI version
run: |
out="$(npx --prefix studio tauri --version)"
echo "$out"
[ "$out" = "tauri-cli 2.10.1" ] || { echo "::error::expected tauri-cli 2.10.1, got $out"; exit 1; }
- name: Frontend build (npm ci, vite)
working-directory: studio/frontend
run: |
npm ci --no-fund --no-audit
npm run build
test -f dist/index.html
- name: Tauri debug build (Linux, no bundle, no codesign)
# `--debug` + `--no-bundle` keeps this lean: compiles the Rust crate,
# confirms the frontend dist is wired into Tauri, but skips the AppImage
# / .deb production. Code signing is irrelevant because we never produce
# a distributable artifact.
env:
TAURI_SIGNING_PRIVATE_KEY: ''
TAURI_SIGNING_PRIVATE_KEY_PASSWORD: ''
run: npx --prefix studio tauri build --debug --no-bundle
- name: Inspect produced binary
run: |
BIN=$(find studio/src-tauri/target/debug -maxdepth 1 -type f -executable 2>/dev/null \
| grep -Ev '\.(d|so|dylib|dll)$' \
| grep -Ev '/(deps|build|examples)$' \
| head -1)
echo "binary: $BIN"
if [ -z "$BIN" ]; then
echo "::error::Tauri debug binary not produced"
ls -la studio/src-tauri/target/debug/ || true
exit 1
fi
file "$BIN"
du -h "$BIN"
- uses: actions/upload-artifact@v4
if: failure()
with:
name: tauri-debug-build
path: |
studio/src-tauri/target/debug
studio/frontend/dist
retention-days: 3

124
.github/workflows/wheel-smoke.yml vendored Normal file
View file

@ -0,0 +1,124 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
# Builds the PyPI wheel from the PR branch, then verifies the built wheel
# actually contains what we expect to ship and does NOT contain the broken
# Studio bundle that 2026.5.1 published. This is the single workflow that
# would have blocked the 2026.5.1 release before twine upload.
#
# Verified locally end-to-end against this branch:
# - python -m build produces unsloth-<version>-py3-none-any.whl in 13s
# - wheel content sanity passes:
# lockfile shipped, frontend dist shipped,
# no node_modules in wheel, no bun.lock in wheel,
# main bundle has unstable_Provider hits=1 (assistant-ui internals only).
# - Studio backend imports cleanly from the installed wheel with the
# lightweight dep set below.
name: Wheel CI
on:
pull_request:
paths:
- 'pyproject.toml'
- 'studio/**'
- 'unsloth/**'
- 'unsloth_cli/**'
- '.github/workflows/wheel-smoke.yml'
push:
branches: [main, pip]
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
cancel-in-progress: true
jobs:
wheel:
name: Wheel build + content sanity + import smoke
runs-on: ubuntu-latest
timeout-minutes: 15
steps:
- uses: actions/checkout@v4
- uses: actions/setup-node@v4
with:
node-version: '22'
cache: 'npm'
cache-dependency-path: studio/frontend/package-lock.json
- uses: actions/setup-python@v5
with:
python-version: '3.12'
- name: Build frontend
run: |
cd studio/frontend
npm ci --no-fund --no-audit
npm run build
- name: Build wheel + sdist
run: |
python -m pip install --upgrade pip build
rm -rf dist build ./*.egg-info
python -m build
- name: Wheel content sanity
run: |
python - <<'PY'
import zipfile, glob, sys
w = glob.glob("dist/unsloth-*.whl")
if not w:
print("FAIL: no wheel produced"); sys.exit(2)
w = w[0]
print(f"wheel: {w}")
with zipfile.ZipFile(w) as z:
n = z.namelist()
checks = {
"lockfile shipped": any(s.endswith("studio/frontend/package-lock.json") for s in n),
"frontend dist shipped": any(s.endswith("studio/frontend/dist/index.html") for s in n),
"no node_modules": not any("studio/frontend/node_modules/" in s for s in n),
"no bun.lock": not any(s.endswith("studio/frontend/bun.lock") for s in n),
}
js = [s for s in n
if "studio/frontend/dist/assets/" in s
and s.endswith(".js")
and "/index-" in s]
if not js:
print("FAIL: no main bundle index-*.js in wheel"); sys.exit(2)
data = z.read(js[0]).decode("utf-8", "replace")
hits = data.count("unstable_Provider:")
print(f"main bundle: {js[0]}")
print(f"unstable_Provider hits: {hits} (>=4 indicates 2026.5.1 regression)")
checks["bundle has no Studio unstable_Provider call site"] = (hits < 4)
print()
for k, v in checks.items():
print(f" [{'PASS' if v else 'FAIL'}] {k}")
sys.exit(0 if all(checks.values()) else 1)
PY
- name: Studio backend import smoke
# Imports `studio.backend.main:app` from the freshly-installed wheel in
# a clean venv. This catches the class of bug that 2026.5.1 shipped with:
# frontend dist missing, package-lock.json missing, or the wheel's Python
# source tree broken in a way that surfaces only at app construction time.
run: |
python -m venv /tmp/v
/tmp/v/bin/pip install --upgrade pip
/tmp/v/bin/pip install -r studio/backend/requirements/studio.txt
/tmp/v/bin/pip install \
python-multipart aiofiles sqlalchemy cryptography \
pyyaml jinja2 mammoth unpdf requests \
'numpy<3'
/tmp/v/bin/pip install --no-deps dist/unsloth-*.whl
# Run from /tmp so Python imports the installed package, not the source tree.
cd /tmp
/tmp/v/bin/python -c "from studio.backend.main import app; print('Studio backend OK:', app.title)"
- name: Upload wheel on failure
if: failure()
uses: actions/upload-artifact@v4
with:
name: unsloth-wheel
path: dist/
retention-days: 7

5
.gitignore vendored
View file

@ -24,8 +24,8 @@ dist/
downloads/
eggs/
.eggs/
lib/
lib64/
/lib/
/lib64/
parts/
sdist/
var/
@ -228,3 +228,4 @@ setup_leo.sh
server.pid
*.log
package-lock.json
llama.cpp/

View file

@ -3,6 +3,11 @@
# Local: Set-ExecutionPolicy -Scope Process -ExecutionPolicy Bypass; .\install.ps1 --local
# NoTorch: .\install.ps1 --no-torch (skip PyTorch, GGUF-only mode)
# Test: .\install.ps1 --package roland-sloth
#
# Env vars (priority: UNSLOTH_STUDIO_HOME > STUDIO_HOME > USERPROFILE-redirect > default):
# UNSLOTH_STUDIO_HOME / STUDIO_HOME = path -> install under that path
# (DataDir nests inside; user PATH not modified persistently).
# Default ($USERPROFILE\.unsloth\studio) is preserved when no env var is set.
function Install-UnslothStudio {
$ErrorActionPreference = "Stop"
@ -126,7 +131,94 @@ function Install-UnslothStudio {
}
$PythonVersion = "3.13"
$StudioHome = Join-Path $env:USERPROFILE ".unsloth\studio"
# Resolve install destinations. Priority: UNSLOTH_STUDIO_HOME, then
# STUDIO_HOME alias, then USERPROFILE-redirect, then default.
# Reject whitespace-only values so " " is treated as unset (matches the
# Python resolvers' .strip()), preventing install/runtime layout drift.
$envOverrideVar = $null
$envOverride = $null
if (-not [string]::IsNullOrWhiteSpace($env:UNSLOTH_STUDIO_HOME)) {
$envOverrideVar = "UNSLOTH_STUDIO_HOME"
$envOverride = $env:UNSLOTH_STUDIO_HOME.Trim()
} elseif (-not [string]::IsNullOrWhiteSpace($env:STUDIO_HOME)) {
$envOverrideVar = "STUDIO_HOME"
$envOverride = $env:STUDIO_HOME.Trim()
}
# Custom Studio roots are not supported with --tauri (desktop app still
# resolves %USERPROFILE%\.unsloth\studio). Pass through if override == legacy.
if ($TauriMode -and $envOverride) {
$_tauriOverride = $envOverride
if ($_tauriOverride -eq "~" -or $_tauriOverride -like "~/*" -or $_tauriOverride -like "~\*") {
$_tauriOverride = (Join-Path $env:USERPROFILE $_tauriOverride.Substring(1).TrimStart('/','\'))
}
try {
$_tauriOverride = [System.IO.Path]::GetFullPath($_tauriOverride)
} catch {}
$_legacyTauriRoot = Join-Path $env:USERPROFILE ".unsloth\studio"
try {
$_legacyTauriRoot = [System.IO.Path]::GetFullPath($_legacyTauriRoot)
} catch {}
# Strip trailing separators so ".../studio\" matches ".../studio".
$_trimSeps = @(
[System.IO.Path]::DirectorySeparatorChar,
[System.IO.Path]::AltDirectorySeparatorChar
)
$_tauriOverride = $_tauriOverride.TrimEnd($_trimSeps)
$_legacyTauriRoot = $_legacyTauriRoot.TrimEnd($_trimSeps)
if ($_tauriOverride -ne $_legacyTauriRoot) {
Write-Host "ERROR: $envOverrideVar is not supported with --tauri." -ForegroundColor Red
Write-Host " The desktop app still uses the legacy %USERPROFILE%\.unsloth\studio root." -ForegroundColor Red
Write-Host " Run install.ps1 without --tauri for custom-root shell installs," -ForegroundColor Yellow
Write-Host " or unset the env var for default desktop installs." -ForegroundColor Yellow
throw "$envOverrideVar is not supported with --tauri."
}
}
$defaultProfile = $null
try { $defaultProfile = [Environment]::GetFolderPath("UserProfile") } catch {}
# LOCALAPPDATA may be unset in service / CI contexts; Join-Path would abort
# under ErrorActionPreference=Stop without this guard.
$defaultDataDir = if ($env:LOCALAPPDATA -and -not [string]::IsNullOrWhiteSpace($env:LOCALAPPDATA)) {
Join-Path $env:LOCALAPPDATA "Unsloth Studio"
} else { $null }
if ($envOverride) {
# Tilde expansion: env vars aren't subject to it when quoted on assignment.
if ($envOverride -eq "~" -or $envOverride -like "~/*" -or $envOverride -like "~\*") {
$envOverride = (Join-Path $env:USERPROFILE $envOverride.Substring(1).TrimStart('/','\'))
}
try {
# .NET API: New-Item -Path treats brackets as wildcards and has no
# -LiteralPath in PS 5.1, so a root like C:\studio[abc] would fail.
[System.IO.Directory]::CreateDirectory($envOverride) | Out-Null
$StudioHome = (Resolve-Path -LiteralPath $envOverride).Path
} catch {
Write-Host "ERROR: $envOverrideVar=$envOverride cannot be created or accessed." -ForegroundColor Red
throw "$envOverrideVar=$envOverride cannot be created or accessed."
}
$probe = Join-Path $StudioHome (".unsloth-write-probe-" + [guid]::NewGuid())
try {
# WriteAllText: literal-path safe + closes handle so Remove-Item works.
[System.IO.File]::WriteAllText($probe, "")
Remove-Item -LiteralPath $probe -Force -ErrorAction SilentlyContinue
} catch {
Write-Host "ERROR: $envOverrideVar=$StudioHome is not writable." -ForegroundColor Red
throw "$envOverrideVar=$StudioHome is not writable."
}
$StudioDataDir = Join-Path $StudioHome "share"
$StudioRedirectMode = 'env'
} elseif ($defaultProfile -and $env:USERPROFILE -and ($env:USERPROFILE -ne $defaultProfile)) {
$StudioHome = Join-Path $env:USERPROFILE ".unsloth\studio"
$StudioDataDir = $defaultDataDir
$StudioRedirectMode = 'profile'
} else {
$StudioHome = Join-Path $env:USERPROFILE ".unsloth\studio"
$StudioDataDir = $defaultDataDir
$StudioRedirectMode = 'default'
}
$VenvDir = Join-Path $StudioHome "unsloth_studio"
$Rule = [string]::new([char]0x2500, 52)
@ -378,24 +470,24 @@ function Install-UnslothStudio {
[Parameter(Mandatory = $true)][string]$UnslothExePath
)
if (-not (Test-Path $UnslothExePath)) {
if (-not (Test-Path -LiteralPath $UnslothExePath)) {
substep "cannot create shortcuts, unsloth.exe not found at $UnslothExePath" "Yellow"
return
}
try {
# Persist an absolute path in launcher scripts so shortcut working
# directory changes do not break process startup.
$UnslothExePath = (Resolve-Path $UnslothExePath).Path
$UnslothExePath = (Resolve-Path -LiteralPath $UnslothExePath).Path
# Escape for single-quoted embedding in generated launcher script.
# This prevents runtime variable expansion for paths containing '$'.
$SingleQuotedExePath = $UnslothExePath -replace "'", "''"
$localAppDataDir = $env:LOCALAPPDATA
if (-not $localAppDataDir -or [string]::IsNullOrWhiteSpace($localAppDataDir)) {
substep "LOCALAPPDATA path unavailable; skipped shortcut creation" "Yellow"
# $StudioDataDir = LOCALAPPDATA\Unsloth Studio, or $StudioHome\share in env-mode.
if (-not $StudioDataDir -or [string]::IsNullOrWhiteSpace($StudioDataDir)) {
substep "DataDir path unavailable; skipped shortcut creation" "Yellow"
return
}
$appDir = Join-Path $localAppDataDir "Unsloth Studio"
$appDir = $StudioDataDir
$launcherPs1 = Join-Path $appDir "launch-studio.ps1"
$launcherVbs = Join-Path $appDir "launch-studio.vbs"
$desktopDir = [Environment]::GetFolderPath("Desktop")
@ -427,23 +519,89 @@ function Install-UnslothStudio {
}
$iconUrl = "https://raw.githubusercontent.com/unslothai/unsloth/main/studio/frontend/public/unsloth.ico"
if (-not (Test-Path $appDir)) {
New-Item -ItemType Directory -Path $appDir -Force | Out-Null
if (-not (Test-Path -LiteralPath $appDir)) {
[System.IO.Directory]::CreateDirectory($appDir) | Out-Null
}
# Same-install discriminator: per-install opaque id written once at
# install time and read by both this launcher and the backend
# (/api/health). Replaces the older sha256(resolved $StudioHome)
# scheme to (a) avoid leaking the install path on -H 0.0.0.0
# deployments and (b) sidestep launcher/backend canonicalization
# drift (Resolve-Path vs Path.resolve() junction handling). Lives
# at $StudioHome\share\ (not $appDir) so the backend can find it
# via _STUDIO_ROOT_RESOLVED / "share" / "studio_install_id"
# regardless of mode. 32 bytes of crypto random -> 64 hex chars.
$_studioIdDir = Join-Path $StudioHome "share"
if (-not (Test-Path -LiteralPath $_studioIdDir)) {
[System.IO.Directory]::CreateDirectory($_studioIdDir) | Out-Null
}
$_studioIdFile = Join-Path $_studioIdDir "studio_install_id"
$_studioRootId = ""
if ((Test-Path -LiteralPath $_studioIdFile) -and `
((Get-Item -LiteralPath $_studioIdFile).Length -gt 0)) {
$_studioRootId = ([System.IO.File]::ReadAllText($_studioIdFile)).Trim()
}
if (-not $_studioRootId) {
$_idBytes = New-Object byte[] 32
[Security.Cryptography.RandomNumberGenerator]::Create().GetBytes($_idBytes)
$_studioRootId = -join ($_idBytes | ForEach-Object { $_.ToString('x2') })
# Atomic write: write to a temp sibling then rename, so a partial
# install cannot leave a half-written id.
$_idTmp = $_studioIdFile + ".$PID.tmp"
[System.IO.File]::WriteAllText($_idTmp, $_studioRootId)
Move-Item -LiteralPath $_idTmp -Destination $_studioIdFile -Force
}
# Env-mode: persist UNSLOTH_STUDIO_HOME (and llama path) so fresh
# shells don't need to re-export, and bake per-install $portFile /
# $mutexName so concurrent custom-root launchers cannot serialize
# through one global mutex on 8888..8908. Default installs get an
# empty prefix to match pre-PR behavior.
$studioHomeExport = if ($StudioRedirectMode -eq 'env') {
# When override == legacy default, llama.cpp stays at
# ~/.unsloth/llama.cpp (one shared build). Canonicalize the
# legacy side so the comparison survives path normalization.
$_legacyStudio = Join-Path $env:USERPROFILE ".unsloth\studio"
if (Test-Path -LiteralPath $_legacyStudio -PathType Container) {
$_legacyStudio = (Resolve-Path -LiteralPath $_legacyStudio).Path
}
$_llamaPath = if ($StudioHome -eq $_legacyStudio) {
Join-Path $env:USERPROFILE ".unsloth\llama.cpp"
} else {
Join-Path $StudioHome "llama.cpp"
}
$_sq = $StudioHome -replace "'", "''"
$_llama = $_llamaPath -replace "'", "''"
$_appDirSq = $appDir -replace "'", "''"
$_appBytes = [Text.Encoding]::UTF8.GetBytes($appDir)
$_appHash = ([BitConverter]::ToString(
[Security.Cryptography.SHA256]::Create().ComputeHash($_appBytes)
) -replace '-', '').Substring(0, 16)
# UNSLOTH_LLAMA_CPP_PATH is a pre-existing user override; only default if unset.
"`$env:UNSLOTH_STUDIO_HOME = '$_sq'`nif (-not `$env:UNSLOTH_LLAMA_CPP_PATH) {`n `$env:UNSLOTH_LLAMA_CPP_PATH = '$_llama'`n}`n`$portFile = '$_appDirSq\studio.port'`n`$mutexName = 'Local\UnslothStudioLauncher-$_appHash'`n"
} else {
"`$portFile = `$null`n`$mutexName = 'Local\UnslothStudioLauncher'`n"
}
$launcherContent = @"
`$ErrorActionPreference = 'Stop'
$studioHomeExport`$ErrorActionPreference = 'Stop'
`$basePort = 8888
`$maxPortOffset = 20
`$timeoutSec = 60
`$pollIntervalMs = 1000
`$_ExpectedStudioRootId = '$_studioRootId'
function Test-StudioHealth {
param([Parameter(Mandatory = `$true)][int]`$Port)
try {
`$url = "http://127.0.0.1:`$Port/api/health"
`$resp = Invoke-RestMethod -Uri `$url -TimeoutSec 1 -Method Get
return (`$resp -and `$resp.status -eq 'healthy' -and `$resp.service -eq 'Unsloth UI Backend')
if (-not (`$resp -and `$resp.status -eq 'healthy' -and `$resp.service -eq 'Unsloth UI Backend')) { return `$false }
# why: verify the backend belongs to THIS install via the install-time
# hex digest; raw path is not leaked over /api/health.
if (`$_ExpectedStudioRootId -and `$resp.studio_root_id -ne `$_ExpectedStudioRootId) { return `$false }
return `$true
} catch {
return `$false
}
@ -469,6 +627,17 @@ function Get-CandidatePorts {
}
function Find-HealthyStudioPort {
if (`$portFile) {
if (Test-Path -LiteralPath `$portFile) {
`$cached = Get-Content -LiteralPath `$portFile -ErrorAction SilentlyContinue | Select-Object -First 1
if (`$cached -match '^\d+`$') {
`$cachedPort = [int]`$cached
if (Test-StudioHealth -Port `$cachedPort) { return `$cachedPort }
Remove-Item -LiteralPath `$portFile -Force -ErrorAction SilentlyContinue
}
}
return `$null
}
foreach (`$candidate in (Get-CandidatePorts)) {
if (Test-StudioHealth -Port `$candidate) {
return `$candidate
@ -522,7 +691,7 @@ if (`$existingPort) {
exit 0
}
`$launchMutex = [System.Threading.Mutex]::new(`$false, 'Local\UnslothStudioLauncher')
`$launchMutex = [System.Threading.Mutex]::new(`$false, `$mutexName)
`$haveMutex = `$false
try {
try {
@ -552,7 +721,9 @@ try {
} catch {}
exit 1
}
`$studioCommand = '& "' + `$studioExe + '" studio -p ' + `$launchPort
# Single-quote the path in the child -Command so `$` / backtick in custom
# roots don't get reparsed; double any apostrophes so 'O''Brien' survives.
`$studioCommand = "& '" + (`$studioExe -replace "'", "''") + "' studio -p " + `$launchPort
`$launchArgs = @(
'-NoExit',
'-NoProfile',
@ -576,9 +747,13 @@ try {
`$browserOpened = `$false
`$deadline = (Get-Date).AddSeconds(`$timeoutSec)
while ((Get-Date) -lt `$deadline) {
`$healthyPort = Find-HealthyStudioPort
if (`$healthyPort) {
Start-Process "http://localhost:`$healthyPort"
if (Test-StudioHealth -Port `$launchPort) {
if (`$portFile) {
try {
[System.IO.File]::WriteAllText(`$portFile, "`$launchPort`n")
} catch {}
}
Start-Process "http://localhost:`$launchPort"
`$browserOpened = `$true
break
}
@ -613,19 +788,19 @@ cmd = "powershell -NoProfile -ExecutionPolicy Bypass -WindowStyle Hidden -File "
shell.Run cmd, 0, False
"@
# WSH handles UTF-16LE reliably for .vbs files with non-ASCII paths.
Set-Content -Path $launcherVbs -Value $vbsContent -Encoding Unicode -Force
Set-Content -LiteralPath $launcherVbs -Value $vbsContent -Encoding Unicode -Force
# Prefer bundled icon from local clone/dev installs.
# If not available, best-effort download from raw GitHub.
# We only attach the icon if the resulting file has a valid ICO header.
$hasValidIcon = $false
if ($bundledIcon -and (Test-Path $bundledIcon)) {
if ($bundledIcon -and (Test-Path -LiteralPath $bundledIcon)) {
try {
Copy-Item -Path $bundledIcon -Destination $iconPath -Force
Copy-Item -LiteralPath $bundledIcon -Destination $iconPath -Force
} catch {
Write-Host "[DEBUG] Error copying bundled icon: $($_.Exception.Message)" -ForegroundColor DarkGray
}
} elseif (-not (Test-Path $iconPath)) {
} elseif (-not (Test-Path -LiteralPath $iconPath)) {
try {
Invoke-WebRequest -Uri $iconUrl -OutFile $iconPath -UseBasicParsing
} catch {
@ -633,7 +808,7 @@ shell.Run cmd, 0, False
}
}
if (Test-Path $iconPath) {
if (Test-Path -LiteralPath $iconPath) {
try {
$bytes = [System.IO.File]::ReadAllBytes($iconPath)
if (
@ -645,14 +820,21 @@ shell.Run cmd, 0, False
) {
$hasValidIcon = $true
} else {
Remove-Item $iconPath -Force -ErrorAction SilentlyContinue
Remove-Item -LiteralPath $iconPath -Force -ErrorAction SilentlyContinue
}
} catch {
Write-Host "[DEBUG] Error validating or removing icon: $($_.Exception.Message)" -ForegroundColor DarkGray
Remove-Item $iconPath -Force -ErrorAction SilentlyContinue
Remove-Item -LiteralPath $iconPath -Force -ErrorAction SilentlyContinue
}
}
# Env-mode: skip persistent Desktop / Start Menu .lnk shortcuts
# that may point at a deleted workspace; launcher + icon stay.
if ($StudioRedirectMode -eq 'env') {
substep "wrote launcher at $launcherPs1 (persistent shortcuts skipped in env-override mode)"
return
}
$wscriptExe = Join-Path $env:SystemRoot "System32\wscript.exe"
$shortcutArgs = "//B //Nologo `"$launcherVbs`""
@ -850,8 +1032,9 @@ shell.Run cmd, 0, False
# Pass the resolved executable path to uv so it does not re-resolve
# a version string back to a conda interpreter.
Write-TauriLog "STEP" "Creating virtual environment"
if (-not (Test-Path $StudioHome)) {
New-Item -ItemType Directory -Path $StudioHome -Force | Out-Null
if (-not (Test-Path -LiteralPath $StudioHome)) {
# .NET API: New-Item -Path treats brackets as wildcards.
[System.IO.Directory]::CreateDirectory($StudioHome) | Out-Null
}
$VenvPython = Join-Path $VenvDir "Scripts\python.exe"
@ -865,11 +1048,13 @@ shell.Run cmd, 0, False
$stamp = Get-Date -Format "yyyyMMddHHmmss"
$candidate = Join-Path $StudioHome "unsloth_studio.rollback.$stamp.$PID"
$suffix = 0
while (Test-Path $candidate) {
# -LiteralPath: a custom $StudioHome may contain [ ] * ? which
# plain Test-Path / Move-Item would interpret as wildcards.
while (Test-Path -LiteralPath $candidate) {
$suffix++
$candidate = Join-Path $StudioHome "unsloth_studio.rollback.$stamp.$PID.$suffix"
}
Move-Item -Path $ExistingDir -Destination $candidate -ErrorAction Stop
Move-Item -LiteralPath $ExistingDir -Destination $candidate -ErrorAction Stop
$script:StudioVenvRollbackDir = $candidate
$script:StudioVenvRollbackTarget = $ExistingDir
$script:StudioVenvRollbackActive = $true
@ -880,16 +1065,16 @@ shell.Run cmd, 0, False
if (-not $script:StudioVenvRollbackActive) { return }
$backup = $script:StudioVenvRollbackDir
$target = $script:StudioVenvRollbackTarget
if (-not $backup -or -not (Test-Path $backup)) {
if (-not $backup -or -not (Test-Path -LiteralPath $backup)) {
$script:StudioVenvRollbackActive = $false
return
}
substep "restoring previous environment after failed install..." "Yellow"
try {
if (Test-Path $target) {
Remove-Item -Recurse -Force $target -ErrorAction SilentlyContinue
if (Test-Path -LiteralPath $target) {
Remove-Item -LiteralPath $target -Recurse -Force -ErrorAction SilentlyContinue
}
Move-Item -Path $backup -Destination $target -Force -ErrorAction Stop
Move-Item -LiteralPath $backup -Destination $target -Force -ErrorAction Stop
substep "restored previous environment"
$script:StudioVenvRollbackActive = $false
$script:StudioVenvRollbackDir = $null
@ -902,14 +1087,29 @@ shell.Run cmd, 0, False
function Complete-StudioVenvRollback {
if (-not $script:StudioVenvRollbackActive) { return }
$backup = $script:StudioVenvRollbackDir
if ($backup -and (Test-Path $backup)) {
Remove-Item -Recurse -Force $backup -ErrorAction SilentlyContinue
if ($backup -and (Test-Path -LiteralPath $backup)) {
Remove-Item -LiteralPath $backup -Recurse -Force -ErrorAction SilentlyContinue
}
$script:StudioVenvRollbackActive = $false
$script:StudioVenvRollbackDir = $null
}
if (Test-Path $VenvPython) {
if (Test-Path -LiteralPath $VenvPython) {
# why: matching guard to the .venv branch below -- in env-mode
# $StudioHome is a user-chosen workspace, so refuse to nuke an
# existing $StudioHome\unsloth_studio that lacks Studio sentinels.
# -PathType Leaf rejects a directory at the sentinel path. Accept the
# in-VENV ownership marker so partial-install retries are not blocked.
if (
$StudioRedirectMode -eq 'env' -and
-not (Test-Path -LiteralPath (Join-Path $VenvDir ".unsloth-studio-owned") -PathType Leaf) -and
-not (Test-Path -LiteralPath (Join-Path $StudioHome "share\studio.conf") -PathType Leaf) -and
-not (Test-Path -LiteralPath (Join-Path $StudioHome "bin\unsloth.exe") -PathType Leaf)
) {
Write-Host "[ERROR] $VenvDir already exists but does not look like an Unsloth Studio install." -ForegroundColor Red
Write-Host " Move it aside or choose an empty UNSLOTH_STUDIO_HOME." -ForegroundColor Yellow
throw "Refusing to delete non-Studio venv at $VenvDir"
}
# New layout already exists -- replace only after preserving rollback copy.
substep "preserving existing environment for rollback..."
try {
@ -918,8 +1118,13 @@ shell.Run cmd, 0, False
Write-Host "[ERROR] Could not prepare existing environment for reinstall: $($_.Exception.Message)" -ForegroundColor Red
return (Exit-InstallFailure "Could not prepare existing environment for reinstall")
}
} elseif (Test-Path (Join-Path $StudioHome ".venv\Scripts\python.exe")) {
# Old layout (~/.unsloth/studio/.venv) exists -- validate before migrating
} elseif (
$StudioRedirectMode -ne 'env' `
-and (Test-Path -LiteralPath (Join-Path $StudioHome ".venv\Scripts\python.exe"))
) {
# Old layout (~/.unsloth/studio/.venv) exists -- validate before migrating.
# Skip in env-mode so we don't blow away an unrelated .venv at the
# workspace root (e.g. user's existing project Python venv).
$OldVenv = Join-Path $StudioHome ".venv"
$OldPy = Join-Path $OldVenv "Scripts\python.exe"
substep "found legacy Studio environment, validating..."
@ -936,24 +1141,29 @@ shell.Run cmd, 0, False
$ErrorActionPreference = $prevEAP2
if ($legacyOk) {
substep "legacy environment is healthy -- migrating..."
Move-Item -Path $OldVenv -Destination $VenvDir -Force
Move-Item -LiteralPath $OldVenv -Destination $VenvDir -Force
substep "moved .venv -> unsloth_studio"
$_Migrated = $true
} else {
substep "legacy environment failed validation -- creating fresh environment" "Yellow"
$invalidVenv = Join-Path $StudioHome (".venv.invalid.{0}.{1}" -f (Get-Date -Format "yyyyMMddHHmmss"), $PID)
Move-Item -Path $OldVenv -Destination $invalidVenv -Force -ErrorAction SilentlyContinue
Move-Item -LiteralPath $OldVenv -Destination $invalidVenv -Force -ErrorAction SilentlyContinue
}
} elseif (Test-Path (Join-Path $env:USERPROFILE "unsloth_studio\Scripts\python.exe")) {
# CWD-relative venv from old install.ps1 -- migrate to absolute path
} elseif (
$StudioRedirectMode -ne 'env' `
-and (Test-Path -LiteralPath (Join-Path $env:USERPROFILE "unsloth_studio\Scripts\python.exe"))
) {
# CWD-relative venv from old install.ps1 -> migrate to absolute path.
# Skip in env-mode so we don't relocate the default-install venv into
# the workspace root.
$CwdVenv = Join-Path $env:USERPROFILE "unsloth_studio"
substep "found CWD-relative Studio environment, migrating to $VenvDir..."
Move-Item -Path $CwdVenv -Destination $VenvDir -Force
Move-Item -LiteralPath $CwdVenv -Destination $VenvDir -Force
substep "moved ~/unsloth_studio -> ~/.unsloth/studio/unsloth_studio"
$_Migrated = $true
}
if (-not (Test-Path $VenvPython)) {
if (-not (Test-Path -LiteralPath $VenvPython)) {
step "venv" "creating Python $($DetectedPython.Version) virtual environment"
substep "$VenvDir"
$venvExit = Invoke-InstallCommand { uv venv $VenvDir --python "$($DetectedPython.Path)" }
@ -966,6 +1176,13 @@ shell.Run cmd, 0, False
substep "$VenvDir"
}
# Mark the freshly-created venv as Studio-owned so a partial install can be
# repaired by re-running install.ps1; the env-mode deletion guard above
# accepts this marker as the primary sentinel.
if (Test-Path -LiteralPath $VenvDir -PathType Container) {
try { [System.IO.File]::WriteAllText((Join-Path $VenvDir ".unsloth-studio-owned"), "") } catch {}
}
# ── Detect GPU (robust: PATH + hardcoded fallback paths, mirrors setup.ps1) ──
$HasNvidiaSmi = $false
$NvidiaSmiExe = $null
@ -1054,7 +1271,7 @@ shell.Run cmd, 0, False
if ($StudioLocalInstall -and (Test-Path (Join-Path $RepoRoot "studio\backend\requirements\no-torch-runtime.txt"))) {
return Join-Path $RepoRoot "studio\backend\requirements\no-torch-runtime.txt"
}
$installed = Get-ChildItem -Path $VenvDir -Recurse -Filter "no-torch-runtime.txt" -ErrorAction SilentlyContinue |
$installed = Get-ChildItem -LiteralPath $VenvDir -Recurse -Filter "no-torch-runtime.txt" -ErrorAction SilentlyContinue |
Where-Object { $_.FullName -like "*studio*backend*requirements*no-torch-runtime.txt" } |
Select-Object -ExpandProperty FullName -First 1
return $installed
@ -1192,23 +1409,25 @@ shell.Run cmd, 0, False
foreach ($rel in $overlayMap.Keys) {
$src = Join-Path $scriptDir $rel
$dst = Join-Path $VenvDir $overlayMap[$rel]
if (-not (Test-Path $src)) { continue }
# -LiteralPath: $VenvDir derives from $StudioHome which may
# contain [ ] * ? when the user overrode UNSLOTH_STUDIO_HOME.
if (-not (Test-Path -LiteralPath $src)) { continue }
$dstParent = Split-Path -Parent $dst
if (-not (Test-Path $dstParent)) {
if (-not (Test-Path -LiteralPath $dstParent)) {
Write-Host "[WARN] Overlay target dir missing: $dstParent; studio setup may use stale bundled file" -ForegroundColor Yellow
continue
}
try {
if (-not (Test-Path $dst)) {
if (-not (Test-Path -LiteralPath $dst)) {
# Backfill: target file missing but parent dir exists.
Copy-Item $src $dst -Force
Copy-Item -LiteralPath $src -Destination $dst -Force
substep ("backfilled bundled " + (Split-Path -Leaf $rel))
} else {
# Hash-compare so re-runs are no-ops when files already match.
$srcHash = (Get-FileHash $src -Algorithm SHA256).Hash
$dstHash = (Get-FileHash $dst -Algorithm SHA256).Hash
$srcHash = (Get-FileHash -LiteralPath $src -Algorithm SHA256).Hash
$dstHash = (Get-FileHash -LiteralPath $dst -Algorithm SHA256).Hash
if ($srcHash -ne $dstHash) {
Copy-Item $src $dst -Force
Copy-Item -LiteralPath $src -Destination $dst -Force
substep ("applied bundled " + (Split-Path -Leaf $rel))
}
}
@ -1225,7 +1444,8 @@ shell.Run cmd, 0, False
Write-TauriLog "STEP" "Running studio setup"
step "setup" "running unsloth studio setup..."
$UnslothExe = Join-Path $VenvDir "Scripts\unsloth.exe"
if (-not (Test-Path $UnslothExe)) {
if (-not (Test-Path -LiteralPath $UnslothExe)) {
Write-TauriLog "ERROR" "unsloth CLI was not installed correctly"
Write-Host "[ERROR] unsloth CLI was not installed correctly." -ForegroundColor Red
Write-Host " Expected: $UnslothExe" -ForegroundColor Yellow
Write-Host " This usually means an older unsloth version was installed that does not include the Studio CLI." -ForegroundColor Yellow
@ -1250,6 +1470,15 @@ shell.Run cmd, 0, False
# Use 'studio setup' (not 'studio update') because 'update' pops
# SKIP_STUDIO_BASE, which would cause redundant package reinstallation
# and bypass the fast-path version check from PR #4667.
# Propagate UNSLOTH_STUDIO_HOME only for env-override installs; otherwise
# an inherited value would put llama.cpp in the wrong place.
$previousUnslothStudioHome = $env:UNSLOTH_STUDIO_HOME
$hadPreviousUnslothStudioHome = ($null -ne $previousUnslothStudioHome)
if ($StudioRedirectMode -eq 'env') {
$env:UNSLOTH_STUDIO_HOME = $StudioHome
} else {
Remove-Item Env:UNSLOTH_STUDIO_HOME -ErrorAction SilentlyContinue
}
$studioArgs = @('studio', 'setup')
if ($script:UnslothVerbose) { $studioArgs += '--verbose' }
$env:UNSLOTH_INSTALL_ROLLBACK_MANAGED = "1"
@ -1257,6 +1486,11 @@ shell.Run cmd, 0, False
& $UnslothExe @studioArgs
$setupExit = $LASTEXITCODE
} finally {
if ($hadPreviousUnslothStudioHome) {
$env:UNSLOTH_STUDIO_HOME = $previousUnslothStudioHome
} else {
Remove-Item Env:UNSLOTH_STUDIO_HOME -ErrorAction SilentlyContinue
}
Remove-Item Env:UNSLOTH_INSTALL_ROLLBACK_MANAGED -ErrorAction SilentlyContinue
}
if ($setupExit -ne 0) {
@ -1301,20 +1535,32 @@ shell.Run cmd, 0, False
}
} catch { }
$ShimDir = Join-Path $StudioHome "bin"
New-Item -ItemType Directory -Force -Path $ShimDir | Out-Null
[System.IO.Directory]::CreateDirectory($ShimDir) | Out-Null
$ShimExe = Join-Path $ShimDir "unsloth.exe"
# Fatal preflight outside the lock-handling try/catch -- a directory at
# the shim path must not be downgraded to "Continuing with the existing
# launcher", or the install finishes with no usable shim.
if (Test-Path -LiteralPath $ShimExe -PathType Container) {
Write-Host "[ERROR] Cannot create unsloth launcher: $ShimExe is a directory." -ForegroundColor Red
Write-Host " Move or remove it manually, then re-run the installer." -ForegroundColor Yellow
throw "Cannot create unsloth launcher: $ShimExe is a directory."
}
# try/catch: if unsloth.exe is locked (Studio running), keep the old shim.
$shimUpdated = $false
try {
if (Test-Path $ShimExe) { Remove-Item $ShimExe -Force -ErrorAction Stop }
if (Test-Path -LiteralPath $ShimExe) { Remove-Item -LiteralPath $ShimExe -Force -ErrorAction Stop }
try {
# New-Item -ItemType HardLink does NOT accept -LiteralPath in any
# PowerShell version, so use -Path. Wildcards in $ShimExe (e.g.
# brackets in custom roots) glob-expand here and fall through to
# the Copy-Item -LiteralPath fallback below.
New-Item -ItemType HardLink -Path $ShimExe -Target $UnslothExe -ErrorAction Stop | Out-Null
} catch {
Copy-Item -Path $UnslothExe -Destination $ShimExe -Force -ErrorAction Stop # fallback: copy
Copy-Item -LiteralPath $UnslothExe -Destination $ShimExe -Force -ErrorAction Stop # fallback: copy
}
$shimUpdated = $true
} catch {
if (Test-Path $ShimExe) {
if (Test-Path -LiteralPath $ShimExe) {
Write-Host "[WARN] Could not refresh unsloth launcher at $ShimExe." -ForegroundColor Yellow
Write-Host " This usually means a running 'unsloth studio' process still holds the file open." -ForegroundColor Yellow
Write-Host " Close Studio and re-run the installer to pick up the latest launcher." -ForegroundColor Yellow
@ -1325,10 +1571,13 @@ shell.Run cmd, 0, False
Write-Host " Launch unsloth studio directly via '$UnslothExe' until the next successful install." -ForegroundColor Yellow
}
}
# Only add to PATH when the launcher actually exists on disk.
# Add to PATH only when launcher exists. Env-mode: session-only export,
# no registry change (workspace path may be deleted later).
$pathAdded = $false
if (Test-Path $ShimExe) {
$pathAdded = Add-ToUserPath -Directory $ShimDir -Position 'Prepend'
if (Test-Path -LiteralPath $ShimExe) {
if ($StudioRedirectMode -ne 'env') {
$pathAdded = Add-ToUserPath -Directory $ShimDir -Position 'Prepend'
}
}
if ($shimUpdated -and $pathAdded) {
step "path" "added unsloth launcher to PATH"
@ -1336,12 +1585,20 @@ shell.Run cmd, 0, False
Refresh-SessionPath # sync current session with registry
Complete-StudioVenvRollback
# Env-mode session export AFTER Refresh-SessionPath; otherwise a legacy
# User PATH entry (Machine > User > current $env:Path) would win.
if ($StudioRedirectMode -eq 'env' -and (Test-Path -LiteralPath $ShimExe)) {
$env:Path = "$ShimDir;$env:Path"
step "path" "exported $ShimDir for this session (no registry PATH change in env-override mode)"
}
# ── Tauri mode: done, skip shortcuts and auto-launch ──
if ($TauriMode) {
Write-TauriLog "DONE" ""
return
}
# New-StudioShortcuts gates the .lnk shortcuts on env-mode internally.
New-StudioShortcuts -UnslothExePath $UnslothExe
# In interactive terminals, ask the user before starting Studio.
@ -1360,8 +1617,21 @@ shell.Run cmd, 0, False
}
} else {
step "launch" "manual commands:"
substep "& `"$VenvDir\Scripts\Activate.ps1`""
substep "unsloth studio -p 8888"
# Single-quote the printed paths so $-vars / backticks in custom roots
# do not reparse when the user pastes the command.
$_actLiteral = "'" + ((Join-Path $VenvDir "Scripts\Activate.ps1") -replace "'", "''") + "'"
if ($StudioRedirectMode -eq 'env') {
# Env-mode skips registry PATH; print the absolute shim path.
$_shim = Join-Path $StudioHome "bin\unsloth.exe"
$_shimLiteral = "'" + ($_shim -replace "'", "''") + "'"
substep "& $_shimLiteral studio -p 8888"
substep "or activate env first:"
substep "& $_actLiteral"
substep "unsloth studio -p 8888"
} else {
substep "& $_actLiteral"
substep "unsloth studio -p 8888"
}
substep "(add -H 0.0.0.0 to allow network / cloud access)"
Write-Host ""
}

View file

@ -6,6 +6,12 @@
# Usage (no-torch): ./install.sh --no-torch (skip PyTorch, GGUF-only mode)
# Usage (test): ./install.sh --package roland-sloth (install a different package name)
# Usage (py): ./install.sh --python 3.12 (override auto-detected Python version)
#
# Env vars (priority: UNSLOTH_STUDIO_HOME > STUDIO_HOME > HOME-redirect > default):
# UNSLOTH_STUDIO_HOME=/abs/path -> install under that path
# STUDIO_HOME=/abs/path -> alias, same effect (UNSLOTH_STUDIO_HOME wins)
# (DATA_DIR + unsloth CLI shim nest inside; no shell rc-file append.)
# Default ($HOME/.unsloth/studio) is preserved when no env var is set.
set -e
# ── Output style (aligned with studio/setup.sh) ──
@ -66,6 +72,56 @@ if [ "$_VERBOSE" = true ]; then
export UNSLOTH_VERBOSE=1
fi
# Custom Studio roots are not supported with --tauri (desktop app still
# resolves ~/.unsloth/studio). Pass through if the override == legacy default.
if [ "$TAURI_MODE" = true ]; then
_tauri_override_var=""
_tauri_override="${UNSLOTH_STUDIO_HOME:-}"
if [ -n "$_tauri_override" ]; then
_tauri_override_var="UNSLOTH_STUDIO_HOME"
else
_tauri_override="${STUDIO_HOME:-}"
[ -n "$_tauri_override" ] && _tauri_override_var="STUDIO_HOME"
fi
# Strip whitespace so " " is treated as unset (matches Python .strip()).
_tauri_override=$(printf '%s' "$_tauri_override" | sed -e 's/^[[:space:]]*//' -e 's/[[:space:]]*$//')
if [ -n "$_tauri_override" ]; then
case "$_tauri_override" in
"~") _tauri_override="$HOME" ;;
"~/"*) _tauri_override="$HOME/${_tauri_override#'~/'}" ;;
esac
# Canonicalize both sides (CDPATH=, -P) so a CDPATH-set env or
# symlinked $HOME doesn't break the legacy-equality comparison.
if [ -d "$_tauri_override" ]; then
_tauri_override_abs=$(CDPATH= cd -P -- "$_tauri_override" 2>/dev/null && pwd -P) \
|| _tauri_override_abs="$_tauri_override"
else
_tauri_override_abs="$_tauri_override"
fi
# Strip trailing separators so ".../studio/" matches ".../studio".
while [ "$_tauri_override_abs" != "/" ] \
&& [ "${_tauri_override_abs%/}" != "$_tauri_override_abs" ]; do
_tauri_override_abs=${_tauri_override_abs%/}
done
_tauri_legacy_root="$HOME/.unsloth/studio"
if [ -d "$_tauri_legacy_root" ]; then
_tauri_legacy_root=$(CDPATH= cd -P -- "$_tauri_legacy_root" 2>/dev/null && pwd -P) \
|| _tauri_legacy_root="$HOME/.unsloth/studio"
fi
while [ "$_tauri_legacy_root" != "/" ] \
&& [ "${_tauri_legacy_root%/}" != "$_tauri_legacy_root" ]; do
_tauri_legacy_root=${_tauri_legacy_root%/}
done
if [ "$_tauri_override_abs" != "$_tauri_legacy_root" ]; then
echo "ERROR: $_tauri_override_var is not supported with --tauri." >&2
echo " The desktop app still uses the legacy ~/.unsloth/studio root." >&2
echo " Run install.sh without --tauri for custom-root shell installs," >&2
echo " or unset the env var for default desktop installs." >&2
exit 1
fi
fi
fi
_is_verbose() {
[ "${UNSLOTH_VERBOSE:-0}" = "1" ]
}
@ -219,7 +275,67 @@ _tauri_gpu_branch() {
}
PYTHON_VERSION="" # resolved after platform detection
STUDIO_HOME="$HOME/.unsloth/studio"
# Resolve install destinations: env override, HOME-redirect (best-effort
# via getent/dscl), or default. Env-var priority: UNSLOTH_STUDIO_HOME wins
# over STUDIO_HOME (the more specific signal beats the generic alias).
_resolve_studio_destinations() {
_override_var=""
_override="${UNSLOTH_STUDIO_HOME:-}"
if [ -n "$_override" ]; then
_override_var="UNSLOTH_STUDIO_HOME"
else
_override="${STUDIO_HOME:-}"
[ -n "$_override" ] && _override_var="STUDIO_HOME"
fi
# Strip surrounding whitespace so " " is treated as unset (matches the
# Python resolvers' .strip()), preventing install/runtime layout drift.
_override=$(printf '%s' "$_override" | sed -e 's/^[[:space:]]*//' -e 's/[[:space:]]*$//')
# Tilde expansion: env vars are not subject to it when quoted on assignment.
case "$_override" in
"~") _override="$HOME" ;;
"~/"*) _override="$HOME/${_override#'~/'}" ;;
esac
if [ -n "$_override" ]; then
mkdir -p -- "$_override" 2>/dev/null || { echo "ERROR: $_override_var=$_override cannot be created." >&2; exit 1; }
[ -w "$_override" ] || { echo "ERROR: $_override_var=$_override is not writable." >&2; exit 1; }
STUDIO_HOME="$(CDPATH= cd -P -- "$_override" && pwd -P)" || exit 1
DATA_DIR="$STUDIO_HOME/share"
_LOCAL_BIN="$STUDIO_HOME/bin"
_STUDIO_HOME_REDIRECT=env
substep "custom $_override_var=$STUDIO_HOME"
return 0
fi
_default_home=""
if command -v getent >/dev/null 2>&1; then
_default_home=$(getent passwd "${USER:-$(whoami)}" 2>/dev/null | cut -d: -f6)
elif [ "$(uname)" = "Darwin" ] && command -v dscl >/dev/null 2>&1; then
_default_home=$(dscl . -read "/Users/${USER:-$(whoami)}" NFSHomeDirectory 2>/dev/null | awk '{print $2}')
fi
# Canonicalize both sides so a trailing slash on $HOME (or symlink mismatch
# with passwd-DB output) doesn't misfire the redirection branch.
_home_canon="$HOME"
if [ -d "$_home_canon" ]; then
_home_canon=$(CDPATH= cd -P -- "$_home_canon" 2>/dev/null && pwd -P) || _home_canon="$HOME"
fi
_default_home_canon="$_default_home"
if [ -n "$_default_home_canon" ] && [ -d "$_default_home_canon" ]; then
_default_home_canon=$(CDPATH= cd -P -- "$_default_home_canon" 2>/dev/null && pwd -P) || _default_home_canon="$_default_home"
fi
if [ -n "$_default_home_canon" ] && [ "$_home_canon" != "$_default_home_canon" ]; then
STUDIO_HOME="$HOME/.unsloth/studio"
DATA_DIR="$HOME/.local/share/unsloth"
_LOCAL_BIN="$HOME/.local/bin"
_STUDIO_HOME_REDIRECT=home
substep "HOME redirected ($HOME); install follows \$HOME"
return 0
fi
STUDIO_HOME="$HOME/.unsloth/studio"
DATA_DIR="$HOME/.local/share/unsloth"
_LOCAL_BIN="$HOME/.local/bin"
_STUDIO_HOME_REDIRECT=default
}
_resolve_studio_destinations
VENV_DIR="$STUDIO_HOME/unsloth_studio"
_VENV_ROLLBACK_DIR=""
_VENV_ROLLBACK_TARGET="$VENV_DIR"
@ -383,23 +499,65 @@ create_studio_shortcuts() {
_css_exe_dir=$(cd "$(dirname "$_css_exe")" && pwd)
_css_exe="$_css_exe_dir/$(basename "$_css_exe")"
_css_data_dir="$HOME/.local/share/unsloth"
_css_data_dir="$DATA_DIR"
_css_launcher="$_css_data_dir/launch-studio.sh"
_css_icon_png="$_css_data_dir/unsloth-studio.png"
_css_gem_png="$_css_data_dir/unsloth-gem.png"
mkdir -p "$_css_data_dir"
# Same-install discriminator: per-install opaque id written once at install
# time and read by both this launcher and the backend (/api/health). Replaces
# the older sha256(canonical $STUDIO_HOME) scheme to (a) avoid leaking the
# install path on -H 0.0.0.0 deployments and (b) sidestep launcher/backend
# canonicalization drift (cd -P vs Path.resolve() symlink/junction handling).
# Lives at $STUDIO_HOME/share/ (not $DATA_DIR) so the backend can find it
# via _STUDIO_ROOT_RESOLVED / "share" / "studio_install_id" regardless of
# mode (in env-mode $STUDIO_HOME/share == $DATA_DIR; in default mode they
# diverge but the backend only knows the studio_root). 32 bytes of urandom
# -> 64 hex chars, byte-compatible with the prior digest so launcher
# placeholder, _check_health, and tests stay length-agnostic.
_css_id_dir="$STUDIO_HOME/share"
mkdir -p "$_css_id_dir"
_css_id_file="$_css_id_dir/studio_install_id"
if [ ! -s "$_css_id_file" ]; then
if [ -r /dev/urandom ]; then
_css_new_id=$(od -An -N32 -tx1 /dev/urandom 2>/dev/null | tr -d ' \n')
fi
if [ -z "${_css_new_id:-}" ] && command -v python3 >/dev/null 2>&1; then
_css_new_id=$(python3 -c 'import secrets; print(secrets.token_hex(32))' 2>/dev/null)
fi
if [ -z "${_css_new_id:-}" ]; then
echo "[WARN] Cannot create launcher: no entropy source for studio_install_id" >&2
return 1
fi
# Atomic write so a partial install can't leave a half-written id.
_css_id_tmp="$_css_id_file.$$.tmp"
printf '%s' "$_css_new_id" > "$_css_id_tmp" \
&& mv "$_css_id_tmp" "$_css_id_file"
chmod 600 "$_css_id_file" 2>/dev/null || true
unset _css_new_id _css_id_tmp
fi
_css_studio_root_id=$(cat "$_css_id_file" 2>/dev/null)
if [ -z "$_css_studio_root_id" ]; then
echo "[WARN] Cannot create launcher: failed to read $_css_id_file" >&2
return 1
fi
_css_is_env_mode=false
[ "$_STUDIO_HOME_REDIRECT" = "env" ] && _css_is_env_mode=true
# ── Write launcher script ──
# The launcher is Bash (not POSIX sh).
# We write it with a placeholder and substitute the exe path via sed.
# Single-quoted heredoc; @@DATA_DIR@@, @@STUDIO_ROOT_ID@@, and
# @@INSTALLED_IS_ENV_MODE@@ are substituted via sed below.
cat > "$_css_launcher" << 'LAUNCHER_EOF'
#!/usr/bin/env bash
# Unsloth Studio Launcher
# Auto-generated by install.sh -- do not edit manually.
set -euo pipefail
DATA_DIR="$HOME/.local/share/unsloth"
DATA_DIR='@@DATA_DIR@@'
_EXPECTED_STUDIO_ROOT_ID='@@STUDIO_ROOT_ID@@'
_INSTALLED_IS_ENV_MODE='@@INSTALLED_IS_ENV_MODE@@'
# Read exe path from config written at install time.
# Sourcing is safe: the config file is written by install.sh, not user input.
@ -416,7 +574,23 @@ MAX_PORT_OFFSET=20
TIMEOUT_SEC=60
POLL_INTERVAL_SEC=1
LOG_FILE="$DATA_DIR/studio.log"
# why: in env-override mode multiple installs share an OS user; namespace the
# lock and remember our own healthy port so we never attach to an unrelated
# Studio listening on the global 8888..8908 range.
LOCK_DIR="${XDG_RUNTIME_DIR:-/tmp}/unsloth-studio-launcher-$(id -u).lock"
PORT_FILE=""
# why: gate on the install-time mode (baked above) instead of the runtime env
# var; sourcing a custom-root studio.conf in shell must not flip a default-mode
# launcher into env-mode behavior with stale state.
if [ "$_INSTALLED_IS_ENV_MODE" = "true" ]; then
if command -v cksum >/dev/null 2>&1; then
_LOCK_KEY=$(printf '%s' "$DATA_DIR" | cksum | awk '{print $1}')
else
_LOCK_KEY=""
fi
[ -n "$_LOCK_KEY" ] && LOCK_DIR="${XDG_RUNTIME_DIR:-/tmp}/unsloth-studio-launcher-$(id -u)-${_LOCK_KEY}.lock"
PORT_FILE="$DATA_DIR/studio.port"
fi
# ── HTTP GET helper (supports curl and wget) ──
_http_get() {
@ -435,10 +609,20 @@ _check_health() {
_port=$1
_resp=$(_http_get "http://127.0.0.1:$_port/api/health") || return 1
case "$_resp" in
*'"status"'*'"healthy"'*'"service"'*'"Unsloth UI Backend"'*) return 0 ;;
*'"service"'*'"Unsloth UI Backend"'*'"status"'*'"healthy"'*) return 0 ;;
*'"status"'*'"healthy"'*'"service"'*'"Unsloth UI Backend"'*) ;;
*'"service"'*'"Unsloth UI Backend"'*'"status"'*'"healthy"'*) ;;
*) return 1 ;;
esac
return 1
# why: verify the backend belongs to THIS install. Baked hex digest avoids
# JSON-escape mismatches on paths with `\`/`"` and avoids leaking the raw
# install path to unauthenticated callers.
if [ -n "$_EXPECTED_STUDIO_ROOT_ID" ]; then
case "$_resp" in
*"\"studio_root_id\":\"$_EXPECTED_STUDIO_ROOT_ID\""*|*"\"studio_root_id\": \"$_EXPECTED_STUDIO_ROOT_ID\""*) return 0 ;;
*) return 1 ;;
esac
fi
return 0
}
# ── Port scanning ──
@ -461,6 +645,25 @@ _candidate_ports() {
}
_find_healthy_port() {
if [ -n "$PORT_FILE" ] && [ -f "$PORT_FILE" ]; then
# why: env-mode installs only attach to a port we previously launched
# ourselves; never to a sibling Studio that happens to be healthy.
_p=$(cat "$PORT_FILE" 2>/dev/null || true)
case "$_p" in
''|*[!0-9]*) ;;
*)
if _check_health "$_p"; then
echo "$_p"
return 0
fi
rm -f "$PORT_FILE"
;;
esac
return 1
fi
if [ -n "$PORT_FILE" ]; then
return 1
fi
for _p in $(_candidate_ports | sort -un); do
if _check_health "$_p"; then
echo "$_p"
@ -611,6 +814,7 @@ if [ -t 1 ]; then
_obwr_deadline=$(($(date +%s) + TIMEOUT_SEC))
while [ "$(date +%s)" -lt "$_obwr_deadline" ]; do
if _check_health "$_launch_port"; then
[ -n "$PORT_FILE" ] && printf '%s\n' "$_launch_port" > "$PORT_FILE" 2>/dev/null || true
_release_lock
_open_browser "http://localhost:$_launch_port"
exit 0
@ -634,6 +838,7 @@ else
_deadline=$(($(date +%s) + TIMEOUT_SEC))
while [ "$(date +%s)" -lt "$_deadline" ]; do
if _check_health "$_launch_port"; then
[ -n "$PORT_FILE" ] && printf '%s\n' "$_launch_port" > "$PORT_FILE" 2>/dev/null || true
_open_browser "http://localhost:$_launch_port"
exit 0
fi
@ -646,13 +851,62 @@ else
fi
LAUNCHER_EOF
# why: bake non-user-controlled placeholders FIRST so a literal
# `@@STUDIO_ROOT_ID@@` inside $DATA_DIR cannot be rewritten below.
sed -e "s|@@STUDIO_ROOT_ID@@|$_css_studio_root_id|g" \
-e "s|@@INSTALLED_IS_ENV_MODE@@|$_css_is_env_mode|g" \
"$_css_launcher" > "$_css_launcher.tmp" \
&& mv "$_css_launcher.tmp" "$_css_launcher"
# Env-mode bakes an absolute DATA_DIR (root fixed at install time);
# default / HOME-redirect keeps the literal $HOME/.local/share/unsloth
# so behavior is byte-identical to pre-override.
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
# Two-stage escape: (1) `'` -> `'\''` for shell single-quote embedding,
# (2) backslash/&/| escape so the value survives the s|...|VALUE| sed
# below. Verified end-to-end with apostrophes, spaces, &, |, $.
_sq_escaped=$(printf '%s' "$DATA_DIR" | sed "s/'/'\\\\''/g")
_sed_safe=$(printf '%s' "$_sq_escaped" | sed 's/[\\&|]/\\&/g')
sed "s|@@DATA_DIR@@|$_sed_safe|g" "$_css_launcher" > "$_css_launcher.tmp" \
&& mv "$_css_launcher.tmp" "$_css_launcher"
else
sed "s|DATA_DIR='@@DATA_DIR@@'|DATA_DIR=\"\$HOME/.local/share/unsloth\"|" \
"$_css_launcher" > "$_css_launcher.tmp" \
&& mv "$_css_launcher.tmp" "$_css_launcher"
fi
chmod +x "$_css_launcher"
# Write the exe path to a separate conf file sourced by the launcher.
# Using single-quote wrapping with the standard '\'' escape for any
# embedded apostrophes. This avoids all sed metacharacter issues.
# studio.conf: exe path + (env-mode only) persisted env vars so fresh
# shells launch the right install without re-exporting.
_css_quoted_exe=$(printf '%s' "$_css_exe" | sed "s/'/'\\\\''/g")
printf '%s\n' "UNSLOTH_EXE='$_css_quoted_exe'" > "$_css_data_dir/studio.conf"
{
printf '%s\n' "UNSLOTH_EXE='$_css_quoted_exe'"
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
# When an override resolves to the legacy default, llama.cpp
# still lives at ~/.unsloth/llama.cpp (one shared build).
# Canonicalize the legacy side so a symlinked $HOME doesn't
# break the comparison.
_css_legacy_studio="$HOME/.unsloth/studio"
if [ -d "$_css_legacy_studio" ]; then
_css_legacy_studio=$(CDPATH= cd -P -- "$_css_legacy_studio" 2>/dev/null && pwd -P) \
|| _css_legacy_studio="$HOME/.unsloth/studio"
fi
if [ "$STUDIO_HOME" = "$_css_legacy_studio" ]; then
_css_llama_path="$HOME/.unsloth/llama.cpp"
else
_css_llama_path="$STUDIO_HOME/llama.cpp"
fi
_css_quoted_home=$(printf '%s' "$STUDIO_HOME" | sed "s/'/'\\\\''/g")
_css_quoted_llama=$(printf '%s' "$_css_llama_path" | sed "s/'/'\\\\''/g")
printf '%s\n' "export UNSLOTH_STUDIO_HOME='$_css_quoted_home'"
# UNSLOTH_LLAMA_CPP_PATH is a pre-existing user-controlled
# llama.cpp dir override; only default it if unset.
printf '%s\n' 'if [ -z "${UNSLOTH_LLAMA_CPP_PATH:-}" ]; then'
printf '%s\n' " export UNSLOTH_LLAMA_CPP_PATH='$_css_quoted_llama'"
printf '%s\n' 'fi'
fi
} > "$_css_data_dir/studio.conf"
# ── Icon: try bundled, then download ──
# rounded-512.png used for both Linux and macOS icons
@ -698,6 +952,14 @@ LAUNCHER_EOF
fi
# ── Platform-specific shortcuts ──
# Env-mode installs are workspace-scoped: skip persistent desktop /
# Start-Menu / dock launchers that may point at a deleted workspace.
# Runtime launcher + studio.conf + icon are still written above.
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
substep "wrote launcher at $_css_launcher (persistent shortcuts skipped in env-override mode)"
return 0
fi
_css_created=0
if [ "$_css_os" = "linux" ]; then
@ -775,11 +1037,18 @@ DESKTOP_EOF
</plist>
PLIST_EOF
# Executable stub
cat > "$_css_macos_dir/launch-studio" << STUB_EOF
# Executable stub: same single-quoted-heredoc + sed-substitute
# pattern as launch-studio.sh so $-vars in $_css_data_dir don't
# expand at .app launch time.
_css_sq_dir=$(printf '%s' "$_css_data_dir" | sed "s/'/'\\\\''/g")
_css_sed_dir=$(printf '%s' "$_css_sq_dir" | sed 's/[\\&|]/\\&/g')
cat > "$_css_macos_dir/launch-studio" << 'STUB_EOF'
#!/bin/sh
exec "$HOME/.local/share/unsloth/launch-studio.sh" "\$@"
exec '@@DATA_DIR@@/launch-studio.sh' "$@"
STUB_EOF
sed "s|@@DATA_DIR@@|$_css_sed_dir|g" "$_css_macos_dir/launch-studio" \
> "$_css_macos_dir/launch-studio.tmp" \
&& mv "$_css_macos_dir/launch-studio.tmp" "$_css_macos_dir/launch-studio"
chmod +x "$_css_macos_dir/launch-studio"
# Build AppIcon.icns from unsloth-gem.png (2240x2240)
@ -1079,11 +1348,28 @@ mkdir -p "$STUDIO_HOME"
_MIGRATED=false
if [ -x "$VENV_DIR/bin/python" ]; then
# why: matching guard to the .venv branch below -- in env-mode
# $STUDIO_HOME is a user-chosen workspace, so refuse to nuke an
# existing $STUDIO_HOME/unsloth_studio that lacks Studio sentinels.
# Accept the in-VENV ownership marker so partial-install retries are
# not blocked. Sentinels must be regular files: -f follows symlinks
# to files (the legitimate ln -s shim shape) but rejects directories
# and broken/dir-targeted symlinks.
if [ "$_STUDIO_HOME_REDIRECT" = "env" ] \
&& [ ! -f "$VENV_DIR/.unsloth-studio-owned" ] \
&& [ ! -f "$STUDIO_HOME/share/studio.conf" ] \
&& [ ! -f "$STUDIO_HOME/bin/unsloth" ]; then
echo "ERROR: $VENV_DIR already exists but does not look like an Unsloth Studio install." >&2
echo " Move it aside or choose an empty UNSLOTH_STUDIO_HOME." >&2
exit 1
fi
# New layout already exists — replace only after preserving rollback copy.
substep "preserving existing environment for rollback..."
_start_studio_venv_replacement "$VENV_DIR"
elif [ -x "$STUDIO_HOME/.venv/bin/python" ]; then
elif [ "$_STUDIO_HOME_REDIRECT" != "env" ] && [ -x "$STUDIO_HOME/.venv/bin/python" ]; then
# Old layout exists — validate before migrating.
# Skip in env-mode so we don't rm -rf an unrelated .venv at the
# workspace root (e.g. user's existing project Python venv).
# In no-torch mode, a missing torch package is expected; validate Python only.
substep "found legacy Studio environment, validating..."
_legacy_ok=false
@ -1132,6 +1418,13 @@ if [ ! -x "$VENV_DIR/bin/python" ]; then
run_install_cmd "create venv" uv venv "$VENV_DIR" --python "$PYTHON_VERSION"
fi
# Mark the freshly-created venv as Studio-owned so a partial install can be
# repaired by re-running install.sh; the env-mode deletion guard above accepts
# this marker as the primary sentinel.
if [ -x "$VENV_DIR/bin/python" ]; then
: > "$VENV_DIR/.unsloth-studio-owned" 2>/dev/null || true
fi
# Guard against Python 3.13.8 torch import bug on Apple Silicon
# (skip when the user explicitly chose a version via --python)
if [ -z "$_USER_PYTHON" ] && [ "$OS" = "macos" ] && [ "$_ARCH" = "arm64" ]; then
@ -1143,6 +1436,9 @@ if [ -z "$_USER_PYTHON" ] && [ "$OS" = "macos" ] && [ "$_ARCH" = "arm64" ]; then
rm -rf "$VENV_DIR"
PYTHON_VERSION="3.12"
run_install_cmd "recreate venv" uv venv "$VENV_DIR" --python "$PYTHON_VERSION"
if [ -x "$VENV_DIR/bin/python" ]; then
: > "$VENV_DIR/.unsloth-studio-owned" 2>/dev/null || true
fi
fi
fi
@ -1721,6 +2017,12 @@ else
fi
fi
# ── Install mlx-vlm on Apple Silicon (optional, for VLM training) ──
if [ "$OS" = "macos" ] && [ "$_ARCH" = "arm64" ]; then
substep "installing mlx-vlm (VLM training support)..."
run_install_cmd "install mlx-vlm" uv pip install --python "$_VENV_PY" mlx-vlm
fi
# ── Run studio setup ──
tauri_log "STEP" "Running Studio setup"
# When --local, use the repo's own setup.sh directly.
@ -1768,7 +2070,17 @@ _SKIP_FRONTEND=0
if [ "$TAURI_MODE" = true ]; then
_SKIP_FRONTEND=1
fi
# Prepend UNSLOTH_STUDIO_HOME=$STUDIO_HOME to "$@" for env-override installs
# without word-splitting on whitespace paths.
_run_setup_with_studio_home() {
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
UNSLOTH_STUDIO_HOME="$STUDIO_HOME" "$@"
else
"$@"
fi
}
if [ "$STUDIO_LOCAL_INSTALL" = true ]; then
_run_setup_with_studio_home env \
SKIP_STUDIO_BASE="$_SKIP_BASE" \
SKIP_STUDIO_FRONTEND="$_SKIP_FRONTEND" \
STUDIO_PACKAGE_NAME="$PACKAGE_NAME" \
@ -1782,6 +2094,7 @@ else
# the same session) does not silently flip a normal install onto the
# local-dev path in setup.sh and install_python_stack.py. Mirrors the
# reset already done in install.ps1 for PowerShell.
_run_setup_with_studio_home env \
SKIP_STUDIO_BASE="$_SKIP_BASE" \
SKIP_STUDIO_FRONTEND="$_SKIP_FRONTEND" \
STUDIO_PACKAGE_NAME="$PACKAGE_NAME" \
@ -1791,36 +2104,53 @@ else
bash "$SETUP_SH" </dev/null || _SETUP_EXIT=$?
fi
# ── Make 'unsloth' available globally via ~/.local/bin ──
mkdir -p "$HOME/.local/bin"
ln -sf "$VENV_DIR/bin/unsloth" "$HOME/.local/bin/unsloth"
# ── Make 'unsloth' available via $_LOCAL_BIN (resolved earlier) ──
# Env-mode: $_LOCAL_BIN is $STUDIO_HOME/bin; skip shell-rc PATH append so we
# don't pollute the user's profile with a workspace-scoped path.
mkdir -p "$_LOCAL_BIN"
# ln -sf into an existing dir creates link inside it. Refuse to delete a
# real directory at the shim path -- that could destroy unrelated user data.
_shim_path="$_LOCAL_BIN/unsloth"
if [ -d "$_shim_path" ] && [ ! -L "$_shim_path" ]; then
echo "ERROR: $_shim_path is a directory; refusing to delete it." >&2
echo " Move or remove it manually, then re-run the installer." >&2
exit 1
fi
# why: -sfn is atomic and -n prevents descent into a symlink-to-directory at
# the shim path (the directory guard above already rejects a real directory).
ln -sfn "$VENV_DIR/bin/unsloth" "$_shim_path"
_LOCAL_BIN="$HOME/.local/bin"
case ":$PATH:" in
*":$_LOCAL_BIN:"*) ;; # already on PATH
*)
_SHELL_PROFILE=""
if [ -n "${ZSH_VERSION:-}" ] || [ "$(basename "${SHELL:-}")" = "zsh" ]; then
_SHELL_PROFILE="$HOME/.zshrc"
elif [ -f "$HOME/.bashrc" ]; then
_SHELL_PROFILE="$HOME/.bashrc"
elif [ -f "$HOME/.profile" ]; then
_SHELL_PROFILE="$HOME/.profile"
fi
if [ -n "$_SHELL_PROFILE" ]; then
if ! grep -q '\.local/bin' "$_SHELL_PROFILE" 2>/dev/null; then
echo '' >> "$_SHELL_PROFILE"
echo '# Added by Unsloth installer' >> "$_SHELL_PROFILE"
echo 'export PATH="$HOME/.local/bin:$PATH"' >> "$_SHELL_PROFILE"
step "path" "added ~/.local/bin to PATH in $_SHELL_PROFILE"
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
export PATH="$_LOCAL_BIN:$PATH"
step "path" "exported $_LOCAL_BIN for this session (no rc-file append in env-override mode)"
else
_SHELL_PROFILE=""
if [ -n "${ZSH_VERSION:-}" ] || [ "$(basename "${SHELL:-}")" = "zsh" ]; then
_SHELL_PROFILE="$HOME/.zshrc"
elif [ -f "$HOME/.bashrc" ]; then
_SHELL_PROFILE="$HOME/.bashrc"
elif [ -f "$HOME/.profile" ]; then
_SHELL_PROFILE="$HOME/.profile"
fi
if [ -n "$_SHELL_PROFILE" ]; then
if ! grep -q '\.local/bin' "$_SHELL_PROFILE" 2>/dev/null; then
echo '' >> "$_SHELL_PROFILE"
echo '# Added by Unsloth installer' >> "$_SHELL_PROFILE"
echo 'export PATH="$HOME/.local/bin:$PATH"' >> "$_SHELL_PROFILE"
step "path" "added ~/.local/bin to PATH in $_SHELL_PROFILE"
fi
fi
export PATH="$_LOCAL_BIN:$PATH"
fi
export PATH="$_LOCAL_BIN:$PATH"
;;
esac
# Non-Tauri installs keep shortcuts even if setup reports failure.
# create_studio_shortcuts gates persistent menu shortcuts on env-mode;
# launcher + studio.conf + icon are always written.
if [ "$TAURI_MODE" != true ]; then
create_studio_shortcuts "$VENV_ABS_BIN/unsloth" "$OS"
fi
@ -1883,10 +2213,21 @@ if [ -t 1 ]; then
esac
else
step "launch" "manual commands:"
substep "unsloth studio -p 8888"
substep "or activate env first:"
substep "source ${VENV_DIR}/bin/activate"
substep "unsloth studio -p 8888"
# Single-quote-escape so paths with spaces / apostrophes copy-paste cleanly.
_li_shim_q="'$(printf '%s' "${_LOCAL_BIN}/unsloth" | sed "s/'/'\\\\''/g")'"
_li_act_q="'$(printf '%s' "${VENV_DIR}/bin/activate" | sed "s/'/'\\\\''/g")'"
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
# Env-mode skips the rc PATH append, so print the absolute shim path.
substep "$_li_shim_q studio -p 8888"
substep "or activate env first:"
substep "source $_li_act_q"
substep "unsloth studio -p 8888"
else
substep "unsloth studio -p 8888"
substep "or activate env first:"
substep "source $_li_act_q"
substep "unsloth studio -p 8888"
fi
substep "(add -H 0.0.0.0 to allow network / cloud access)"
echo ""
fi

View file

@ -89,7 +89,7 @@ huggingfacenotorch = [
]
huggingface = [
"unsloth[huggingfacenotorch]",
"unsloth_zoo>=2026.5.1",
"unsloth_zoo>=2026.4.8",
"torchvision",
"unsloth[triton]",
]
@ -579,7 +579,7 @@ colab-ampere-torch220 = [
"flash-attn>=2.6.3 ; ('linux' in sys_platform)",
]
colab-new = [
"unsloth_zoo>=2026.5.1",
"unsloth_zoo>=2026.4.8",
"packaging",
"tyro",
"transformers>=4.51.3,!=4.52.0,!=4.52.1,!=4.52.2,!=4.52.3,!=4.53.0,!=4.54.0,!=4.55.0,!=4.55.1,!=4.57.0,!=4.57.4,!=4.57.5,!=5.0.0,!=5.1.0,<=5.5.0",

View file

@ -9,16 +9,14 @@ Export backend - handles model exporting in various formats
import glob
import json
import structlog
import tempfile
from loggers import get_logger
import os
import shutil
from pathlib import Path
from typing import Optional, Tuple, List
from peft import PeftModel, PeftModelForCausalLM
from unsloth import FastLanguageModel, FastVisionModel
from unsloth import FastLanguageModel, FastVisionModel, _IS_MLX
from huggingface_hub import HfApi, ModelCard
from transformers.modeling_utils import PushToHubMixin
import torch
from utils.hardware import clear_gpu_cache
from utils.models import is_vision_model, get_base_model_from_lora
@ -26,6 +24,12 @@ from utils.models.model_config import detect_audio_type
from utils.paths import ensure_dir, outputs_root, resolve_export_dir, resolve_output_dir
from core.inference import get_inference_backend
# GPU-only imports — guarded for Apple Silicon where these aren't needed
if not _IS_MLX:
from peft import PeftModel, PeftModelForCausalLM
from transformers.modeling_utils import PushToHubMixin
import torch
logger = get_logger(__name__)
_LLAMA_CPP_SCRIPTS_WARNING_EMITTED = False
@ -225,7 +229,7 @@ class ExportBackend:
model, tokenizer = FastModel.from_pretrained(
model_name = checkpoint_path,
max_seq_length = max_seq_length,
dtype = torch.float32,
dtype = None if _IS_MLX else torch.float32,
load_in_4bit = False,
trust_remote_code = trust_remote_code,
)
@ -262,8 +266,12 @@ class ExportBackend:
trust_remote_code = trust_remote_code,
)
# Check if PEFT model
self.is_peft = isinstance(model, (PeftModel, PeftModelForCausalLM))
# Check if PEFT / LoRA model
if _IS_MLX:
# MLX doesn't use PeftModel — detect LoRA via adapter_config.json
self.is_peft = adapter_config.exists()
else:
self.is_peft = isinstance(model, (PeftModel, PeftModelForCausalLM))
# Store loaded model
self.current_model = model
@ -325,9 +333,7 @@ class ExportBackend:
private: Whether to make the repo private
Returns:
Tuple of (success, message, output_path). output_path is the
resolved absolute on-disk directory of the saved model when
``save_directory`` was set, else None.
Tuple of (success: bool, message: str, output_path: Optional[str])
"""
if not self.current_model or not self.current_tokenizer:
return False, "No model loaded. Please select a checkpoint first.", None
@ -341,14 +347,17 @@ class ExportBackend:
output_path: Optional[str] = None
try:
# Determine save method
if format_type == "4-bit (FP4)":
save_method = "merged_4bit_forced"
elif self._audio_type == "whisper":
# Whisper uses save_method=None for local 16-bit merged save
save_method = None
else: # 16-bit (FP16)
save_method = "merged_16bit"
if _IS_MLX:
mlx_save_method = (
"merged_4bit" if format_type == "4-bit (FP4)" else "merged_16bit"
)
else:
if format_type == "4-bit (FP4)":
save_method = "merged_4bit_forced"
elif self._audio_type == "whisper":
save_method = None
else:
save_method = "merged_16bit"
# Save locally if requested
if save_directory:
@ -356,11 +365,17 @@ class ExportBackend:
logger.info(f"Saving merged model locally to: {save_directory}")
ensure_dir(Path(save_directory))
self.current_model.save_pretrained_merged(
save_directory, self.current_tokenizer, save_method = save_method
)
if _IS_MLX:
self.current_model.save_pretrained_merged(
save_directory,
self.current_tokenizer,
save_method = mlx_save_method,
)
else:
self.current_model.save_pretrained_merged(
save_directory, self.current_tokenizer, save_method = save_method
)
# Write export metadata so the Chat page can identify the base model
self._write_export_metadata(save_directory)
logger.info(f"Model saved successfully to {save_directory}")
output_path = str(Path(save_directory).resolve())
@ -376,17 +391,40 @@ class ExportBackend:
logger.info(f"Pushing merged model to Hub: {repo_id}")
# Whisper uses save_method=None for local but "merged_16bit" for hub push
hub_save_method = (
save_method if save_method is not None else "merged_16bit"
)
self.current_model.push_to_hub_merged(
repo_id,
self.current_tokenizer,
save_method = hub_save_method,
token = hf_token,
private = private,
)
if _IS_MLX:
if save_directory:
self.current_model.push_to_hub_merged(
repo_id,
self.current_tokenizer,
save_directory = save_directory,
token = hf_token,
private = private,
)
else:
with tempfile.TemporaryDirectory() as tmp_dir:
self.current_model.save_pretrained_merged(
tmp_dir,
self.current_tokenizer,
save_method = mlx_save_method,
)
self.current_model.push_to_hub_merged(
repo_id,
self.current_tokenizer,
save_directory = tmp_dir,
token = hf_token,
private = private,
)
else:
hub_save_method = (
save_method if save_method is not None else "merged_16bit"
)
self.current_model.push_to_hub_merged(
repo_id,
self.current_tokenizer,
save_method = hub_save_method,
token = hf_token,
private = private,
)
logger.info(f"Model pushed successfully to {repo_id}")
return True, "Model exported successfully", output_path
@ -411,9 +449,7 @@ class ExportBackend:
Export base model (for non-PEFT models).
Returns:
Tuple of (success, message, output_path). output_path is the
resolved absolute on-disk directory of the saved model when
``save_directory`` was set, else None.
Tuple of (success: bool, message: str, output_path: Optional[str])
"""
if not self.current_model or not self.current_tokenizer:
return False, "No model loaded. Please select a checkpoint first.", None
@ -433,8 +469,16 @@ class ExportBackend:
logger.info(f"Saving base model locally to: {save_directory}")
ensure_dir(Path(save_directory))
self.current_model.save_pretrained(save_directory)
self.current_tokenizer.save_pretrained(save_directory)
if _IS_MLX:
# MLX: save_pretrained_merged handles non-LoRA models too
# (fuse() is a no-op when there are no LoRA layers)
self.current_model.save_pretrained_merged(
save_directory,
self.current_tokenizer,
)
else:
self.current_model.save_pretrained(save_directory)
self.current_tokenizer.save_pretrained(save_directory)
# Write export metadata so the Chat page can identify the base model
self._write_export_metadata(save_directory)
@ -452,44 +496,73 @@ class ExportBackend:
logger.info(f"Pushing base model to Hub: {repo_id}")
# Get base model name from request or model config
base_model = (
base_model_id
or self.current_model.config._name_or_path
or "unknown"
)
# Create repo
hf_api = HfApi(token = hf_token)
repo_id = PushToHubMixin._create_repo(
PushToHubMixin,
repo_id = repo_id,
private = private,
token = hf_token,
)
username = repo_id.split("/")[0]
# Create and push model card
content = MODEL_CARD.format(
username = username,
base_model = base_model,
model_type = self.current_model.config.model_type,
method = "",
extra = "unsloth",
)
card = ModelCard(content)
card.push_to_hub(
repo_id, token = hf_token, commit_message = "Unsloth Model Card"
)
# Upload model files
if save_directory:
hf_api.upload_folder(
folder_path = save_directory, repo_id = repo_id, repo_type = "model"
)
logger.info(f"Model pushed successfully to {repo_id}")
if _IS_MLX:
if save_directory:
self.current_model.push_to_hub_merged(
repo_id,
self.current_tokenizer,
save_directory = save_directory,
token = hf_token,
private = private,
)
else:
with tempfile.TemporaryDirectory() as tmp_dir:
self.current_model.save_pretrained_merged(
tmp_dir,
self.current_tokenizer,
)
self.current_model.push_to_hub_merged(
repo_id,
self.current_tokenizer,
save_directory = tmp_dir,
token = hf_token,
private = private,
)
else:
return False, "Local save directory required for Hub upload", None
# Get base model name from request or model config
base_model = (
base_model_id
or self.current_model.config._name_or_path
or "unknown"
)
# Create repo
hf_api = HfApi(token = hf_token)
repo_id = PushToHubMixin._create_repo(
PushToHubMixin,
repo_id = repo_id,
private = private,
token = hf_token,
)
username = repo_id.split("/")[0]
# Create and push model card
content = MODEL_CARD.format(
username = username,
base_model = base_model,
model_type = self.current_model.config.model_type,
method = "",
extra = "unsloth",
)
card = ModelCard(content)
card.push_to_hub(
repo_id, token = hf_token, commit_message = "Unsloth Model Card"
)
# Upload model files
if save_directory:
hf_api.upload_folder(
folder_path = save_directory,
repo_id = repo_id,
repo_type = "model",
)
logger.info(f"Model pushed successfully to {repo_id}")
else:
return (
False,
"Local save directory required for Hub upload",
None,
)
return True, "Model exported successfully", output_path
@ -519,9 +592,7 @@ class ExportBackend:
hf_token: Hugging Face token
Returns:
Tuple of (success, message, output_path). output_path is the
resolved absolute on-disk directory containing the .gguf
files when ``save_directory`` was set, else None.
Tuple of (success: bool, message: str, output_path: Optional[str])
"""
if not self.current_model or not self.current_tokenizer:
return False, "No model loaded. Please select a checkpoint first.", None
@ -692,9 +763,7 @@ class ExportBackend:
Export LoRA adapter only (not merged).
Returns:
Tuple of (success, message, output_path). output_path is the
resolved absolute on-disk directory of the saved adapter
when ``save_directory`` was set, else None.
Tuple of (success: bool, message: str, output_path: Optional[str])
"""
if not self.current_model or not self.current_tokenizer:
return False, "No model loaded. Please select a checkpoint first.", None
@ -710,8 +779,13 @@ class ExportBackend:
logger.info(f"Saving LoRA adapter locally to: {save_directory}")
ensure_dir(Path(save_directory))
self.current_model.save_pretrained(save_directory)
self.current_tokenizer.save_pretrained(save_directory)
if _IS_MLX:
# MLX: save adapters.safetensors + tokenizer files
self.current_model.save_lora_adapters(save_directory)
self.current_tokenizer.save_pretrained(save_directory)
else:
self.current_model.save_pretrained(save_directory)
self.current_tokenizer.save_pretrained(save_directory)
logger.info(f"Adapter saved successfully to {save_directory}")
output_path = str(Path(save_directory).resolve())
@ -726,10 +800,24 @@ class ExportBackend:
logger.info(f"Pushing LoRA adapter to Hub: {repo_id}")
self.current_model.push_to_hub(repo_id, token = hf_token, private = private)
self.current_tokenizer.push_to_hub(
repo_id, token = hf_token, private = private
)
if _IS_MLX:
with tempfile.TemporaryDirectory() as tmp_dir:
self.current_model.save_lora_adapters(tmp_dir)
self.current_tokenizer.save_pretrained(tmp_dir)
hf_api = HfApi(token = hf_token)
hf_api.create_repo(repo_id, private = private, exist_ok = True)
hf_api.upload_folder(
folder_path = tmp_dir,
repo_id = repo_id,
repo_type = "model",
)
else:
self.current_model.push_to_hub(
repo_id, token = hf_token, private = private
)
self.current_tokenizer.push_to_hub(
repo_id, token = hf_token, private = private
)
logger.info(f"Adapter pushed successfully to {repo_id}")
return True, "LoRA adapter exported successfully", output_path

View file

@ -732,22 +732,46 @@ class LlamaCppBackend:
if win_bin.is_file():
return str(win_bin)
# 24. ~/.unsloth/llama.cpp (primary — setup.sh / setup.ps1 build here)
unsloth_home = Path.home() / ".unsloth" / "llama.cpp"
# Root dir (make builds copy binaries here)
home_root = unsloth_home / binary_name
if home_root.is_file():
return str(home_root)
# build/bin/ (cmake builds on Linux)
home_linux = unsloth_home / "build" / "bin" / binary_name
if home_linux.is_file():
return str(home_linux)
# 2-4. Match installer layout: env-mode -> $STUDIO_HOME/llama.cpp;
# default/HOME-redirect -> ~/.unsloth/llama.cpp (sibling of studio).
legacy_llama = Path.home() / ".unsloth" / "llama.cpp"
try:
from utils.paths.storage_roots import studio_root as _sr # noqa: WPS433
# 3. Windows MSVC build has Release subdir
if sys.platform == "win32":
home_win = unsloth_home / "build" / "bin" / "Release" / binary_name
if home_win.is_file():
return str(home_win)
_resolved_sr = _sr()
_legacy_studio = Path.home() / ".unsloth" / "studio"
try:
_is_legacy = _resolved_sr.resolve() == _legacy_studio.resolve()
except (OSError, ValueError):
_is_legacy = _resolved_sr == _legacy_studio
if _is_legacy:
search_roots = [legacy_llama]
else:
# why: _kill_orphaned_servers excludes the legacy root in custom
# mode; discovery must match so we never spawn a server we then
# refuse to clean up. UNSLOTH_LLAMA_CPP_PATH (handled earlier)
# is the explicit way to share a build across roots.
search_roots = [_resolved_sr / "llama.cpp"]
except (ImportError, OSError, ValueError):
search_roots = [legacy_llama]
_seen_roots: set[str] = set()
_unique_roots: list[Path] = []
for r in search_roots:
k = str(r)
if k not in _seen_roots:
_seen_roots.add(k)
_unique_roots.append(r)
for unsloth_home in _unique_roots:
home_root = unsloth_home / binary_name
if home_root.is_file():
return str(home_root)
home_linux = unsloth_home / "build" / "bin" / binary_name
if home_linux.is_file():
return str(home_linux)
if sys.platform == "win32":
home_win = unsloth_home / "build" / "bin" / "Release" / binary_name
if home_win.is_file():
return str(home_win)
# 56. Legacy: in-tree build (older setup.sh / setup.ps1 versions)
project_root = Path(__file__).resolve().parents[4]
@ -2592,8 +2616,27 @@ class LlamaCppBackend:
# (binary must be *under* one of these)
install_roots: list[Path] = []
# Primary install dir (setup.sh / prebuilt installer)
install_roots.append(Path.home() / ".unsloth" / "llama.cpp")
# Env-mode custom root (mirrors _find_llama_server_binary).
_is_custom_root = False
try:
from utils.paths.storage_roots import studio_root as _sr # noqa: WPS433
_resolved_sr = _sr()
_legacy_studio = Path.home() / ".unsloth" / "studio"
try:
_is_custom_root = _resolved_sr.resolve() != _legacy_studio.resolve()
except (OSError, ValueError):
_is_custom_root = _resolved_sr != _legacy_studio
if _is_custom_root:
install_roots.append(_resolved_sr / "llama.cpp")
except (ImportError, OSError, ValueError):
pass
# Primary install dir (default mode only). Env-mode skips this so
# a custom-root Studio cannot kill a concurrent default-install
# Studio's llama-server (same OS user, different install).
if not _is_custom_root:
install_roots.append(Path.home() / ".unsloth" / "llama.cpp")
# Legacy in-tree build dirs (older setup.sh versions)
project_root = Path(__file__).resolve().parents[4]

View file

@ -0,0 +1,395 @@
# SPDX-License-Identifier: AGPL-3.0-only
"""MLX inference backend for Apple Silicon.
Drop-in replacement for InferenceBackend same interface, uses mlx-lm/mlx-vlm
instead of torch/transformers for model loading and generation.
"""
import threading
from typing import Optional, Generator
from loggers import get_logger
logger = get_logger(__name__)
class MLXInferenceBackend:
def __init__(self):
self.models = {}
self.active_model_name = None
self.loading_models = set()
self.loaded_local_models = []
self.device = "mlx"
self._generation_lock = threading.Lock()
# MLX state
self._model = None
self._tokenizer = None
self._processor = None
self._is_vlm = False
self._config = {}
# Recorded for unload to release pinned memory back to the OS.
self._memory_limits_applied = {}
def _configure_memory_limits(self):
"""Apply Metal memory caps before loading a model.
Mirrors MLXTrainer._configure_memory_limits's defaults:
memory_limit = 85% of recommended working-set,
wired_limit = min(recommended, memory_limit). Recorded so unload
can lower wired_limit back to release pinned RAM.
"""
import mlx.core as mx
if not mx.metal.is_available():
return
info = mx.device_info()
rec_bytes = info.get("max_recommended_working_set_size")
if not rec_bytes or rec_bytes <= 0:
return
rec_gb = rec_bytes / 1e9
memory_limit_gb = rec_gb * 0.85
wired_limit_gb = min(rec_gb, memory_limit_gb)
mx.set_memory_limit(int(memory_limit_gb * 1e9))
mx.set_wired_limit(int(wired_limit_gb * 1e9))
self._memory_limits_applied = {
"memory_limit_gb": memory_limit_gb,
"wired_limit_gb": wired_limit_gb,
"recommended_gb": rec_gb,
}
logger.info(
"MLX memory caps: memory_limit=%.2f GB, wired_limit=%.2f GB",
memory_limit_gb,
wired_limit_gb,
)
def load_model(
self,
config,
max_seq_length = 2048,
load_in_4bit = True,
hf_token = None,
trust_remote_code = False,
gpu_ids = None,
dtype = None,
) -> bool:
import mlx.core as mx
model_name = config.identifier if hasattr(config, "identifier") else str(config)
is_vision = getattr(config, "is_vision", False)
if hf_token:
import os
os.environ["HF_TOKEN"] = hf_token
self._configure_memory_limits()
is_lora = getattr(config, "is_lora", False)
logger.info(
"Loading %s via %s (is_lora=%s)",
model_name,
"mlx-vlm" if is_vision else "mlx-lm",
is_lora,
)
try:
from unsloth_zoo.mlx_loader import FastMLXModel
except ImportError as e:
raise ImportError(
"Unsloth: MLX inference requires unsloth-zoo with the MLX modules "
"(unsloth_zoo.mlx_loader). Reinstall via install.sh on Apple Silicon."
) from e
model, tokenizer_or_processor = FastMLXModel.from_pretrained(
model_name,
max_seq_length = max_seq_length,
dtype = dtype,
load_in_4bit = load_in_4bit,
token = hf_token,
trust_remote_code = trust_remote_code,
text_only = False if is_vision else True,
)
if is_vision:
processor = tokenizer_or_processor
self._model = model
self._processor = processor
self._tokenizer = getattr(processor, "tokenizer", processor)
self._is_vlm = True
else:
tokenizer = tokenizer_or_processor
self._model = model
self._tokenizer = tokenizer
self._processor = None
self._is_vlm = False
self.active_model_name = model_name
self.models[model_name] = {
"model": self._model,
"tokenizer": self._tokenizer,
"processor": self._processor,
"is_vision": is_vision,
"is_lora": getattr(config, "is_lora", False),
"is_audio": False,
"audio_type": None,
"has_audio_input": False,
}
logger.info("Model %s loaded successfully", model_name)
return True
def unload_model(self, model_name: str) -> bool:
import mlx.core as mx
import gc
if model_name in self.models:
del self.models[model_name]
self._model = None
self._tokenizer = None
self._processor = None
if self.active_model_name == model_name:
self.active_model_name = None
gc.collect()
mx.clear_cache()
if mx.metal.is_available() and self._memory_limits_applied and not self.models:
try:
mx.set_wired_limit(0)
logger.info("MLX wired_limit released back to OS on unload")
except Exception as e:
logger.warning("Failed to release wired_limit: %s", e)
self._memory_limits_applied = {}
logger.info("Model %s unloaded", model_name)
return True
def generate_chat_response(
self,
messages,
system_prompt = "",
image = None,
temperature = 0.7,
top_p = 0.9,
top_k = 40,
min_p = 0.0,
max_new_tokens = 256,
repetition_penalty = 1.0,
cancel_event = None,
) -> Generator[str, None, None]:
if self._model is None:
raise RuntimeError("No model loaded")
# Build messages with system prompt
full_messages = []
if system_prompt:
full_messages.append({"role": "system", "content": system_prompt})
full_messages.extend(messages)
# Inject image into the last user message for VLM
if self._is_vlm and image is not None:
for msg in reversed(full_messages):
if msg.get("role") == "user":
content = msg.get("content", "")
if isinstance(content, str):
msg["content"] = [
{"type": "image"},
{"type": "text", "text": content},
]
elif isinstance(content, list):
# Prepend image if not already there
has_image = any(
p.get("type") == "image"
for p in content
if isinstance(p, dict)
)
if not has_image:
content.insert(0, {"type": "image"})
break
if self._is_vlm:
yield from self._generate_vlm(
full_messages,
image,
temperature,
top_p,
top_k,
min_p,
max_new_tokens,
repetition_penalty,
cancel_event,
)
else:
yield from self._generate_text(
full_messages,
temperature,
top_p,
top_k,
min_p,
max_new_tokens,
repetition_penalty,
cancel_event,
)
def _generate_text(
self,
messages,
temperature,
top_p,
top_k,
min_p,
max_new_tokens,
repetition_penalty,
cancel_event,
):
from mlx_lm import stream_generate
from mlx_lm.sample_utils import make_sampler, make_logits_processors
prompt = self._tokenizer.apply_chat_template(
messages,
tokenize = False,
add_generation_prompt = True,
)
if prompt is None:
raise RuntimeError(
"apply_chat_template returned None — tokenizer may be incompatible"
)
sampler = make_sampler(
temp = temperature,
top_p = top_p,
top_k = int(top_k or 0),
min_p = float(min_p or 0.0),
min_tokens_to_keep = 1,
)
# Only build a logits processor when we actually have a non-trivial
# repetition penalty (1.0 is the no-op value).
logits_processors = None
if repetition_penalty is not None and float(repetition_penalty) not in (
0.0,
1.0,
):
logits_processors = make_logits_processors(
repetition_penalty = float(repetition_penalty),
)
token_ids = []
logger.info(
"Generating: prompt_len=%d, max_tokens=%d, model=%s, tokenizer=%s",
len(prompt),
max_new_tokens,
type(self._model).__name__,
type(self._tokenizer).__name__,
)
with self._generation_lock:
try:
gen_kwargs = dict(
prompt = prompt,
max_tokens = max_new_tokens,
sampler = sampler,
)
if logits_processors is not None:
gen_kwargs["logits_processors"] = logits_processors
for response in stream_generate(
self._model,
self._tokenizer,
**gen_kwargs,
):
token_ids.append(response.token)
# Decode full sequence with skip_special_tokens — same as GPU
cumulative = self._tokenizer.decode(
token_ids,
skip_special_tokens = True,
)
yield cumulative
if cancel_event and cancel_event.is_set():
break
except Exception as e:
import traceback
logger.error("stream_generate failed:\n%s", traceback.format_exc())
raise
def _generate_vlm(
self,
messages,
image,
temperature,
top_p,
top_k,
min_p,
max_new_tokens,
repetition_penalty,
cancel_event,
):
from mlx_vlm import stream_generate as vlm_stream
# Apply chat template
chat_fn = getattr(self._processor, "apply_chat_template", None)
if (
chat_fn is None
or not hasattr(self._processor, "chat_template")
or self._processor.chat_template is None
):
tok = getattr(self._processor, "tokenizer", self._processor)
chat_fn = tok.apply_chat_template
prompt = chat_fn(messages, tokenize = False, add_generation_prompt = True)
# For VLM: always use mlx_vlm's stream_generate which handles
# pixel_values properly (passes None for text-only, image for VLM)
images = [image] if image is not None else None
cumulative = ""
logger.info(
"VLM generating: prompt_len=%d, has_image=%s",
len(prompt),
image is not None,
)
# mlx_vlm.stream_generate forwards **kwargs into generate_step, which
# accepts temp/top_p/top_k/repetition_penalty (and builds the sampler
# + logits_processors internally). Pass them through.
# NOTE: mlx_vlm.generate_step expects ``temperature=`` (long form) —
# passing ``temp=`` silently falls into **kwargs and is ignored,
# leaving generation stuck at the default 0.0 (greedy).
vlm_kwargs = dict(
max_tokens = max_new_tokens,
temperature = temperature,
top_p = top_p,
top_k = int(top_k or 0),
min_p = float(min_p or 0.0),
)
if repetition_penalty is not None and float(repetition_penalty) not in (
0.0,
1.0,
):
vlm_kwargs["repetition_penalty"] = float(repetition_penalty)
with self._generation_lock:
for response in vlm_stream(
self._model,
self._processor,
prompt,
images,
**vlm_kwargs,
):
token_text = (
response.text if hasattr(response, "text") else str(response)
)
cumulative += token_text
yield cumulative
if cancel_event and cancel_event.is_set():
break
def generate_with_adapter_control(
self, use_adapter = None, cancel_event = None, **gen_kwargs
) -> Generator[str, None, None]:
# MLX LoRA adapter toggling not yet supported — generate normally
yield from self.generate_chat_response(cancel_event = cancel_event, **gen_kwargs)
def reset_generation_state(self):
import mlx.core as mx
import gc
gc.collect()
mx.clear_cache()

View file

@ -663,6 +663,98 @@ def run_inference_process(
model_name = config["model_name"]
# ── 0. MLX fast-path — skip torch/transformers entirely ──
backend_path = str(Path(__file__).resolve().parent.parent.parent)
if backend_path not in sys.path:
sys.path.insert(0, backend_path)
from utils.hardware import hardware as _hw
_hw.detect_hardware()
if _hw.DEVICE == _hw.DeviceType.MLX:
try:
_activate_transformers_version(model_name)
except Exception:
pass
try:
from core.inference.mlx_inference import MLXInferenceBackend
backend = MLXInferenceBackend()
_send_response(
resp_queue,
{"type": "status", "message": "Loading model...", "ts": time.time()},
)
_handle_load(backend, config, resp_queue)
except Exception as exc:
_send_response(
resp_queue,
{
"type": "error",
"error": f"MLX inference init failed: {exc}",
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
},
)
return
# Enter same command loop as GPU path
logger.info("MLX inference subprocess ready, entering command loop")
while True:
try:
cmd = cmd_queue.get(timeout = 1.0)
except _queue.Empty:
continue
except (EOFError, OSError):
return
if cmd is None:
continue
cmd_type = cmd.get("type", "")
try:
if cmd_type == "generate":
cancel_event.clear()
_handle_generate(backend, cmd, resp_queue, cancel_event)
elif cmd_type == "load":
if backend.active_model_name:
backend.unload_model(backend.active_model_name)
_handle_load(backend, cmd, resp_queue)
elif cmd_type == "unload":
_handle_unload(backend, cmd, resp_queue)
elif cmd_type == "cancel":
cancel_event.set()
elif cmd_type == "reset":
cancel_event.set()
backend.reset_generation_state()
_send_response(resp_queue, {"type": "reset_ack", "ts": time.time()})
elif cmd_type == "status":
_send_response(
resp_queue,
{
"type": "status_response",
"active_model": backend.active_model_name,
"models": {
k: {kk: vv for kk, vv in v.items() if kk != "model"}
for k, v in backend.models.items()
},
"loading": list(backend.loading_models),
"ts": time.time(),
},
)
elif cmd_type == "shutdown":
return
except Exception as exc:
logger.error("MLX command error (%s): %s", cmd_type, exc)
_send_response(
resp_queue,
{
"type": "gen_error" if cmd_type == "generate" else "error",
"request_id": cmd.get("request_id"),
"error": str(exc),
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
},
)
return
# ── 1. Activate correct transformers version BEFORE any ML imports ──
try:
_activate_transformers_version(model_name)

View file

@ -62,6 +62,7 @@ from datasets import Dataset, load_dataset
from utils.models import is_vision_model, detect_audio_type
from utils.datasets import format_and_template_dataset
from utils.datasets import MODEL_TO_TEMPLATE_MAPPER, TEMPLATE_TO_RESPONSES_MAPPER
from utils.datasets.raw_text import prepare_raw_text_dataset
from utils.paths import (
ensure_dir,
resolve_dataset_path,
@ -125,6 +126,7 @@ class UnslothTrainer:
self.load_in_4bit = True # Track quantization mode for metadata
# Model state tracking
self.is_cpt = False # Set to True for Continued Pretraining
self.is_vlm = False
self.is_audio = False
self.is_audio_vlm = (
@ -925,6 +927,7 @@ class UnslothTrainer:
use_gradient_checkpointing: str = "unsloth",
use_rslora: bool = False,
use_loftq: bool = False,
modules_to_save: list = None,
) -> bool:
"""
Prepare model for training (with optional LoRA).
@ -1121,11 +1124,14 @@ class UnslothTrainer:
loftq_config = {"loftq_bits": 4, "loftq_iter": 1}
if use_loftq
else None,
modules_to_save = modules_to_save,
)
else:
# Text model LoRA
logger.info(f"Text model LoRA configuration:")
logger.info(f" - Target modules: {target_modules}\n")
if modules_to_save:
logger.info(f" - Modules to save: {modules_to_save}\n")
self.model = FastLanguageModel.get_peft_model(
self.model,
@ -1140,6 +1146,7 @@ class UnslothTrainer:
loftq_config = {"loftq_bits": 4, "loftq_iter": 1}
if use_loftq
else None,
modules_to_save = modules_to_save,
)
# Check if stopped during LoRA preparation
@ -2342,6 +2349,7 @@ class UnslothTrainer:
eval_steps: float = 0.00,
dataset_slice_start: int = None,
dataset_slice_end: int = None,
is_cpt: bool = False,
) -> Optional[tuple]:
"""
Load and prepare dataset for training.
@ -2360,6 +2368,35 @@ class UnslothTrainer:
False # True if eval comes from a separate HF split
)
eval_enabled = eval_steps is not None and eval_steps > 0
raw_text_mode = is_cpt or format_type == "raw"
def _raw_mode_label() -> str:
return "CPT" if is_cpt else "raw text"
def _apply_raw_text_prep(ds: Dataset, split_name: str) -> Dataset:
try:
result = prepare_raw_text_dataset(
ds,
mode_label = _raw_mode_label(),
split_name = split_name,
eos_token = getattr(self.tokenizer, "eos_token", None),
append_eos = True,
)
except ValueError as exc:
error_msg = str(exc)
logger.error(error_msg)
self._update_progress(error = error_msg)
raise
for notice in result.notices:
if notice.level == "warning":
logger.warning(notice.message)
if notice.update_status:
self._update_progress(status_message = notice.message)
else:
logger.info(f"{notice.message}\n")
return result.dataset
if local_datasets:
# Load local datasets using load_dataset() so the result is
@ -2534,6 +2571,48 @@ class UnslothTrainer:
processed = self._preprocess_dac_dataset(dataset, custom_format_mapping)
return ({"dataset": processed, "final_format": "audio_dac"}, None)
# ========== RAW TEXT BYPASS ==========
if raw_text_mode:
logger.info(
f"{_raw_mode_label().capitalize()} mode: bypassing chat template, "
"using raw text\n"
)
dataset = _apply_raw_text_prep(dataset, "train")
if has_separate_eval_source and eval_dataset is not None:
eval_dataset = _apply_raw_text_prep(eval_dataset, "eval")
dataset_info = {
"dataset": dataset,
"detected_format": "raw_text",
"final_format": "raw_text",
"success": True,
}
if has_separate_eval_source and eval_dataset is not None:
logger.info(
f"{_raw_mode_label().capitalize()}: eval dataset "
f"({len(eval_dataset)} rows) kept as raw text\n"
)
elif eval_enabled and not has_separate_eval_source:
split_result = self._resolve_eval_split_from_dataset(dataset)
if split_result is not None:
train_portion, eval_dataset = split_result
dataset_info["dataset"] = train_portion
train_dataset = dataset_info["dataset"]
n = len(train_dataset) if hasattr(train_dataset, "__len__") else None
n_display = f"{n:,}" if isinstance(n, int) else "streaming"
self._update_progress(
status_message = f"Dataset ready ({n_display} samples, raw text)"
)
logger.info(f"Raw-text dataset ready ({n_display} samples)\n")
if "text" not in train_dataset.column_names:
raise ValueError(
f"Raw-text dataset missing 'text' column: {train_dataset.column_names}"
)
return (dataset_info, eval_dataset)
elif self.is_audio_vlm:
formatted = self._format_audio_vlm_dataset(
dataset, custom_format_mapping
@ -2676,6 +2755,7 @@ class UnslothTrainer:
output_dir: str | None = None,
num_epochs: int = 3,
learning_rate: float = 2e-4,
embedding_learning_rate: float | None = None,
batch_size: int = 2,
gradient_accumulation_steps: int = 4,
warmup_steps: int = None,
@ -2728,6 +2808,7 @@ class UnslothTrainer:
"output_dir": output_dir,
"num_epochs": num_epochs,
"learning_rate": learning_rate,
"embedding_learning_rate": embedding_learning_rate,
"batch_size": batch_size,
"gradient_accumulation_steps": gradient_accumulation_steps,
"warmup_steps": warmup_steps,
@ -2945,6 +3026,13 @@ class UnslothTrainer:
logger.info("Configuring data collator...\n")
dataset_final_format = (
str(dataset.get("final_format", "")).lower()
if isinstance(dataset, dict)
else ""
)
raw_text_mode = dataset_final_format == "raw_text"
data_collator = None # Default to built-in data collator
if is_deepseek_ocr:
# Special DeepSeek OCR collator - auto-install if needed
@ -2984,7 +3072,7 @@ class UnslothTrainer:
self._update_progress(error = error_msg, is_training = False)
return
elif self.is_audio_vlm:
elif self.is_audio_vlm and not raw_text_mode:
# Audio VLM collator (e.g. Gemma 3N with audio data)
# Mirrors the collate_fn from Gemma3N_(4B)-Audio notebook
logger.info("Configuring audio VLM data collator...\n")
@ -3026,7 +3114,7 @@ class UnslothTrainer:
data_collator = audio_vlm_collate_fn
logger.info("Audio VLM data collator configured\n")
elif self.is_vlm:
elif self.is_vlm and not raw_text_mode:
# Standard VLM collator (images)
logger.info("Using UnslothVisionDataCollator for vision model\n")
from unsloth.trainer import UnslothVisionDataCollator
@ -3137,8 +3225,9 @@ class UnslothTrainer:
optim_value = training_args.get("optim", "adamw_8bit")
lr_scheduler_type_value = training_args.get("lr_scheduler_type", "linear")
if self.is_vlm or self.is_audio_vlm:
if (self.is_vlm or self.is_audio_vlm) and not raw_text_mode:
# Vision / audio VLM config (both need skip_prepare_dataset + remove_unused_columns)
# Raw-text runs on VLM-capable models are routed to the text path below.
label = "audio VLM" if self.is_audio_vlm else "vision"
logger.info(f"Configuring {label} model training parameters\n")
# Use provided values or defaults for vision models
@ -3160,7 +3249,14 @@ class UnslothTrainer:
}
)
else:
logger.info("Configuring text model training parameters\n")
is_cpt = training_args.get("is_cpt", False)
self.is_cpt = is_cpt
if is_cpt:
logger.info("Configuring Continued Pretraining (CPT) parameters\n")
elif raw_text_mode:
logger.info("Configuring raw-text training parameters\n")
else:
logger.info("Configuring text model training parameters\n")
config_args.update(
{
"optim": optim_value,
@ -3189,9 +3285,10 @@ class UnslothTrainer:
logger.info("Training configuration prepared\n")
# ========== TRAINER INITIALIZATION ==========
if self.is_audio_vlm:
if self.is_audio_vlm and not raw_text_mode:
# Audio VLM (e.g. Gemma 3N + audio): raw Dataset from _format_audio_vlm_dataset
# Notebook uses processing_class=processor.tokenizer (text tokenizer only)
# Raw-text runs are routed to the text path below.
train_dataset = (
dataset if isinstance(dataset, Dataset) else dataset["dataset"]
)
@ -3210,8 +3307,9 @@ class UnslothTrainer:
if eval_dataset is not None:
trainer_kwargs["eval_dataset"] = eval_dataset
self.trainer = SFTTrainer(**trainer_kwargs)
elif self.is_vlm:
elif self.is_vlm and not raw_text_mode:
# Image VLM: dataset is dict wrapper from format_and_template_dataset
# Raw-text runs are routed to the text path below.
train_dataset = (
dataset["dataset"] if isinstance(dataset, dict) else dataset
)
@ -3242,16 +3340,48 @@ class UnslothTrainer:
)
sft_tokenizer = self.tokenizer.tokenizer
trainer_kwargs = {
"model": self.model,
"tokenizer": sft_tokenizer,
"train_dataset": dataset["dataset"],
"data_collator": data_collator,
"args": SFTConfig(**config_args),
}
if eval_dataset is not None:
trainer_kwargs["eval_dataset"] = eval_dataset
self.trainer = SFTTrainer(**trainer_kwargs)
if is_cpt:
try:
from unsloth import (
UnslothTrainer as _UnslothCPTTrainer,
UnslothTrainingArguments as _UnslothTrainingArguments,
)
except ImportError as exc:
raise RuntimeError(
"CPT requires a newer Unsloth install that exports "
"`UnslothTrainer` and `UnslothTrainingArguments` "
"(for embedding_learning_rate support). "
"Upgrade with: `pip install -U unsloth unsloth_zoo`."
) from exc
embedding_lr = training_args.get("embedding_learning_rate")
logger.info(
f"CPT: using UnslothTrainer with embedding_learning_rate={embedding_lr}\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,
),
}
if eval_dataset is not None:
trainer_kwargs["eval_dataset"] = eval_dataset
self.trainer = _UnslothCPTTrainer(**trainer_kwargs)
else:
trainer_kwargs = {
"model": self.model,
"tokenizer": sft_tokenizer,
"train_dataset": dataset["dataset"],
"data_collator": data_collator,
"args": SFTConfig(**config_args),
}
if eval_dataset is not None:
trainer_kwargs["eval_dataset"] = eval_dataset
self.trainer = SFTTrainer(**trainer_kwargs)
# Restore the full processor as processing_class so checkpoint
# saves include preprocessor_config.json (needed for GGUF export).
if sft_tokenizer is not self.tokenizer:
@ -3260,19 +3390,32 @@ class UnslothTrainer:
# ========== TRAIN ON RESPONSES ONLY ==========
# Determine if we should train on responses only
# Raw-text datasets always train on all tokens.
instruction_part = None
response_part = None
train_on_responses_enabled = training_args.get(
"train_on_completions", False
is_cpt = training_args.get("is_cpt", False)
train_on_responses_enabled = (
False
if (is_cpt or raw_text_mode)
else training_args.get("train_on_completions", False)
)
if is_cpt:
logger.info(
"CPT mode: skipping train_on_responses_only — training on all tokens\n"
)
elif raw_text_mode:
logger.info(
"Raw-text mode: skipping train_on_responses_only — training on all tokens\n"
)
# DeepSeek OCR handles this internally in its collator, so skip
# Audio VLM handles label masking in its collator, so skip
if (
train_on_responses_enabled
and not self.is_audio_vlm
and not self.is_audio
and not (is_deepseek_ocr or dataset["final_format"].lower() == "alpaca")
and not (is_deepseek_ocr or dataset_final_format == "alpaca")
):
try:
logger.info("Configuring train on responses only...\n")
@ -3318,7 +3461,7 @@ class UnslothTrainer:
and response_part
and not self.is_audio_vlm
and not self.is_audio
and not (is_deepseek_ocr or dataset["final_format"].lower() == "alpaca")
and not (is_deepseek_ocr or dataset_final_format == "alpaca")
):
try:
from unsloth.chat_templates import train_on_responses_only
@ -3451,7 +3594,9 @@ class UnslothTrainer:
config = json.load(f)
# Determine the training method
if self.load_in_4bit:
if self.is_cpt:
method = "CPT"
elif self.load_in_4bit:
method = "qlora"
else:
method = "lora"

View file

@ -62,6 +62,7 @@ class TrainingProgress:
grad_norm: Optional[float] = None
num_tokens: Optional[int] = None
eval_loss: Optional[float] = None
peak_memory_gb: Optional[float] = None
class TrainingBackend:
@ -158,6 +159,7 @@ class TrainingBackend:
"is_embedding": kwargs.get("is_embedding", False),
"num_epochs": kwargs.get("num_epochs", 3),
"learning_rate": kwargs.get("learning_rate", "2e-4"),
"embedding_learning_rate": kwargs.get("embedding_learning_rate"),
"batch_size": kwargs.get("batch_size", 2),
"gradient_accumulation_steps": kwargs.get("gradient_accumulation_steps", 4),
"warmup_steps": kwargs.get("warmup_steps"),
@ -194,26 +196,33 @@ class TrainingBackend:
"gpu_ids": kwargs.get("gpu_ids"),
}
# Derive load_in_4bit from training_type
if config["training_type"] != "LoRA/QLoRA":
# Full finetuning always runs in 16-bit. LoRA/QLoRA and CPT preserve the
# explicit request so 4-bit adapter/raw-text runs remain possible.
if config["training_type"] == "Full Finetuning":
config["load_in_4bit"] = False
# Spawn subprocess — use locals so state is untouched on failure
resolved_gpu_ids, gpu_selection = prepare_gpu_selection(
kwargs.get("gpu_ids"),
model_name = config["model_name"],
hf_token = config["hf_token"] or None,
training_type = config["training_type"],
load_in_4bit = config["load_in_4bit"],
batch_size = config.get("batch_size", 4),
max_seq_length = config.get("max_seq_length", 2048),
lora_rank = config.get("lora_r", 16),
target_modules = config.get("target_modules"),
gradient_checkpointing = config.get("gradient_checkpointing", "unsloth"),
optimizer = config.get("optim", "adamw_8bit"),
)
config["resolved_gpu_ids"] = resolved_gpu_ids
config["gpu_selection"] = gpu_selection
from utils.hardware import hardware as _hw
if _hw.DEVICE == _hw.DeviceType.MLX:
config["resolved_gpu_ids"] = None
config["gpu_selection"] = None
else:
resolved_gpu_ids, gpu_selection = prepare_gpu_selection(
kwargs.get("gpu_ids"),
model_name = config["model_name"],
hf_token = config["hf_token"] or None,
training_type = config["training_type"],
load_in_4bit = config["load_in_4bit"],
batch_size = config.get("batch_size", 4),
max_seq_length = config.get("max_seq_length", 2048),
lora_rank = config.get("lora_r", 16),
target_modules = config.get("target_modules"),
gradient_checkpointing = config.get("gradient_checkpointing", "unsloth"),
optimizer = config.get("optim", "adamw_8bit"),
)
config["resolved_gpu_ids"] = resolved_gpu_ids
config["gpu_selection"] = gpu_selection
from .worker import run_training_process
@ -512,6 +521,12 @@ class TrainingBackend:
self._progress.grad_norm = event.get("grad_norm")
self._progress.num_tokens = event.get("num_tokens")
self._progress.eval_loss = event.get("eval_loss")
_peak = event.get("peak_memory_gb")
if _peak is not None:
try:
self._progress.peak_memory_gb = float(_peak)
except (TypeError, ValueError):
pass
self._progress.is_training = True
status = event.get("status_message", "")
if status:

View file

@ -15,6 +15,7 @@ from __future__ import annotations
import structlog
from loggers import get_logger
import math
import os
import shutil
import sys
@ -338,6 +339,594 @@ def _activate_transformers_version(model_name: str) -> None:
activate_transformers_for_subprocess(model_name)
def _adapt_for_mlx_vlm(items):
"""Adapt GPU-path VLM dataset output for mlx-vlm consumption.
The GPU path embeds PIL images inside messages content as
{"type": "image", "image": PIL_Image}. mlx-vlm's prepare_inputs
needs images at top-level to produce pixel_values regardless of
model type. Extract them and leave bare {"type": "image"} placeholders.
"""
adapted = []
for item in items:
images = []
messages = []
for msg in item.get("messages", []):
content = msg.get("content", "")
if isinstance(content, list):
new_content = []
for part in content:
if isinstance(part, dict) and part.get("type") == "image":
img = part.get("image")
if img is not None:
images.append(img)
new_content.append({"type": "image"})
else:
new_content.append(part)
messages.append({"role": msg["role"], "content": new_content})
else:
messages.append(msg)
out = {"messages": messages}
if images:
out["image"] = images[0] if len(images) == 1 else images
elif "image" in item:
out["image"] = item["image"]
elif "images" in item:
out["images"] = item["images"]
adapted.append(out)
return adapted
_MLX_STUDIO_OPTIM_MAP = {
"adamw_8bit": "adamw",
"paged_adamw_8bit": "adamw",
"adamw_bnb_8bit": "adamw",
"paged_adamw_32bit": "adamw",
"adamw_torch": "adamw",
"adamw_torch_fused": "adamw",
"adamw": "adamw",
"adafactor": "adafactor",
"sgd": "sgd",
"adam": "adam",
"muon": "muon",
"lion": "lion",
}
_MLX_STUDIO_LR_SCHEDULERS = {"linear", "cosine", "constant"}
def _normalize_mlx_studio_optimizer(value):
raw = str(value or "adamw_8bit").strip().lower()
try:
return _MLX_STUDIO_OPTIM_MAP[raw]
except KeyError:
supported = ", ".join(sorted(_MLX_STUDIO_OPTIM_MAP))
raise ValueError(
f"Unsupported optimizer for MLX training: {value!r}. "
f"Supported values: {supported}."
)
def _normalize_mlx_studio_scheduler(value):
raw = str(value or "linear").strip().lower()
if raw not in _MLX_STUDIO_LR_SCHEDULERS:
supported = ", ".join(sorted(_MLX_STUDIO_LR_SCHEDULERS))
raise ValueError(
f"Unsupported LR scheduler for MLX training: {value!r}. "
f"Supported values: {supported}."
)
return raw
def _run_mlx_training(event_queue, stop_queue, config):
"""Self-contained MLX training path for Apple Silicon.
Uses MLXTrainer from unsloth_zoo directly -- no torch/SFTTrainer needed.
Mirrors the event_queue protocol so the parent process pump works unchanged.
"""
import time
import gc
import math
import threading
import queue as _queue
from pathlib import Path
def _send(event_type, **kwargs):
if event_type == "status" and "message" not in kwargs:
sm = kwargs.get("status_message")
if sm is not None:
kwargs["message"] = sm
event_queue.put({"type": event_type, "ts": time.time(), **kwargs})
_send("status", status_message = "Loading MLX libraries...")
import mlx.core as mx
try:
from unsloth_zoo.mlx_loader import FastMLXModel
from unsloth_zoo.mlx_trainer import (
MLXTrainer,
MLXTrainingConfig,
train_on_responses_only,
)
except ImportError as e:
raise ImportError(
"Unsloth: MLX training requires unsloth-zoo with the MLX modules "
"(unsloth_zoo.mlx_loader / unsloth_zoo.mlx_trainer). Reinstall via "
"install.sh on Apple Silicon."
) from e
from datasets import load_dataset
if mx.metal.is_available():
info = mx.device_info()
rec_bytes = info.get("max_recommended_working_set_size", 0) or 0
if rec_bytes > 0:
memory_cap = int(rec_bytes * 0.85)
wired_cap = min(int(rec_bytes), memory_cap)
mx.set_memory_limit(memory_cap)
mx.set_wired_limit(wired_cap)
model_name = config["model_name"]
hf_token = config.get("hf_token") or None
if hf_token:
os.environ["HF_TOKEN"] = hf_token
if config.get("use_loftq"):
message = "LoftQ is not supported for MLX training yet."
_send("error", error = message)
raise NotImplementedError(message)
optim_name = _normalize_mlx_studio_optimizer(config.get("optim", "adamw_8bit"))
lr_scheduler_type = _normalize_mlx_studio_scheduler(
config.get("lr_scheduler_type", "linear")
)
# ── 1. Load model ──
# Force text-only if the dataset is not an image dataset, even if the model
# has vision capabilities (e.g. Qwen3.5-VL trained on plain alpaca text).
_send("status", status_message = f"Loading {model_name}...")
is_dataset_image = bool(config.get("is_dataset_image", False))
training_type = config.get("training_type", "LoRA/QLoRA")
use_lora = training_type == "LoRA/QLoRA"
model, tokenizer = FastMLXModel.from_pretrained(
model_name,
load_in_4bit = config.get("load_in_4bit", True),
full_finetuning = not use_lora,
text_only = None if is_dataset_image else True,
token = hf_token,
trust_remote_code = bool(config.get("trust_remote_code", False)),
random_state = config.get("random_seed", 3407),
)
is_vlm = bool(is_dataset_image and getattr(model, "_is_vlm_model", False))
model._is_vlm_model = is_vlm
# ── 2. Apply LoRA / full FT ──
# Pass gradient_checkpointing as string ("mlx"/"unsloth"/"none"/etc.)
# get_peft_model and MLXTrainer both accept strings and handle them.
gc_setting = config.get("gradient_checkpointing", "mlx")
if isinstance(gc_setting, str):
use_grad_checkpoint = (
gc_setting if gc_setting.lower() not in ("false", "") else False
)
else:
use_grad_checkpoint = gc_setting
if use_lora:
_send("status", status_message = "Configuring LoRA adapters...")
peft_kwargs = dict(
r = config.get("lora_r", 16),
lora_alpha = config.get("lora_alpha", 16),
lora_dropout = config.get("lora_dropout", 0.0),
use_rslora = config.get("use_rslora", False),
init_lora_weights = config.get("init_lora_weights", True),
random_state = config.get("random_seed", 3407),
target_modules = config.get("target_modules")
or [
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"gate_proj",
"up_proj",
"down_proj",
],
use_gradient_checkpointing = use_grad_checkpoint,
)
finetune_language = config.get("finetune_language_layers", True)
finetune_attention = config.get("finetune_attention_modules", True)
finetune_mlp = config.get("finetune_mlp_modules", True)
finetune_vision = (
config.get("finetune_vision_layers", False) if is_vlm else False
)
if (
(finetune_attention or finetune_mlp)
and not finetune_language
and not finetune_vision
):
finetune_language = True
peft_kwargs["finetune_language_layers"] = finetune_language
peft_kwargs["finetune_attention_modules"] = finetune_attention
peft_kwargs["finetune_mlp_modules"] = finetune_mlp
if is_vlm:
peft_kwargs["finetune_vision_layers"] = finetune_vision
model = FastMLXModel.get_peft_model(model, **peft_kwargs)
# ── 3. Load dataset ──
_send("status", status_message = "Loading dataset...")
hf_dataset = config.get("hf_dataset", "")
subset = config.get("subset")
train_split = config.get("train_split", "train") or "train"
eval_split = config.get("eval_split")
slice_start = config.get("dataset_slice_start")
slice_end = config.get("dataset_slice_end")
def _slice(ds):
if slice_start is not None or slice_end is not None:
start = slice_start if slice_start is not None else 0
end = slice_end if slice_end is not None else len(ds) - 1
if end < start:
return ds.select([])
ds = ds.select(range(start, min(end + 1, len(ds))))
return ds
def _load_local(file_paths):
from core.training.trainer import UnslothTrainer
from datasets import load_from_disk
if len(file_paths) == 1:
p = Path(file_paths[0])
if p.is_dir() and (
(p / "dataset_info.json").exists() or (p / "state.json").exists()
):
return load_from_disk(str(p))
all_files = UnslothTrainer._resolve_local_files(file_paths)
if not all_files:
raise ValueError("No local dataset files found")
loader = UnslothTrainer._loader_for_files(all_files)
return load_dataset(loader, data_files = all_files, split = "train")
if hf_dataset:
load_kwargs = {"split": train_split, "token": hf_token}
if subset:
load_kwargs["name"] = subset
dataset = load_dataset(hf_dataset, **load_kwargs)
dataset = _slice(dataset)
elif config.get("local_datasets"):
dataset = _load_local(config["local_datasets"])
dataset = _slice(dataset)
else:
raise ValueError("No dataset specified")
# Eval dataset (separate split or local file)
eval_dataset = None
if eval_split and hf_dataset:
eval_kwargs = {"split": eval_split, "token": hf_token}
if subset:
eval_kwargs["name"] = subset
try:
eval_dataset = load_dataset(hf_dataset, **eval_kwargs)
except Exception as e:
_send("status", status_message = f"Eval split load failed: {e}")
eval_dataset = None
elif config.get("local_eval_datasets"):
eval_dataset = _load_local(config["local_eval_datasets"])
# ── 3b. Format dataset (VLM or text) ──
# Reuse the GPU path's format pipeline for both VLM (auto-detects OCR/caption/
# llava/sharegpt+images) and text (alpaca/sharegpt/chatml → "text" column).
format_type = config.get("format_type", "")
try:
from utils.datasets import format_and_template_dataset
def _fmt_progress(status_message = "", **_kw):
_send("status", status_message = status_message)
if is_vlm:
_send("status", status_message = "Formatting VLM dataset...")
vlm_info = format_and_template_dataset(
dataset,
model_name = model_name,
tokenizer = tokenizer,
is_vlm = True,
dataset_name = hf_dataset or "local",
progress_callback = _fmt_progress,
)
if vlm_info.get("success"):
dataset = _adapt_for_mlx_vlm(vlm_info["dataset"])
else:
errors = vlm_info.get("errors", [])
raise ValueError(
f"VLM dataset format conversion failed: {'; '.join(errors)}"
)
if eval_dataset is not None:
ev_info = format_and_template_dataset(
eval_dataset,
model_name = model_name,
tokenizer = tokenizer,
is_vlm = True,
dataset_name = hf_dataset or "local",
)
if ev_info.get("success"):
eval_dataset = _adapt_for_mlx_vlm(ev_info["dataset"])
elif format_type:
_send("status", status_message = f"Formatting dataset ({format_type})...")
info = format_and_template_dataset(
dataset,
model_name = model_name,
tokenizer = tokenizer,
is_vlm = False,
format_type = format_type,
dataset_name = hf_dataset or "local",
)
if info.get("success", True):
dataset = info.get("dataset", dataset)
if eval_dataset is not None:
ev = format_and_template_dataset(
eval_dataset,
model_name = model_name,
tokenizer = tokenizer,
is_vlm = False,
format_type = format_type,
dataset_name = hf_dataset or "local",
)
if ev.get("success", True):
eval_dataset = ev.get("dataset", eval_dataset)
except ImportError:
_send("status", status_message = "Format helper unavailable, using raw dataset")
# ── 4. Resolve training steps ──
max_steps = config.get("max_steps", 0) or 0
num_epochs = config.get("num_epochs", 3)
max_seq_length = config.get("max_seq_length", 2048)
batch_size = config.get("batch_size", 4)
grad_accum = config.get("gradient_accumulation_steps", 4)
if max_steps <= 0:
max_steps = max(
1,
math.ceil(len(dataset) / batch_size / grad_accum) * num_epochs,
)
lr_value = float(config.get("learning_rate", "2e-4"))
# Warmup: prefer warmup_steps; fall back to warmup_ratio
warmup_steps = config.get("warmup_steps")
warmup_ratio = config.get("warmup_ratio")
if warmup_steps is None and warmup_ratio is not None:
warmup_steps = int(round(warmup_ratio * max_steps))
if warmup_steps is None:
warmup_steps = 5
# ── 5. Build output dir ──
output_dir = config.get("output_dir", "")
if not output_dir:
output_dir = f"{model_name.replace('/', '_')}_{int(time.time())}"
# Resolve to ~/.unsloth/studio/outputs/ so the export page can find it
from utils.paths import resolve_output_dir, ensure_dir
output_dir = str(resolve_output_dir(output_dir))
ensure_dir(Path(output_dir))
# ── 6. Create trainer ──
eval_steps_val = config.get("eval_steps", 0) or 0
if isinstance(eval_steps_val, float) and 0 < eval_steps_val < 1:
# Studio sometimes sends fraction-of-total-steps
eval_steps_val = max(1, int(eval_steps_val * max_steps))
else:
eval_steps_val = int(eval_steps_val)
trainer = MLXTrainer(
model = model,
tokenizer = tokenizer,
train_dataset = dataset,
eval_dataset = eval_dataset,
args = MLXTrainingConfig(
per_device_train_batch_size = batch_size,
gradient_accumulation_steps = grad_accum,
max_steps = max_steps,
learning_rate = lr_value,
warmup_steps = warmup_steps,
lr_scheduler_type = lr_scheduler_type,
optim = optim_name,
weight_decay = float(config.get("weight_decay", 0.001) or 0.001),
logging_steps = 1,
max_seq_length = max_seq_length,
seed = config.get("random_seed", 3407),
use_cce = True,
compile = True,
gradient_checkpointing = use_grad_checkpoint,
streaming = is_vlm,
packing = bool(config.get("packing", False)),
output_dir = output_dir,
save_steps = int(config.get("save_steps", 0) or 0),
eval_steps = eval_steps_val,
),
)
# Tell the parent that eval is configured so the frontend shows the eval chart
if eval_dataset is not None and eval_steps_val > 0:
_send("eval_configured")
# ── 7. Apply train_on_responses_only if requested ──
if config.get("train_on_completions", False):
_send("status", status_message = "Configuring response-only training...")
try:
from utils.datasets import (
MODEL_TO_TEMPLATE_MAPPER,
TEMPLATE_TO_RESPONSES_MAPPER,
)
template_name = MODEL_TO_TEMPLATE_MAPPER.get(model_name.lower())
markers = (
TEMPLATE_TO_RESPONSES_MAPPER.get(template_name)
if template_name
else None
)
if markers:
trainer = train_on_responses_only(
trainer,
instruction_part = markers["instruction"],
response_part = markers["response"],
)
else:
_send(
"status",
status_message = f"train_on_completions skipped (no template for {model_name})",
)
except Exception as e:
_send("status", status_message = f"train_on_completions failed: {e}")
# ── 8. Setup wandb / tensorboard ──
wandb_run = None
tb_writer = None
if config.get("enable_wandb", False):
try:
import wandb as _wandb
wandb_token = config.get("wandb_token")
if wandb_token:
os.environ["WANDB_API_KEY"] = wandb_token
_wandb_sensitive = {"hf_token", "wandb_token"}
wandb_run = _wandb.init(
project = config.get("wandb_project") or "unsloth-mlx",
config = {k: v for k, v in config.items() if k not in _wandb_sensitive},
reinit = True,
)
except Exception as e:
_send("status", status_message = f"wandb init failed: {e}")
if config.get("enable_tensorboard", False):
try:
from tensorboardX import SummaryWriter
except ImportError:
try:
from torch.utils.tensorboard import SummaryWriter
except ImportError:
SummaryWriter = None
if SummaryWriter is not None:
try:
tb_dir = config.get("tensorboard_dir") or f"{output_dir}/runs"
tb_writer = SummaryWriter(log_dir = tb_dir)
except Exception as e:
_send("status", status_message = f"tensorboard init failed: {e}")
else:
_send(
"status",
status_message = "tensorboard unavailable (install tensorboardX)",
)
# ── 9. Real-time progress callback ──
_send("status", status_message = f"Training {model_name}...")
def _on_step(step, total, loss, lr, tok_s, peak_gb, elapsed, num_tokens):
eta = (elapsed / step * (total - step)) if step > 0 else 0
_send(
"progress",
step = step,
epoch = round(step / total * num_epochs, 2) if total > 0 else 0,
loss = loss,
learning_rate = lr,
total_steps = total,
elapsed_seconds = elapsed,
eta_seconds = max(0, eta),
grad_norm = None,
num_tokens = num_tokens,
eval_loss = None,
status_message = None,
peak_memory_gb = peak_gb,
)
if wandb_run is not None:
try:
wandb_run.log(
{
"train/loss": loss,
"train/learning_rate": lr,
"train/tokens_per_sec": tok_s,
"train/peak_gb": peak_gb,
"train/num_tokens": num_tokens,
},
step = step,
)
except Exception:
pass
if tb_writer is not None:
try:
tb_writer.add_scalar("train/loss", loss, step)
tb_writer.add_scalar("train/learning_rate", lr, step)
tb_writer.add_scalar("train/tokens_per_sec", tok_s, step)
tb_writer.add_scalar("train/peak_gb", peak_gb, step)
except Exception:
pass
trainer.add_step_callback(_on_step)
def _on_eval(step, eval_loss, perplexity):
_send("progress", step = step, eval_loss = eval_loss)
if wandb_run is not None:
try:
wandb_run.log(
{"eval/loss": eval_loss, "eval/perplexity": perplexity}, step = step
)
except Exception:
pass
if tb_writer is not None:
try:
tb_writer.add_scalar("eval/loss", eval_loss, step)
tb_writer.add_scalar("eval/perplexity", perplexity, step)
except Exception:
pass
trainer.add_eval_callback(_on_eval)
# ── 10. Stop signal polling ──
_stop_save = [True] # mutable so thread can update; [save_flag]
def _poll_stop():
while True:
try:
msg = stop_queue.get(timeout = 1.0)
if msg and msg.get("type") == "stop":
_stop_save[0] = msg.get("save", True)
trainer.stop_requested = True
return
except _queue.Empty:
continue
except (EOFError, OSError):
# why safe: pipe permanently broken, no further messages can arrive
return
stop_thread = threading.Thread(target = _poll_stop, daemon = True)
stop_thread.start()
# ── 11. Run training ──
gc.collect()
mx.synchronize()
trainer.train()
# ── 12. Save and finalize ──
if trainer.stop_requested and not _stop_save[0]:
# User clicked "Cancel" (save=False) — skip saving
_send("complete", output_dir = None, status_message = "Training cancelled")
else:
_send("status", status_message = "Saving model...")
mx.synchronize()
trainer.save_model(output_dir)
_send("complete", output_dir = output_dir, status_message = "Training completed")
if tb_writer is not None:
try:
tb_writer.close()
except Exception:
pass
if wandb_run is not None:
try:
wandb_run.finish()
except Exception:
pass
def run_training_process(
*,
event_queue: Any,
@ -371,6 +960,46 @@ def run_training_process(
model_name = config["model_name"]
# ── 0. MLX FAST-PATH (must run before any torch/transformers imports) ──
# Apple Silicon uses MLXTrainer directly -- skip transformers version
# activation, causal-conv1d install, and torch imports entirely.
backend_path = str(Path(__file__).resolve().parent.parent.parent)
if backend_path not in sys.path:
sys.path.insert(0, backend_path)
from utils.hardware import hardware as _hw
_hw.detect_hardware()
if _hw.DEVICE == _hw.DeviceType.MLX:
if config.get("is_dataset_audio"):
event_queue.put(
{
"type": "error",
"error": "Audio dataset training is not yet supported on Apple Silicon.",
"stack": "",
"ts": time.time(),
}
)
return
# Activate correct transformers version (Gemma-4 needs 5.5.0, etc.)
# Must happen before any transformers/mlx-lm imports in _run_mlx_training.
try:
_activate_transformers_version(model_name)
except Exception:
pass # Non-fatal: fall through with whatever version is installed
try:
_run_mlx_training(event_queue, stop_queue, config)
except Exception as exc:
event_queue.put(
{
"type": "error",
"error": str(exc),
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
}
)
return
# ── 1. Activate correct transformers version BEFORE any ML imports ──
try:
_activate_transformers_version(model_name)
@ -580,6 +1209,8 @@ def run_training_process(
# ── 4b. Load and format dataset (LLM helper may use VRAM briefly) ──
_send_status(event_queue, "Loading and formatting dataset...")
hf_dataset = config.get("hf_dataset", "")
training_type = config.get("training_type", "LoRA/QLoRA")
_is_cpt_for_dataset = training_type == "Continued Pretraining"
dataset_result = trainer.load_and_format_dataset(
dataset_source = hf_dataset if hf_dataset and hf_dataset.strip() else None,
format_type = config.get("format_type", ""),
@ -592,6 +1223,7 @@ def run_training_process(
eval_steps = config.get("eval_steps", 0.00),
dataset_slice_start = config.get("dataset_slice_start"),
dataset_slice_end = config.get("dataset_slice_end"),
is_cpt = _is_cpt_for_dataset,
)
if isinstance(dataset_result, tuple):
@ -677,7 +1309,9 @@ def run_training_process(
_tqdm_thread.start()
training_type = config.get("training_type", "LoRA/QLoRA")
use_lora = training_type == "LoRA/QLoRA"
is_cpt = training_type == "Continued Pretraining"
use_lora = training_type in ("LoRA/QLoRA", "Continued Pretraining")
cpt_trains_embeddings = False
# ── 4c. Load training model (uses VRAM — dataset already formatted) ──
_send_status(event_queue, "Loading model...")
@ -709,8 +1343,41 @@ def run_training_process(
)
return
# ── 4d. Prepare model (LoRA or full finetuning) ──
if use_lora:
# ── 4d. Prepare model (LoRA, full finetuning, or CPT) ──
if is_cpt:
_send_status(event_queue, "Configuring LoRA for continued pretraining...")
# embed_tokens (if the user included it) goes to modules_to_save —
# trained full-precision at embedding_learning_rate. lm_head stays as
# a LoRA target for merge compatibility (see unsloth PR #4106).
_user_modules = config.get("target_modules") or []
wants_embed = "embed_tokens" in _user_modules
cpt_trains_embeddings = wants_embed
cpt_target_modules = [m for m in _user_modules if m != "embed_tokens"]
if not cpt_target_modules:
cpt_target_modules = [
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"gate_proj",
"up_proj",
"down_proj",
"lm_head",
]
success = trainer.prepare_model_for_training(
use_lora = True,
target_modules = cpt_target_modules,
modules_to_save = ["embed_tokens"] if wants_embed else None,
lora_r = config.get("lora_r", 128),
lora_alpha = config.get("lora_alpha", 32),
lora_dropout = config.get("lora_dropout", 0.0),
use_gradient_checkpointing = config.get(
"gradient_checkpointing", "unsloth"
),
use_rslora = config.get("use_rslora", False),
use_loftq = config.get("use_loftq", False),
)
elif use_lora:
_send_status(event_queue, "Configuring LoRA adapters...")
success = trainer.prepare_model_for_training(
use_lora = True,
@ -751,9 +1418,9 @@ def run_training_process(
)
return
# Convert learning rate
lr_default = "5e-5" if is_cpt else "2e-4"
try:
lr_value = float(config.get("learning_rate", "2e-4"))
lr_value = float(config.get("learning_rate", lr_default))
except ValueError:
event_queue.put(
{
@ -765,6 +1432,25 @@ def run_training_process(
)
return
# embedding_learning_rate is validated by the Pydantic model (Optional[float],
# gt=0, lt=1.0); if present it is already a finite float in range.
embedding_lr_value = config.get("embedding_learning_rate")
if is_cpt:
if cpt_trains_embeddings:
if embedding_lr_value is None:
# Default embedding_learning_rate = lr/10 per Unsloth's CPT notebook.
embedding_lr_value = lr_value / 10.0
logger.info(
f"CPT: using default embedding_learning_rate={embedding_lr_value:.1e} "
f"(lr/10). Set explicitly to override.\n"
)
elif embedding_lr_value is not None:
logger.warning(
"CPT: embedding_learning_rate was provided but embed_tokens is "
"not being trained; ignoring the override.\n"
)
embedding_lr_value = None
# Generate output dir
resume_from_checkpoint = config.get("resume_from_checkpoint")
output_dir = config.get("output_dir") or _output_dir_from_resume_checkpoint(
@ -797,6 +1483,7 @@ def run_training_process(
output_dir = output_dir,
num_epochs = config.get("num_epochs", 3),
learning_rate = lr_value,
embedding_learning_rate = embedding_lr_value,
batch_size = config.get("batch_size", 2),
gradient_accumulation_steps = config.get("gradient_accumulation_steps", 4),
warmup_steps = config.get("warmup_steps"),
@ -806,7 +1493,9 @@ def run_training_process(
weight_decay = config.get("weight_decay", 0.001),
random_seed = config.get("random_seed", 3407),
packing = config.get("packing", False),
train_on_completions = config.get("train_on_completions", False),
train_on_completions = False
if is_cpt
else config.get("train_on_completions", False),
enable_wandb = config.get("enable_wandb", False),
wandb_project = config.get("wandb_project", "unsloth-training"),
wandb_token = config.get("wandb_token"),
@ -817,6 +1506,7 @@ def run_training_process(
max_seq_length = config.get("max_seq_length", 2048),
optim = config.get("optim", "adamw_8bit"),
lr_scheduler_type = config.get("lr_scheduler_type", "linear"),
is_cpt = is_cpt,
resume_from_checkpoint = resume_from_checkpoint,
)

View file

@ -23,12 +23,67 @@ if _backend_dir not in sys.path:
# See: https://github.com/python/cpython/issues/102396
import _platform_compat # noqa: F401
# Direct `uvicorn main:app` launches bypass run.py, so re-export here too
# (mirrors run.py). Required BEFORE the unsloth-zoo import below, since
# its LLAMA_CPP_DEFAULT_DIR binding is import-time.
from utils.paths.storage_roots import studio_root as _studio_root
try:
_LEGACY_STUDIO_ROOT = (_Path.home() / ".unsloth" / "studio").resolve()
except (OSError, ValueError):
_LEGACY_STUDIO_ROOT = _Path.home() / ".unsloth" / "studio"
try:
_STUDIO_ROOT_RESOLVED = _studio_root().resolve()
except (OSError, ValueError):
_STUDIO_ROOT_RESOLVED = _studio_root()
if _STUDIO_ROOT_RESOLVED != _LEGACY_STUDIO_ROOT:
if not os.environ.get("UNSLOTH_STUDIO_HOME"):
os.environ["UNSLOTH_STUDIO_HOME"] = str(_STUDIO_ROOT_RESOLVED)
if not os.environ.get("UNSLOTH_LLAMA_CPP_PATH"):
os.environ["UNSLOTH_LLAMA_CPP_PATH"] = str(_STUDIO_ROOT_RESOLVED / "llama.cpp")
import mimetypes
import re as _re
import shutil
import warnings
from contextlib import asynccontextmanager
from importlib.metadata import PackageNotFoundError, version as package_version
_STUDIO_INSTALL_ID_RE = _re.compile(r"^[0-9a-f]{64}$")
def _read_studio_install_id() -> str:
"""Per-install opaque id written by install.sh / install.ps1 at
$STUDIO_HOME/share/studio_install_id. Returns "" when the file is
absent (pre-PR install, fresh tree never run through the installer)
or contains anything other than a 64-char lowercase-hex token --
in which case /api/health emits "" and the launcher's _check_health
falls back to the existing "no baked id, accept any healthy
Unsloth backend" path. This intentionally replaces a previous
sha256(resolved_install_path) so the field carries no install-path
information for callers reaching /api/health (relevant when Studio
is run with -H 0.0.0.0)."""
try:
token = (
(_STUDIO_ROOT_RESOLVED / "share" / "studio_install_id").read_text().strip()
)
except (OSError, ValueError):
return ""
return token if _STUDIO_INSTALL_ID_RE.fullmatch(token) else ""
_STUDIO_ROOT_ID_CACHE: str = _read_studio_install_id()
def _studio_root_id() -> str:
"""Same-install discriminator for /api/health: a per-install opaque
token written once by the installer and read once at module import.
Empty when no installer-written token is present; the launcher
contract treats "" as "no baked id, accept any healthy backend"."""
return _STUDIO_ROOT_ID_CACHE
# Fix broken Windows registry MIME types. Some Windows installs map .js to
# "text/plain" in the registry (HKCR\.js\Content Type). Python's mimetypes
# module reads from the registry, and FastAPI/Starlette's StaticFiles uses
@ -245,6 +300,10 @@ async def health_check():
"chat_only": _hw_module.CHAT_ONLY,
"desktop_protocol_version": 1,
"supports_desktop_auth": True,
# why: launchers compare against an install-time hash so a sibling
# Studio on the same port is rejected; hex digest avoids leaking the
# raw install path on -H 0.0.0.0.
"studio_root_id": _studio_root_id(),
"native_path_leases_supported": native_path_leases_supported(),
}

View file

@ -16,8 +16,11 @@ class TrainingStartRequest(BaseModel):
model_name: str = Field(
..., description = "Model identifier (e.g., 'unsloth/llama-3-8b-bnb-4bit')"
)
training_type: str = Field(
..., description = "Training type: 'LoRA/QLoRA' or 'Full Finetuning'"
training_type: Literal["LoRA/QLoRA", "Full Finetuning", "Continued Pretraining"] = (
Field(
...,
description = "Training type: 'LoRA/QLoRA', 'Full Finetuning', or 'Continued Pretraining'",
)
)
hf_token: Optional[str] = Field(None, description = "HuggingFace token")
load_in_4bit: bool = Field(True, description = "Load model in 4-bit quantization")
@ -86,6 +89,13 @@ class TrainingStartRequest(BaseModel):
packing: bool = Field(False, description = "Enable sequence packing")
optim: str = Field("adamw_8bit", description = "Optimizer")
lr_scheduler_type: str = Field("linear", description = "Learning rate scheduler type")
embedding_learning_rate: Optional[float] = Field(
None,
gt = 0,
lt = 1.0,
description = "Separate learning rate for embedding matrices (CPT). "
"Must be in (0, 1). Should be 2-10x smaller than the main learning rate.",
)
# LoRA parameters
use_lora: bool = Field(True, description = "Use LoRA (derived from training_type)")

View file

@ -207,6 +207,7 @@ async def start_training(
"custom_format_mapping": request.custom_format_mapping,
"num_epochs": request.num_epochs,
"learning_rate": request.learning_rate,
"embedding_learning_rate": request.embedding_learning_rate,
"batch_size": request.batch_size,
"gradient_accumulation_steps": request.gradient_accumulation_steps,
"warmup_steps": request.warmup_steps,

View file

@ -159,7 +159,27 @@ def _find_free_port(host: str, start: int, max_attempts: int = 20) -> int:
)
_PID_FILE = Path.home() / ".unsloth" / "studio" / "studio.pid"
from utils.paths.storage_roots import studio_root as _studio_root
_PID_FILE = _studio_root() / "studio.pid"
# Direct backend launches bypass the CLI's env re-export; do it here for
# real custom roots so unsloth-zoo's import-time LLAMA_CPP_DEFAULT_DIR
# picks up the custom build. Skip for legacy-default to avoid flipping
# default-mode installs into env-override.
try:
_LEGACY_STUDIO_ROOT = (Path.home() / ".unsloth" / "studio").resolve()
except (OSError, ValueError):
_LEGACY_STUDIO_ROOT = Path.home() / ".unsloth" / "studio"
try:
_STUDIO_ROOT_RESOLVED = _studio_root().resolve()
except (OSError, ValueError):
_STUDIO_ROOT_RESOLVED = _studio_root()
if _STUDIO_ROOT_RESOLVED != _LEGACY_STUDIO_ROOT:
if not os.environ.get("UNSLOTH_STUDIO_HOME"):
os.environ["UNSLOTH_STUDIO_HOME"] = str(_STUDIO_ROOT_RESOLVED)
if not os.environ.get("UNSLOTH_LLAMA_CPP_PATH"):
os.environ["UNSLOTH_LLAMA_CPP_PATH"] = str(_STUDIO_ROOT_RESOLVED / "llama.cpp")
def _write_pid_file():

View file

@ -0,0 +1,157 @@
# SPDX-License-Identifier: AGPL-3.0-only
import sys
import types
from types import SimpleNamespace
class _DummyMetal:
@staticmethod
def is_available():
return False
class _DummyMX:
metal = _DummyMetal()
@staticmethod
def set_wired_limit(_limit):
return None
@staticmethod
def device_info():
return {"max_recommended_working_set_size": 1024}
class _DummyTokenizer:
pass
class _DummyProcessor:
tokenizer = _DummyTokenizer()
class _DummyModel:
pass
def _install_fake_mlx(monkeypatch):
mlx_pkg = types.ModuleType("mlx")
mlx_core = types.ModuleType("mlx.core")
mlx_core.metal = _DummyMetal()
mlx_core.set_wired_limit = _DummyMX.set_wired_limit
mlx_core.device_info = _DummyMX.device_info
mlx_pkg.core = mlx_core
monkeypatch.setitem(sys.modules, "mlx", mlx_pkg)
monkeypatch.setitem(sys.modules, "mlx.core", mlx_core)
def _install_fake_fast_mlx(monkeypatch, calls):
class _FastMLXModel:
@staticmethod
def from_pretrained(*args, **kwargs):
calls.append((args, kwargs))
if kwargs["text_only"] is False:
return _DummyModel(), _DummyProcessor()
return _DummyModel(), _DummyTokenizer()
unsloth_zoo_pkg = types.ModuleType("unsloth_zoo")
mlx_loader = types.ModuleType("unsloth_zoo.mlx_loader")
mlx_loader.FastMLXModel = _FastMLXModel
unsloth_zoo_pkg.mlx_loader = mlx_loader
monkeypatch.setitem(sys.modules, "unsloth_zoo", unsloth_zoo_pkg)
monkeypatch.setitem(sys.modules, "unsloth_zoo.mlx_loader", mlx_loader)
def test_mlx_inference_text_load_forwards_studio_settings(monkeypatch):
_install_fake_mlx(monkeypatch)
calls = []
_install_fake_fast_mlx(monkeypatch, calls)
from core.inference.mlx_inference import MLXInferenceBackend
backend = MLXInferenceBackend()
config = SimpleNamespace(identifier = "fake/text", is_vision = False, is_lora = False)
assert backend.load_model(
config,
max_seq_length = 4096,
load_in_4bit = False,
hf_token = "hf-token",
trust_remote_code = True,
dtype = "float16",
)
assert calls == [
(
("fake/text",),
{
"max_seq_length": 4096,
"dtype": "float16",
"load_in_4bit": False,
"token": "hf-token",
"trust_remote_code": True,
"text_only": True,
},
)
]
assert backend._is_vlm is False
assert isinstance(backend._tokenizer, _DummyTokenizer)
def test_mlx_inference_vlm_lora_uses_unsloth_loader_without_native_adapter_rewrite(
monkeypatch,
tmp_path,
):
_install_fake_mlx(monkeypatch)
calls = []
_install_fake_fast_mlx(monkeypatch, calls)
def _native_vlm_load(*_args, **_kwargs):
raise AssertionError("Studio MLX VLM inference must use FastMLXModel")
mlx_vlm = types.ModuleType("mlx_vlm")
mlx_vlm.load = _native_vlm_load
monkeypatch.setitem(sys.modules, "mlx_vlm", mlx_vlm)
adapter_dir = tmp_path / "adapter"
adapter_dir.mkdir()
cfg_path = adapter_dir / "adapter_config.json"
original_cfg = '{"base_model_name_or_path": "fake/base", "rank": 8}\n'
cfg_path.write_text(original_cfg)
from core.inference.mlx_inference import MLXInferenceBackend
backend = MLXInferenceBackend()
config = SimpleNamespace(
identifier = str(adapter_dir),
is_vision = True,
is_lora = True,
base_model = "fake/base",
)
assert backend.load_model(
config,
max_seq_length = 8192,
load_in_4bit = True,
hf_token = "hf-token",
trust_remote_code = True,
)
assert calls == [
(
(str(adapter_dir),),
{
"max_seq_length": 8192,
"dtype": None,
"load_in_4bit": True,
"token": "hf-token",
"trust_remote_code": True,
"text_only": False,
},
)
]
assert cfg_path.read_text() == original_cfg
assert backend._is_vlm is True
assert isinstance(backend._processor, _DummyProcessor)
assert isinstance(backend._tokenizer, _DummyTokenizer)

View file

@ -0,0 +1,83 @@
# SPDX-License-Identifier: AGPL-3.0-only
import importlib.util
import sys
import types
from pathlib import Path
import pytest
def _load_worker_module():
stub_names = (
"structlog",
"loggers",
"utils",
"utils.hardware",
"utils.wheel_utils",
)
previous_modules = {name: sys.modules.get(name) for name in stub_names}
try:
sys.modules["structlog"] = types.ModuleType("structlog")
loggers = types.ModuleType("loggers")
loggers.get_logger = lambda *_args, **_kwargs: None
sys.modules["loggers"] = loggers
utils = types.ModuleType("utils")
utils.__path__ = []
sys.modules["utils"] = utils
hardware = types.ModuleType("utils.hardware")
hardware.apply_gpu_ids = lambda *_args, **_kwargs: None
sys.modules["utils.hardware"] = hardware
wheel_utils = types.ModuleType("utils.wheel_utils")
for name in (
"direct_wheel_url",
"flash_attn_wheel_url",
"install_wheel",
"probe_torch_wheel_env",
"url_exists",
):
setattr(wheel_utils, name, lambda *_args, **_kwargs: None)
sys.modules["utils.wheel_utils"] = wheel_utils
worker_path = (
Path(__file__).resolve().parents[1] / "core" / "training" / "worker.py"
)
spec = importlib.util.spec_from_file_location(
"mlx_training_worker_under_test", worker_path
)
module = importlib.util.module_from_spec(spec)
assert spec.loader is not None
spec.loader.exec_module(module)
return module
finally:
for name, module in previous_modules.items():
if module is None:
sys.modules.pop(name, None)
else:
sys.modules[name] = module
_worker = _load_worker_module()
_normalize_mlx_studio_optimizer = _worker._normalize_mlx_studio_optimizer
_normalize_mlx_studio_scheduler = _worker._normalize_mlx_studio_scheduler
def test_mlx_studio_optimizer_aliases_are_explicit():
assert _normalize_mlx_studio_optimizer("adamw_8bit") == "adamw"
assert _normalize_mlx_studio_optimizer("paged_adamw_8bit") == "adamw"
assert _normalize_mlx_studio_optimizer("adafactor") == "adafactor"
def test_mlx_studio_rejects_unknown_optimizer():
with pytest.raises(ValueError, match = "Unsupported optimizer for MLX training"):
_normalize_mlx_studio_optimizer("adamw_typo")
def test_mlx_studio_rejects_unknown_scheduler():
with pytest.raises(ValueError, match = "Unsupported LR scheduler for MLX training"):
_normalize_mlx_studio_scheduler("linear_typo")

View file

@ -0,0 +1,183 @@
# 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 asyncio
import importlib.util
import unittest
from pathlib import Path
from unittest.mock import patch
from datasets import Dataset
from core.training.training import TrainingBackend
from models.training import TrainingStartRequest
from utils.datasets import format_dataset, format_and_template_dataset
from utils.datasets.raw_text import prepare_raw_text_dataset
_BACKEND_ROOT = Path(__file__).resolve().parent.parent
def _load_route_module(name: str, relative_path: str):
spec = importlib.util.spec_from_file_location(name, _BACKEND_ROOT / relative_path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module
class TestTrainingRawSupport(unittest.TestCase):
def test_training_backend_preserves_cpt_4bit_and_embedding_lr(self):
backend = TrainingBackend()
class DummyProcess:
pid = 12345
def start(self):
return None
class DummyThread:
def start(self):
return None
dummy_queue = object()
with (
patch(
"core.training.training.prepare_gpu_selection",
return_value = ([0], {"selection_mode": "auto"}),
),
patch(
"core.training.training._CTX.Queue",
side_effect = [dummy_queue, dummy_queue],
),
patch(
"core.training.training._CTX.Process", return_value = DummyProcess()
) as mock_process,
patch(
"core.training.training.threading.Thread",
return_value = DummyThread(),
),
):
backend.start_training(
job_id = "test-cpt-raw",
model_name = "unsloth/test-bnb-4bit",
training_type = "Continued Pretraining",
format_type = "raw",
load_in_4bit = True,
embedding_learning_rate = 1e-5,
)
config = mock_process.call_args.kwargs["kwargs"]["config"]
self.assertTrue(config["load_in_4bit"])
self.assertEqual(config["embedding_learning_rate"], 1e-5)
def test_training_route_forwards_embedding_learning_rate(self):
training_route = _load_route_module(
"training_route_module_raw_support",
"routes/training.py",
)
captured: dict = {}
class DummyBackend:
current_job_id = None
def is_training_active(self):
return False
def start_training(self, **kwargs):
captured.update(kwargs)
return True
request = TrainingStartRequest(
model_name = "unsloth/test-bnb-4bit",
training_type = "Continued Pretraining",
format_type = "raw",
load_in_4bit = True,
embedding_learning_rate = 1e-5,
)
with (
patch.object(
training_route,
"get_training_backend",
return_value = DummyBackend(),
),
patch.object(training_route, "load_model_defaults", return_value = {}),
patch(
"core.inference.get_inference_backend",
return_value = type(
"InferenceBackend",
(),
{"active_model_name": None},
)(),
),
patch(
"core.export.get_export_backend",
return_value = type(
"ExportBackend",
(),
{"current_checkpoint": None},
)(),
),
):
response = asyncio.run(
training_route.start_training(request, current_subject = "test-user")
)
self.assertEqual(response.status, "queued")
self.assertEqual(captured["embedding_learning_rate"], 1e-5)
self.assertTrue(captured["load_in_4bit"])
def test_format_dataset_supports_raw_text(self):
dataset = Dataset.from_dict(
{
"body": ["hello", "world"],
"title": ["a", "b"],
"id": [1, 2],
}
)
result = format_dataset(dataset, format_type = "raw")
self.assertEqual(result["final_format"], "raw_text")
self.assertIn("text", result["dataset"].column_names)
self.assertEqual(result["dataset"][0]["text"], "hello")
self.assertFalse(result["requires_manual_mapping"])
def test_format_and_template_dataset_supports_raw_text_without_template(self):
dataset = Dataset.from_dict({"body": ["hello raw world"]})
result = format_and_template_dataset(
dataset,
model_name = "unsloth/test",
tokenizer = None,
format_type = "raw",
)
self.assertTrue(result["success"])
self.assertEqual(result["final_format"], "raw_text")
self.assertEqual(result["dataset"][0]["text"], "hello raw world")
def test_prepare_raw_text_dataset_drops_null_rows_before_appending_eos(self):
dataset = Dataset.from_dict({"text": ["hello", None, "world"]})
result = prepare_raw_text_dataset(
dataset,
mode_label = "CPT",
split_name = "train",
eos_token = "<eos>",
append_eos = True,
)
self.assertEqual(len(result.dataset), 2)
self.assertEqual(result.dataset[0]["text"], "hello<eos>")
self.assertEqual(result.dataset[1]["text"], "world<eos>")
self.assertTrue(
any(
"null or non-string 'text' values" in notice.message
for notice in result.notices
)
)
if __name__ == "__main__":
unittest.main()

View file

@ -41,6 +41,7 @@ from .chat_templates import (
get_tokenizer_chat_template,
DEFAULT_ALPACA_TEMPLATE,
)
from .raw_text import prepare_raw_text_dataset
from .vlm_processing import generate_smart_vlm_instruction
from .data_collators import DeepSeekOCRDataCollator, VLMDataCollator
from .model_mappings import TEMPLATE_TO_MODEL_MAPPER
@ -437,6 +438,20 @@ def format_dataset(
# Detect multimodal first (needed for all flows)
multimodal_info = detect_multimodal_dataset(dataset)
if format_type == "raw":
raw_result = prepare_raw_text_dataset(dataset)
return {
"dataset": raw_result.dataset,
"detected_format": "raw_text",
"final_format": "raw_text",
"chat_column": "text",
"is_standardized": True,
"requires_manual_mapping": False,
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": [notice.message for notice in raw_result.notices],
}
# If user provided explicit mapping, skip detection and apply in the requested format
if custom_format_mapping:
try:
@ -1105,6 +1120,21 @@ def format_and_template_dataset(
num_proc = num_proc,
)
if dataset_info["final_format"] == "raw_text":
summary = get_dataset_info_summary(dataset_info)
return {
"dataset": dataset_info["dataset"],
"detected_format": dataset_info["detected_format"],
"final_format": dataset_info["final_format"],
"chat_column": dataset_info.get("chat_column"),
"is_vlm": False,
"success": True,
"requires_manual_mapping": False,
"warnings": dataset_info.get("warnings", []),
"errors": [],
"summary": summary,
}
# Step 2: Apply chat template
detected = dataset_info.get("detected_format", "unknown")
if progress_callback and n_rows:

View file

@ -0,0 +1,142 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""
Shared helpers for raw-text dataset preparation.
"""
from dataclasses import dataclass
from typing import Literal
from datasets import Dataset
@dataclass(frozen = True)
class RawTextNotice:
message: str
level: Literal["info", "warning"]
update_status: bool = False
@dataclass(frozen = True)
class RawTextPreparationResult:
dataset: Dataset
notices: list[RawTextNotice]
def _string_columns(dataset: Dataset) -> list[str]:
feature_map = getattr(dataset, "features", {}) or {}
string_cols: list[str] = []
for col in dataset.column_names:
feature = feature_map.get(col)
dtype = str(getattr(feature, "dtype", ""))
if dtype in {"string", "large_string"}:
string_cols.append(col)
return string_cols
def _split_scope(split_name: str | None) -> str:
return f"the {split_name} split" if split_name else "this dataset"
def _drop_invalid_text_rows(
dataset: Dataset,
*,
mode_title: str,
split_scope: str,
) -> tuple[Dataset, list[RawTextNotice]]:
filtered_dataset = dataset.filter(lambda ex: isinstance(ex["text"], str))
dropped_rows = len(dataset) - len(filtered_dataset)
if not dropped_rows:
return filtered_dataset, []
if len(filtered_dataset) == 0:
raise ValueError(
f"{mode_title} training requires at least one string 'text' value "
f"in {split_scope}; all {dropped_rows} rows were null or non-string."
)
return filtered_dataset, [
RawTextNotice(
message = (
f"{mode_title}: dropped {dropped_rows:,} row(s) with null or "
f"non-string 'text' values from {split_scope}"
),
level = "warning",
update_status = True,
)
]
def prepare_raw_text_dataset(
dataset: Dataset,
*,
mode_label: str = "raw text",
split_name: str | None = None,
eos_token: str | None = None,
append_eos: bool = False,
) -> RawTextPreparationResult:
notices: list[RawTextNotice] = []
mode_title = mode_label.capitalize()
split_scope = _split_scope(split_name)
if "text" not in dataset.column_names:
string_cols = _string_columns(dataset)
if not string_cols:
raise ValueError(
f"{mode_title} training requires a string 'text' column but none "
f"was found in {split_scope} (columns: {dataset.column_names})."
)
renamed_col = string_cols[0]
if len(string_cols) > 1:
notices.append(
RawTextNotice(
message = (
f"{mode_title}: dataset has {len(string_cols)} string "
f"columns ({string_cols}); auto-selecting '{renamed_col}' "
"as the training text. Rename the intended column to "
"'text' to override."
),
level = "warning",
update_status = True,
)
)
notices.append(
RawTextNotice(
message = (
f"{mode_title}: renaming column '{renamed_col}' -> 'text' "
f"for {split_scope}"
),
level = "info",
)
)
dataset = dataset.rename_column(renamed_col, "text")
dataset, invalid_row_notices = _drop_invalid_text_rows(
dataset,
mode_title = mode_title,
split_scope = split_scope,
)
notices.extend(invalid_row_notices)
if append_eos:
if not eos_token:
notices.append(
RawTextNotice(
message = (
f"{mode_title}: tokenizer has no eos_token; skipping EOS "
"append. Model will not learn document boundaries."
),
level = "warning",
)
)
else:
def _append_eos(ex, _eos = eos_token):
text = ex["text"]
return {"text": text if text.endswith(_eos) else text + _eos}
dataset = dataset.map(_append_eos)
return RawTextPreparationResult(dataset = dataset, notices = notices)

View file

@ -143,6 +143,7 @@ def detect_hardware() -> DeviceType:
# --- MLX: Apple Silicon ---
if is_apple_silicon() and _has_mlx():
DEVICE = DeviceType.MLX
CHAT_ONLY = False
chip = platform.processor() or platform.machine()
print(f"Hardware detected: MLX — Apple Silicon ({chip})")
return DEVICE
@ -270,19 +271,30 @@ def get_gpu_memory_info() -> Dict[str, Any]:
import mlx.core as mx
import psutil
# MLX uses unified memory — report system memory as the pool
# MLX uses unified memory. Total = system RAM. GPU memory used
# comes from IORegistry's AGXAccelerator (system-wide, no sudo).
total = psutil.virtual_memory().total
# MLX doesn't expose per-process GPU allocation; report 0 as allocated
allocated = 0
agx = _read_apple_gpu_stats()
allocated = agx.get("vram_used_bytes", 0) if agx else 0
try:
info = mx.device_info()
gpu_name = (
info.get("device_name")
or platform.processor()
or platform.machine()
)
except Exception:
gpu_name = platform.processor() or platform.machine()
return {
"available": True,
"backend": _backend_label(device),
"device": 0,
"device_name": f"Apple Silicon ({platform.processor() or platform.machine()})",
"device_name": f"Apple Silicon ({gpu_name})",
"total_gb": total / (1024**3),
"allocated_gb": allocated / (1024**3),
"reserved_gb": 0,
"reserved_gb": allocated / (1024**3),
"free_gb": (total - allocated) / (1024**3),
"utilization_pct": (allocated / total) * 100 if total else 0,
}
@ -460,6 +472,39 @@ def _smi_query(func_name: str, *args, **kwargs) -> Optional[Dict[str, Any]]:
return None
def _read_apple_gpu_stats() -> Dict[str, Any]:
"""Query macOS IORegistry for AGX (Apple GPU) live stats. No sudo needed.
Returns dict with utilization_pct, vram_used_bytes (system-wide GPU memory).
Returns empty dict on failure.
"""
import subprocess
import re
try:
result = subprocess.run(
["ioreg", "-r", "-c", "AGXAccelerator"],
capture_output = True,
timeout = 2,
)
text = result.stdout.decode("utf-8", errors = "replace")
except Exception:
return {}
# PerformanceStatistics block has GPU utilization and in-use memory
m = re.search(r'"PerformanceStatistics" = \{([^}]+)\}', text)
if not m:
return {}
stats_str = m.group(1)
pairs = re.findall(r'"([^"]+)"=(\d+)', stats_str)
stats = {k: int(v) for k, v in pairs}
return {
"utilization_pct": stats.get("Device Utilization %", 0),
"vram_used_bytes": stats.get("In use system memory", 0),
}
def get_gpu_utilization() -> Dict[str, Any]:
"""Return a live snapshot of device utilization information."""
device = get_device()
@ -480,6 +525,50 @@ def get_gpu_utilization() -> Dict[str, Any]:
)
return result
# MLX path: single _read_apple_gpu_stats() call carries both VRAM-used
# bytes and GPU utilization %. psutil for unified-memory total is cheap.
if device == DeviceType.MLX:
try:
import psutil
agx = _read_apple_gpu_stats()
total_bytes = psutil.virtual_memory().total
except Exception as e:
logger.error(f"Error getting MLX GPU utilization: {e}")
return {"available": False, "backend": device.value, "error": str(e)}
if not agx:
return {"available": False, "backend": device.value}
allocated_bytes = agx.get("vram_used_bytes", 0) or 0
vram_used_gb = allocated_bytes / (1024**3)
total_gb = total_bytes / (1024**3)
try:
from core.training import get_training_backend
tb = get_training_backend()
tb_progress = getattr(tb, "_progress", None)
if tb_progress is not None and getattr(tb_progress, "is_training", False):
tb_peak = getattr(tb_progress, "peak_memory_gb", None)
if tb_peak is not None and tb_peak > 0:
vram_used_gb = float(tb_peak)
except Exception:
pass
return {
"available": True,
"backend": device.value,
"gpu_utilization_pct": agx.get("utilization_pct") if agx else None,
"temperature_c": None,
"vram_used_gb": round(vram_used_gb, 2),
"vram_total_gb": round(total_gb, 2),
"vram_utilization_pct": (
round((vram_used_gb / total_gb) * 100, 1) if total_gb > 0 else None
),
"power_draw_w": None,
"power_limit_w": None,
"power_utilization_pct": None,
}
mem = get_gpu_memory_info()
if device != DeviceType.CPU and mem.get("available"):
return {

View file

@ -500,7 +500,9 @@ _VLM_MODEL_TYPES = {
# Pre-computed .venv_t5 paths and backend dir for subprocess version switching.
# Vision check uses 5.5.0 (newest, recognizes all architectures).
_VENV_T5_DIR = str(Path.home() / ".unsloth" / "studio" / ".venv_t5_550")
from utils.paths.storage_roots import studio_root as _studio_root # noqa: E402
_VENV_T5_DIR = str(_studio_root() / ".venv_t5_550")
_BACKEND_DIR = str(Path(__file__).resolve().parent.parent.parent)
# Inline script executed in a subprocess with transformers 5.x activated.

View file

@ -5,17 +5,59 @@ from __future__ import annotations
import json
import os
import sys
from pathlib import Path
import tempfile
def _infer_studio_home_from_venv() -> Path | None:
"""Return parent dir of sys.prefix as STUDIO_HOME if running from an
installer-managed unsloth_studio venv. Sentinel-gated (share/studio.conf
or bin shim) so a developer venv named unsloth_studio is not misidentified.
"""
try:
prefix = Path(sys.prefix).resolve()
except (OSError, ValueError):
return None
if prefix.name != "unsloth_studio":
return None
candidate = prefix.parent
shim_name = "unsloth.exe" if os.name == "nt" else "unsloth"
try:
has_sentinel = (candidate / "share" / "studio.conf").is_file() or (
candidate / "bin" / shim_name
).is_file()
except OSError:
return None
if has_sentinel:
return candidate
return None
def studio_root() -> Path:
"""Studio install root.
Priority: UNSLOTH_STUDIO_HOME, then STUDIO_HOME alias, then sys.prefix
inference, then legacy ~/.unsloth/studio. UNSLOTH_STUDIO_HOME wins when
both are set (the more specific signal beats the generic alias).
"""
override = (os.environ.get("UNSLOTH_STUDIO_HOME") or "").strip()
if not override:
override = (os.environ.get("STUDIO_HOME") or "").strip()
if override:
try:
return Path(override).expanduser().resolve()
except (OSError, ValueError):
return Path(override).expanduser()
inferred = _infer_studio_home_from_venv()
if inferred is not None:
return inferred
return Path.home() / ".unsloth" / "studio"
def cache_root() -> Path:
"""Central cache directory for all studio downloads (models, datasets, etc.)."""
return Path.home() / ".unsloth" / "studio" / "cache"
return studio_root() / "cache"
def assets_root() -> Path:

View file

@ -95,9 +95,11 @@ TRANSFORMERS_DEFAULT_VERSION = "4.57.6"
# Consumers should prefer TRANSFORMERS_530_VERSION / TRANSFORMERS_550_VERSION.
TRANSFORMERS_5_VERSION = TRANSFORMERS_550_VERSION
# Pre-installed directories — created by setup.sh / setup.ps1
_VENV_T5_530_DIR = str(Path.home() / ".unsloth" / "studio" / ".venv_t5_530")
_VENV_T5_550_DIR = str(Path.home() / ".unsloth" / "studio" / ".venv_t5_550")
# Pre-installed directories — created by setup.sh / setup.ps1.
from utils.paths.storage_roots import studio_root as _studio_root # noqa: E402
_VENV_T5_530_DIR = str(_studio_root() / ".venv_t5_530")
_VENV_T5_550_DIR = str(_studio_root() / ".venv_t5_550")
# Backwards-compat alias
_VENV_T5_DIR = _VENV_T5_550_DIR

View file

@ -38,11 +38,11 @@ import {
Download03Icon,
GemIcon,
Globe02Icon,
HelpCircleIcon,
Search01Icon,
PowerIcon,
PencilEdit02Icon,
LayoutAlignLeftIcon,
HelpCircleIcon,
Settings02Icon,
ZapIcon,
} from "@hugeicons/core-free-icons";

View file

@ -316,15 +316,28 @@ const ReasoningGroupImpl: ReasoningGroupComponent = ({
if (message.status?.type !== "running") {
return false;
}
const lastIndex = message.parts.length - 1;
if (lastIndex < 0) {
const parts = message.parts;
const len = parts.length;
if (len === 0) {
return false;
}
const lastType = message.parts[lastIndex]?.type;
if (lastType !== "reasoning") {
let groupHasReasoning = false;
for (let i = startIndex; i <= endIndex && i < len; i += 1) {
if (parts[i]?.type === "reasoning") {
groupHasReasoning = true;
break;
}
}
if (!groupHasReasoning) {
return false;
}
return lastIndex >= startIndex && lastIndex <= endIndex;
for (let i = endIndex + 1; i < len; i += 1) {
if (parts[i]?.type !== "tool-call") {
return false;
}
}
return true;
});
const persistedDuration = useAuiState(({ message }) => {

View file

@ -50,7 +50,7 @@ export async function fetchDeviceType(): Promise<DeviceType> {
if (res.ok) {
const data = (await res.json()) as { device_type?: string; chat_only?: boolean };
const deviceType = data.device_type ?? detectLocalPlatform();
const chatOnly = data.chat_only ?? deviceType === "mac";
const chatOnly = data.chat_only ?? false;
usePlatformStore.setState({ deviceType, chatOnly, fetched: true });
return deviceType;
}

View file

@ -76,6 +76,13 @@ export const TARGET_MODULES = [
"down_proj",
];
/** CPT requires embed_tokens and lm_head in addition to standard LoRA modules. */
export const CPT_TARGET_MODULES = [
...TARGET_MODULES,
"embed_tokens",
"lm_head",
];
export const OPTIMIZER_OPTIONS: ReadonlyArray<{ value: string; label: string }> = [
{ value: "adamw_8bit", label: "AdamW 8-bit" },
{ value: "paged_adamw_8bit", label: "Paged AdamW 8-bit" },
@ -96,11 +103,14 @@ export const LR_SCHEDULER_OPTIONS: ReadonlyArray<{ value: string; label: string
*/
export const LR_DEFAULT_LORA = 2e-4;
export const LR_DEFAULT_FULL = 2e-5;
export const LR_DEFAULT_CPT = 5e-5;
export const DEFAULT_HYPERPARAMS = {
epochs: 3,
contextLength: 2048,
learningRate: LR_DEFAULT_LORA,
// null = let backend auto-compute (lr/10 per Unsloth CPT recipe). Only used by CPT.
embeddingLearningRate: null as number | null,
optimizerType: "adamw_8bit",
lrSchedulerType: "linear",
loraRank: 16,

View file

@ -74,6 +74,7 @@ export const METHOD_LABELS: Record<TrainingMethod, string> = {
qlora: "QLoRA",
lora: "LoRA",
full: "Full Fine-tune",
cpt: "Continued Pretraining",
};
export const GUIDE_STEPS = [

View file

@ -63,6 +63,7 @@ const FORMAT_OPTIONS: { value: DatasetFormat; label: string }[] = [
{ value: "alpaca", label: "Alpaca" },
{ value: "chatml", label: "ChatML" },
{ value: "sharegpt", label: "ShareGPT" },
{ value: "raw", label: "Raw Text" },
];
export function DatasetStep() {

View file

@ -366,6 +366,7 @@ export function ModelSelectionStep() {
<SelectItem value="qlora">QLoRA (4-bit)</SelectItem>
<SelectItem value="lora">LoRA (16-bit)</SelectItem>
<SelectItem value="full">Full Fine-tune</SelectItem>
<SelectItem value="cpt">Continued Pretraining</SelectItem>
</SelectContent>
</Select>
</div>

View file

@ -5,6 +5,7 @@ import { Badge } from "@/components/ui/badge";
import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
import { Separator } from "@/components/ui/separator";
import { useTrainingConfigStore } from "@/features/training";
import { getTrainingMethodLabel } from "@/features/training/lib/training-methods";
import { useHardwareInfo } from "@/hooks";
import { isAdapterMethod } from "@/types/training";
import { ChipIcon, Database02Icon, GpuIcon, Settings04Icon } from "@hugeicons/core-free-icons";
@ -102,6 +103,7 @@ export function SummaryStep() {
const showLoraParams = isAdapterMethod(trainingMethod);
const datasetName = datasetSource === "upload" ? uploadedFile : dataset;
const trainingMethodLabel = getTrainingMethodLabel(trainingMethod);
return (
<div className="grid grid-cols-2 gap-3">
@ -150,7 +152,7 @@ export function SummaryStep() {
<Separator className="my-2" />
<div className="space-y-1 text-sm">
<Row label="Type" value={modelType} capitalize />
<Row label="Method" value={trainingMethod === "qlora" ? "QLoRA" : trainingMethod === "lora" ? "LoRA" : "Full"} />
<Row label="Method" value={trainingMethodLabel} />
</div>
</CardContent>
</Card>
@ -199,7 +201,7 @@ export function SummaryStep() {
<div className="flex flex-1 flex-col">
<span className="text-xs text-muted-foreground">Training</span>
<span className="text-sm font-medium">
{trainingMethod === "qlora" ? "QLoRA" : trainingMethod === "lora" ? "LoRA" : "Full"}
{trainingMethodLabel}
</span>
</div>
</div>

View file

@ -4,6 +4,7 @@
import type { TrainingViewData } from "@/features/training";
import { getTrainingRun } from "@/features/training";
import type { TrainingRunDetailResponse } from "@/features/training";
import { parseBackendTrainingMethod } from "@/features/training/lib/training-methods";
import { type ReactElement, useEffect, useState } from "react";
import { ChartsSection } from "./sections/charts-section";
import { ProgressSection } from "./sections/progress-section";
@ -12,15 +13,6 @@ interface HistoricalTrainingViewProps {
runId: string;
}
function normalizeTrainingMethod(config: Record<string, unknown>): string {
const type = config?.training_type as string | undefined;
if (!type || type === "Full Finetuning") return "full";
if (type === "LoRA/QLoRA") {
return config?.load_in_4bit ? "qlora" : "lora";
}
return "full";
}
function mapToViewData(detail: TrainingRunDetailResponse): TrainingViewData {
const { run, metrics } = detail;
@ -79,7 +71,10 @@ function mapToViewData(detail: TrainingRunDetailResponse): TrainingViewData {
error: run.status === "error" ? run.error_message : null,
isTrainingRunning: false,
modelName: run.model_name,
trainingMethod: normalizeTrainingMethod(detail.config),
trainingMethod: parseBackendTrainingMethod(
detail.config?.training_type,
detail.config?.load_in_4bit,
),
lossHistory,
lrHistory,
gradNormHistory,
@ -143,7 +138,11 @@ export function HistoricalTrainingView({
loraRank: detail.config.lora_r as number | undefined,
loraAlpha: detail.config.lora_alpha as number | undefined,
loraDropout: detail.config.lora_dropout as number | undefined,
loraVariant: detail.config.use_rslora ? "rsLoRA" : undefined,
loraVariant: detail.config.use_rslora
? "rslora"
: detail.config.use_loftq
? "loftq"
: "lora",
}
: undefined;

View file

@ -15,6 +15,7 @@ import { Badge } from "@/components/ui/badge";
import { Spinner } from "@/components/ui/spinner";
import { useTrainingActions, useTrainingConfigStore } from "@/features/training";
import { checkDatasetFormat } from "@/features/training/api/datasets-api";
import { isRawTextDatasetFormat } from "@/features/training/lib/training-methods";
import type { CheckFormatResponse } from "@/features/training/types/datasets";
import { Database02Icon, AlertCircleIcon } from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
@ -90,10 +91,11 @@ export function DatasetPreviewDialog({
const effectiveIsAudio = !!data?.is_audio;
const effectiveIsVlm = isVlm || !!data?.is_image;
const isRawFormat = isRawTextDatasetFormat(datasetFormat);
const hasHeuristicMapping = !data?.requires_manual_mapping && !!data?.suggested_mapping;
const mappingEnabled = !!data?.requires_manual_mapping || hasHeuristicMapping;
const mappingEnabled = !isRawFormat && (!!data?.requires_manual_mapping || hasHeuristicMapping);
const showMappingFooter = mode === "mapping" && mappingEnabled;
const mappingOk = isMappingComplete(manualMapping, effectiveIsVlm, datasetFormat, effectiveIsAudio);
const mappingOk = isRawFormat || isMappingComplete(manualMapping, effectiveIsVlm, datasetFormat, effectiveIsAudio);
const availableRoles = getAvailableRoles(effectiveIsVlm, datasetFormat, effectiveIsAudio);
const isHfDataset = datasetSource === "huggingface";
@ -413,7 +415,7 @@ export function DatasetPreviewDialog({
<MetaRow label="Source" value={sourceLabel} />
<MetaRow
label="Format"
value={data.detected_format || "--"}
value={isRawFormat ? "Raw Text" : (data.detected_format || "--")}
/>
<MetaRow
label="Total Rows"
@ -441,7 +443,7 @@ export function DatasetPreviewDialog({
/>
</div>
{data.warning && (
{data.warning && !isRawFormat && (
<div className="rounded-lg border border-amber-200 bg-amber-50 px-4 py-3 text-xs text-amber-700 dark:border-amber-800 dark:bg-amber-950 dark:text-amber-400 mb-4 flex items-start gap-2.5">
<HugeiconsIcon icon={AlertCircleIcon} className="size-4 shrink-0 mt-0.5" />
<span>{data.warning}</span>

View file

@ -913,6 +913,7 @@ export function DatasetSection() {
<SelectItem value="alpaca">Alpaca</SelectItem>
<SelectItem value="chatml">ChatML</SelectItem>
<SelectItem value="sharegpt">ShareGPT</SelectItem>
<SelectItem value="raw">Raw Text</SelectItem>
</SelectContent>
</Select>
</div>

View file

@ -67,6 +67,7 @@ const METHOD_DOTS: Record<string, string> = {
qlora: "bg-emerald-400",
lora: "bg-blue-400",
full: "bg-amber-400",
cpt: "bg-purple-400",
};
const DARK_TRIGGER =
@ -570,7 +571,9 @@ export function ModelSection() {
</TooltipTrigger>
<TooltipContent className="max-w-xs">
QLoRA uses 4-bit quantization for lowest VRAM. LoRA uses
16-bit. Full updates all weights.{" "}
16-bit. Full updates all weights. CPT (Continued Pretraining)
trains on raw text to adapt the model to a new domain without
chat formatting.{" "}
<a
href="https://unsloth.ai/docs/get-started/fine-tuning-llms-guide/lora-hyperparameters-guide"
target="_blank"
@ -617,6 +620,14 @@ export function ModelSection() {
Full Fine-tune
</span>
</SelectItem>
<SelectItem value="cpt">
<span className="flex items-center gap-2">
<span
className={`size-2 shrink-0 rounded-full ${METHOD_DOTS.cpt}`}
/>
Continued Pretraining
</span>
</SelectItem>
</SelectContent>
</Select>
</div>

View file

@ -1,6 +1,7 @@
// 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 { usePlatformStore } from "@/config/env";
import { SectionCard } from "@/components/section-card";
import { Checkbox } from "@/components/ui/checkbox";
import {
@ -33,11 +34,14 @@ import {
} from "@/components/ui/tooltip";
import {
CONTEXT_LENGTHS,
CPT_TARGET_MODULES,
LR_SCHEDULER_OPTIONS,
OPTIMIZER_OPTIONS,
TARGET_MODULES,
} from "@/config/training";
import { useMaxStepsEpochsToggle, useTrainingConfigStore } from "@/features/training";
import { isRawTextDatasetFormat } from "@/features/training/lib/training-methods";
import { isAdapterMethod } from "@/types/training";
import type { GradientCheckpointing } from "@/types/training";
import {
ArrowDown01Icon,
@ -124,10 +128,14 @@ function SliderRow({
export function ParamsSection(): ReactElement {
const store = useTrainingConfigStore();
const isLora = store.trainingMethod !== "full";
const platformDeviceType = usePlatformStore((s) => s.deviceType);
const isLora = isAdapterMethod(store.trainingMethod);
const isCpt = store.trainingMethod === "cpt";
const isRawText = isRawTextDatasetFormat(store.datasetFormat);
const showVisionLora = store.isVisionModel && store.isDatasetImage === true;
const [loraOpen, setLoraOpen] = useState(false);
const [hyperOpen, setHyperOpen] = useState(false);
const needsExpandedHeight = isCpt || (isLora && loraOpen) || hyperOpen;
const [ctxInput, setCtxInput] = useState(String(store.contextLength));
const ctxAnchorRef = useRef<HTMLDivElement>(null);
const ctxItems = CONTEXT_LENGTHS.map(String);
@ -166,7 +174,7 @@ export function ParamsSection(): ReactElement {
title="Parameters"
description="Configure training hyperparameters"
accent="orange"
className={`${(isLora && loraOpen) || hyperOpen
className={`${needsExpandedHeight
? "min-h-studio-config-column"
: "h-studio-config-column"} duration-150`}
>
@ -376,10 +384,62 @@ export function ParamsSection(): ReactElement {
className="w-full font-mono"
/>
<p className="text-[10px] text-muted-foreground">
Recommended: 2e-4 for LoRA, 2e-5 for full fine-tune
Recommended: 2e-4 for LoRA, 5e-5 for CPT, 2e-5 for full fine-tune
</p>
</div>
{/* Embedding Learning Rate (CPT only) */}
{isCpt && (
<div className="flex flex-col gap-2">
<span className="flex items-center gap-1.5 text-xs font-medium text-muted-foreground">
Embedding Learning Rate
<Tooltip>
<TooltipTrigger asChild={true}>
<button
type="button"
className="text-foreground/70 hover:text-foreground"
>
<HugeiconsIcon
icon={InformationCircleIcon}
className="size-3"
/>
</button>
</TooltipTrigger>
<TooltipContent>
Only used when CPT is training <code>embed_tokens</code>.
Embeddings are easier to destabilize than LoRA weights, so
they usually need a smaller LR. Leave blank to use
<code>lr/10</code>; typical working range is 2x-10x smaller
than the main LR. Increase it only if vocabulary or
domain-token adaptation is too slow.
</TooltipContent>
</Tooltip>
</span>
<Input
type="number"
step="0.00001"
min="0"
max="1"
placeholder={`auto (${(store.learningRate / 10).toExponential(1)})`}
value={store.embeddingLearningRate ?? ""}
onChange={(e) => {
const raw = e.target.value;
if (raw === "") {
store.setEmbeddingLearningRate(null);
return;
}
const n = Number(raw);
store.setEmbeddingLearningRate(Number.isFinite(n) ? n : null);
}}
className="w-full font-mono"
/>
<p className="text-[10px] text-muted-foreground">
Leave blank to use lr/10 (recommended). Typical range is
2x-10x smaller than the main learning rate.
</p>
</div>
)}
{/* LoRA Settings */}
{isLora && (
<Collapsible open={loraOpen} onOpenChange={setLoraOpen}>
@ -514,7 +574,7 @@ export function ParamsSection(): ReactElement {
Target Modules
</span>
<div className="flex flex-wrap gap-1.5">
{TARGET_MODULES.map((mod) => {
{(isCpt ? CPT_TARGET_MODULES : TARGET_MODULES).map((mod) => {
const active = store.targetModules.includes(mod);
return (
<button
@ -883,7 +943,11 @@ export function ParamsSection(): ReactElement {
<SelectContent>
<SelectItem value="none">None</SelectItem>
<SelectItem value="true">Standard</SelectItem>
<SelectItem value="unsloth">Unsloth</SelectItem>
{platformDeviceType === "mac" ? (
<SelectItem value="mlx">MLX</SelectItem>
) : (
<SelectItem value="unsloth">Unsloth</SelectItem>
)}
</SelectContent>
</Select>
</Row>
@ -902,7 +966,7 @@ export function ParamsSection(): ReactElement {
</label>
</div>
)}
{!store.isEmbeddingModel && (
{!store.isEmbeddingModel && !isCpt && !isRawText && (
<div className="flex items-center gap-2">
<Checkbox
id="trainOnCompletions"

View file

@ -26,6 +26,7 @@ import {
useTrainingConfigStore,
useTrainingRuntimeStore,
} from "@/features/training";
import { getTrainingMethodLabel } from "@/features/training/lib/training-methods";
import type { TrainingViewData } from "@/features/training";
import { useGpuUtilization } from "@/hooks";
import { cn } from "@/lib/utils";
@ -86,6 +87,7 @@ export function ProgressSection({
configOverride,
}: ProgressSectionProps): ReactElement {
const navigate = useNavigate();
const trainingMethodLabel = getTrainingMethodLabel(data.trainingMethod);
const config = useTrainingConfigStore(
useShallow((state) => ({
@ -272,7 +274,7 @@ export function ProgressSection({
{data.modelName || "--"}
</MetricStat>
<MetricStat label="Method">
{data.trainingMethod === "qlora" ? "QLoRA" : data.trainingMethod === "lora" ? "LoRA" : "Full"}
{trainingMethodLabel}
</MetricStat>
</div>

View file

@ -3,9 +3,10 @@
import type { TrainingConfigState } from "../types/config";
import type { TrainingStartRequest } from "../types/api";
const BACKEND_LORA_TYPE = "LoRA/QLoRA";
const BACKEND_FULL_TYPE = "Full Finetuning";
import {
isRawTextDatasetFormat,
toBackendTrainingType,
} from "../lib/training-methods";
function parseSliceValue(value: string | null): number | null {
if (value == null) return null;
@ -16,16 +17,15 @@ function parseSliceValue(value: string | null): number | null {
return num;
}
export function toBackendTrainingType(trainingMethod: string): string {
return trainingMethod === "full" ? BACKEND_FULL_TYPE : BACKEND_LORA_TYPE;
}
export function buildTrainingStartPayload(
config: TrainingConfigState,
): TrainingStartRequest {
const isCpt = config.trainingMethod === "cpt";
const adapterMethod = config.trainingMethod !== "full";
const isQloraMethod = config.trainingMethod === "qlora";
const isFourBitModel = (config.selectedModel ?? "").toLowerCase().includes("4bit");
const isEmbedding = config.isEmbeddingModel;
const isRawText = isRawTextDatasetFormat(config.datasetFormat);
const hfDataset = config.datasetSource === "huggingface" ? config.dataset : null;
const localDatasets =
config.datasetSource === "upload" && config.uploadedFile
@ -53,7 +53,7 @@ export function buildTrainingStartPayload(
model_name: config.selectedModel ?? "",
training_type: toBackendTrainingType(config.trainingMethod),
hf_token: config.hfToken.trim() || null,
load_in_4bit: adapterMethod ? isQloraMethod : false,
load_in_4bit: (adapterMethod && isQloraMethod) || (isCpt && isFourBitModel),
max_seq_length: config.contextLength,
trust_remote_code: config.trustRemoteCode ?? false,
hf_dataset: hfDataset,
@ -71,6 +71,10 @@ export function buildTrainingStartPayload(
custom_format_mapping: customFormatMapping,
num_epochs: config.epochs,
learning_rate: String(config.learningRate),
embedding_learning_rate:
isCpt && config.embeddingLearningRate != null
? config.embeddingLearningRate
: null,
batch_size: config.batchSize,
gradient_accumulation_steps: config.gradientAccumulation,
warmup_steps: isEmbedding ? null : config.warmupSteps,
@ -91,7 +95,8 @@ export function buildTrainingStartPayload(
gradient_checkpointing: config.gradientCheckpointing,
use_rslora: config.loraVariant === "rslora",
use_loftq: config.loraVariant === "loftq",
train_on_completions: isEmbedding ? false : config.trainOnCompletions,
// CPT always trains on full sequences (no chat format masking)
train_on_completions: (isEmbedding || isCpt || isRawText) ? false : config.trainOnCompletions,
finetune_vision_layers: config.finetuneVisionLayers,
finetune_language_layers: config.finetuneLanguageLayers,
finetune_attention_modules: config.finetuneAttentionModules,

View file

@ -8,6 +8,7 @@ import { checkDatasetFormat } from "../api/datasets-api";
import { getTrainingRun } from "../api/history-api";
import { buildTrainingStartPayload } from "../api/mappers";
import { resetTraining, startTraining, stopTraining } from "../api/train-api";
import { isRawTextDatasetFormat } from "../lib/training-methods";
import { syncTrainingRuntimeFromBackend } from "../lib/sync-runtime";
import { validateTrainingConfig } from "../lib/validation";
import { useDatasetPreviewDialogStore } from "../stores/dataset-preview-dialog-store";
@ -88,7 +89,10 @@ export function useTrainingActions() {
});
}
const needsReview = check.requires_manual_mapping || check.detected_format === "custom_heuristic";
const isRawFormat = isRawTextDatasetFormat(config.datasetFormat);
const needsReview =
!isRawFormat &&
(check.requires_manual_mapping || check.detected_format === "custom_heuristic");
if (needsReview && !hasManualMapping(config, isVlm, isAudio)) {
// Pre-fill from suggested_mapping or VLM detected columns
const hint: Record<string, string> = {};

View file

@ -3,6 +3,7 @@
import type { BackendModelConfig } from "../api/models-api";
import type { TrainingConfigState } from "../types/config";
import { usePlatformStore } from "@/config/env";
type ModelDefaultsPatch = Partial<
Pick<
@ -69,7 +70,13 @@ function toStringArray(value: unknown): string[] | undefined {
function toGradientCheckpointing(
value: unknown,
): TrainingConfigState["gradientCheckpointing"] | undefined {
if (value === "none" || value === "true" || value === "unsloth") return value;
if (value === "none" || value === "true" || value === "unsloth" || value === "mlx") {
// On Mac, map "unsloth" → "mlx" since Unsloth GC is GPU-only
if (usePlatformStore.getState().deviceType === "mac" && value === "unsloth") {
return "mlx";
}
return value;
}
return undefined;
}

View file

@ -0,0 +1,48 @@
// 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 type { DatasetFormat, TrainingMethod } from "@/types/training";
const BACKEND_TRAINING_TYPE: Record<TrainingMethod, string> = {
qlora: "LoRA/QLoRA",
lora: "LoRA/QLoRA",
full: "Full Finetuning",
cpt: "Continued Pretraining",
};
const TRAINING_METHOD_LABELS: Record<TrainingMethod, string> = {
qlora: "QLoRA",
lora: "LoRA",
full: "Full",
cpt: "CPT",
};
export function toBackendTrainingType(trainingMethod: TrainingMethod): string {
return BACKEND_TRAINING_TYPE[trainingMethod];
}
export function getTrainingMethodLabel(
trainingMethod: TrainingMethod | string,
): string {
if (Object.prototype.hasOwnProperty.call(TRAINING_METHOD_LABELS, trainingMethod)) {
return TRAINING_METHOD_LABELS[trainingMethod as TrainingMethod];
}
return TRAINING_METHOD_LABELS.full;
}
export function parseBackendTrainingMethod(
trainingType: unknown,
loadIn4Bit: unknown,
): TrainingMethod {
if (trainingType === "Continued Pretraining") return "cpt";
if (trainingType === "LoRA/QLoRA") {
return loadIn4Bit ? "qlora" : "lora";
}
return "full";
}
export function isRawTextDatasetFormat(
datasetFormat: DatasetFormat,
): boolean {
return datasetFormat === "raw";
}

View file

@ -1,15 +1,17 @@
// 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 { DEFAULT_HYPERPARAMS, LR_DEFAULT_FULL, LR_DEFAULT_LORA, STEPS } from "@/config/training";
import { CPT_TARGET_MODULES, DEFAULT_HYPERPARAMS, LR_DEFAULT_CPT, LR_DEFAULT_FULL, LR_DEFAULT_LORA, STEPS, TARGET_MODULES } from "@/config/training";
import { authFetch } from "@/features/auth";
import { isAdapterMethod } from "@/types/training";
import type { DatasetFormat } from "@/types/training";
import type { ModelType, StepNumber, TrainingMethod } from "@/types/training";
import { create } from "zustand";
import { persist } from "zustand/middleware";
import { checkDatasetFormat } from "../api/datasets-api";
import { checkVisionModel, getModelConfig } from "../api/models-api";
import { mapBackendModelConfigToTrainingPatch } from "../lib/model-defaults";
import { isRawTextDatasetFormat } from "../lib/training-methods";
import type { BackendModelConfig } from "../api/models-api";
import type { TrainingConfigState, TrainingConfigStore } from "../types/config";
@ -108,6 +110,11 @@ let _learningRateManuallySet = false;
// setTrainingMethod can restore it when switching back from full to adapter.
let _yamlLearningRate: number | undefined = undefined;
// Track whether entering CPT auto-forced datasetFormat="raw" so that
// leaving CPT can restore the prior user-visible format.
let _datasetFormatBeforeCpt: DatasetFormat | null = null;
let _datasetFormatAutoForcedByCpt = false;
const NON_PERSISTED_STATE_KEYS: ReadonlySet<keyof TrainingConfigState> = new Set([
"modelType",
"isCheckingVision",
@ -156,6 +163,123 @@ function canProceedForStep(state: TrainingConfigState): boolean {
}
}
type TrainingMethodStatePatch = Partial<
Pick<
TrainingConfigState,
| "trainingMethod"
| "learningRate"
| "loraRank"
| "loraAlpha"
| "loraVariant"
| "targetModules"
| "datasetFormat"
| "trainOnCompletions"
>
>;
function getCptTrainingPatch(): TrainingMethodStatePatch {
return {
loraRank: 128,
loraAlpha: 32,
loraVariant: "rslora",
targetModules: CPT_TARGET_MODULES,
datasetFormat: "raw",
trainOnCompletions: false,
};
}
function getCptModelDefaultsPatch(): TrainingMethodStatePatch {
return {
...getCptTrainingPatch(),
learningRate: LR_DEFAULT_CPT,
};
}
function getRestoreFromCptPatch(): TrainingMethodStatePatch {
return {
loraRank: DEFAULT_HYPERPARAMS.loraRank,
loraAlpha: DEFAULT_HYPERPARAMS.loraAlpha,
loraVariant: DEFAULT_HYPERPARAMS.loraVariant,
targetModules: TARGET_MODULES,
};
}
function clearCptDatasetFormatTracking(): void {
_datasetFormatBeforeCpt = null;
_datasetFormatAutoForcedByCpt = false;
}
function recordCptDatasetFormatOverride(currentDatasetFormat: DatasetFormat): void {
if (isRawTextDatasetFormat(currentDatasetFormat)) {
clearCptDatasetFormatTracking();
return;
}
_datasetFormatBeforeCpt = currentDatasetFormat;
_datasetFormatAutoForcedByCpt = true;
}
function getRestoreDatasetFormatFromCptPatch(): TrainingMethodStatePatch {
if (!_datasetFormatAutoForcedByCpt || _datasetFormatBeforeCpt == null) {
clearCptDatasetFormatTracking();
return {};
}
const previousDatasetFormat = _datasetFormatBeforeCpt;
clearCptDatasetFormatTracking();
return { datasetFormat: previousDatasetFormat };
}
function resolveTrainingMethodLearningRate(
prevMethod: TrainingMethod,
nextMethod: TrainingMethod,
): number | undefined {
if (_learningRateManuallySet) {
return undefined;
}
const wasCpt = prevMethod === "cpt";
const wasAdapter = isAdapterMethod(prevMethod);
const nowAdapter = isAdapterMethod(nextMethod);
if (nextMethod === "cpt") {
return LR_DEFAULT_CPT;
}
if (wasCpt && nowAdapter) {
return _yamlLearningRate ?? LR_DEFAULT_LORA;
}
if (wasAdapter && nowAdapter) {
return undefined;
}
return nowAdapter ? _yamlLearningRate ?? LR_DEFAULT_LORA : LR_DEFAULT_FULL;
}
function buildTrainingMethodPatch(
prevMethod: TrainingMethod,
nextMethod: TrainingMethod,
currentDatasetFormat: DatasetFormat,
): TrainingMethodStatePatch {
const patch: TrainingMethodStatePatch = { trainingMethod: nextMethod };
if (prevMethod !== "cpt" && nextMethod === "cpt") {
recordCptDatasetFormatOverride(currentDatasetFormat);
Object.assign(patch, getCptTrainingPatch());
}
if (prevMethod === "cpt" && nextMethod !== "cpt") {
Object.assign(
patch,
getRestoreFromCptPatch(),
getRestoreDatasetFormatFromCptPatch(),
);
}
const learningRate = resolveTrainingMethodLearningRate(prevMethod, nextMethod);
if (learningRate !== undefined) {
patch.learningRate = learningRate;
}
return patch;
}
export const useTrainingConfigStore = create<TrainingConfigStore>()(
persist(
(set, get) => {
@ -216,11 +340,14 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
// Auto-select training method based on model size vs GPU memory.
// If model_size * 1.5 * context_scale fits in free VRAM, use LoRA 16-bit.
// Otherwise use QLoRA 4-bit.
// Auto-select LoRA vs QLoRA based on GPU memory.
// Skip if user has manually chosen CPT -- don't override it.
const modelSizeBytes = modelDetails.model_size_bytes;
if (modelSizeBytes && modelSizeBytes > 0) {
if (modelSizeBytes && modelSizeBytes > 0 && get().trainingMethod !== "cpt") {
void autoSelectTrainingMethod(modelSizeBytes, patch.contextLength ?? get().contextLength)
.then((method) => {
if (get().selectedModel !== modelName) return;
if (get().trainingMethod === "cpt") return;
if (method) {
const lrPatch = !_learningRateManuallySet && !modelConfigHasLR
? { learningRate: method === "full" ? LR_DEFAULT_FULL : LR_DEFAULT_LORA }
@ -230,8 +357,16 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
});
}
// Preserve CPT hyperparams: YAML adapter defaults (r/alpha/targets/LR)
// are tuned for standard LoRA and would otherwise clobber CPT settings.
const cptOverrides =
get().trainingMethod === "cpt"
? getCptModelDefaultsPatch()
: {};
set({
...patch,
...cptOverrides,
modelType: inferredModelType,
isVisionModel: modelDetails.is_vision,
isEmbeddingModel: isEmbedding,
@ -396,29 +531,14 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
void loadAndApplyModelDefaults(state.selectedModel);
},
setTrainingMethod: (trainingMethod) => {
if (_learningRateManuallySet) {
set({ trainingMethod });
return;
}
const prev = get().trainingMethod;
const wasAdapter = isAdapterMethod(prev);
const nowAdapter = isAdapterMethod(trainingMethod);
// qlora <-> lora: same LR range, don't touch learning rate
if (wasAdapter && nowAdapter) {
set({ trainingMethod });
return;
}
// Category changed (adapter <-> full)
if (nowAdapter) {
// Switching TO adapter: restore YAML LR if available
set({ trainingMethod, learningRate: _yamlLearningRate ?? LR_DEFAULT_LORA });
} else {
// Switching TO full: no YAML full-LR exists, use constant
set({ trainingMethod, learningRate: LR_DEFAULT_FULL });
}
const state = get();
set(
buildTrainingMethodPatch(
state.trainingMethod,
trainingMethod,
state.datasetFormat,
),
);
},
setHfToken: (hfToken) =>
set({ hfToken: hfToken.trim().replace(/^["']+|["']+$/g, "") }),
@ -448,7 +568,26 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
runDatasetCheck(uploadedFile, "train");
}
},
setDatasetFormat: (datasetFormat) => set({ datasetFormat }),
setDatasetFormat: (datasetFormat) =>
set((state) => {
if (state.trainingMethod === "cpt") {
if (isRawTextDatasetFormat(datasetFormat)) {
clearCptDatasetFormatTracking();
}
return {
datasetFormat: "raw",
trainOnCompletions: false,
};
}
return {
datasetFormat,
trainOnCompletions:
isRawTextDatasetFormat(datasetFormat)
? false
: state.trainOnCompletions,
};
}),
setDataset: (dataset) => {
_datasetCheckController?.abort();
_datasetCheckController = null;
@ -566,6 +705,8 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
_learningRateManuallySet = true;
set({ learningRate });
},
setEmbeddingLearningRate: (embeddingLearningRate) =>
set({ embeddingLearningRate }),
setOptimizerType: (optimizerType) => set({ optimizerType }),
setLrSchedulerType: (lrSchedulerType) => set({ lrSchedulerType }),
setLoraRank: (loraRank) => set({ loraRank }),
@ -608,6 +749,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
_trainOnCompletionsManuallySet = false;
_learningRateManuallySet = false;
_yamlLearningRate = undefined;
clearCptDatasetFormatTracking();
set(initialState);
},
resetToModelDefaults: () => {
@ -629,7 +771,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
},
{
name: "unsloth_training_config_v1",
version: 9,
version: 10,
migrate: (persisted, version) => {
const s = persisted as Record<string, unknown>;
if (version < 2 && s.datasetSubset == null && s.datasetConfig != null) {
@ -665,6 +807,17 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
s.weightDecay = DEFAULT_HYPERPARAMS.weightDecay;
}
}
if (version < 10 && s.trainingMethod === "cpt") {
// Backfill CPT defaults for state persisted before they existed.
s.loraRank = 128;
s.loraAlpha = 32;
s.loraVariant = "rslora";
s.targetModules = CPT_TARGET_MODULES;
s.datasetFormat = "raw";
if (s.learningRate == null || s.learningRate === LR_DEFAULT_LORA) {
s.learningRate = LR_DEFAULT_CPT;
}
}
return s as unknown as TrainingConfigStore;
},
partialize: partializePersistedState,

View file

@ -21,6 +21,8 @@ export interface TrainingStartRequest {
custom_format_mapping?: Record<string, unknown> | null;
num_epochs: number;
learning_rate: string;
/** Optional CPT embedding LR. If omitted, backend uses lr/10; typical range is 2x-10x smaller than main LR. */
embedding_learning_rate?: number | null;
batch_size: number;
gradient_accumulation_steps: number;
warmup_steps: number | null;

View file

@ -41,6 +41,7 @@ export interface TrainingConfigState {
epochs: number;
contextLength: number;
learningRate: number;
embeddingLearningRate: number | null;
optimizerType: string;
lrSchedulerType: string;
loraRank: number;
@ -115,6 +116,7 @@ export interface TrainingConfigActions {
setEpochs: (epochs: number) => void;
setContextLength: (length: number) => void;
setLearningRate: (rate: number) => void;
setEmbeddingLearningRate: (rate: number | null) => void;
setOptimizerType: (value: string) => void;
setLrSchedulerType: (value: string) => void;
setLoraRank: (rank: number) => void;

View file

@ -6,6 +6,7 @@ import { listModels } from "@huggingface/hub";
import { type CachedResult, cachedModelInfo, primeCacheFromListing } from "@/lib/hf-cache";
import { useCallback, useMemo } from "react";
import { useHfPaginatedSearch } from "./use-hf-paginated-search";
import { usePlatformStore } from "@/config/env";
export interface HfModelResult {
id: string;
@ -16,7 +17,8 @@ export interface HfModelResult {
isGguf: boolean;
}
const EXCLUDED_TAGS = new Set([
/** Tags to exclude on GPU (CUDA/ROCm) — MLX models won't load on GPU. */
const EXCLUDED_TAGS_GPU = new Set([
"gptq",
"awq",
"exl2",
@ -28,6 +30,18 @@ const EXCLUDED_TAGS = new Set([
"ctranslate2",
]);
/** Tags to exclude on MLX (Mac) — GPU-only quant formats won't load on MLX. */
const EXCLUDED_TAGS_MLX = new Set([
"gptq",
"awq",
"exl2",
"onnx",
"openvino",
"coreml",
"tflite",
"ctranslate2",
]);
// Embedding / sentence-transformer models ship with onnx/openvino as additional
// export formats — they should not be excluded by the tag check above.
const EMBEDDING_TAGS = new Set([
@ -77,7 +91,7 @@ function estimateSizeFromDtypes(
return total > 0 ? total : undefined;
}
function makeMapModel(excludeGguf: boolean) {
function makeMapModel(excludeGguf: boolean, excludedTags: Set<string>) {
return (raw: unknown): HfModelResult | null => {
const m = raw as {
name: string;
@ -87,7 +101,7 @@ function makeMapModel(excludeGguf: boolean) {
tags?: string[];
};
const isEmbedding = m.tags?.some((t) => EMBEDDING_TAGS.has(t));
if (!isEmbedding && m.tags?.some((t) => EXCLUDED_TAGS.has(t))) {
if (!isEmbedding && m.tags?.some((t) => excludedTags.has(t))) {
return null;
}
const isGguf =
@ -314,7 +328,9 @@ export function useHfModelSearch(
[trimmed, searchQuery, pinnedId, task, accessToken, priorityIds],
);
const mapModel = useMemo(() => makeMapModel(excludeGguf), [excludeGguf]);
const deviceType = usePlatformStore((s) => s.deviceType);
const excludedTags = deviceType === "mac" ? EXCLUDED_TAGS_MLX : EXCLUDED_TAGS_GPU;
const mapModel = useMemo(() => makeMapModel(excludeGguf, excludedTags), [excludeGguf, excludedTags]);
const search = useHfPaginatedSearch(createIter, mapModel);
// Secondary sort guarantee: unsloth models always float to the top.

View file

@ -57,7 +57,12 @@ export type VramFitStatus = "fits" | "tight" | "exceeds";
*/
export const FP16_LOADING_BYTES = 2.0;
export type TrainingMethod = "qlora" | "lora" | "full";
export type TrainingMethod = "qlora" | "lora" | "full" | "cpt";
function usesQuantizedLoading(method: TrainingMethod, modelId?: string): boolean {
if (method === "qlora") return true;
return method === "cpt" && (modelId ?? "").toLowerCase().includes("4bit");
}
/**
* Estimate VRAM (GB) needed to load a model with Unsloth.
@ -66,15 +71,18 @@ export type TrainingMethod = "qlora" | "lora" | "full";
* - QLoRA : 4-bit quantized via bnb -> 0.90 bytes/param (calibrated)
* - LoRA : fp16 -> 2.0 bytes/param (theoretical)
* - Full : fp16 -> 2.0 bytes/param (theoretical)
* - CPT : fp16 LoRA (16-bit base) -> 2.0 bytes/param (theoretical)
*
* Formula: totalParams * bytesPerParam + 1.4 GB overhead
*/
export function estimateLoadingVram(
totalParams: number,
method: TrainingMethod = "qlora",
modelId?: string,
): number {
const bytesPerParam =
method === "qlora" ? BNB_4BIT_LOADING_BYTES : FP16_LOADING_BYTES;
const bytesPerParam = usesQuantizedLoading(method, modelId)
? BNB_4BIT_LOADING_BYTES
: FP16_LOADING_BYTES;
const gb = (totalParams / 1e9) * bytesPerParam + LOADING_OVERHEAD_GB;
return Math.round(gb * 10) / 10;
}
@ -119,7 +127,7 @@ export function buildModelVramMap(
continue;
}
const est = estimateLoadingVram(model.totalParams, method);
const est = estimateLoadingVram(model.totalParams, method, model.id);
const status = gpu.available ? checkVramFit(est, gpu.memoryTotalGb) : null;
map.set(model.id, { est, status });
}

View file

@ -2,15 +2,15 @@
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
export type ModelType = "vision" | "audio" | "embeddings" | "text";
export type TrainingMethod = "qlora" | "lora" | "full";
export type TrainingMethod = "qlora" | "lora" | "full" | "cpt";
export function isAdapterMethod(method: TrainingMethod): boolean {
return method === "lora" || method === "qlora";
return method === "lora" || method === "qlora" || method === "cpt";
}
export type StepNumber = 1 | 2 | 3 | 4 | 5;
export type DatasetSource = "huggingface" | "upload";
export type DatasetFormat = "auto" | "alpaca" | "chatml" | "sharegpt";
export type GradientCheckpointing = "none" | "true" | "unsloth";
export type DatasetFormat = "auto" | "alpaca" | "chatml" | "sharegpt" | "raw";
export type GradientCheckpointing = "none" | "true" | "unsloth" | "mlx";
export interface WizardState {
currentStep: StepNumber;

View file

@ -1492,9 +1492,79 @@ if (-not $PythonCmd) {
substep "Using $PythonCmd ($(& $PythonCmd --version 2>&1))"
# The venv must already exist (created by install.ps1).
# This script (setup.ps1 / "unsloth studio update") only updates packages.
$VenvDir = Join-Path $env:USERPROFILE ".unsloth\studio\unsloth_studio"
# The venv must already exist (created by install.ps1); this script only
# updates packages. UNSLOTH_STUDIO_HOME (or STUDIO_HOME alias) overrides the
# root. UNSLOTH_STUDIO_HOME wins when both are set. Whitespace-only values
# are treated as unset to match Python .strip() semantics.
$_studioOverrideVar = $null
$_studioOverride = $null
if (-not [string]::IsNullOrWhiteSpace($env:UNSLOTH_STUDIO_HOME)) {
$_studioOverrideVar = "UNSLOTH_STUDIO_HOME"
$_studioOverride = $env:UNSLOTH_STUDIO_HOME.Trim()
} elseif (-not [string]::IsNullOrWhiteSpace($env:STUDIO_HOME)) {
$_studioOverrideVar = "STUDIO_HOME"
$_studioOverride = $env:STUDIO_HOME.Trim()
}
if ($_studioOverride) {
if ($_studioOverride -eq "~" -or $_studioOverride -like "~/*" -or $_studioOverride -like "~\*") {
$_studioOverride = (Join-Path $env:USERPROFILE $_studioOverride.Substring(1).TrimStart('/','\'))
}
if (Test-Path -LiteralPath $_studioOverride -PathType Container) {
$StudioHome = (Resolve-Path -LiteralPath $_studioOverride).Path
# why: mirror setup.sh:417 and install.ps1:130 -- fail fast when the
# custom root is read-only instead of erroring later while creating
# sidecar venvs / installing packages.
$_setupWriteProbe = Join-Path $StudioHome (".unsloth-write-probe-" + [guid]::NewGuid())
try {
[System.IO.File]::WriteAllText($_setupWriteProbe, "")
Remove-Item -LiteralPath $_setupWriteProbe -Force -ErrorAction SilentlyContinue
} catch {
Write-Host "ERROR: $_studioOverrideVar=$StudioHome is not writable." -ForegroundColor Red
exit 1
}
} else {
Write-Host "ERROR: $_studioOverrideVar=$_studioOverride does not exist." -ForegroundColor Red
Write-Host " Run install.ps1 to create the install root before 'unsloth studio update'." -ForegroundColor Red
exit 1
}
} else {
$StudioHome = Join-Path $env:USERPROFILE ".unsloth\studio"
}
$VenvDir = Join-Path $StudioHome "unsloth_studio"
# why: in env-override mode $StudioHome is user-chosen; require the
# ownership marker before Remove-Item so unrelated dirs survive. Gated on
# the canonical comparison so an override pointing at the legacy default
# still behaves like a default install.
$StudioOwnedMarker = ".unsloth-studio-owned"
$LegacyStudioHome = Join-Path $env:USERPROFILE ".unsloth\studio"
$_studioHomeCanon = $StudioHome
if (Test-Path -LiteralPath $_studioHomeCanon -PathType Container) {
$_studioHomeCanon = (Resolve-Path -LiteralPath $_studioHomeCanon).Path
}
if (Test-Path -LiteralPath $LegacyStudioHome -PathType Container) {
$LegacyStudioHome = (Resolve-Path -LiteralPath $LegacyStudioHome).Path
}
$StudioHomeIsCustom = ($_studioHomeCanon -ne $LegacyStudioHome)
function Assert-StudioOwnedOrAbsent {
param(
[Parameter(Mandatory = $true)][string]$Path,
[Parameter(Mandatory = $true)][string]$Label
)
if (-not (Test-Path -LiteralPath $Path -PathType Container)) { return }
if ($StudioHomeIsCustom -and -not (Test-Path -LiteralPath (Join-Path $Path $StudioOwnedMarker) -PathType Leaf)) {
Write-Host "[ERROR] $Path already exists and is not marked as a Studio-owned $Label." -ForegroundColor Red
Write-Host " Move it aside or choose an empty UNSLOTH_STUDIO_HOME before re-running." -ForegroundColor Yellow
exit 1
}
}
function Mark-StudioOwned {
param([Parameter(Mandatory = $true)][string]$Path)
if (-not (Test-Path -LiteralPath $Path -PathType Container)) { return }
try {
[System.IO.File]::WriteAllText((Join-Path $Path $StudioOwnedMarker), "")
} catch {}
}
# Stale-venv detection: if the venv exists but its torch flavor no longer
# matches the current machine, repair according to invocation context.
@ -1504,12 +1574,12 @@ $VenvDir = Join-Path $env:USERPROFILE ".unsloth\studio\unsloth_studio"
# In no-torch mode, a missing torch package is expected.
$NoTorchMode = $env:UNSLOTH_NO_TORCH -match '^(?i:true|1|yes)$'
$InstallerManagedSetup = $env:UNSLOTH_INSTALL_ROLLBACK_MANAGED -match '^(?i:true|1|yes)$'
if ((Test-Path $VenvDir -PathType Container) -and -not $NoTorchMode) {
if ((Test-Path -LiteralPath $VenvDir -PathType Container) -and -not $NoTorchMode) {
$VenvPyExe = Join-Path $VenvDir "Scripts\python.exe"
$installedTorchTag = $null
$shouldRebuild = $false
if (Test-Path $VenvPyExe) {
if (Test-Path -LiteralPath $VenvPyExe) {
try {
$psi = New-Object System.Diagnostics.ProcessStartInfo
$psi.FileName = $VenvPyExe
@ -1558,8 +1628,21 @@ if ((Test-Path $VenvDir -PathType Container) -and -not $NoTorchMode) {
exit 1
}
substep "Stale venv detected ($reason) -- rebuilding..." "Yellow"
# why: mirror install.ps1 env-mode guard so an update against a custom
# UNSLOTH_STUDIO_HOME never wipes an unrelated unsloth_studio venv;
# -PathType Leaf rejects a directory masquerading as the sentinel.
if (
$StudioHomeIsCustom -and
-not (Test-Path -LiteralPath (Join-Path $VenvDir $StudioOwnedMarker) -PathType Leaf) -and
-not (Test-Path -LiteralPath (Join-Path $StudioHome "share\studio.conf") -PathType Leaf) -and
-not (Test-Path -LiteralPath (Join-Path $StudioHome "bin\unsloth.exe") -PathType Leaf)
) {
Write-Host "[ERROR] $VenvDir already exists but does not look like an Unsloth Studio install." -ForegroundColor Red
Write-Host " Move it aside or choose an empty UNSLOTH_STUDIO_HOME before re-running." -ForegroundColor Yellow
exit 1
}
try {
Remove-Item $VenvDir -Recurse -Force -ErrorAction Stop
Remove-Item -LiteralPath $VenvDir -Recurse -Force -ErrorAction Stop
} catch {
Write-Host " [ERROR] Could not remove stale venv: $($_.Exception.Message)" -ForegroundColor Red
Write-Host " Close any running Studio/Python processes and re-run setup." -ForegroundColor Red
@ -1568,7 +1651,7 @@ if ((Test-Path $VenvDir -PathType Container) -and -not $NoTorchMode) {
}
}
if (-not (Test-Path $VenvDir)) {
if (-not (Test-Path -LiteralPath $VenvDir)) {
Write-Host "[ERROR] Virtual environment not found at $VenvDir" -ForegroundColor Red
Write-Host " Run install.ps1 first to create the environment:" -ForegroundColor Yellow
Write-Host " irm https://unsloth.ai/install.ps1 | iex" -ForegroundColor Yellow
@ -1759,17 +1842,19 @@ if ($stackExit -ne 0) {
# ── Pre-install transformers 5.x into .venv_t5_530/ and .venv_t5_550/ ──
# Runs outside the deps fast-path gate so that upgrades from the legacy
# single .venv_t5 are always migrated to the tiered layout.
$VenvT5_530Dir = Join-Path $env:USERPROFILE ".unsloth\studio\.venv_t5_530"
$VenvT5_550Dir = Join-Path $env:USERPROFILE ".unsloth\studio\.venv_t5_550"
$VenvT5Legacy = Join-Path $env:USERPROFILE ".unsloth\studio\.venv_t5"
# T5 sidecar venvs live under the resolved $StudioHome so custom installs are self-contained.
$VenvT5_530Dir = Join-Path $StudioHome ".venv_t5_530"
$VenvT5_550Dir = Join-Path $StudioHome ".venv_t5_550"
$VenvT5Legacy = Join-Path $StudioHome ".venv_t5"
$_NeedT5Install = $false
if (Test-Path $VenvT5Legacy) {
Remove-Item -Recurse -Force $VenvT5Legacy
if (Test-Path -LiteralPath $VenvT5Legacy) {
Assert-StudioOwnedOrAbsent -Path $VenvT5Legacy -Label "legacy transformers sidecar venv"
Remove-Item -LiteralPath $VenvT5Legacy -Recurse -Force
$_NeedT5Install = $true
}
if (-not (Test-Path $VenvT5_530Dir)) { $_NeedT5Install = $true }
if (-not (Test-Path $VenvT5_550Dir)) { $_NeedT5Install = $true }
if (-not (Test-Path -LiteralPath $VenvT5_530Dir)) { $_NeedT5Install = $true }
if (-not (Test-Path -LiteralPath $VenvT5_550Dir)) { $_NeedT5Install = $true }
# Also reinstall when python deps were updated
if (-not $SkipPythonDeps) { $_NeedT5Install = $true }
@ -1781,8 +1866,10 @@ $ErrorActionPreference = "Continue"
# --- .venv_t5_530 (transformers 5.3.0) ---
substep "pre-installing transformers 5.3.0 for newer model support..."
if (Test-Path $VenvT5_530Dir) { Remove-Item -Recurse -Force $VenvT5_530Dir }
New-Item -ItemType Directory -Path $VenvT5_530Dir -Force | Out-Null
Assert-StudioOwnedOrAbsent -Path $VenvT5_530Dir -Label "transformers 5.3 sidecar venv"
if (Test-Path -LiteralPath $VenvT5_530Dir) { Remove-Item -LiteralPath $VenvT5_530Dir -Recurse -Force }
[System.IO.Directory]::CreateDirectory($VenvT5_530Dir) | Out-Null
Mark-StudioOwned -Path $VenvT5_530Dir
foreach ($pkg in @("transformers==5.3.0", "huggingface_hub==1.8.0", "hf_xet==1.4.2")) {
if ($script:UnslothVerbose) {
Fast-Install --target $VenvT5_530Dir --no-deps $pkg
@ -1814,8 +1901,10 @@ step "transformers" "5.3.0 pre-installed"
# --- .venv_t5_550 (transformers 5.5.0) ---
substep "pre-installing transformers 5.5.0 for Gemma 4 support..."
if (Test-Path $VenvT5_550Dir) { Remove-Item -Recurse -Force $VenvT5_550Dir }
New-Item -ItemType Directory -Path $VenvT5_550Dir -Force | Out-Null
Assert-StudioOwnedOrAbsent -Path $VenvT5_550Dir -Label "transformers 5.5 sidecar venv"
if (Test-Path -LiteralPath $VenvT5_550Dir) { Remove-Item -LiteralPath $VenvT5_550Dir -Recurse -Force }
[System.IO.Directory]::CreateDirectory($VenvT5_550Dir) | Out-Null
Mark-StudioOwned -Path $VenvT5_550Dir
foreach ($pkg in @("transformers==5.5.0", "huggingface_hub==1.8.0", "hf_xet==1.4.2")) {
if ($script:UnslothVerbose) {
Fast-Install --target $VenvT5_550Dir --no-deps $pkg
@ -1851,8 +1940,15 @@ step "transformers" "5.5.0 pre-installed"
# ==========================================================================
# PHASE 3.4: Prefer prebuilt llama.cpp bundles before source build
# ==========================================================================
$UnslothHome = Join-Path $env:USERPROFILE ".unsloth"
if (-not (Test-Path $UnslothHome)) { New-Item -ItemType Directory -Force $UnslothHome | Out-Null }
# Nest llama.cpp under $StudioHome only for real env-overrides, never the
# legacy default. Reuses $StudioHomeIsCustom from the canonical comparison
# computed above so the llama.cpp nest matches ownership-guard semantics.
if ($StudioHomeIsCustom) {
$UnslothHome = $StudioHome
} else {
$UnslothHome = Join-Path $env:USERPROFILE ".unsloth"
}
if (-not (Test-Path -LiteralPath $UnslothHome)) { [System.IO.Directory]::CreateDirectory($UnslothHome) | Out-Null }
$LlamaCppDir = Join-Path $UnslothHome "llama.cpp"
$NeedLlamaSourceBuild = $false
$SkipPrebuiltInstall = $false
@ -1954,9 +2050,15 @@ if ($env:UNSLOTH_LLAMA_FORCE_COMPILE -eq "1") {
} else {
Write-Host ""
substep "installing prebuilt llama.cpp bundle (preferred path)..."
if (Test-Path $LlamaCppDir) {
if (Test-Path -LiteralPath $LlamaCppDir) {
substep "Existing llama.cpp install detected -- validating staged prebuilt update before replacement"
}
# why: install_llama_prebuilt.py uses os.replace(), which would displace
# an unrelated $env:UNSLOTH_STUDIO_HOME\llama.cpp before the source-build
# ownership check below ever runs.
if ($StudioHomeIsCustom) {
Assert-StudioOwnedOrAbsent -Path $LlamaCppDir -Label "llama.cpp install"
}
$prebuiltArgs = @(
"$PSScriptRoot\install_llama_prebuilt.py",
"--install-dir", $LlamaCppDir,
@ -2001,6 +2103,9 @@ if ($env:UNSLOTH_LLAMA_FORCE_COMPILE -eq "1") {
} else {
step "llama.cpp" "prebuilt installed and validated"
}
if ($StudioHomeIsCustom -and (Test-Path -LiteralPath $LlamaCppDir -PathType Container)) {
Mark-StudioOwned -Path $LlamaCppDir
}
$installedRelease = Get-InstalledLlamaPrebuiltRelease -InstallDir $LlamaCppDir
if ($installedRelease) {
substep $installedRelease
@ -2008,7 +2113,7 @@ if ($env:UNSLOTH_LLAMA_FORCE_COMPILE -eq "1") {
} elseif ($prebuiltExit -eq 3) {
step "llama.cpp" "install blocked by active llama.cpp process" "Yellow"
Write-LlamaFailureLog -Output $prebuiltOutput
if (Test-Path $LlamaCppDir) {
if (Test-Path -LiteralPath $LlamaCppDir) {
substep "Existing install was restored" "Yellow"
}
substep "Close Studio or other llama.cpp users and retry" "Yellow"
@ -2016,7 +2121,7 @@ if ($env:UNSLOTH_LLAMA_FORCE_COMPILE -eq "1") {
} else {
step "llama.cpp" "prebuilt install failed (continuing)" "Yellow"
Write-LlamaFailureLog -Output $prebuiltOutput
if (Test-Path $LlamaCppDir) {
if (Test-Path -LiteralPath $LlamaCppDir) {
substep "Prebuilt update failed; existing install was restored or cleaned before source build fallback" "Yellow"
}
substep "Prebuilt llama.cpp path unavailable or failed validation -- falling back to source build" "Yellow"
@ -2092,10 +2197,10 @@ $HasCmakeForBuild = $null -ne (Get-Command cmake -ErrorAction SilentlyContinue)
# Check if existing llama-server matches current GPU mode. A CUDA-built binary
# on a now-CPU-only machine (or vice versa) needs to be rebuilt.
$NeedRebuild = $false
if (Test-Path $LlamaServerBin) {
if (Test-Path -LiteralPath $LlamaServerBin) {
$CmakeCacheFile = Join-Path $BuildDir "CMakeCache.txt"
if (Test-Path $CmakeCacheFile) {
$cachedCuda = Select-String -Path $CmakeCacheFile -Pattern 'GGML_CUDA:BOOL=ON' -Quiet
if (Test-Path -LiteralPath $CmakeCacheFile) {
$cachedCuda = Select-String -LiteralPath $CmakeCacheFile -Pattern 'GGML_CUDA:BOOL=ON' -Quiet
if ($HasNvidiaSmi -and -not $cachedCuda) {
Write-Host " Existing llama-server is CPU-only but GPU is available -- rebuilding" -ForegroundColor Yellow
$NeedRebuild = $true
@ -2109,7 +2214,7 @@ if (Test-Path $LlamaServerBin) {
if (-not $NeedLlamaSourceBuild) {
Write-Host ""
step "llama.cpp" "prebuilt (validated)"
} elseif ((Test-Path $LlamaServerBin) -and -not $NeedRebuild -and $RequestedLlamaTag -ne "master") {
} elseif ((Test-Path -LiteralPath $LlamaServerBin) -and -not $NeedRebuild -and $RequestedLlamaTag -ne "master") {
# Skip rebuild only for pinned tags (e.g. b8635). When the requested
# tag is "master" (a moving target), always rebuild so the binary picks
# up new model architecture support (e.g. Gemma 4).
@ -2211,7 +2316,13 @@ if (-not $NeedLlamaSourceBuild) {
$UseConcreteRef = ($ResolvedSourceRef -ne "latest" -and -not [string]::IsNullOrWhiteSpace($ResolvedSourceRef))
if (Test-Path (Join-Path $LlamaCppDir ".git")) {
if (Test-Path -LiteralPath (Join-Path $LlamaCppDir ".git")) {
# why: in-place git mutation (remote set-url, checkout -B, clean -fdx)
# rewrites $LlamaCppDir; mirror the prebuilt and temp-dir-swap guards
# so an unrelated workspace .git tree is never silently overwritten.
if ($StudioHomeIsCustom) {
Assert-StudioOwnedOrAbsent -Path $LlamaCppDir -Label "llama.cpp install"
}
Write-Host " Syncing llama.cpp to $ResolvedSourceRef..." -ForegroundColor Gray
# Always sync the remote URL so switching between default/fork sources works
Invoke-SetupCommand -AlwaysQuiet { git -C $LlamaCppDir remote set-url origin "$ResolvedSourceUrl.git" } | Out-Null
@ -2282,24 +2393,30 @@ if (-not $NeedLlamaSourceBuild) {
}
}
}
# why: in-place git-sync (the temp-dir clone path calls Mark-StudioOwned
# at swap-time) must mark the existing tree so a subsequent prebuilt
# update path's Assert-StudioOwnedOrAbsent does not exit on the same root.
if ($BuildOk -and $StudioHomeIsCustom) {
Mark-StudioOwned -Path $LlamaCppDir
}
} else {
Write-Host " Cloning llama.cpp @ $ResolvedSourceRef..." -ForegroundColor Gray
$buildTmp = "$LlamaCppDir.build.$PID"
$null = New-Item -ItemType Directory -Force -Path (Split-Path $LlamaCppDir -Parent)
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
$null = [System.IO.Directory]::CreateDirectory((Split-Path -LiteralPath $LlamaCppDir))
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
if ($LlamaPr) {
$cloneExit = Invoke-SetupCommand -AlwaysQuiet { git clone --depth 1 "$LlamaSource.git" $buildTmp }
if ($cloneExit -ne 0) {
$BuildOk = $false
$FailedStep = "git clone"
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
}
if ($BuildOk) {
$fetchExit = Invoke-SetupCommand -AlwaysQuiet { git -C $buildTmp fetch --depth 1 origin "pull/$LlamaPr/head:pr-$LlamaPr" }
if ($fetchExit -ne 0) {
$BuildOk = $false
$FailedStep = "git fetch PR #$LlamaPr"
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
}
}
if ($BuildOk) {
@ -2307,7 +2424,7 @@ if (-not $NeedLlamaSourceBuild) {
if ($checkoutExit -ne 0) {
$BuildOk = $false
$FailedStep = "git checkout PR #$LlamaPr"
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
}
}
} elseif ($ResolvedSourceRefKind -eq "pull") {
@ -2315,14 +2432,14 @@ if (-not $NeedLlamaSourceBuild) {
if ($cloneExit -ne 0) {
$BuildOk = $false
$FailedStep = "git clone"
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
}
if ($BuildOk) {
$fetchExit = Invoke-SetupCommand -AlwaysQuiet { git -C $buildTmp fetch --depth 1 origin $ResolvedSourceRef }
if ($fetchExit -ne 0) {
$BuildOk = $false
$FailedStep = "git fetch source PR ref"
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
}
}
if ($BuildOk) {
@ -2330,7 +2447,7 @@ if (-not $NeedLlamaSourceBuild) {
if ($checkoutExit -ne 0) {
$BuildOk = $false
$FailedStep = "git checkout source PR ref"
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
}
}
} elseif ($ResolvedSourceRefKind -eq "commit") {
@ -2338,14 +2455,14 @@ if (-not $NeedLlamaSourceBuild) {
if ($cloneExit -ne 0) {
$BuildOk = $false
$FailedStep = "git clone"
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
}
if ($BuildOk) {
$fetchExit = Invoke-SetupCommand -AlwaysQuiet { git -C $buildTmp fetch --depth 1 origin $ResolvedSourceRef }
if ($fetchExit -ne 0) {
$BuildOk = $false
$FailedStep = "git fetch source commit"
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
}
}
if ($BuildOk) {
@ -2353,7 +2470,7 @@ if (-not $NeedLlamaSourceBuild) {
if ($checkoutExit -ne 0) {
$BuildOk = $false
$FailedStep = "git checkout source commit"
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
}
}
} else {
@ -2366,7 +2483,7 @@ if (-not $NeedLlamaSourceBuild) {
if ($cloneExit -ne 0) {
$BuildOk = $false
$FailedStep = "git clone"
if (Test-Path $buildTmp) { Remove-Item -Recurse -Force $buildTmp }
if (Test-Path -LiteralPath $buildTmp) { Remove-Item -LiteralPath $buildTmp -Recurse -Force }
}
}
# Use temp dir for build; swap into $LlamaCppDir only after build succeeds
@ -2482,14 +2599,16 @@ if (-not $NeedLlamaSourceBuild) {
# Swap temp build dir into final location (only if we built in a temp dir)
if ($BuildOk -and $LlamaCppDir -ne $OriginalLlamaCppDir) {
if (Test-Path $OriginalLlamaCppDir) { Remove-Item -Recurse -Force $OriginalLlamaCppDir }
Move-Item $LlamaCppDir $OriginalLlamaCppDir
Assert-StudioOwnedOrAbsent -Path $OriginalLlamaCppDir -Label "llama.cpp install"
if (Test-Path -LiteralPath $OriginalLlamaCppDir) { Remove-Item -LiteralPath $OriginalLlamaCppDir -Recurse -Force }
Move-Item -LiteralPath $LlamaCppDir -Destination $OriginalLlamaCppDir
$LlamaCppDir = $OriginalLlamaCppDir
$BuildDir = Join-Path $LlamaCppDir "build"
$LlamaServerBin = Join-Path $BuildDir "bin\Release\llama-server.exe"
Mark-StudioOwned -Path $LlamaCppDir
} elseif (-not $BuildOk -and $LlamaCppDir -ne $OriginalLlamaCppDir) {
# Build failed -- clean up temp dir, preserve existing install
if (Test-Path $LlamaCppDir) { Remove-Item -Recurse -Force $LlamaCppDir }
if (Test-Path -LiteralPath $LlamaCppDir) { Remove-Item -LiteralPath $LlamaCppDir -Recurse -Force }
$LlamaCppDir = $OriginalLlamaCppDir
$BuildDir = Join-Path $LlamaCppDir "build"
$LlamaServerBin = Join-Path $BuildDir "bin\Release\llama-server.exe"
@ -2504,16 +2623,16 @@ if (-not $NeedLlamaSourceBuild) {
$totalSec = [math]::Round($totalSw.Elapsed.TotalSeconds % 60, 1)
# -- Summary --
if ($BuildOk -and (Test-Path $LlamaServerBin)) {
if ($BuildOk -and (Test-Path -LiteralPath $LlamaServerBin)) {
step "llama.cpp" "built"
$QuantizeBin = Join-Path $BuildDir "bin\Release\llama-quantize.exe"
if (Test-Path $QuantizeBin) {
if (Test-Path -LiteralPath $QuantizeBin) {
step "llama-quantize" "built"
}
step "build time" "${totalMin}m ${totalSec}s" "DarkGray"
} else {
$altBin = Join-Path $BuildDir "bin\llama-server.exe"
if ($BuildOk -and (Test-Path $altBin)) {
if ($BuildOk -and (Test-Path -LiteralPath $altBin)) {
step "llama.cpp" "built"
step "build time" "${totalMin}m ${totalSec}s" "DarkGray"
} else {

View file

@ -417,7 +417,36 @@ if [ -d "$SCRIPT_DIR/backend/core/data_recipe/oxc-validator" ] && command -v npm
fi
# ── Python venv + deps ──
STUDIO_HOME="$HOME/.unsloth/studio"
# UNSLOTH_STUDIO_HOME (or STUDIO_HOME alias) overrides the install root
# (mirrors install.sh). UNSLOTH_STUDIO_HOME wins when both are set.
_studio_override_var=""
_studio_override="${UNSLOTH_STUDIO_HOME:-}"
if [ -n "$_studio_override" ]; then
_studio_override_var="UNSLOTH_STUDIO_HOME"
else
_studio_override="${STUDIO_HOME:-}"
[ -n "$_studio_override" ] && _studio_override_var="STUDIO_HOME"
fi
# Strip whitespace so " " is treated as unset (matches Python .strip()).
_studio_override=$(printf '%s' "$_studio_override" | sed -e 's/^[[:space:]]*//' -e 's/[[:space:]]*$//')
case "$_studio_override" in
"~") _studio_override="$HOME" ;;
"~/"*) _studio_override="$HOME/${_studio_override#'~/'}" ;;
esac
if [ -n "$_studio_override" ]; then
# setup.sh runs against an existing install (via 'unsloth studio update');
# a typo in the override must fail fast instead of materializing an
# empty workspace dir. Mirrors setup.ps1 behavior.
if [ ! -d "$_studio_override" ]; then
echo "ERROR: $_studio_override_var=$_studio_override does not exist." >&2
echo " Run install.sh to create the install root before 'unsloth studio update'." >&2
exit 1
fi
[ -w "$_studio_override" ] || { echo "ERROR: $_studio_override_var=$_studio_override is not writable." >&2; exit 1; }
STUDIO_HOME="$(CDPATH= cd -P -- "$_studio_override" && pwd -P)" || exit 1
else
STUDIO_HOME="$HOME/.unsloth/studio"
fi
VENV_DIR="$STUDIO_HOME/unsloth_studio"
VENV_T5_530_DIR="$STUDIO_HOME/.venv_t5_530"
VENV_T5_550_DIR="$STUDIO_HOME/.venv_t5_550"
@ -542,9 +571,39 @@ fi
#
# Runs outside the _SKIP_PYTHON_DEPS gate so that upgrades from legacy
# single .venv_t5 are always migrated to the tiered layout.
# why: in env-override mode $STUDIO_HOME is user-chosen; require the
# ownership marker before rm -rf so unrelated dirs survive. Gated on the
# canonical comparison so an override pointing at the legacy default still
# behaves like a default install.
_STUDIO_OWNED_MARKER=".unsloth-studio-owned"
_LEGACY_STUDIO_HOME="$HOME/.unsloth/studio"
_studio_home_canon="$STUDIO_HOME"
if [ -d "$_studio_home_canon" ]; then
_studio_home_canon=$(CDPATH= cd -P -- "$_studio_home_canon" 2>/dev/null && pwd -P) \
|| _studio_home_canon="$STUDIO_HOME"
fi
if [ -d "$_LEGACY_STUDIO_HOME" ]; then
_LEGACY_STUDIO_HOME=$(CDPATH= cd -P -- "$_LEGACY_STUDIO_HOME" 2>/dev/null && pwd -P) \
|| _LEGACY_STUDIO_HOME="$HOME/.unsloth/studio"
fi
_STUDIO_HOME_IS_CUSTOM=false
if [ "$_studio_home_canon" != "$_LEGACY_STUDIO_HOME" ]; then
_STUDIO_HOME_IS_CUSTOM=true
fi
_assert_studio_owned_or_absent() {
_aso_dir="$1"
_aso_label="$2"
[ -d "$_aso_dir" ] || return 0
if [ "$_STUDIO_HOME_IS_CUSTOM" = true ] && [ ! -f "$_aso_dir/$_STUDIO_OWNED_MARKER" ]; then
echo "ERROR: $_aso_dir already exists and is not marked as a Studio-owned $_aso_label." >&2
echo " Move it aside or choose an empty UNSLOTH_STUDIO_HOME before re-running." >&2
exit 1
fi
}
_NEED_T5_INSTALL=false
if [ -d "$STUDIO_HOME/.venv_t5" ]; then
# Legacy layout — migrate
_assert_studio_owned_or_absent "$STUDIO_HOME/.venv_t5" "legacy transformers sidecar venv"
rm -rf "$STUDIO_HOME/.venv_t5"
_NEED_T5_INSTALL=true
fi
@ -554,16 +613,20 @@ fi
[ "$_SKIP_PYTHON_DEPS" = false ] && _NEED_T5_INSTALL=true
if [ "$_NEED_T5_INSTALL" = true ]; then
_assert_studio_owned_or_absent "$VENV_T5_530_DIR" "transformers 5.3 sidecar venv"
[ -d "$VENV_T5_530_DIR" ] && rm -rf "$VENV_T5_530_DIR"
mkdir -p "$VENV_T5_530_DIR"
: > "$VENV_T5_530_DIR/$_STUDIO_OWNED_MARKER" 2>/dev/null || true
run_quiet "install transformers 5.3.0" fast_install --target "$VENV_T5_530_DIR" --no-deps "transformers==5.3.0"
run_quiet "install huggingface_hub for t5_530" fast_install --target "$VENV_T5_530_DIR" --no-deps "huggingface_hub==1.8.0"
run_quiet "install hf_xet for t5_530" fast_install --target "$VENV_T5_530_DIR" --no-deps "hf_xet==1.4.2"
run_quiet "install tiktoken for t5_530" fast_install --target "$VENV_T5_530_DIR" "tiktoken"
step "transformers" "5.3.0 pre-installed"
_assert_studio_owned_or_absent "$VENV_T5_550_DIR" "transformers 5.5 sidecar venv"
[ -d "$VENV_T5_550_DIR" ] && rm -rf "$VENV_T5_550_DIR"
mkdir -p "$VENV_T5_550_DIR"
: > "$VENV_T5_550_DIR/$_STUDIO_OWNED_MARKER" 2>/dev/null || true
run_quiet "install transformers 5.5.0" fast_install --target "$VENV_T5_550_DIR" --no-deps "transformers==5.5.0"
run_quiet "install huggingface_hub for t5_550" fast_install --target "$VENV_T5_550_DIR" --no-deps "huggingface_hub==1.8.0"
run_quiet "install hf_xet for t5_550" fast_install --target "$VENV_T5_550_DIR" --no-deps "hf_xet==1.4.2"
@ -573,7 +636,13 @@ fi
fi
# ── 7. Prefer prebuilt llama.cpp bundles before any source build path ──
UNSLOTH_HOME="$HOME/.unsloth"
# Nest llama.cpp under $STUDIO_HOME only for real env-overrides; legacy
# default keeps ~/.unsloth/llama.cpp so pre-PR builds are still discovered.
if [ "$_STUDIO_HOME_IS_CUSTOM" = true ]; then
UNSLOTH_HOME="$STUDIO_HOME"
else
UNSLOTH_HOME="$HOME/.unsloth"
fi
mkdir -p "$UNSLOTH_HOME"
LLAMA_CPP_DIR="$UNSLOTH_HOME/llama.cpp"
LLAMA_SERVER_BIN="$LLAMA_CPP_DIR/build/bin/llama-server"
@ -582,11 +651,30 @@ _LLAMA_CPP_DEGRADED=false
_LLAMA_FORCE_COMPILE="${UNSLOTH_LLAMA_FORCE_COMPILE:-0}"
_REQUESTED_LLAMA_TAG="${UNSLOTH_LLAMA_TAG:-${_DEFAULT_LLAMA_TAG}}"
_HOST_SYSTEM="$(uname -s 2>/dev/null || true)"
_HOST_MACHINE="$(uname -m 2>/dev/null || true)"
# Pick the release repo install_llama_prebuilt.py plans against.
# unslothai/llama.cpp ships only Linux CUDA bundles, so CPU-only Linux
# x86_64 routes to ggml-org for bin-ubuntu-x64.tar.gz. Anything with a
# GPU tool installed stays on unslothai (CUDA bundle / ROCm source build).
_LINUX_HAS_GPU=false
for _GPU_TOOL in nvidia-smi rocminfo amd-smi hipconfig hipinfo; do
if command -v "$_GPU_TOOL" >/dev/null 2>&1; then
_LINUX_HAS_GPU=true
break
fi
done
if [ "$_HOST_SYSTEM" = "Darwin" ]; then
_HELPER_RELEASE_REPO="ggml-org/llama.cpp"
elif [ "$_HOST_SYSTEM" = "Linux" ] \
&& [ "$_HOST_MACHINE" = "x86_64" ] \
&& [ "$_LINUX_HAS_GPU" = false ]; then
_HELPER_RELEASE_REPO="ggml-org/llama.cpp"
else
_HELPER_RELEASE_REPO="unslothai/llama.cpp"
fi
unset _GPU_TOOL
_LLAMA_PR="${UNSLOTH_LLAMA_PR:-}"
_SKIP_PREBUILT_INSTALL=false
_LLAMA_PR_FORCE="${UNSLOTH_LLAMA_PR_FORCE:-${_DEFAULT_LLAMA_PR_FORCE}}"
@ -635,6 +723,12 @@ else
if [ -d "$LLAMA_CPP_DIR" ]; then
substep "existing install detected -- validating update"
fi
# why: install_llama_prebuilt.py uses os.replace(), which would displace
# an unrelated $UNSLOTH_STUDIO_HOME/llama.cpp before the source-build
# ownership check below ever runs.
if [ "$_STUDIO_HOME_IS_CUSTOM" = true ]; then
_assert_studio_owned_or_absent "$LLAMA_CPP_DIR" "llama.cpp install"
fi
_PREBUILT_CMD=(
python "$SCRIPT_DIR/install_llama_prebuilt.py"
--install-dir "$LLAMA_CPP_DIR"
@ -662,6 +756,9 @@ else
else
step "llama.cpp" "prebuilt installed and validated"
fi
if [ "$_STUDIO_HOME_IS_CUSTOM" = true ] && [ -d "$LLAMA_CPP_DIR" ]; then
: > "$LLAMA_CPP_DIR/$_STUDIO_OWNED_MARKER" 2>/dev/null || true
fi
print_installed_llama_prebuilt_release "$LLAMA_CPP_DIR"
verbose_substep "llama.cpp install dir: $LLAMA_CPP_DIR"
rm -f "$_PREBUILT_LOG"
@ -1032,8 +1129,10 @@ else
# Swap only after build succeeds -- preserves existing install on failure
if [ "$BUILD_OK" = true ]; then
_assert_studio_owned_or_absent "$LLAMA_CPP_DIR" "llama.cpp install"
rm -rf "$LLAMA_CPP_DIR"
mv "$_BUILD_TMP" "$LLAMA_CPP_DIR"
: > "$LLAMA_CPP_DIR/$_STUDIO_OWNED_MARKER" 2>/dev/null || true
# Symlink to llama.cpp root -- check_llama_cpp() looks for the binary there
QUANTIZE_BIN="$LLAMA_CPP_DIR/build/bin/llama-quantize"
if [ -f "$QUANTIZE_BIN" ]; then

View file

@ -60,6 +60,11 @@ pub async fn check_install_status() -> bool {
cmd.env_remove("PYTHONPATH");
}
// Tauri uses the legacy root regardless of UNSLOTH_STUDIO_HOME / STUDIO_HOME;
// probe subprocesses must follow the same isolation as process.rs.
cmd.env_remove("UNSLOTH_STUDIO_HOME");
cmd.env_remove("STUDIO_HOME");
let mut child = match cmd.spawn() {
Ok(c) => c,
Err(e) => {

View file

@ -203,6 +203,11 @@ async fn provision_desktop_auth() -> Result<(), String> {
cmd.env_remove("PYTHONHOME");
cmd.env_remove("PYTHONPATH");
}
// Tauri uses the legacy root regardless of UNSLOTH_STUDIO_HOME / STUDIO_HOME.
// Scrub so provisioning writes match what the Rust auth code reads.
cmd.env_remove("UNSLOTH_STUDIO_HOME");
cmd.env_remove("STUDIO_HOME");
#[cfg(windows)]
{
use std::os::windows::process::CommandExt;

View file

@ -196,6 +196,11 @@ fn spawn_script(
cmd.env_remove("PYTHONPATH");
}
// Tauri only does default-root installs; install.sh / install.ps1 reject
// these under --tauri. Scrub so an inherited value can't trip the guard.
cmd.env_remove("UNSLOTH_STUDIO_HOME");
cmd.env_remove("STUDIO_HOME");
// On Windows, launch the installer directly with CREATE_NO_WINDOW.
// The app process is assigned to a KILL_ON_JOB_CLOSE job in main.rs, so
// child cleanup on crash comes from inherited job membership instead.

View file

@ -102,6 +102,11 @@ async fn run_cli_probe(bin: &std::path::Path, args: &[&str]) -> bool {
cmd.env_remove("PYTHONPATH");
}
// Tauri uses the legacy root regardless of UNSLOTH_STUDIO_HOME / STUDIO_HOME;
// probe subprocesses must follow the same isolation as process.rs.
cmd.env_remove("UNSLOTH_STUDIO_HOME");
cmd.env_remove("STUDIO_HOME");
#[cfg(windows)]
{
use std::os::windows::process::CommandExt;
@ -135,6 +140,11 @@ async fn probe_cli_capability(bin: &std::path::Path) -> Option<DesktopCapability
cmd.env_remove("PYTHONPATH");
}
// Tauri uses the legacy root regardless of UNSLOTH_STUDIO_HOME / STUDIO_HOME;
// probe subprocesses must follow the same isolation as process.rs.
cmd.env_remove("UNSLOTH_STUDIO_HOME");
cmd.env_remove("STUDIO_HOME");
#[cfg(windows)]
{
use std::os::windows::process::CommandExt;

View file

@ -316,6 +316,12 @@ pub fn start_backend(
cmd.env_remove("PYTHONPATH");
}
// Tauri uses the legacy root regardless of UNSLOTH_STUDIO_HOME / STUDIO_HOME;
// scrub so the spawned Python backend can't diverge. UNSLOTH_LLAMA_CPP_PATH
// is a pre-existing user-controlled llama.cpp dir override; keep it.
cmd.env_remove("UNSLOTH_STUDIO_HOME");
cmd.env_remove("STUDIO_HOME");
// On Windows, launch the backend directly with hidden-window flags.
// The app process is assigned to a KILL_ON_JOB_CLOSE job in main.rs, so
// children inherit crash-safe cleanup without the buggy per-child JobObject wrapper.

View file

@ -61,6 +61,11 @@ fn spawn_update(
cmd.env_remove("PYTHONPATH");
}
// Tauri manages the legacy root; scrub so 'unsloth studio update' targets
// the same install the desktop app uses, not an inherited custom root.
cmd.env_remove("UNSLOTH_STUDIO_HOME");
cmd.env_remove("STUDIO_HOME");
#[cfg(windows)]
let mut child: Box<dyn ChildWrapper + Send> = {
use std::os::windows::process::CommandExt;

141
tests/conftest.py Normal file
View file

@ -0,0 +1,141 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
"""GPU-free test harness.
unsloth's import chain hits unsloth_zoo.device_type, which calls
get_device_type() at import time and raises NotImplementedError on CI
runners with no CUDA / XPU / HIP visible. Pre-load the real
unsloth_zoo.device_type under a temporarily-mocked
torch.cuda.is_available() so its @cache permanently captures "cuda".
On a real accelerator the pre-load is skipped and detection runs
normally.
Mirrors the conftest harness in unslothai/unsloth-zoo PR #624.
"""
from __future__ import annotations
import importlib.util
import os
import sys
import types
def _has_real_accelerator() -> bool:
try:
import torch
except Exception:
return False
for probe in (
lambda: hasattr(torch, "cuda") and torch.cuda.is_available(),
lambda: hasattr(torch, "xpu") and torch.xpu.is_available(),
lambda: hasattr(torch, "accelerator") and torch.accelerator.is_available(),
):
try:
if probe():
return True
except Exception:
pass
return False
def _preload_device_type(package: str, prereqs: tuple[str, ...] = ()) -> bool:
"""Pre-load <package>.device_type under a mocked
torch.cuda.is_available() == True so its @cache permanently
captures "cuda". prereqs lists submodule names of <package> that
must be loaded first (e.g. 'utils' for unsloth_zoo). Returns False
if the package or any prerequisite cannot be imported, in which
case the caller falls back to a stub."""
target = f"{package}.device_type"
if target in sys.modules:
return True
pkg_spec = importlib.util.find_spec(package)
if pkg_spec is None or not pkg_spec.submodule_search_locations:
return False
pkg_path = pkg_spec.submodule_search_locations[0]
skeleton_already = package in sys.modules
if not skeleton_already:
skel = types.ModuleType(package)
skel.__path__ = [pkg_path]
skel.__spec__ = pkg_spec
skel.__package__ = package
sys.modules[package] = skel
try:
for prereq in prereqs:
full = f"{package}.{prereq}"
if full in sys.modules:
continue
prereq_path = os.path.join(pkg_path, f"{prereq}.py")
prereq_spec = importlib.util.spec_from_file_location(full, prereq_path)
prereq_mod = importlib.util.module_from_spec(prereq_spec)
sys.modules[full] = prereq_mod
prereq_spec.loader.exec_module(prereq_mod)
device_type_path = os.path.join(pkg_path, "device_type.py")
dt_spec = importlib.util.spec_from_file_location(target, device_type_path)
dt_mod = importlib.util.module_from_spec(dt_spec)
sys.modules[target] = dt_mod
import torch
_orig_is_avail = torch.cuda.is_available
torch.cuda.is_available = lambda: True # type: ignore[assignment]
try:
dt_spec.loader.exec_module(dt_mod)
finally:
torch.cuda.is_available = _orig_is_avail
except Exception:
sys.modules.pop(target, None)
return False
finally:
if not skeleton_already:
sys.modules.pop(package, None)
return True
def _patch_torch_cuda_for_import() -> None:
"""Stub torch.cuda.* probes that fire at IMPORT time of unsloth /
unsloth_zoo when DEVICE_TYPE was forced to "cuda" above. These are
queries, not real GPU work, so returning plausible Ampere values
lets the import chain finish; tests that touch real tensors run on
CPU like normal."""
try:
import torch.cuda.memory as _cuda_memory # type: ignore
_cuda_memory.mem_get_info = lambda *a, **k: (0, 80 * 1024**3)
except Exception:
pass
try:
import torch
torch.cuda.get_device_capability = lambda *a, **k: (8, 0)
torch.cuda.is_bf16_supported = lambda *a, **k: True
except Exception:
pass
def _install_device_type_stub(name: str) -> None:
stub = types.ModuleType(name)
stub.DEVICE_TYPE = "cuda"
stub.DEVICE_TYPE_TORCH = "cuda"
stub.DEVICE_COUNT = 1
stub.ALLOW_PREQUANTIZED_MODELS = False
stub.is_hip = lambda: False
stub.get_device_type = lambda: "cuda"
stub.get_device_count = lambda: 1
stub.device_synchronize = lambda *a, **k: None
stub.device_empty_cache = lambda *a, **k: None
stub.device_is_bf16_supported = lambda *a, **k: False
sys.modules[name] = stub
if not _has_real_accelerator():
if not _preload_device_type("unsloth_zoo", prereqs = ("utils",)):
_install_device_type_stub("unsloth_zoo.device_type")
if not _preload_device_type("unsloth"):
_install_device_type_stub("unsloth.device_type")
_patch_torch_cuda_for_import()

View file

@ -8,11 +8,45 @@
from __future__ import annotations
import importlib.util
import os
import pathlib
import sys
import types
import pytest
def _stub_module(name: str) -> types.ModuleType:
# __spec__ must be set so importlib.util.find_spec(name) does not raise
# ValueError if a downstream test imports the real package.
mod = types.ModuleType(name)
mod.__spec__ = importlib.util.spec_from_loader(name, loader = None)
return mod
_STUB_KEYS = (
"transformers",
"sentence_transformers",
"sentence_transformers.models",
)
@pytest.fixture(autouse = True)
def _restore_sys_modules():
"""Snapshot the entries we shadow with stubs and restore them after each
test so a downstream test that does `import transformers` for real does
not pick up our non-package stub."""
saved = {k: sys.modules.get(k) for k in _STUB_KEYS}
try:
yield
finally:
for k, v in saved.items():
if v is None:
sys.modules.pop(k, None)
else:
sys.modules[k] = v
class _FakeAuto:
def __init__(self, name):
@ -45,14 +79,14 @@ class _RaisingTransformer:
def _build_driver(transformer_class):
transformers_mod = types.ModuleType("transformers")
transformers_mod = _stub_module("transformers")
transformers_mod.AutoModel = _FakeAuto("AutoModel")
transformers_mod.AutoProcessor = _FakeAuto("AutoProcessor")
transformers_mod.AutoTokenizer = _FakeAuto("AutoTokenizer")
sys.modules["transformers"] = transformers_mod
st_root = types.ModuleType("sentence_transformers")
st_models = types.ModuleType("sentence_transformers.models")
st_root = _stub_module("sentence_transformers")
st_models = _stub_module("sentence_transformers.models")
st_models.Transformer = transformer_class
sys.modules["sentence_transformers"] = st_root
sys.modules["sentence_transformers.models"] = st_models

View file

@ -0,0 +1,46 @@
import ast
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[2]
GPU_INIT = REPO_ROOT / "unsloth" / "_gpu_init.py"
def _find_geteuid_guard(tree: ast.AST):
for node in ast.walk(tree):
if not isinstance(node, ast.If):
continue
for sub in ast.walk(node.test):
if isinstance(sub, ast.Call) and isinstance(sub.func, ast.Attribute):
if sub.func.attr == "geteuid":
return node
return None
def test_gpu_init_has_geteuid_guard():
tree = ast.parse(GPU_INIT.read_text())
guard = _find_geteuid_guard(tree)
assert (
guard is not None
), "_gpu_init.py must guard ldconfig recovery on os.geteuid()"
def test_ldconfig_calls_only_inside_geteuid_guard():
src = GPU_INIT.read_text()
tree = ast.parse(src)
guard = _find_geteuid_guard(tree)
assert guard is not None
guard_src = ast.get_source_segment(src, guard) or ""
ldconfig_lines = [
line for line in src.splitlines() if "ldconfig" in line and "os.system" in line
]
for line in ldconfig_lines:
assert line.strip() in guard_src, (
"os.system('ldconfig ...') must live inside the geteuid guard, "
f"but found unguarded: {line!r}"
)
def test_non_root_branch_warns_when_bnb_present():
src = GPU_INIT.read_text()
assert "elif bnb is not None" in src
assert "sudo ldconfig" in src

View file

@ -754,6 +754,9 @@ def write_linux_install_shape(install_dir: Path) -> None:
(install_dir / "llama-quantize").write_text("#!/bin/sh\n", encoding = "utf-8")
(runtime_dir / "llama-server").write_text("#!/bin/sh\n", encoding = "utf-8")
(runtime_dir / "llama-quantize").write_text("#!/bin/sh\n", encoding = "utf-8")
# Mirror the runtime payload health groups in install_llama_prebuilt.py:
# libllama-common.so* was added by PR #5135 and is required.
(runtime_dir / "libllama-common.so.0").write_bytes(b"DLL")
(runtime_dir / "libllama.so.0").write_bytes(b"DLL")
(runtime_dir / "libggml.so.0").write_bytes(b"DLL")
(runtime_dir / "libggml-base.so.0").write_bytes(b"DLL")

View file

@ -53,6 +53,32 @@ _has_usable_nvidia_gpu = stack_mod._has_usable_nvidia_gpu
_ROCM_TORCH_INDEX = stack_mod._ROCM_TORCH_INDEX
def _extract_sh_function_body(source: str, name: str) -> str:
"""Return the body of a shell function from `source` by brace matching.
Used by structural tests that need to assert ordering of helper
calls inside a specific function rather than across the whole
install.sh file.
"""
needle = f"{name}() {{"
start = source.find(needle)
if start < 0:
return ""
depth = 0
i = start + len(needle) - 1 # land on the opening brace
n = len(source)
while i < n:
ch = source[i]
if ch == "{":
depth += 1
elif ch == "}":
depth -= 1
if depth == 0:
return source[start : i + 1]
i += 1
return source[start:]
# ── Helper: build HostInfo for different scenarios ──────────────────────────
@ -561,12 +587,13 @@ class TestEnsureRocmTorch:
_ensure_rocm_torch()
mock_pip.assert_not_called()
@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_cpu_torch_gets_rocm_reinstall(
self, mock_ver, mock_gpu, mock_nvidia, mock_pip
self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""CPU-only torch on ROCm host should trigger reinstall."""
mock_probe = MagicMock()
@ -575,12 +602,11 @@ class TestEnsureRocmTorch:
with patch("os.path.isdir", return_value = True):
with patch("subprocess.run", return_value = mock_probe):
_ensure_rocm_torch()
# Should call pip_install twice: once for torch, once for bitsandbytes
assert mock_pip.call_count == 2
torch_call = mock_pip.call_args_list[0]
assert "rocm7.1" in str(torch_call)
bnb_call = mock_pip.call_args_list[1]
assert "bitsandbytes" in str(bnb_call)
# Should install torch via pip_install and bitsandbytes via pip_install_try.
assert mock_pip.call_count == 1
assert "rocm7.1" in str(mock_pip.call_args_list[0])
assert mock_pip_try.call_count >= 1
assert "bitsandbytes" in str(mock_pip_try.call_args_list[0])
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@ -642,12 +668,13 @@ class TestEnsureRocmTorch:
torch_call = mock_pip.call_args_list[0]
assert "rocm7.1" in str(torch_call)
@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_probe_timeout_triggers_reinstall(
self, mock_ver, mock_gpu, mock_nvidia, mock_pip
self, mock_ver, mock_gpu, mock_nvidia, mock_pip, mock_pip_try
):
"""Probe subprocess timeout should not crash; should proceed to reinstall."""
with patch("os.path.isdir", return_value = True):
@ -656,8 +683,10 @@ class TestEnsureRocmTorch:
):
_ensure_rocm_torch()
# If probe times out, the function should treat torch as unusable and reinstall
assert mock_pip.call_count == 2
# both torch (via pip_install) and bitsandbytes (via pip_install_try).
assert mock_pip.call_count == 1
assert "rocm7.1" in str(mock_pip.call_args_list[0])
assert mock_pip_try.call_count >= 1
@patch.object(stack_mod, "pip_install")
@patch.object(stack_mod, "_has_usable_nvidia_gpu", return_value = False)
@ -857,15 +886,33 @@ class TestInstallShStructure:
assert "rocm" in source.lower()
def test_cuda_precedence(self):
"""ROCm detection should only run when nvidia-smi is absent."""
"""ROCm detection should only run when nvidia-smi is absent.
install.sh defines _has_amd_rocm_gpu and _has_usable_nvidia_gpu
helpers near each other (file-position order has no semantic
meaning), so check the runtime ordering inside
get_torch_index_url instead: NVIDIA branch runs first and the
AMD/ROCm branch only fires inside the `if [ -z "$_smi" ]`
block.
"""
sh_path = PACKAGE_ROOT / "install.sh"
source = sh_path.read_text()
# The ROCm block should be inside the "if [ -z "$_smi" ]" branch
smi_block_start = source.find('if [ -z "$_smi" ]')
rocm_block_start = source.find("amd-smi")
body = _extract_sh_function_body(source, "get_torch_index_url")
nvidia_call = body.find("_has_usable_nvidia_gpu")
no_nvidia_branch = body.find('if [ -z "$_smi" ]')
rocm_call = body.find("_has_amd_rocm_gpu")
assert (
smi_block_start < rocm_block_start
), "ROCm detection should be inside the 'no nvidia-smi' branch"
nvidia_call >= 0
), "get_torch_index_url should call _has_usable_nvidia_gpu"
assert (
no_nvidia_branch >= 0
), "get_torch_index_url should gate ROCm on no-nvidia-smi"
assert (
rocm_call > no_nvidia_branch
), "ROCm detection should sit inside the 'no nvidia-smi' branch"
assert (
nvidia_call < no_nvidia_branch
), "NVIDIA detection should run before the no-nvidia-smi branch"
def test_bitsandbytes_amd_install(self):
"""install.sh should install bitsandbytes for AMD when ROCm detected."""
@ -963,16 +1010,32 @@ class TestLiveRegression:
if not shutil.which("nvidia-smi"):
pytest.skip("No nvidia-smi available")
sh_path = PACKAGE_ROOT / "install.sh"
# Extract just the function (don't source the whole installer)
result = subprocess.run(
# Skip if nvidia-smi exists but does not actually list a GPU on this
# host (containers occasionally ship the binary without a driver).
check = subprocess.run(
[
"bash",
"-c",
f"eval \"$(sed -n '/^get_torch_index_url()/,/^}}/p' '{sh_path}')\"; "
"get_torch_index_url",
"nvidia-smi -L 2>/dev/null | "
"awk '/^GPU[[:space:]]+[0-9]+:/{f=1} END{exit !f}'",
],
capture_output = True,
)
if check.returncode != 0:
pytest.skip("nvidia-smi is on PATH but no GPU is listed")
sh_path = PACKAGE_ROOT / "install.sh"
# get_torch_index_url calls _has_usable_nvidia_gpu and
# _has_amd_rocm_gpu, so all three function definitions must be
# in scope when we eval the extract.
extract_cmd = (
f"sed -n '/^_has_amd_rocm_gpu()/,/^}}$/p; "
f"/^_has_usable_nvidia_gpu()/,/^}}$/p; "
f"/^get_torch_index_url()/,/^}}$/p' '{sh_path}'"
)
result = subprocess.run(
["bash", "-c", f'eval "$({extract_cmd})"; get_torch_index_url'],
capture_output = True,
text = True,
timeout = 30,
)
@ -988,19 +1051,23 @@ class TestLiveRegression:
# Load worker.py module
_WORKER_PATH = PACKAGE_ROOT / "studio" / "backend" / "core" / "training" / "worker.py"
# The wheel-probe subprocess was hoisted out of worker.py into wheel_utils
# during the wheel-resolver refactor; the probe script literal lives there.
_WHEEL_UTILS_PATH = PACKAGE_ROOT / "studio" / "backend" / "utils" / "wheel_utils.py"
class TestWorkerRocmMambaSsm:
"""Verify worker.py Mamba/SSM install logic on ROCm."""
def test_probe_returns_hip_version_field(self):
"""_probe_causal_conv1d_env probe script should include hip_version."""
source = _WORKER_PATH.read_text()
assert "hip_version" in source
"""The wheel probe should include hip_version, and worker.py should
consume it."""
assert "hip_version" in _WHEEL_UTILS_PATH.read_text()
assert "hip_version" in _WORKER_PATH.read_text()
def test_probe_script_has_getattr_hip(self):
"""Probe script should use getattr for torch.version.hip (safe on CUDA)."""
source = _WORKER_PATH.read_text()
source = _WHEEL_UTILS_PATH.read_text()
assert "getattr(torch.version, 'hip', None)" in source
def test_direct_wheel_url_returns_none_without_cuda_major(self):
@ -1121,7 +1188,7 @@ class TestAmdGpuMonitoring:
assert metrics["vram_utilization_pct"] is not None
assert metrics["power_utilization_pct"] is not None
def test_amd_primary_gpu_with_mock(self):
def test_amd_primary_gpu_with_mock(self, monkeypatch):
"""get_primary_gpu_utilization returns correct dict with mocked amd-smi."""
amd_path = PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "amd.py"
_amd_spec = importlib.util.spec_from_file_location("test_amd2", amd_path)
@ -1136,6 +1203,17 @@ class TestAmdGpuMonitoring:
except Exception:
pytest.skip("Could not load amd module")
# _first_visible_amd_gpu_id() short-circuits to None when any of
# HIP / ROCR / CUDA_VISIBLE_DEVICES is set to "" or "-1". CI runners
# often unset CUDA at the env level by setting CUDA_VISIBLE_DEVICES
# to "" so the test must not inherit that.
for var in (
"HIP_VISIBLE_DEVICES",
"ROCR_VISIBLE_DEVICES",
"CUDA_VISIBLE_DEVICES",
):
monkeypatch.delenv(var, raising = False)
mock_json = json.dumps(
[
{
@ -1216,27 +1294,45 @@ class TestHardwareAmdBranching:
assert "from . import amd" in source
def test_hardware_branches_on_is_rocm_for_utilization(self):
"""get_gpu_utilization should check IS_ROCM before choosing backend."""
"""get_gpu_utilization should dispatch to amd.py via _smi_query
when IS_ROCM, and the dispatcher itself must check IS_ROCM and
import the amd backend."""
hw_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
)
source = hw_path.read_text()
# Find the get_gpu_utilization function
func_start = source.find("def get_gpu_utilization")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "IS_ROCM" in func_body
assert "amd.get_primary_gpu_utilization" in func_body
assert '_smi_query("get_primary_gpu_utilization"' in func_body
smi = source[
source.find("def _smi_query") : source.find(
"\ndef ", source.find("def _smi_query") + 1
)
]
assert "IS_ROCM" in smi
assert "from . import amd" in smi
def test_hardware_branches_on_is_rocm_for_visible(self):
"""get_visible_gpu_utilization should check IS_ROCM."""
"""get_visible_gpu_utilization should dispatch to amd.py via
_smi_query when IS_ROCM."""
hw_path = (
PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py"
)
source = hw_path.read_text()
func_start = source.find("def get_visible_gpu_utilization")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "IS_ROCM" in func_body
assert "amd.get_visible_gpu_utilization" in func_body
# The dispatcher call may wrap onto multiple lines; allow whitespace
# between the open paren and the literal func name argument.
import re as _re
assert _re.search(r'_smi_query\(\s*"get_visible_gpu_utilization"', func_body)
smi = source[
source.find("def _smi_query") : source.find(
"\ndef ", source.find("def _smi_query") + 1
)
]
assert "IS_ROCM" in smi
assert "from . import amd" in smi
def test_hardware_branches_on_is_rocm_for_physical_count(self):
"""get_physical_gpu_count should try amd.py when IS_ROCM."""
@ -1247,7 +1343,7 @@ class TestHardwareAmdBranching:
func_start = source.find("def get_physical_gpu_count")
func_body = source[func_start : source.find("\ndef ", func_start + 1)]
assert "IS_ROCM" in func_body
assert "amd.get_physical_gpu_count" in func_body
assert "from . import amd" in func_body
# =============================================================================

View file

@ -0,0 +1,122 @@
# SPDX-License-Identifier: AGPL-3.0-only
import ast
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[2]
EXPORT = REPO_ROOT / "studio" / "backend" / "core" / "export" / "export.py"
EXPORT_FNS = (
"export_merged_model",
"export_base_model",
"export_gguf",
"export_lora_adapter",
)
def _find_method(tree, cls_name, method_name):
for cls in ast.walk(tree):
if isinstance(cls, ast.ClassDef) and cls.name == cls_name:
for item in cls.body:
if isinstance(item, ast.FunctionDef) and item.name == method_name:
return item
return None
def _return_tuple_arity(fn):
arities = []
for node in ast.walk(fn):
if isinstance(node, ast.Return) and isinstance(node.value, ast.Tuple):
arities.append(len(node.value.elts))
return arities
def test_export_methods_return_three_tuple_annotation():
tree = ast.parse(EXPORT.read_text())
for fn_name in EXPORT_FNS:
fn = _find_method(tree, "ExportBackend", fn_name)
assert fn is not None, f"missing ExportBackend.{fn_name}"
ret = fn.returns
assert isinstance(ret, ast.Subscript), f"{fn_name} return must be Tuple[...]"
slc = ret.slice
elts = slc.elts if isinstance(slc, ast.Tuple) else None
assert (
elts is not None and len(elts) == 3
), f"{fn_name} return annotation must be a 3-tuple, got {ast.dump(ret)}"
def test_export_methods_return_three_element_tuples():
tree = ast.parse(EXPORT.read_text())
for fn_name in EXPORT_FNS:
fn = _find_method(tree, "ExportBackend", fn_name)
assert fn is not None
arities = _return_tuple_arity(fn)
assert arities, f"{fn_name} has no tuple-return statements"
for arity in arities:
assert arity == 3, f"{fn_name} return tuple arity {arity}, expected 3"
def test_local_save_assigns_output_path():
tree = ast.parse(EXPORT.read_text())
for fn_name in EXPORT_FNS:
fn = _find_method(tree, "ExportBackend", fn_name)
assert fn is not None
assigns = []
for node in ast.walk(fn):
if isinstance(node, ast.Assign):
for tgt in node.targets:
if isinstance(tgt, ast.Name) and tgt.id == "output_path":
assigns.append(node)
non_none = [
a
for a in assigns
if not (isinstance(a.value, ast.Constant) and a.value.value is None)
]
assert non_none, f"{fn_name} never assigns a non-None output_path"
def test_gpu_save_method_bound_for_hub_only():
tree = ast.parse(EXPORT.read_text())
fn = _find_method(tree, "ExportBackend", "export_merged_model")
assert fn is not None
found_pre_save_method = False
for node in ast.walk(fn):
if isinstance(node, ast.Try):
for stmt in node.body:
if isinstance(stmt, ast.If):
test = stmt.test
if isinstance(test, ast.Name) and test.id == "_IS_MLX":
for sub in ast.walk(
ast.Module(body = stmt.orelse, type_ignores = [])
):
if isinstance(sub, ast.Assign) and any(
isinstance(t, ast.Name) and t.id == "save_method"
for t in sub.targets
):
found_pre_save_method = True
break
if found_pre_save_method:
break
if found_pre_save_method:
break
assert found_pre_save_method, (
"GPU save_method must be assigned at the top of the try block, "
"before the `if save_directory:` guard, so Hub-only export does not "
"raise UnboundLocalError."
)
def test_mlx_hub_only_uses_temp_directory():
src = EXPORT.read_text()
assert (
src.count("tempfile.TemporaryDirectory") >= 3
), "expected TemporaryDirectory in merged, base, and lora hub-push paths"
assert "import tempfile" in src.split("class ExportBackend")[0]
def test_is_mlx_imported_from_unsloth():
src = EXPORT.read_text()
assert "from unsloth import" in src
head = src.split("class ExportBackend")[0]
assert "_IS_MLX" in head
assert "_IS_MLX = platform.system()" not in src

View file

@ -0,0 +1,395 @@
# SPDX-License-Identifier: AGPL-3.0-only
"""
Comprehensive hardware dispatch matrix for Studio.
Drives every supported hardware profile from a single test host by
spoofing platform / torch.cuda / torch.xpu / sys.modules['mlx'] so we
can exercise the CUDA, ROCm, XPU, MLX, and CPU dispatch paths
deterministically without real hardware.
Profiles checked:
nvidia_cuda Linux x86_64 + torch.cuda.is_available()=True,
torch.version.hip=None
amd_rocm Linux x86_64 + torch.cuda.is_available()=True,
torch.version.hip="6.1" (PyTorch ROCm aliases
torch.cuda.* over HIP)
intel_xpu Linux x86_64 + torch.cuda off, torch.xpu.is_available()=True
apple_silicon_mlx Darwin arm64 + cuda off + xpu off + mlx importable
apple_silicon_no_mlx Darwin arm64 + everything off (no mlx pkg)
linux_arm64_with_mlx Linux arm64 + mlx importable -- gate must NOT activate
(canary against accidental Linux-arm64 hijack)
cpu_only Linux x86_64 + nothing -- pure CPU fallback
For each profile we assert three contracts:
1. ``unsloth._IS_MLX`` (re-evaluated under the spoof).
2. ``utils.hardware.detect_hardware()`` ``DeviceType`` and ``IS_ROCM``.
3. ``utils.hardware.is_apple_silicon()``.
Add a row to ``PROFILES`` to extend coverage; tests parametrize over it
automatically. No real hardware required.
"""
from __future__ import annotations
import importlib
import importlib.machinery
import importlib.util
import sys
import types
from dataclasses import dataclass, field
from pathlib import Path
from typing import Optional
import pytest
REPO_ROOT = Path(__file__).resolve().parents[2]
STUDIO_BACKEND = REPO_ROOT / "studio" / "backend"
# ---------------------------------------------------------------------------
# Profile definition
# ---------------------------------------------------------------------------
@dataclass
class HardwareProfile:
name: str
system: str # platform.system() value
machine: str # platform.machine() value
cuda_available: bool # torch.cuda.is_available() value
hip_version: Optional[
str
] # torch.version.hip; None for NVIDIA, "6.1" etc. for ROCm
xpu_available: bool # torch.xpu.is_available() value
has_mlx: bool # whether to inject a fake mlx into sys.modules
mps_available: bool # torch.backends.mps.is_available() value
expect_is_mlx: bool # unsloth._IS_MLX
expect_device_type: (
str # Studio DeviceType (uppercased name: "CUDA"/"XPU"/"MLX"/"CPU")
)
expect_is_rocm: bool # Studio IS_ROCM
expect_apple_silicon: bool # Studio is_apple_silicon()
extra_notes: str = ""
PROFILES = [
HardwareProfile(
name = "nvidia_cuda",
system = "Linux",
machine = "x86_64",
cuda_available = True,
hip_version = None,
xpu_available = False,
has_mlx = False,
mps_available = False,
expect_is_mlx = False,
expect_device_type = "CUDA",
expect_is_rocm = False,
expect_apple_silicon = False,
),
HardwareProfile(
name = "amd_rocm",
system = "Linux",
machine = "x86_64",
cuda_available = True,
hip_version = "6.1",
xpu_available = False,
has_mlx = False,
mps_available = False,
expect_is_mlx = False,
expect_device_type = "CUDA",
expect_is_rocm = True,
expect_apple_silicon = False,
extra_notes = "PyTorch ROCm reuses torch.cuda.* over HIP; "
"Studio still uses DeviceType.CUDA but flips IS_ROCM=True.",
),
HardwareProfile(
name = "intel_xpu",
system = "Linux",
machine = "x86_64",
cuda_available = False,
hip_version = None,
xpu_available = True,
has_mlx = False,
mps_available = False,
expect_is_mlx = False,
expect_device_type = "XPU",
expect_is_rocm = False,
expect_apple_silicon = False,
),
HardwareProfile(
name = "apple_silicon_mlx",
system = "Darwin",
machine = "arm64",
cuda_available = False,
hip_version = None,
xpu_available = False,
has_mlx = True,
mps_available = True,
expect_is_mlx = True,
expect_device_type = "MLX",
expect_is_rocm = False,
expect_apple_silicon = True,
),
HardwareProfile(
name = "apple_silicon_no_mlx",
system = "Darwin",
machine = "arm64",
cuda_available = False,
hip_version = None,
xpu_available = False,
has_mlx = False,
mps_available = True,
expect_is_mlx = False,
expect_device_type = "CPU",
expect_is_rocm = False,
expect_apple_silicon = True,
extra_notes = "Mac without mlx falls through to CPU (chat-only).",
),
HardwareProfile(
name = "linux_arm64_with_mlx",
system = "Linux",
machine = "arm64",
cuda_available = False,
hip_version = None,
xpu_available = False,
has_mlx = True,
mps_available = False,
expect_is_mlx = False,
expect_device_type = "CPU",
expect_is_rocm = False,
expect_apple_silicon = False,
extra_notes = "Canary: Linux ARM64 with mlx package installed must NOT "
"trigger MLX dispatch; the system check is what guards it.",
),
HardwareProfile(
name = "cpu_only",
system = "Linux",
machine = "x86_64",
cuda_available = False,
hip_version = None,
xpu_available = False,
has_mlx = False,
mps_available = False,
expect_is_mlx = False,
expect_device_type = "CPU",
expect_is_rocm = False,
expect_apple_silicon = False,
),
]
PROFILE_IDS = [p.name for p in PROFILES]
# ---------------------------------------------------------------------------
# Spoofing helpers
# ---------------------------------------------------------------------------
@pytest.fixture
def spoof_hardware(monkeypatch):
"""Return a function that applies a HardwareProfile to the live process.
Idempotent: each call re-applies the profile. Cleanup happens
automatically when the test exits via monkeypatch.
"""
def _apply(profile: HardwareProfile) -> None:
import platform
import torch
# platform spoof (used by both the unsloth gate and Studio's helpers)
monkeypatch.setattr(platform, "system", lambda: profile.system)
monkeypatch.setattr(platform, "machine", lambda: profile.machine)
# torch.cuda.is_available
monkeypatch.setattr(torch.cuda, "is_available", lambda: profile.cuda_available)
# detect_hardware reads torch.cuda.get_device_properties(0).name when
# cuda_available is True. On a CPU CI runner that triggers _cuda_init
# and crashes with "No CUDA GPUs are available". Stub it so the
# dispatch path under test runs end-to-end.
if profile.cuda_available:
stub_props = types.SimpleNamespace(
name = "Stub GPU" if not profile.hip_version else "Stub AMD GPU",
)
monkeypatch.setattr(
torch.cuda,
"get_device_properties",
lambda i = 0: stub_props,
raising = False,
)
# torch.version.hip — None on NVIDIA, "6.1" etc. on ROCm
torch_version = torch.version
monkeypatch.setattr(torch_version, "hip", profile.hip_version, raising = False)
# torch.xpu.is_available + get_device_name -- detect_hardware reads both.
# Real torch.xpu.get_device_name requires the XPU-compiled torch build,
# so always stub it under the spoof to keep tests hardware-agnostic.
if hasattr(torch, "xpu"):
monkeypatch.setattr(
torch.xpu, "is_available", lambda: profile.xpu_available
)
monkeypatch.setattr(
torch.xpu,
"get_device_name",
lambda i = 0: "Intel XPU (stub)",
raising = False,
)
elif profile.xpu_available:
xpu_stub = types.SimpleNamespace(
is_available = lambda: True,
get_device_name = lambda i = 0: "Intel XPU (stub)",
)
monkeypatch.setattr(torch, "xpu", xpu_stub, raising = False)
# torch.backends.mps.is_available
if hasattr(torch.backends, "mps"):
monkeypatch.setattr(
torch.backends.mps, "is_available", lambda: profile.mps_available
)
# mlx + mlx.core in sys.modules
if profile.has_mlx:
fake_mlx = types.ModuleType("mlx")
fake_mlx.__spec__ = importlib.machinery.ModuleSpec("mlx", loader = None)
fake_mlx.__path__ = []
fake_mlx_core = types.ModuleType("mlx.core")
fake_mlx.core = fake_mlx_core
monkeypatch.setitem(sys.modules, "mlx", fake_mlx)
monkeypatch.setitem(sys.modules, "mlx.core", fake_mlx_core)
else:
monkeypatch.delitem(sys.modules, "mlx", raising = False)
monkeypatch.delitem(sys.modules, "mlx.core", raising = False)
real_find_spec = importlib.util.find_spec
def _no_mlx(name, *args, **kwargs):
if name == "mlx":
return None
return real_find_spec(name, *args, **kwargs)
monkeypatch.setattr(importlib.util, "find_spec", _no_mlx)
return _apply
def _evaluate_unsloth_is_mlx_gate() -> bool:
"""Re-evaluate the exact expression from unsloth/__init__.py:20-24."""
import importlib.util
import platform
return (
platform.system() == "Darwin"
and platform.machine() == "arm64"
and importlib.util.find_spec("mlx") is not None
)
def _import_studio_hardware_module():
"""Lazy-load Studio's hardware module under the bare-imports layout."""
if str(STUDIO_BACKEND) not in sys.path:
sys.path.insert(0, str(STUDIO_BACKEND))
# Force a fresh import so detect_hardware re-runs under the current spoofs.
sys.modules.pop("utils.hardware.hardware", None)
sys.modules.pop("utils.hardware", None)
from utils.hardware import hardware as hw # type: ignore
return hw
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("profile", PROFILES, ids = PROFILE_IDS)
def test_unsloth_is_mlx_gate_matches_profile(profile, spoof_hardware):
"""The _IS_MLX expression in unsloth/__init__.py flips correctly per profile."""
spoof_hardware(profile)
actual = _evaluate_unsloth_is_mlx_gate()
assert actual is profile.expect_is_mlx, (
f"profile {profile.name}: expected _IS_MLX={profile.expect_is_mlx}, "
f"got {actual}. {profile.extra_notes}"
)
@pytest.mark.parametrize("profile", PROFILES, ids = PROFILE_IDS)
def test_studio_detect_hardware_matches_profile(profile, spoof_hardware):
"""Studio's detect_hardware() routes to the right DeviceType per profile."""
spoof_hardware(profile)
hw = _import_studio_hardware_module()
detected = hw.detect_hardware()
expected = getattr(hw.DeviceType, profile.expect_device_type)
assert detected == expected, (
f"profile {profile.name}: expected {profile.expect_device_type}, "
f"got {detected!r}. {profile.extra_notes}"
)
assert hw.IS_ROCM is profile.expect_is_rocm, (
f"profile {profile.name}: expected IS_ROCM={profile.expect_is_rocm}, "
f"got {hw.IS_ROCM}"
)
@pytest.mark.parametrize("profile", PROFILES, ids = PROFILE_IDS)
def test_studio_is_apple_silicon_matches_profile(profile, spoof_hardware):
"""Studio's is_apple_silicon() helper agrees with platform spoof."""
spoof_hardware(profile)
hw = _import_studio_hardware_module()
assert hw.is_apple_silicon() is profile.expect_apple_silicon, (
f"profile {profile.name}: expected is_apple_silicon={profile.expect_apple_silicon}, "
f"got {hw.is_apple_silicon()}"
)
# ---------------------------------------------------------------------------
# Negative-space tests: catch regressions where the dispatch order changes.
# ---------------------------------------------------------------------------
def test_cuda_takes_priority_over_mlx_when_both_available(spoof_hardware):
"""If both CUDA and MLX are available, Studio MUST pick CUDA. This is the
canary that protects every existing GPU user from being silently routed
to MLX after future refactors.
"""
profile = HardwareProfile(
name = "cuda_plus_mlx",
system = "Darwin",
machine = "arm64",
cuda_available = True,
hip_version = None,
xpu_available = False,
has_mlx = True,
mps_available = True,
expect_is_mlx = True,
expect_device_type = "CUDA",
expect_is_rocm = False,
expect_apple_silicon = True,
)
spoof_hardware(profile)
hw = _import_studio_hardware_module()
assert hw.detect_hardware() == hw.DeviceType.CUDA
def test_xpu_takes_priority_over_mlx_when_both_available(spoof_hardware):
"""XPU is selected over MLX in the dispatch order."""
profile = HardwareProfile(
name = "xpu_plus_mlx",
system = "Darwin",
machine = "arm64",
cuda_available = False,
hip_version = None,
xpu_available = True,
has_mlx = True,
mps_available = True,
expect_is_mlx = True,
expect_device_type = "XPU",
expect_is_rocm = False,
expect_apple_silicon = True,
)
spoof_hardware(profile)
hw = _import_studio_hardware_module()
assert hw.detect_hardware() == hw.DeviceType.XPU

View file

@ -0,0 +1,213 @@
# SPDX-License-Identifier: AGPL-3.0-only
"""
Regression tests for the CUDA-vs-MLX dispatch gates Studio relies on.
Two gates drive every dispatch decision in Studio's MLX path:
1. ``unsloth._IS_MLX`` at the top of ``unsloth/__init__.py`` -- evaluated
once at import time and read by Studio worker code to choose between
the GPU and MLX trainer / inference / export paths. Defined as
``Darwin AND arm64 AND find_spec("mlx") is not None``.
2. ``utils.hardware.detect_hardware()`` -- runtime probe in the Studio
backend. Priority order: CUDA -> XPU -> MLX -> CPU. The MLX branch is
reached only when both CUDA and XPU are unavailable AND the host is
Apple Silicon AND mlx is importable.
These gates are the canaries for "MLX support accidentally hijacks
CUDA/AMD/Intel users". The tests here:
* verify the source-level structure of the ``_IS_MLX`` expression so an
accidental rewrite (e.g. dropping the ``arm64`` check) is caught,
* exercise the runtime gate logic under a spoofed Darwin+arm64 platform
with a fake ``mlx`` module in ``sys.modules`` to confirm both gates
flip True together,
* confirm that on the actual Linux+CUDA test host both gates remain in
their CUDA-side state.
No real MLX install is required; uses the same ``monkeypatch.setitem``
fake-mlx pattern as ``test_mlx_inference_backend.py``.
"""
import ast
import importlib
import sys
import types
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[2]
UNSLOTH_INIT = REPO_ROOT / "unsloth" / "__init__.py"
# ---------------------------------------------------------------------------
# 1. Source-level structure check on _IS_MLX (no platform dependencies).
# ---------------------------------------------------------------------------
def test_is_mlx_gate_uses_three_required_predicates():
"""The _IS_MLX assignment must AND together exactly the three checks
that Studio depends on: Darwin OS, arm64 machine, and an importable
mlx package. Dropping any one of them silently breaks dispatch.
"""
tree = ast.parse(UNSLOTH_INIT.read_text())
target = None
for node in ast.walk(tree):
if (
isinstance(node, ast.Assign)
and len(node.targets) == 1
and isinstance(node.targets[0], ast.Name)
and node.targets[0].id == "_IS_MLX"
):
target = node.value
break
assert target is not None, "_IS_MLX assignment not found in unsloth/__init__.py"
assert isinstance(target, ast.BoolOp) and isinstance(
target.op, ast.And
), "_IS_MLX must be a BoolOp(And) of platform + mlx checks"
expr_src = ast.unparse(target)
assert (
"platform.system()" in expr_src and "Darwin" in expr_src
), "_IS_MLX must check platform.system() == 'Darwin'"
assert (
"platform.machine()" in expr_src and "arm64" in expr_src
), "_IS_MLX must check platform.machine() == 'arm64'"
assert (
"find_spec" in expr_src and "'mlx'" in expr_src
), "_IS_MLX must check importlib.util.find_spec('mlx')"
# ---------------------------------------------------------------------------
# 2. Runtime gate behavior with the platform spoofed to Apple Silicon and a
# fake mlx module in sys.modules. Re-evaluates the same expression
# rather than reloading unsloth (which would cascade-reload torch).
# ---------------------------------------------------------------------------
def _evaluate_is_mlx_gate(platform_module, importlib_util):
"""Re-evaluate the _IS_MLX expression using injected dependencies.
Mirrors the assignment in unsloth/__init__.py exactly.
"""
return (
platform_module.system() == "Darwin"
and platform_module.machine() == "arm64"
and importlib_util.find_spec("mlx") is not None
)
def test_is_mlx_gate_true_on_apple_silicon_with_mlx_present(monkeypatch):
import platform
import importlib.util
# Inject a fake mlx package so find_spec returns a non-None ModuleSpec.
fake_mlx = types.ModuleType("mlx")
fake_mlx.__spec__ = importlib.machinery.ModuleSpec("mlx", loader = None)
fake_mlx.__path__ = []
monkeypatch.setitem(sys.modules, "mlx", fake_mlx)
monkeypatch.setattr(platform, "system", lambda: "Darwin")
monkeypatch.setattr(platform, "machine", lambda: "arm64")
assert _evaluate_is_mlx_gate(platform, importlib.util) is True
def test_is_mlx_gate_false_when_mlx_missing(monkeypatch):
import platform
import importlib.util
# Apple Silicon platform but no mlx package -> gate must be False.
monkeypatch.delitem(sys.modules, "mlx", raising = False)
monkeypatch.setattr(platform, "system", lambda: "Darwin")
monkeypatch.setattr(platform, "machine", lambda: "arm64")
real_find_spec = importlib.util.find_spec
def _no_mlx(name, *args, **kwargs):
if name == "mlx":
return None
return real_find_spec(name, *args, **kwargs)
monkeypatch.setattr(importlib.util, "find_spec", _no_mlx)
assert _evaluate_is_mlx_gate(platform, importlib.util) is False
def test_is_mlx_gate_false_on_non_apple_silicon():
"""On the real Linux+CUDA / AMD / Intel test host, the gate stays False."""
import platform
import importlib.util
if platform.system() == "Darwin" and platform.machine() == "arm64":
# On a Mac CI runner this assertion would not apply; skip there.
import pytest
pytest.skip("Test host is Apple Silicon; CUDA-side canary doesn't apply.")
assert _evaluate_is_mlx_gate(platform, importlib.util) is False
# ---------------------------------------------------------------------------
# 3. Studio's runtime detect_hardware() picks MLX only when CUDA + XPU are
# both unavailable AND the host is Apple Silicon AND mlx is importable.
# ---------------------------------------------------------------------------
def _import_studio_hardware():
"""Lazy import for the Studio hardware module, with the bare-imports
convention that Studio uses (studio/backend on sys.path).
"""
studio_backend = REPO_ROOT / "studio" / "backend"
if str(studio_backend) not in sys.path:
sys.path.insert(0, str(studio_backend))
from utils.hardware import hardware as hw # type: ignore
return hw
def test_detect_hardware_picks_mlx_when_only_apple_silicon_available(monkeypatch):
hw = _import_studio_hardware()
# Force CUDA + XPU paths off so detect_hardware falls through to MLX.
import torch
monkeypatch.setattr(torch.cuda, "is_available", lambda: False)
if hasattr(torch, "xpu"):
monkeypatch.setattr(torch.xpu, "is_available", lambda: False)
# Spoof Apple Silicon and provide an importable mlx.core for _has_mlx().
import platform
monkeypatch.setattr(platform, "system", lambda: "Darwin")
monkeypatch.setattr(platform, "machine", lambda: "arm64")
fake_mlx = types.ModuleType("mlx")
fake_mlx_core = types.ModuleType("mlx.core")
fake_mlx.core = fake_mlx_core
monkeypatch.setitem(sys.modules, "mlx", fake_mlx)
monkeypatch.setitem(sys.modules, "mlx.core", fake_mlx_core)
detected = hw.detect_hardware()
assert detected == hw.DeviceType.MLX, f"expected MLX, got {detected!r}"
def test_detect_hardware_picks_cuda_on_real_host():
"""Canary: on a real CUDA host the MLX branch must NOT be taken even
if mlx happens to be importable. Protects CUDA/AMD/Intel users from
accidental MLX dispatch when MLX support is added.
"""
import torch
if not torch.cuda.is_available():
import pytest
pytest.skip("No CUDA available on this host; canary not applicable.")
hw = _import_studio_hardware()
detected = hw.detect_hardware()
assert (
detected == hw.DeviceType.CUDA
), f"CUDA host must dispatch to CUDA, got {detected!r}"

View file

@ -0,0 +1,90 @@
# SPDX-License-Identifier: AGPL-3.0-only
import ast
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[2]
WORKER = REPO_ROOT / "studio" / "backend" / "core" / "training" / "worker.py"
def _find_func(tree, name):
for node in ast.walk(tree):
if isinstance(node, ast.FunctionDef) and node.name == name:
return node
return None
def test_run_mlx_training_passes_token_to_from_pretrained():
tree = ast.parse(WORKER.read_text())
fn = _find_func(tree, "_run_mlx_training")
assert fn is not None
found = False
for node in ast.walk(fn):
if (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute)
and node.func.attr == "from_pretrained"
and isinstance(node.func.value, ast.Name)
and node.func.value.id == "FastMLXModel"
):
kwarg_names = {kw.arg for kw in node.keywords if kw.arg}
assert (
"token" in kwarg_names
), f"FastMLXModel.from_pretrained must forward token=hf_token; got {kwarg_names!r}"
found = True
assert found, "FastMLXModel.from_pretrained call not found in _run_mlx_training"
def test_wandb_init_strips_secret_keys():
src = WORKER.read_text()
assert "_wandb_sensitive" in src, "expected a sensitive-key set near wandb.init"
assert '"hf_token"' in src and '"wandb_token"' in src
assert (
"config = dict(config)" not in src
), "wandb.init received raw config dict; secrets would leak"
def test_local_dataset_loader_uses_load_dataset_path():
src = WORKER.read_text()
assert "_resolve_local_files" in src
assert "_loader_for_files" in src
assert "data_files = all_files" in src or "data_files=all_files" in src
def test_send_aliases_status_message_to_message():
src = WORKER.read_text()
assert 'kwargs["message"] = sm' in src or 'kwargs["message"]=sm' in src
def test_slice_uses_inclusive_end_and_handles_zero():
src = WORKER.read_text()
assert "min(end + 1, len(ds))" in src or "min(end+1, len(ds))" in src
assert "slice_start if slice_start is not None else 0" in src
assert "slice_end if slice_end is not None else len(ds) - 1" in src
def test_poll_stop_returns_on_broken_pipe():
src = WORKER.read_text()
assert "except (EOFError, OSError)" in src
lines = src.splitlines()
for i, line in enumerate(lines):
if "except (EOFError, OSError)" in line:
for j in range(i + 1, min(i + 6, len(lines))):
stripped = lines[j].strip()
if not stripped or stripped.startswith("#"):
continue
assert stripped.startswith(
"return"
), f"expected return after EOFError/OSError, got {stripped!r}"
break
break
else:
raise AssertionError("EOFError/OSError handler not found in worker.py")
def test_unsloth_zoo_mlx_imports_have_friendly_error():
src = WORKER.read_text()
assert "from unsloth_zoo.mlx_loader import FastMLXModel" in src
assert "from unsloth_zoo.mlx_trainer import" in src
assert "raise ImportError" in src
assert "install.sh" in src

View file

@ -35,8 +35,11 @@ class MockDataset:
return cls(data_dict)
# Mock datasets module
# Mock datasets module. __spec__ must be set so importlib.util.find_spec
# does not raise ValueError when transformers' import_utils probes for
# the real `datasets` package later in the test session.
datasets_mock = type(sys)("datasets")
datasets_mock.__spec__ = importlib.util.spec_from_loader("datasets", loader = None)
datasets_mock.Dataset = MockDataset
sys.modules["datasets"] = datasets_mock

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,154 @@
"""Resilience checks for Studio install-root inference under hostile
filesystem conditions:
- _infer_studio_home_from_venv must NOT propagate PermissionError /
OSError out through studio_root() (it would crash module import in
run.py / main.py / transformers_version.py / model_config.py).
- _kill_orphaned_servers must catch (ImportError, OSError, ValueError)
on the studio_root() probe so a transient resolve / sentinel failure
cannot crash server startup.
- _find_llama_server_binary must keep the custom-root in search_roots
when the inner resolve() comparison itself fails."""
from __future__ import annotations
import importlib.util
import re
import sys
import textwrap
from pathlib import Path
from unittest import mock
import pytest
REPO_ROOT = Path(__file__).resolve().parents[1]
STORAGE_ROOTS = (
REPO_ROOT / "studio" / "backend" / "utils" / "paths" / "storage_roots.py"
)
LLAMA_CPP = REPO_ROOT / "studio" / "backend" / "core" / "inference" / "llama_cpp.py"
def _load(name: str, path: Path):
spec = importlib.util.spec_from_file_location(name, path)
assert spec is not None and spec.loader is not None
mod = importlib.util.module_from_spec(spec)
sys.modules[name] = mod
spec.loader.exec_module(mod)
return mod
def test_infer_studio_home_swallows_permission_error(tmp_path, monkeypatch):
candidate = tmp_path / "fake_root"
venv = candidate / "unsloth_studio"
venv.mkdir(parents = True)
monkeypatch.setattr(sys, "prefix", str(venv))
sys.modules.pop("sr_perm", None)
mod = _load("sr_perm", STORAGE_ROOTS)
with mock.patch.object(Path, "is_file", side_effect = PermissionError("denied")):
# Must NOT raise.
assert mod._infer_studio_home_from_venv() is None
def test_studio_root_does_not_crash_on_permission_error(tmp_path, monkeypatch):
"""studio_root() must remain callable even when the venv inference
encounters a restricted filesystem; it should fall through to the
legacy default."""
candidate = tmp_path / "fake_root"
venv = candidate / "unsloth_studio"
venv.mkdir(parents = True)
monkeypatch.setattr(sys, "prefix", str(venv))
monkeypatch.delenv("UNSLOTH_STUDIO_HOME", raising = False)
monkeypatch.delenv("STUDIO_HOME", raising = False)
sys.modules.pop("sr_studio_perm", None)
mod = _load("sr_studio_perm", STORAGE_ROOTS)
with mock.patch.object(Path, "is_file", side_effect = OSError("ebusy")):
result = mod.studio_root()
assert result == Path.home() / ".unsloth" / "studio"
def test_kill_orphan_catches_oserror_from_studio_root():
"""_kill_orphaned_servers must catch (ImportError, OSError, ValueError)
on the studio_root() probe specifically; the sister function
_find_llama_server_binary uses the same broader catch on its own probe."""
src = LLAMA_CPP.read_text()
fn_start = src.index("def _kill_orphaned_servers")
fn_body = src[fn_start : fn_start + 4000]
# The studio_root() probe in this fn is the one that imports as `_sr`
# and assigns `_resolved_sr = _sr()`. Find the except that closes it.
probe_idx = fn_body.index("storage_roots import studio_root as _sr")
# The matching except is the next `except ...:` after the inner
# OSError/ValueError block that wraps resolve().
after = fn_body[probe_idx:]
# Skip over the inner `except (OSError, ValueError):` that wraps resolve().
inner_idx = after.index("except (OSError, ValueError):")
after_inner = after[inner_idx + len("except (OSError, ValueError):") :]
outer_match = re.search(r"except\s*\(?[^)]*?\)?:", after_inner)
assert outer_match, "outer except for studio_root probe missing"
clause = outer_match.group(0)
assert (
"OSError" in clause and "ValueError" in clause
), f"_kill_orphaned_servers studio_root probe catch too narrow: {clause!r}"
def _exec_search_roots_block(
home: Path, studio_root_value: Path, resolve_raises: bool
) -> list[Path]:
"""Extract _find_llama_server_binary's env-mode search_roots block
and execute it with controlled inputs."""
src = LLAMA_CPP.read_text()
block_start = src.index('legacy_llama = Path.home() / ".unsloth" / "llama.cpp"')
block_end = src.index("_seen_roots: set[str]", block_start)
raw = src[block_start:block_end]
indent = " " * 8
block = textwrap.dedent(indent + raw)
fake_module = type(sys)("fake_storage_roots")
fake_module.studio_root = lambda: studio_root_value
sys.modules["utils.paths.storage_roots"] = fake_module
try:
original_resolve = Path.resolve
def _resolve(self, *a, **k):
if resolve_raises:
raise OSError("ebusy")
return original_resolve(self, *a, **k)
with (
mock.patch.object(Path, "home", classmethod(lambda cls: home)),
mock.patch.object(Path, "resolve", _resolve),
):
ns: dict = {"Path": Path}
exec(block, ns) # noqa: S102
return ns["search_roots"]
finally:
sys.modules.pop("utils.paths.storage_roots", None)
def test_search_roots_keeps_custom_when_resolve_fails(tmp_path):
home = tmp_path / "home"
home.mkdir()
custom = tmp_path / "custom_studio"
custom.mkdir()
roots = _exec_search_roots_block(
home = home, studio_root_value = custom, resolve_raises = True
)
# On resolve() failure, the inner except falls back to direct equality;
# custom != legacy_studio so the custom root must remain in search_roots.
assert (
custom / "llama.cpp" in roots
), f"custom root dropped on resolve() failure: {roots}"
# custom-mode discovery excludes the legacy tree to match _kill_orphaned_servers.
assert (
(home / ".unsloth" / "llama.cpp") not in roots
), f"legacy llama path must not appear in custom-mode search_roots: {roots}"
def test_search_roots_default_mode_uses_legacy_only(tmp_path):
home = tmp_path / "home"
home.mkdir()
legacy = home / ".unsloth" / "studio"
legacy.mkdir(parents = True)
roots = _exec_search_roots_block(
home = home, studio_root_value = legacy, resolve_raises = False
)
# Default mode: only legacy_llama.
assert roots == [home / ".unsloth" / "llama.cpp"]

View file

@ -12,348 +12,117 @@
# See the License for the specific language governing permissions and
# limitations under the License.
import warnings, importlib, sys
from packaging.version import Version
import os, re, subprocess, inspect, functools
import numpy as np
import os, platform, importlib.util
# Log Unsloth is being used
os.environ["UNSLOTH_IS_PRESENT"] = "1"
# Check if modules that need patching are already imported
critical_modules = ["trl", "transformers", "peft"]
already_imported = [mod for mod in critical_modules if mod in sys.modules]
# Fix some issues before importing other packages
from .import_fixes import (
fix_message_factory_issue,
check_fbgemm_gpu_version,
disable_broken_causal_conv1d,
disable_broken_vllm,
configure_amdgpu_asic_id_table_path,
torchvision_compatibility_check,
fix_diffusers_warnings,
fix_huggingface_hub,
# Detect Apple Silicon + MLX before any torch/numpy imports
_IS_MLX = (
platform.system() == "Darwin"
and platform.machine() == "arm64"
and importlib.util.find_spec("mlx") is not None
)
# Configure libdrm ids table path early so ROCm can resolve AMD GPU names.
configure_amdgpu_asic_id_table_path()
disable_broken_causal_conv1d()
disable_broken_vllm()
fix_message_factory_issue()
check_fbgemm_gpu_version()
torchvision_compatibility_check()
fix_diffusers_warnings()
fix_huggingface_hub()
del configure_amdgpu_asic_id_table_path
del disable_broken_causal_conv1d
del disable_broken_vllm
del fix_message_factory_issue
del check_fbgemm_gpu_version
del torchvision_compatibility_check
del fix_diffusers_warnings
del fix_huggingface_hub
# This check is critical because Unsloth optimizes these libraries by modifying
# their code at import time. If they're imported first, the original (slower,
# more memory-intensive) implementations will be used instead of Unsloth's
# optimized versions, potentially causing OOM errors or slower training.
if already_imported:
# stacklevel=2 makes warning point to user's import line rather than this library code,
# showing them exactly where to fix the import order in their script
warnings.warn(
f"WARNING: Unsloth should be imported before [{', '.join(already_imported)}] "
f"to ensure all optimizations are applied. Your code may run slower or encounter "
f"memory issues without these optimizations.\n\n"
f"Please restructure your imports with 'import unsloth' at the top of your file.",
stacklevel = 2,
)
del already_imported, critical_modules
# Unsloth currently does not work on multi GPU setups - sadly we are a 2 brother team so
# enabling it will require much more work, so we have to prioritize. Please understand!
# We do have a beta version, which you can contact us about!
# Thank you for your understanding and we appreciate it immensely!
# Fixes https://github.com/unslothai/unsloth/issues/1266
os.environ["PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION"] = "python"
# [TODO] Check why some GPUs don't work
# "pinned_use_cuda_host_register:True,"\
# "pinned_num_register_threads:8"
from importlib.metadata import version as importlib_version
from importlib.metadata import PackageNotFoundError
# Check for unsloth_zoo
try:
unsloth_zoo_version = importlib_version("unsloth_zoo")
if Version(unsloth_zoo_version) < Version("2026.3.4"):
print(
"Unsloth: Please update Unsloth and Unsloth-Zoo to the latest version!\n"
"Do this via `pip install --upgrade --force-reinstall --no-cache-dir --no-deps unsloth unsloth_zoo`"
)
# if os.environ.get("UNSLOTH_DISABLE_AUTO_UPDATES", "0") == "0":
# try:
# os.system("pip install --upgrade --no-cache-dir --no-deps unsloth_zoo")
# except:
# try:
# os.system("pip install --upgrade --no-cache-dir --no-deps --user unsloth_zoo")
# except:
# raise ImportError("Unsloth: Please update unsloth_zoo via `pip install --upgrade --no-cache-dir --no-deps unsloth_zoo`")
import unsloth_zoo
except PackageNotFoundError:
raise ImportError(
f"Unsloth: Please install unsloth_zoo via `pip install unsloth_zoo` then retry!"
)
except:
raise
del PackageNotFoundError, importlib_version
# Try importing PyTorch and check version
try:
import torch
except ModuleNotFoundError:
raise ImportError(
"Unsloth: Pytorch is not installed. Go to https://pytorch.org/.\n"
"We have some installation instructions on our Github page."
)
except:
raise
from unsloth_zoo.device_type import (
is_hip,
get_device_type,
DEVICE_TYPE,
DEVICE_TYPE_TORCH,
DEVICE_COUNT,
ALLOW_PREQUANTIZED_MODELS,
)
# Fix other issues
from .import_fixes import (
fix_xformers_performance_issue,
fix_vllm_aimv2_issue,
check_vllm_torch_sm100_compatibility,
fix_vllm_guided_decoding_params,
fix_trl_vllm_ascend,
fix_vllm_pdl_blackwell,
fix_triton_compiled_kernel_missing_attrs,
patch_trunc_normal_precision_issue,
ignore_logger_messages,
patch_ipykernel_hf_xet,
patch_trackio,
patch_datasets,
patch_enable_input_require_grads,
fix_openenv_no_vllm,
patch_openspiel_env_async,
fix_executorch,
patch_vllm_for_notebooks,
patch_torchcodec_audio_decoder,
disable_torchcodec_if_broken,
disable_broken_wandb,
patch_peft_weight_converter_compatibility,
)
fix_xformers_performance_issue()
fix_vllm_aimv2_issue()
# Check vLLM + torch < 2.9.0 + SM100 compatibility BEFORE importing vLLM
check_vllm_torch_sm100_compatibility()
fix_vllm_guided_decoding_params()
fix_trl_vllm_ascend()
fix_vllm_pdl_blackwell()
fix_triton_compiled_kernel_missing_attrs()
patch_trunc_normal_precision_issue()
ignore_logger_messages()
patch_ipykernel_hf_xet()
patch_trackio()
patch_datasets()
patch_enable_input_require_grads()
fix_openenv_no_vllm()
patch_openspiel_env_async()
fix_executorch()
patch_vllm_for_notebooks()
patch_torchcodec_audio_decoder()
disable_torchcodec_if_broken()
disable_broken_wandb()
patch_peft_weight_converter_compatibility()
del fix_xformers_performance_issue
del fix_vllm_aimv2_issue
del check_vllm_torch_sm100_compatibility
del fix_vllm_guided_decoding_params
del fix_trl_vllm_ascend
del fix_vllm_pdl_blackwell
del fix_triton_compiled_kernel_missing_attrs
del patch_trunc_normal_precision_issue
del ignore_logger_messages
del patch_ipykernel_hf_xet
del patch_trackio
del patch_datasets
del patch_enable_input_require_grads
del fix_openenv_no_vllm
del patch_openspiel_env_async
del fix_executorch
del patch_vllm_for_notebooks
del patch_torchcodec_audio_decoder
del disable_torchcodec_if_broken
del disable_broken_wandb
del patch_peft_weight_converter_compatibility
# Torch 2.4 has including_emulation
if DEVICE_TYPE == "cuda":
major_version, minor_version = torch.cuda.get_device_capability()
SUPPORTS_BFLOAT16 = major_version >= 8
old_is_bf16_supported = torch.cuda.is_bf16_supported
if "including_emulation" in str(inspect.signature(old_is_bf16_supported)):
def is_bf16_supported(including_emulation = False):
return old_is_bf16_supported(including_emulation)
torch.cuda.is_bf16_supported = is_bf16_supported
else:
def is_bf16_supported():
return SUPPORTS_BFLOAT16
torch.cuda.is_bf16_supported = is_bf16_supported
del major_version, minor_version
elif DEVICE_TYPE == "hip":
SUPPORTS_BFLOAT16 = torch.cuda.is_bf16_supported()
elif DEVICE_TYPE == "xpu":
# torch.xpu.is_bf16_supported() does not have including_emulation
# set SUPPORTS_BFLOAT16 as torch.xpu.is_bf16_supported()
SUPPORTS_BFLOAT16 = torch.xpu.is_bf16_supported()
# For Gradio HF Spaces?
# if "SPACE_AUTHOR_NAME" not in os.environ and "SPACE_REPO_NAME" not in os.environ:
import triton
if DEVICE_TYPE == "cuda":
libcuda_dirs = lambda: None
if Version(triton.__version__) >= Version("3.0.0"):
try:
from triton.backends.nvidia.driver import libcuda_dirs
except:
pass
else:
from triton.common.build import libcuda_dirs
# Try loading bitsandbytes and triton
if _IS_MLX:
try:
import bitsandbytes as bnb
except:
print(
"Unsloth: `bitsandbytes` is not installed - 4bit QLoRA unallowed, but 16bit and full finetuning works!"
)
bnb = None
import unsloth_zoo
except ImportError as _e:
raise ImportError(
"Unsloth: MLX support requires `unsloth-zoo` with MLX modules. "
"Reinstall with `pip install unsloth-zoo` or rerun install.sh."
) from _e
# The mlx_trainer / mlx_loader submodules ship with unsloth-zoo's MLX
# support. An older installed unsloth-zoo (e.g. from PyPI before the
# MLX release lands) will satisfy `import unsloth_zoo` but be missing
# these submodules. Surface the same friendly install hint instead of
# a raw ImportError on the submodule path.
try:
cdequantize_blockwise_fp32 = bnb.functional.lib.cdequantize_blockwise_fp32
libcuda_dirs()
except:
# Only run the ldconfig recovery when we can actually run
# ldconfig (root). On non-root environments (shared HPC,
# locked-down containers, CI runners, etc.) the recovery would
# shell out to `ldconfig` and fail with "Permission denied",
# which is especially noisy for users who don't even have
# bitsandbytes installed and are just doing 16bit/full
# finetuning. libcuda_dirs() is used by both triton and bnb,
# so we still run the recovery whenever we're root, regardless
# of whether bnb is installed.
if hasattr(os, "geteuid") and os.geteuid() == 0:
warnings.warn("Unsloth: Running `ldconfig /usr/lib64-nvidia` to link CUDA.")
from unsloth_zoo.mlx_trainer import MLXTrainer, MLXTrainingConfig
from unsloth_zoo.mlx_loader import FastMLXModel
except ImportError as _e:
raise ImportError(
"Unsloth: MLX support requires an unsloth-zoo build that includes "
"`unsloth_zoo.mlx_trainer` and `unsloth_zoo.mlx_loader`. Upgrade with "
"`pip install -U unsloth-zoo` or rerun install.sh."
) from _e
if os.path.exists("/usr/lib64-nvidia"):
os.system("ldconfig /usr/lib64-nvidia")
elif os.path.exists("/usr/local"):
# Sometimes bitsandbytes cannot be linked properly in Runpod for example
possible_cudas = (
subprocess.check_output(["ls", "-al", "/usr/local"])
.decode("utf-8")
.split("\n")
)
find_cuda = re.compile(r"[\s](cuda\-[\d\.]{2,})$")
possible_cudas = [find_cuda.search(x) for x in possible_cudas]
possible_cudas = [x.group(1) for x in possible_cudas if x is not None]
# Load raw_text helpers without executing dataprep/__init__.py, which
# imports synthetic.py -> torch and would defeat the torch-free MLX path.
from pathlib import Path as _Path
# Try linking cuda folder, or everything in local
if len(possible_cudas) == 0:
os.system("ldconfig /usr/local/")
else:
find_number = re.compile(r"([\d\.]{2,})")
latest_cuda = np.argsort(
[float(find_number.search(x).group(1)) for x in possible_cudas]
)[::-1][0]
latest_cuda = possible_cudas[latest_cuda]
os.system(f"ldconfig /usr/local/{latest_cuda}")
del find_number, latest_cuda
del possible_cudas, find_cuda
_raw_text_path = _Path(__file__).resolve().parent / "dataprep" / "raw_text.py"
_raw_text_spec = importlib.util.spec_from_file_location(
"unsloth._mlx_raw_text", _raw_text_path
)
if _raw_text_spec is None or _raw_text_spec.loader is None:
raise ImportError("Unsloth: could not load MLX raw_text dataprep helpers.")
_raw_text = importlib.util.module_from_spec(_raw_text_spec)
_raw_text_spec.loader.exec_module(_raw_text)
RawTextDataLoader = _raw_text.RawTextDataLoader
TextPreprocessor = _raw_text.TextPreprocessor
del _raw_text, _raw_text_spec, _raw_text_path, _Path
if bnb is not None:
importlib.reload(bnb)
importlib.reload(triton)
try:
libcuda_dirs = lambda: None
if Version(triton.__version__) >= Version("3.0.0"):
try:
from triton.backends.nvidia.driver import libcuda_dirs
except:
pass
else:
from triton.common.build import libcuda_dirs
cdequantize_blockwise_fp32 = (
bnb.functional.lib.cdequantize_blockwise_fp32
)
libcuda_dirs()
except:
warnings.warn(
"Unsloth: CUDA is not linked properly.\n"
"Try running `python -m bitsandbytes` then `python -m xformers.info`\n"
"We tried running `ldconfig /usr/lib64-nvidia` ourselves, but it didn't work.\n"
"You need to run in your terminal `sudo ldconfig /usr/lib64-nvidia` yourself, then import Unsloth.\n"
"Also try `sudo ldconfig /usr/local/cuda-xx.x` - find the latest cuda version.\n"
"Unsloth will still run for now, but maybe it might crash - let's hope it works!"
)
elif bnb is not None:
# Non-root + bnb installed: we can't run ldconfig ourselves,
# but bnb is going to crash later when the user actually uses
# 4bit quantization - tell them how to fix it manually so
# they're not surprised by an opaque error down the road.
warnings.warn(
"Unsloth: CUDA is not linked properly.\n"
"You need to run in your terminal `sudo ldconfig /usr/lib64-nvidia` yourself, then import Unsloth.\n"
"Also try `sudo ldconfig /usr/local/cuda-xx.x` - find the latest cuda version.\n"
"Unsloth will still run for now, but maybe it might crash - let's hope it works!"
__version__ = unsloth_zoo.__version__
DEVICE_TYPE = "mlx"
class FastLanguageModel:
@staticmethod
def from_pretrained(*args, **kwargs):
return FastMLXModel.from_pretrained(*args, **kwargs)
@staticmethod
def get_peft_model(*args, **kwargs):
return FastMLXModel.get_peft_model(*args, **kwargs)
@staticmethod
def for_inference(*args, **kwargs):
return args[0] if args else None
class FastVisionModel(FastLanguageModel):
@staticmethod
def from_pretrained(*args, **kwargs):
kwargs.setdefault("text_only", False)
return FastMLXModel.from_pretrained(*args, **kwargs)
@staticmethod
def for_training(*args, **kwargs):
return args[0] if args else None
FastTextModel = FastLanguageModel
FastModel = FastLanguageModel
class FastSentenceTransformer:
@staticmethod
def from_pretrained(*args, **kwargs):
raise NotImplementedError(
"Unsloth: FastSentenceTransformer is not yet supported on MLX."
)
del libcuda_dirs
elif DEVICE_TYPE == "hip":
# NO-OP for rocm device
pass
elif DEVICE_TYPE == "xpu":
import bitsandbytes as bnb
# TODO: check triton for intel installed properly.
pass
@staticmethod
def get_peft_model(*args, **kwargs):
raise NotImplementedError(
"Unsloth: FastSentenceTransformer is not yet supported on MLX."
)
from .models import *
from .models import __version__
from .save import *
from .chat_templates import *
from .tokenizer_utils import *
from .trainer import *
def is_bfloat16_supported():
try:
import mlx.core as mx
# Export dataprep utilities for CLI and downstream users
from .dataprep.raw_text import RawTextDataLoader, TextPreprocessor
from unsloth_zoo.rl_environments import (
check_python_modules,
create_locked_down_function,
execute_with_time_limit,
Benchmarker,
is_port_open,
launch_openenv,
)
name = mx.device_info().get("device_name", "") or ""
return not name.startswith(("Apple M1", "Apple M2"))
except Exception:
return True
# Patch TRL trainers for backwards compatibility
_patch_trl_trainer()
is_bf16_supported = is_bfloat16_supported
class UnslothVisionDataCollator:
def __init__(self, *args, **kwargs):
raise NotImplementedError(
"Unsloth: UnslothVisionDataCollator is not used on MLX. "
"Use the MLX trainer/data path instead."
)
else:
# GPU path: load everything from _gpu_init
from ._gpu_init import *
from ._gpu_init import __version__

346
unsloth/_gpu_init.py Normal file
View file

@ -0,0 +1,346 @@
# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import warnings, importlib, sys
from packaging.version import Version
import os, re, subprocess, inspect, functools
import numpy as np
# Log Unsloth is being used
os.environ["UNSLOTH_IS_PRESENT"] = "1"
# Check if modules that need patching are already imported
critical_modules = ["trl", "transformers", "peft"]
already_imported = [mod for mod in critical_modules if mod in sys.modules]
# Fix some issues before importing other packages
from .import_fixes import (
fix_message_factory_issue,
check_fbgemm_gpu_version,
disable_broken_causal_conv1d,
disable_broken_vllm,
configure_amdgpu_asic_id_table_path,
torchvision_compatibility_check,
fix_diffusers_warnings,
fix_huggingface_hub,
)
# Configure libdrm ids table path early so ROCm can resolve AMD GPU names.
configure_amdgpu_asic_id_table_path()
disable_broken_causal_conv1d()
disable_broken_vllm()
fix_message_factory_issue()
check_fbgemm_gpu_version()
torchvision_compatibility_check()
fix_diffusers_warnings()
fix_huggingface_hub()
del configure_amdgpu_asic_id_table_path
del disable_broken_causal_conv1d
del disable_broken_vllm
del fix_message_factory_issue
del check_fbgemm_gpu_version
del torchvision_compatibility_check
del fix_diffusers_warnings
del fix_huggingface_hub
# This check is critical because Unsloth optimizes these libraries by modifying
# their code at import time. If they're imported first, the original (slower,
# more memory-intensive) implementations will be used instead of Unsloth's
# optimized versions, potentially causing OOM errors or slower training.
if already_imported:
# stacklevel=2 makes warning point to user's import line rather than this library code,
# showing them exactly where to fix the import order in their script
warnings.warn(
f"WARNING: Unsloth should be imported before [{', '.join(already_imported)}] "
f"to ensure all optimizations are applied. Your code may run slower or encounter "
f"memory issues without these optimizations.\n\n"
f"Please restructure your imports with 'import unsloth' at the top of your file.",
stacklevel = 2,
)
del already_imported, critical_modules
# Unsloth currently does not work on multi GPU setups - sadly we are a 2 brother team so
# enabling it will require much more work, so we have to prioritize. Please understand!
# We do have a beta version, which you can contact us about!
# Thank you for your understanding and we appreciate it immensely!
# Fixes https://github.com/unslothai/unsloth/issues/1266
os.environ["PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION"] = "python"
# [TODO] Check why some GPUs don't work
# "pinned_use_cuda_host_register:True,"\
# "pinned_num_register_threads:8"
from importlib.metadata import version as importlib_version
from importlib.metadata import PackageNotFoundError
# Check for unsloth_zoo
try:
unsloth_zoo_version = importlib_version("unsloth_zoo")
if Version(unsloth_zoo_version) < Version("2026.3.4"):
print(
"Unsloth: Please update Unsloth and Unsloth-Zoo to the latest version!\n"
"Do this via `pip install --upgrade --force-reinstall --no-cache-dir --no-deps unsloth unsloth_zoo`"
)
# if os.environ.get("UNSLOTH_DISABLE_AUTO_UPDATES", "0") == "0":
# try:
# os.system("pip install --upgrade --no-cache-dir --no-deps unsloth_zoo")
# except:
# try:
# os.system("pip install --upgrade --no-cache-dir --no-deps --user unsloth_zoo")
# except:
# raise ImportError("Unsloth: Please update unsloth_zoo via `pip install --upgrade --no-cache-dir --no-deps unsloth_zoo`")
import unsloth_zoo
except PackageNotFoundError:
raise ImportError(
f"Unsloth: Please install unsloth_zoo via `pip install unsloth_zoo` then retry!"
)
except:
raise
del PackageNotFoundError, importlib_version
# Try importing PyTorch and check version
try:
import torch
except ModuleNotFoundError:
raise ImportError(
"Unsloth: Pytorch is not installed. Go to https://pytorch.org/.\n"
"We have some installation instructions on our Github page."
)
except:
raise
from unsloth_zoo.device_type import (
is_hip,
get_device_type,
DEVICE_TYPE,
DEVICE_TYPE_TORCH,
DEVICE_COUNT,
ALLOW_PREQUANTIZED_MODELS,
)
# Fix other issues
from .import_fixes import (
fix_xformers_performance_issue,
fix_vllm_aimv2_issue,
check_vllm_torch_sm100_compatibility,
fix_vllm_guided_decoding_params,
fix_vllm_pdl_blackwell,
fix_triton_compiled_kernel_missing_attrs,
patch_trunc_normal_precision_issue,
ignore_logger_messages,
patch_ipykernel_hf_xet,
patch_trackio,
patch_datasets,
patch_enable_input_require_grads,
fix_openenv_no_vllm,
patch_openspiel_env_async,
fix_executorch,
patch_vllm_for_notebooks,
patch_torchcodec_audio_decoder,
disable_torchcodec_if_broken,
disable_broken_wandb,
fix_trl_vllm_ascend,
patch_peft_weight_converter_compatibility,
)
fix_xformers_performance_issue()
fix_vllm_aimv2_issue()
# Check vLLM + torch < 2.9.0 + SM100 compatibility BEFORE importing vLLM
check_vllm_torch_sm100_compatibility()
fix_vllm_guided_decoding_params()
fix_trl_vllm_ascend()
fix_vllm_pdl_blackwell()
fix_triton_compiled_kernel_missing_attrs()
patch_trunc_normal_precision_issue()
ignore_logger_messages()
patch_ipykernel_hf_xet()
patch_trackio()
patch_datasets()
patch_enable_input_require_grads()
fix_openenv_no_vllm()
patch_openspiel_env_async()
fix_executorch()
patch_vllm_for_notebooks()
patch_torchcodec_audio_decoder()
disable_torchcodec_if_broken()
disable_broken_wandb()
patch_peft_weight_converter_compatibility()
del fix_xformers_performance_issue
del fix_vllm_aimv2_issue
del check_vllm_torch_sm100_compatibility
del fix_vllm_guided_decoding_params
del fix_trl_vllm_ascend
del fix_vllm_pdl_blackwell
del fix_triton_compiled_kernel_missing_attrs
del patch_trunc_normal_precision_issue
del ignore_logger_messages
del patch_ipykernel_hf_xet
del patch_trackio
del patch_datasets
del patch_enable_input_require_grads
del fix_openenv_no_vllm
del patch_openspiel_env_async
del fix_executorch
del patch_vllm_for_notebooks
del patch_torchcodec_audio_decoder
del disable_torchcodec_if_broken
del disable_broken_wandb
del patch_peft_weight_converter_compatibility
# Torch 2.4 has including_emulation
if DEVICE_TYPE == "cuda":
major_version, minor_version = torch.cuda.get_device_capability()
SUPPORTS_BFLOAT16 = major_version >= 8
old_is_bf16_supported = torch.cuda.is_bf16_supported
if "including_emulation" in str(inspect.signature(old_is_bf16_supported)):
def is_bf16_supported(including_emulation = False):
return old_is_bf16_supported(including_emulation)
torch.cuda.is_bf16_supported = is_bf16_supported
else:
def is_bf16_supported():
return SUPPORTS_BFLOAT16
torch.cuda.is_bf16_supported = is_bf16_supported
del major_version, minor_version
elif DEVICE_TYPE == "hip":
SUPPORTS_BFLOAT16 = torch.cuda.is_bf16_supported()
elif DEVICE_TYPE == "xpu":
# torch.xpu.is_bf16_supported() does not have including_emulation
# set SUPPORTS_BFLOAT16 as torch.xpu.is_bf16_supported()
SUPPORTS_BFLOAT16 = torch.xpu.is_bf16_supported()
# For Gradio HF Spaces?
# if "SPACE_AUTHOR_NAME" not in os.environ and "SPACE_REPO_NAME" not in os.environ:
import triton
if DEVICE_TYPE == "cuda":
libcuda_dirs = lambda: None
if Version(triton.__version__) >= Version("3.0.0"):
try:
from triton.backends.nvidia.driver import libcuda_dirs
except:
pass
else:
from triton.common.build import libcuda_dirs
# Try loading bitsandbytes and triton
try:
import bitsandbytes as bnb
except:
print(
"Unsloth: `bitsandbytes` is not installed - 4bit QLoRA unallowed, but 16bit and full finetuning works!"
)
bnb = None
try:
cdequantize_blockwise_fp32 = bnb.functional.lib.cdequantize_blockwise_fp32
libcuda_dirs()
except:
if hasattr(os, "geteuid") and os.geteuid() == 0:
warnings.warn("Unsloth: Running `ldconfig /usr/lib64-nvidia` to link CUDA.")
if os.path.exists("/usr/lib64-nvidia"):
os.system("ldconfig /usr/lib64-nvidia")
elif os.path.exists("/usr/local"):
# Sometimes bitsandbytes cannot be linked properly in Runpod for example
possible_cudas = (
subprocess.check_output(["ls", "-al", "/usr/local"])
.decode("utf-8")
.split("\n")
)
find_cuda = re.compile(r"[\s](cuda\-[\d\.]{2,})$")
possible_cudas = [find_cuda.search(x) for x in possible_cudas]
possible_cudas = [x.group(1) for x in possible_cudas if x is not None]
# Try linking cuda folder, or everything in local
if len(possible_cudas) == 0:
os.system("ldconfig /usr/local/")
else:
find_number = re.compile(r"([\d\.]{2,})")
latest_cuda = np.argsort(
[float(find_number.search(x).group(1)) for x in possible_cudas]
)[::-1][0]
latest_cuda = possible_cudas[latest_cuda]
os.system(f"ldconfig /usr/local/{latest_cuda}")
del find_number, latest_cuda
del possible_cudas, find_cuda
if bnb is not None:
importlib.reload(bnb)
importlib.reload(triton)
try:
libcuda_dirs = lambda: None
if Version(triton.__version__) >= Version("3.0.0"):
try:
from triton.backends.nvidia.driver import libcuda_dirs
except:
pass
else:
from triton.common.build import libcuda_dirs
cdequantize_blockwise_fp32 = (
bnb.functional.lib.cdequantize_blockwise_fp32
)
libcuda_dirs()
except:
warnings.warn(
"Unsloth: CUDA is not linked properly.\n"
"Try running `python -m bitsandbytes` then `python -m xformers.info`\n"
"We tried running `ldconfig /usr/lib64-nvidia` ourselves, but it didn't work.\n"
"You need to run in your terminal `sudo ldconfig /usr/lib64-nvidia` yourself, then import Unsloth.\n"
"Also try `sudo ldconfig /usr/local/cuda-xx.x` - find the latest cuda version.\n"
"Unsloth will still run for now, but maybe it might crash - let's hope it works!"
)
elif bnb is not None:
warnings.warn(
"Unsloth: CUDA is not linked properly.\n"
"You need to run in your terminal `sudo ldconfig /usr/lib64-nvidia` yourself, then import Unsloth.\n"
"Also try `sudo ldconfig /usr/local/cuda-xx.x` - find the latest cuda version.\n"
"Unsloth will still run for now, but maybe it might crash - let's hope it works!"
)
del libcuda_dirs
elif DEVICE_TYPE == "hip":
# NO-OP for rocm device
pass
elif DEVICE_TYPE == "xpu":
import bitsandbytes as bnb
# TODO: check triton for intel installed properly.
pass
from .models import *
from .models import __version__
from .save import *
from .chat_templates import *
from .tokenizer_utils import *
from .trainer import *
# Export dataprep utilities for CLI and downstream users
from .dataprep.raw_text import RawTextDataLoader, TextPreprocessor
from unsloth_zoo.rl_environments import (
check_python_modules,
create_locked_down_function,
execute_with_time_limit,
Benchmarker,
is_port_open,
launch_openenv,
)
# Patch TRL trainers for backwards compatibility
_patch_trl_trainer()

View file

@ -161,8 +161,14 @@ else:
if DEVICE_TYPE == "xpu":
_gpu_getCurrentRawStream = torch._C._xpu_getCurrentRawStream
# NVIDIA GPU Default Logic
else:
elif hasattr(torch._C, "_cuda_getCurrentRawStream"):
_gpu_getCurrentRawStream = torch._C._cuda_getCurrentRawStream
else:
# CPU-only torch wheel (no compiled CUDA backend). _get_tensor_stream
# is only invoked during real GPU work, so a no-op binding is safe.
def _gpu_getCurrentRawStream(_index = 0):
return 0
c_void_p = ctypes.c_void_p
@ -177,36 +183,49 @@ global XPU_STREAMS
global WEIGHT_BUFFERS
global ABSMAX_BUFFERS
# INTEL GPU Specific Logic
# DEVICE_COUNT == 0 = no visible accelerator (e.g. CPU-only CI runner).
# The consumer functions below only index these arrays during real GPU
# work, so empty containers are safe -- they just need to be defined so
# the module imports cleanly.
if DEVICE_TYPE == "xpu":
_XPU_STREAMS = {
(index := torch.xpu.device(i).idx): ctypes.c_void_p(
torch._C._xpu_getCurrentRawStream(index)
)
for i in range(DEVICE_COUNT)
}
XPU_STREAMS = [None] * (max(_XPU_STREAMS.keys()) + 1)
WEIGHT_BUFFERS = [None] * (max(_XPU_STREAMS.keys()) + 1)
ABSMAX_BUFFERS = [None] * (max(_XPU_STREAMS.keys()) + 1)
for k, v in _XPU_STREAMS.items():
XPU_STREAMS[k] = v
XPU_STREAMS = tuple(XPU_STREAMS)
del _XPU_STREAMS
if DEVICE_COUNT > 0:
_XPU_STREAMS = {
(index := torch.xpu.device(i).idx): ctypes.c_void_p(
torch._C._xpu_getCurrentRawStream(index)
)
for i in range(DEVICE_COUNT)
}
XPU_STREAMS = [None] * (max(_XPU_STREAMS.keys()) + 1)
WEIGHT_BUFFERS = [None] * (max(_XPU_STREAMS.keys()) + 1)
ABSMAX_BUFFERS = [None] * (max(_XPU_STREAMS.keys()) + 1)
for k, v in _XPU_STREAMS.items():
XPU_STREAMS[k] = v
XPU_STREAMS = tuple(XPU_STREAMS)
del _XPU_STREAMS
else:
XPU_STREAMS = ()
WEIGHT_BUFFERS = []
ABSMAX_BUFFERS = []
else:
# NVIDIA GPU Default Logic
_CUDA_STREAMS = {
(index := torch.cuda.device(i).idx): ctypes.c_void_p(
torch._C._cuda_getCurrentRawStream(index)
)
for i in range(DEVICE_COUNT)
}
CUDA_STREAMS = [None] * (max(_CUDA_STREAMS.keys()) + 1)
WEIGHT_BUFFERS = [None] * (max(_CUDA_STREAMS.keys()) + 1)
ABSMAX_BUFFERS = [None] * (max(_CUDA_STREAMS.keys()) + 1)
for k, v in _CUDA_STREAMS.items():
CUDA_STREAMS[k] = v
CUDA_STREAMS = tuple(CUDA_STREAMS)
del _CUDA_STREAMS
if DEVICE_COUNT > 0:
_CUDA_STREAMS = {
(index := torch.cuda.device(i).idx): ctypes.c_void_p(
torch._C._cuda_getCurrentRawStream(index)
)
for i in range(DEVICE_COUNT)
}
CUDA_STREAMS = [None] * (max(_CUDA_STREAMS.keys()) + 1)
WEIGHT_BUFFERS = [None] * (max(_CUDA_STREAMS.keys()) + 1)
ABSMAX_BUFFERS = [None] * (max(_CUDA_STREAMS.keys()) + 1)
for k, v in _CUDA_STREAMS.items():
CUDA_STREAMS[k] = v
CUDA_STREAMS = tuple(CUDA_STREAMS)
del _CUDA_STREAMS
else:
CUDA_STREAMS = ()
WEIGHT_BUFFERS = []
ABSMAX_BUFFERS = []
# Bitsandbytes operations
ctypes_c_int = ctypes.c_int

View file

@ -23,7 +23,9 @@ from functools import wraps
import trl
import inspect
from trl import SFTTrainer
from . import is_bfloat16_supported
# why: bypass partially-initialised unsloth ns during _gpu_init load
from .models._utils import is_bfloat16_supported
from unsloth.utils import (
configure_padding_free,
configure_sample_packing,

View file

@ -20,7 +20,73 @@ import typer
studio_app = typer.Typer(help = "Unsloth Studio commands.")
STUDIO_HOME = Path.home() / ".unsloth" / "studio"
# Resolve install root: UNSLOTH_STUDIO_HOME, then STUDIO_HOME alias, then
# sys.prefix inference (so a direct call to <root>/bin/unsloth resolves after
# the installer's env var has expired), then legacy ~/.unsloth/studio.
# UNSLOTH_STUDIO_HOME wins when both env vars are set.
def _looks_like_installer_managed_studio_home(candidate: Path) -> bool:
"""Sentinel check (studio.conf or bin shim) so a dev venv named
unsloth_studio is not misidentified as a custom Studio root.
"""
shim_name = "unsloth.exe" if platform.system() == "Windows" else "unsloth"
return (candidate / "share" / "studio.conf").is_file() or (
candidate / "bin" / shim_name
).is_file()
def _resolve_studio_home() -> tuple[Path, bool]:
override = (os.environ.get("UNSLOTH_STUDIO_HOME") or "").strip()
if not override:
override = (os.environ.get("STUDIO_HOME") or "").strip()
if override:
try:
return Path(override).expanduser().resolve(), True
except (OSError, ValueError):
return Path(override).expanduser(), True
try:
prefix = Path(sys.prefix).resolve()
if prefix.name == "unsloth_studio":
inferred = prefix.parent
legacy = (Path.home() / ".unsloth" / "studio").resolve()
if inferred != legacy and _looks_like_installer_managed_studio_home(
inferred
):
return inferred, True
except (OSError, ValueError):
pass
return Path.home() / ".unsloth" / "studio", False
STUDIO_HOME, _STUDIO_HOME_IS_CUSTOM = _resolve_studio_home()
def _ensure_studio_env_exported() -> None:
"""Re-export UNSLOTH_STUDIO_HOME / UNSLOTH_LLAMA_CPP_PATH only for real
custom roots so subprocesses inherit the right install. Called from each
studio subcommand entry rather than at import time, to avoid leaking env
state into unrelated importers (tests, --help, CLI introspection).
"""
if not _STUDIO_HOME_IS_CUSTOM:
return
# Truthy-check (not setdefault) so a blank UNSLOTH_STUDIO_HOME= does not
# suppress the inferred custom root.
if not os.environ.get("UNSLOTH_STUDIO_HOME"):
os.environ["UNSLOTH_STUDIO_HOME"] = str(STUDIO_HOME)
# When override == legacy default, llama.cpp stays at ~/.unsloth/llama.cpp.
try:
_legacy_studio = (Path.home() / ".unsloth" / "studio").resolve()
_is_legacy = STUDIO_HOME.resolve() == _legacy_studio
except (OSError, ValueError):
_is_legacy = STUDIO_HOME == (Path.home() / ".unsloth" / "studio")
if _is_legacy:
_llama_dir = Path.home() / ".unsloth" / "llama.cpp"
else:
_llama_dir = STUDIO_HOME / "llama.cpp"
if not os.environ.get("UNSLOTH_LLAMA_CPP_PATH"):
os.environ["UNSLOTH_LLAMA_CPP_PATH"] = str(_llama_dir)
BOOTSTRAP_PASSWORD_FILE = ".bootstrap_password"
DESKTOP_SECRET_FILE = ".desktop_secret"
DEFAULT_ADMIN_USERNAME = "unsloth"
@ -427,6 +493,8 @@ def studio_default(
),
):
"""Launch the Unsloth Studio server."""
# Runs before any subcommand; covers run/setup/update/etc in one place.
_ensure_studio_env_exported()
if ctx.invoked_subcommand is not None:
return