diff --git a/.github/workflows/studio-backend-ci.yml b/.github/workflows/studio-backend-ci.yml new file mode 100644 index 0000000000..5a858888e7 --- /dev/null +++ b/.github/workflows/studio-backend-ci.yml @@ -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 diff --git a/.github/workflows/studio-frontend-ci.yml b/.github/workflows/studio-frontend-ci.yml new file mode 100644 index 0000000000..039bd5dd08 --- /dev/null +++ b/.github/workflows/studio-frontend-ci.yml @@ -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 diff --git a/.github/workflows/studio-inference-smoke.yml b/.github/workflows/studio-inference-smoke.yml new file mode 100644 index 0000000000..8efe072d28 --- /dev/null +++ b/.github/workflows/studio-inference-smoke.yml @@ -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 diff --git a/.github/workflows/studio-tauri-smoke.yml b/.github/workflows/studio-tauri-smoke.yml new file mode 100644 index 0000000000..fcc9c8d963 --- /dev/null +++ b/.github/workflows/studio-tauri-smoke.yml @@ -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 diff --git a/.github/workflows/wheel-smoke.yml b/.github/workflows/wheel-smoke.yml new file mode 100644 index 0000000000..080a6bb261 --- /dev/null +++ b/.github/workflows/wheel-smoke.yml @@ -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--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 diff --git a/.gitignore b/.gitignore index b6786ee655..ae6770bc07 100644 --- a/.gitignore +++ b/.gitignore @@ -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/ diff --git a/install.ps1 b/install.ps1 index d5db1785a1..ef87c5ed08 100644 --- a/install.ps1 +++ b/install.ps1 @@ -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 "" } diff --git a/install.sh b/install.sh index 1d21117d16..9046a9bdf6 100755 --- a/install.sh +++ b/install.sh @@ -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_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" &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 diff --git a/pyproject.toml b/pyproject.toml index c2f884e192..5687ea12f8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/studio/backend/core/export/export.py b/studio/backend/core/export/export.py index 6fee5a38f7..4ab95d896f 100644 --- a/studio/backend/core/export/export.py +++ b/studio/backend/core/export/export.py @@ -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 diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py index f768764c22..8da836de38 100644 --- a/studio/backend/core/inference/llama_cpp.py +++ b/studio/backend/core/inference/llama_cpp.py @@ -732,22 +732,46 @@ class LlamaCppBackend: if win_bin.is_file(): return str(win_bin) - # 2–4. ~/.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) # 5–6. 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] diff --git a/studio/backend/core/inference/mlx_inference.py b/studio/backend/core/inference/mlx_inference.py new file mode 100644 index 0000000000..1d2b03ecb9 --- /dev/null +++ b/studio/backend/core/inference/mlx_inference.py @@ -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() diff --git a/studio/backend/core/inference/worker.py b/studio/backend/core/inference/worker.py index fbcce276ba..085a1ab899 100644 --- a/studio/backend/core/inference/worker.py +++ b/studio/backend/core/inference/worker.py @@ -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) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index fe8d277ac0..a3f063694f 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -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" diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index 5642faa189..72b13c3225 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -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: diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py index 60b9e994ab..ef5cafb175 100644 --- a/studio/backend/core/training/worker.py +++ b/studio/backend/core/training/worker.py @@ -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, ) diff --git a/studio/backend/main.py b/studio/backend/main.py index 0958094ff0..cd901327db 100644 --- a/studio/backend/main.py +++ b/studio/backend/main.py @@ -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(), } diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index a9f4caa1bb..8127af1ee6 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -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)") diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index e5195bb337..19202f3883 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -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, diff --git a/studio/backend/run.py b/studio/backend/run.py index c5b103ff70..1dd1230a17 100644 --- a/studio/backend/run.py +++ b/studio/backend/run.py @@ -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(): diff --git a/studio/backend/tests/test_mlx_inference_backend.py b/studio/backend/tests/test_mlx_inference_backend.py new file mode 100644 index 0000000000..868e537372 --- /dev/null +++ b/studio/backend/tests/test_mlx_inference_backend.py @@ -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) diff --git a/studio/backend/tests/test_mlx_training_worker_config.py b/studio/backend/tests/test_mlx_training_worker_config.py new file mode 100644 index 0000000000..5900af4e3d --- /dev/null +++ b/studio/backend/tests/test_mlx_training_worker_config.py @@ -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") diff --git a/studio/backend/tests/test_training_raw_support.py b/studio/backend/tests/test_training_raw_support.py new file mode 100644 index 0000000000..876ee34686 --- /dev/null +++ b/studio/backend/tests/test_training_raw_support.py @@ -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 = "", + append_eos = True, + ) + + self.assertEqual(len(result.dataset), 2) + self.assertEqual(result.dataset[0]["text"], "hello") + self.assertEqual(result.dataset[1]["text"], "world") + self.assertTrue( + any( + "null or non-string 'text' values" in notice.message + for notice in result.notices + ) + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/studio/backend/utils/datasets/dataset_utils.py b/studio/backend/utils/datasets/dataset_utils.py index fac8c3d295..26378d64ee 100644 --- a/studio/backend/utils/datasets/dataset_utils.py +++ b/studio/backend/utils/datasets/dataset_utils.py @@ -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: diff --git a/studio/backend/utils/datasets/raw_text.py b/studio/backend/utils/datasets/raw_text.py new file mode 100644 index 0000000000..353145fd5a --- /dev/null +++ b/studio/backend/utils/datasets/raw_text.py @@ -0,0 +1,142 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +""" +Shared 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) diff --git a/studio/backend/utils/hardware/hardware.py b/studio/backend/utils/hardware/hardware.py index b800ba0d6a..3b1c2a54dc 100644 --- a/studio/backend/utils/hardware/hardware.py +++ b/studio/backend/utils/hardware/hardware.py @@ -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 { diff --git a/studio/backend/utils/models/model_config.py b/studio/backend/utils/models/model_config.py index 16f6d21edb..dc8dd08315 100644 --- a/studio/backend/utils/models/model_config.py +++ b/studio/backend/utils/models/model_config.py @@ -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. diff --git a/studio/backend/utils/paths/storage_roots.py b/studio/backend/utils/paths/storage_roots.py index b52609b06b..58a4d7967c 100644 --- a/studio/backend/utils/paths/storage_roots.py +++ b/studio/backend/utils/paths/storage_roots.py @@ -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: diff --git a/studio/backend/utils/transformers_version.py b/studio/backend/utils/transformers_version.py index 17af40f663..9075c590ca 100644 --- a/studio/backend/utils/transformers_version.py +++ b/studio/backend/utils/transformers_version.py @@ -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 diff --git a/studio/frontend/src/components/app-sidebar.tsx b/studio/frontend/src/components/app-sidebar.tsx index edcd5120eb..13b8adfa48 100644 --- a/studio/frontend/src/components/app-sidebar.tsx +++ b/studio/frontend/src/components/app-sidebar.tsx @@ -38,11 +38,11 @@ import { Download03Icon, GemIcon, Globe02Icon, + HelpCircleIcon, Search01Icon, PowerIcon, PencilEdit02Icon, LayoutAlignLeftIcon, - HelpCircleIcon, Settings02Icon, ZapIcon, } from "@hugeicons/core-free-icons"; diff --git a/studio/frontend/src/components/assistant-ui/reasoning.tsx b/studio/frontend/src/components/assistant-ui/reasoning.tsx index 387f8cd458..fe913baf2a 100644 --- a/studio/frontend/src/components/assistant-ui/reasoning.tsx +++ b/studio/frontend/src/components/assistant-ui/reasoning.tsx @@ -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 }) => { diff --git a/studio/frontend/src/config/env.ts b/studio/frontend/src/config/env.ts index 72bb3fa815..3839706d25 100644 --- a/studio/frontend/src/config/env.ts +++ b/studio/frontend/src/config/env.ts @@ -50,7 +50,7 @@ export async function fetchDeviceType(): Promise { 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; } diff --git a/studio/frontend/src/config/training.ts b/studio/frontend/src/config/training.ts index 913d612838..e9fe1d679c 100644 --- a/studio/frontend/src/config/training.ts +++ b/studio/frontend/src/config/training.ts @@ -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, diff --git a/studio/frontend/src/features/export/constants.ts b/studio/frontend/src/features/export/constants.ts index e9c3b8c95b..c97d9c1f6e 100644 --- a/studio/frontend/src/features/export/constants.ts +++ b/studio/frontend/src/features/export/constants.ts @@ -74,6 +74,7 @@ export const METHOD_LABELS: Record = { qlora: "QLoRA", lora: "LoRA", full: "Full Fine-tune", + cpt: "Continued Pretraining", }; export const GUIDE_STEPS = [ diff --git a/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx b/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx index 9b0456f67a..2945654d68 100644 --- a/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx +++ b/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx @@ -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() { diff --git a/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx b/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx index 1ff23cb0fc..52f61700e2 100644 --- a/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx +++ b/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx @@ -366,6 +366,7 @@ export function ModelSelectionStep() { QLoRA (4-bit) LoRA (16-bit) Full Fine-tune + Continued Pretraining diff --git a/studio/frontend/src/features/onboarding/components/steps/summary-step.tsx b/studio/frontend/src/features/onboarding/components/steps/summary-step.tsx index 8840983574..1988cfe970 100644 --- a/studio/frontend/src/features/onboarding/components/steps/summary-step.tsx +++ b/studio/frontend/src/features/onboarding/components/steps/summary-step.tsx @@ -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 (
@@ -150,7 +152,7 @@ export function SummaryStep() {
- +
@@ -199,7 +201,7 @@ export function SummaryStep() {
Training - {trainingMethod === "qlora" ? "QLoRA" : trainingMethod === "lora" ? "LoRA" : "Full"} + {trainingMethodLabel}
diff --git a/studio/frontend/src/features/studio/historical-training-view.tsx b/studio/frontend/src/features/studio/historical-training-view.tsx index e461fc5a90..d13ad12727 100644 --- a/studio/frontend/src/features/studio/historical-training-view.tsx +++ b/studio/frontend/src/features/studio/historical-training-view.tsx @@ -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 { - 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; diff --git a/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx b/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx index 5ad2d582c4..d05a8cf242 100644 --- a/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx @@ -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({ - {data.warning && ( + {data.warning && !isRawFormat && (
{data.warning} diff --git a/studio/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx index 11c2321863..80f0a06c3f 100644 --- a/studio/frontend/src/features/studio/sections/dataset-section.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx @@ -913,6 +913,7 @@ export function DatasetSection() { Alpaca ChatML ShareGPT + Raw Text
diff --git a/studio/frontend/src/features/studio/sections/model-section.tsx b/studio/frontend/src/features/studio/sections/model-section.tsx index 775073eb64..a3c737dae5 100644 --- a/studio/frontend/src/features/studio/sections/model-section.tsx +++ b/studio/frontend/src/features/studio/sections/model-section.tsx @@ -67,6 +67,7 @@ const METHOD_DOTS: Record = { 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() { 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.{" "} + + + + Continued Pretraining + + diff --git a/studio/frontend/src/features/studio/sections/params-section.tsx b/studio/frontend/src/features/studio/sections/params-section.tsx index 6566303a81..19b75db140 100644 --- a/studio/frontend/src/features/studio/sections/params-section.tsx +++ b/studio/frontend/src/features/studio/sections/params-section.tsx @@ -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(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" />

- 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

+ {/* Embedding Learning Rate (CPT only) */} + {isCpt && ( +
+ + Embedding Learning Rate + + + + + + Only used when CPT is training embed_tokens. + Embeddings are easier to destabilize than LoRA weights, so + they usually need a smaller LR. Leave blank to use + lr/10; typical working range is 2x-10x smaller + than the main LR. Increase it only if vocabulary or + domain-token adaptation is too slow. + + + + { + 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" + /> +

+ Leave blank to use lr/10 (recommended). Typical range is + 2x-10x smaller than the main learning rate. +

+
+ )} + {/* LoRA Settings */} {isLora && ( @@ -514,7 +574,7 @@ export function ParamsSection(): ReactElement { Target Modules
- {TARGET_MODULES.map((mod) => { + {(isCpt ? CPT_TARGET_MODULES : TARGET_MODULES).map((mod) => { const active = store.targetModules.includes(mod); return (
)} - {!store.isEmbeddingModel && ( + {!store.isEmbeddingModel && !isCpt && !isRawText && (
({ @@ -272,7 +274,7 @@ export function ProgressSection({ {data.modelName || "--"} - {data.trainingMethod === "qlora" ? "QLoRA" : data.trainingMethod === "lora" ? "LoRA" : "Full"} + {trainingMethodLabel}
diff --git a/studio/frontend/src/features/training/api/mappers.ts b/studio/frontend/src/features/training/api/mappers.ts index 561dbe1408..5e68ccd72c 100644 --- a/studio/frontend/src/features/training/api/mappers.ts +++ b/studio/frontend/src/features/training/api/mappers.ts @@ -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, diff --git a/studio/frontend/src/features/training/hooks/use-training-actions.ts b/studio/frontend/src/features/training/hooks/use-training-actions.ts index deaec6c6c2..cff4e929ae 100644 --- a/studio/frontend/src/features/training/hooks/use-training-actions.ts +++ b/studio/frontend/src/features/training/hooks/use-training-actions.ts @@ -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 = {}; diff --git a/studio/frontend/src/features/training/lib/model-defaults.ts b/studio/frontend/src/features/training/lib/model-defaults.ts index c40a1e2282..8bc9c4e064 100644 --- a/studio/frontend/src/features/training/lib/model-defaults.ts +++ b/studio/frontend/src/features/training/lib/model-defaults.ts @@ -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; } diff --git a/studio/frontend/src/features/training/lib/training-methods.ts b/studio/frontend/src/features/training/lib/training-methods.ts new file mode 100644 index 0000000000..9070f2adcf --- /dev/null +++ b/studio/frontend/src/features/training/lib/training-methods.ts @@ -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 = { + qlora: "LoRA/QLoRA", + lora: "LoRA/QLoRA", + full: "Full Finetuning", + cpt: "Continued Pretraining", +}; + +const TRAINING_METHOD_LABELS: Record = { + 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"; +} diff --git a/studio/frontend/src/features/training/stores/training-config-store.ts b/studio/frontend/src/features/training/stores/training-config-store.ts index 8214b0eb2a..ef16f641f5 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -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 = 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()( persist( (set, get) => { @@ -216,11 +340,14 @@ export const useTrainingConfigStore = create()( // 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()( }); } + // 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()( 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()( 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()( _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()( _trainOnCompletionsManuallySet = false; _learningRateManuallySet = false; _yamlLearningRate = undefined; + clearCptDatasetFormatTracking(); set(initialState); }, resetToModelDefaults: () => { @@ -629,7 +771,7 @@ export const useTrainingConfigStore = create()( }, { name: "unsloth_training_config_v1", - version: 9, + version: 10, migrate: (persisted, version) => { const s = persisted as Record; if (version < 2 && s.datasetSubset == null && s.datasetConfig != null) { @@ -665,6 +807,17 @@ export const useTrainingConfigStore = create()( 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, diff --git a/studio/frontend/src/features/training/types/api.ts b/studio/frontend/src/features/training/types/api.ts index ae65d1a53c..fb8a2f899e 100644 --- a/studio/frontend/src/features/training/types/api.ts +++ b/studio/frontend/src/features/training/types/api.ts @@ -21,6 +21,8 @@ export interface TrainingStartRequest { custom_format_mapping?: Record | 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; diff --git a/studio/frontend/src/features/training/types/config.ts b/studio/frontend/src/features/training/types/config.ts index 2d19dea874..5b156316ca 100644 --- a/studio/frontend/src/features/training/types/config.ts +++ b/studio/frontend/src/features/training/types/config.ts @@ -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; diff --git a/studio/frontend/src/hooks/use-hf-model-search.ts b/studio/frontend/src/hooks/use-hf-model-search.ts index 77214f38d9..efe4d726be 100644 --- a/studio/frontend/src/hooks/use-hf-model-search.ts +++ b/studio/frontend/src/hooks/use-hf-model-search.ts @@ -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) { 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. diff --git a/studio/frontend/src/lib/vram.ts b/studio/frontend/src/lib/vram.ts index a5abebb043..fad2c438bd 100644 --- a/studio/frontend/src/lib/vram.ts +++ b/studio/frontend/src/lib/vram.ts @@ -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 }); } diff --git a/studio/frontend/src/types/training.ts b/studio/frontend/src/types/training.ts index d65d14fb83..feca65ebda 100644 --- a/studio/frontend/src/types/training.ts +++ b/studio/frontend/src/types/training.ts @@ -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; diff --git a/studio/setup.ps1 b/studio/setup.ps1 index 3d082aa70d..f2753d5c88 100644 --- a/studio/setup.ps1 +++ b/studio/setup.ps1 @@ -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 { diff --git a/studio/setup.sh b/studio/setup.sh index 3e875eed30..c5beb7ebd3 100755 --- a/studio/setup.sh +++ b/studio/setup.sh @@ -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 diff --git a/studio/src-tauri/src/commands.rs b/studio/src-tauri/src/commands.rs index 48a0af6e48..e9a27644df 100644 --- a/studio/src-tauri/src/commands.rs +++ b/studio/src-tauri/src/commands.rs @@ -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) => { diff --git a/studio/src-tauri/src/desktop_auth.rs b/studio/src-tauri/src/desktop_auth.rs index 483e7c0432..49b19008fb 100644 --- a/studio/src-tauri/src/desktop_auth.rs +++ b/studio/src-tauri/src/desktop_auth.rs @@ -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; diff --git a/studio/src-tauri/src/install.rs b/studio/src-tauri/src/install.rs index 9d672f5e73..024b730735 100644 --- a/studio/src-tauri/src/install.rs +++ b/studio/src-tauri/src/install.rs @@ -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. diff --git a/studio/src-tauri/src/preflight.rs b/studio/src-tauri/src/preflight.rs index d3df06d057..c0bbb07b36 100644 --- a/studio/src-tauri/src/preflight.rs +++ b/studio/src-tauri/src/preflight.rs @@ -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 = { use std::os::windows::process::CommandExt; diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000000..de41c50fc1 --- /dev/null +++ b/tests/conftest.py @@ -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 .device_type under a mocked + torch.cuda.is_available() == True so its @cache permanently + captures "cuda". prereqs lists submodule names of 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() diff --git a/tests/python/test_fast_sentence_transformer_redirect_lifecycle.py b/tests/python/test_fast_sentence_transformer_redirect_lifecycle.py index ff9b91ec23..31d86b09a4 100644 --- a/tests/python/test_fast_sentence_transformer_redirect_lifecycle.py +++ b/tests/python/test_fast_sentence_transformer_redirect_lifecycle.py @@ -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 diff --git a/tests/python/test_gpu_init_ldconfig_guard.py b/tests/python/test_gpu_init_ldconfig_guard.py new file mode 100644 index 0000000000..081a6132b4 --- /dev/null +++ b/tests/python/test_gpu_init_ldconfig_guard.py @@ -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 diff --git a/tests/studio/install/test_install_llama_prebuilt_logic.py b/tests/studio/install/test_install_llama_prebuilt_logic.py index 79dab30129..128b90cfe2 100644 --- a/tests/studio/install/test_install_llama_prebuilt_logic.py +++ b/tests/studio/install/test_install_llama_prebuilt_logic.py @@ -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") diff --git a/tests/studio/install/test_rocm_support.py b/tests/studio/install/test_rocm_support.py index 99bc9c11bc..81aa66c999 100644 --- a/tests/studio/install/test_rocm_support.py +++ b/tests/studio/install/test_rocm_support.py @@ -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 # ============================================================================= diff --git a/tests/studio/test_export_output_path_contract.py b/tests/studio/test_export_output_path_contract.py new file mode 100644 index 0000000000..e99fc42091 --- /dev/null +++ b/tests/studio/test_export_output_path_contract.py @@ -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 diff --git a/tests/studio/test_hardware_dispatch_matrix.py b/tests/studio/test_hardware_dispatch_matrix.py new file mode 100644 index 0000000000..c7a6841936 --- /dev/null +++ b/tests/studio/test_hardware_dispatch_matrix.py @@ -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 diff --git a/tests/studio/test_is_mlx_dispatch_gate.py b/tests/studio/test_is_mlx_dispatch_gate.py new file mode 100644 index 0000000000..fc07a497e7 --- /dev/null +++ b/tests/studio/test_is_mlx_dispatch_gate.py @@ -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}" diff --git a/tests/studio/test_mlx_training_worker_behaviors.py b/tests/studio/test_mlx_training_worker_behaviors.py new file mode 100644 index 0000000000..6c067ea00b --- /dev/null +++ b/tests/studio/test_mlx_training_worker_behaviors.py @@ -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 diff --git a/tests/test_raw_text.py b/tests/test_raw_text.py index d8289fed20..056ea8660f 100644 --- a/tests/test_raw_text.py +++ b/tests/test_raw_text.py @@ -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 diff --git a/tests/test_studio_install_workspace_guard.py b/tests/test_studio_install_workspace_guard.py new file mode 100644 index 0000000000..d077cdb824 --- /dev/null +++ b/tests/test_studio_install_workspace_guard.py @@ -0,0 +1,1021 @@ +"""install.sh / install.ps1 must refuse to rm -rf an existing +$STUDIO_HOME/unsloth_studio in env-override mode unless the directory +carries a Studio sentinel (share/studio.conf or bin/unsloth). Also +asserts studio/setup.ps1 has the matching writability probe that +setup.sh:417 already performs.""" + +from __future__ import annotations + +import re +import subprocess +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parents[1] +INSTALL_SH = REPO_ROOT / "install.sh" +INSTALL_PS1 = REPO_ROOT / "install.ps1" +SETUP_PS1 = REPO_ROOT / "studio" / "setup.ps1" +SETUP_SH = REPO_ROOT / "studio" / "setup.sh" + +# Stubs for helpers that the extracted install.sh guard block calls in real +# installs (`substep` for status output, `_start_studio_venv_replacement` for +# the rollback-managed move). The tests run the block in isolation, so we +# stand in a minimal `mv`-based replacement that exercises the same observable +# effect (venv directory is no longer present at $VENV_DIR after a permitted +# cleanup) without dragging in install.sh's full rollback machinery. +_INSTALL_GUARD_STUBS = ( + "substep() { :; }\n" + "_start_studio_venv_replacement() {\n" + ' mv -- "$1" "$1.replaced"\n' + "}\n" +) + + +def _extract_install_sh_guard_block() -> str: + """Pull the `if [ -x "$VENV_DIR/bin/python" ]; then ... fi` block out + of install.sh as a self-contained snippet. Stops at the first elif so + the block can be paired with a synthetic else and run in isolation.""" + src = INSTALL_SH.read_text() + m = re.search( + r'(if \[ -x "\$VENV_DIR/bin/python" \]; then\n.*?)elif \[ "\$_STUDIO_HOME_REDIRECT" != "env"', + src, + re.DOTALL, + ) + assert m, "install.sh venv guard block not found" + return m.group(1) + "fi\n" + + +def _build_install_guard_script( + studio_home: Path, redirect: str, block: str | None = None +) -> str: + """Build a self-contained bash script that exercises the extracted + guard block. Includes stubs for substep / _start_studio_venv_replacement + so the snippet runs without install.sh's full rollback machinery.""" + if block is None: + block = _extract_install_sh_guard_block() + return ( + _INSTALL_GUARD_STUBS + + f'STUDIO_HOME="{studio_home}"\n' + + f'VENV_DIR="$STUDIO_HOME/unsloth_studio"\n' + + f'_STUDIO_HOME_REDIRECT="{redirect}"\n' + + block + + "echo RESULT=ok\n" + ) + + +def _run_install_guard( + studio_home: Path, + redirect: str, + create_share_conf: bool = False, + create_bin_shim: bool = False, + create_venv_marker: bool = False, +) -> subprocess.CompletedProcess: + venv_dir = studio_home / "unsloth_studio" + (venv_dir / "bin").mkdir(parents = True, exist_ok = True) + py = venv_dir / "bin" / "python" + py.write_text("#!/bin/sh\nexit 0\n") + py.chmod(0o755) + if create_share_conf: + (studio_home / "share").mkdir(parents = True, exist_ok = True) + (studio_home / "share" / "studio.conf").write_text("") + if create_bin_shim: + (studio_home / "bin").mkdir(parents = True, exist_ok = True) + (studio_home / "bin" / "unsloth").write_text("") + if create_venv_marker: + (venv_dir / ".unsloth-studio-owned").write_text("") + script = _build_install_guard_script(studio_home, redirect) + return subprocess.run( + ["bash", "-c", script], + env = {"PATH": "/usr/bin:/bin"}, + text = True, + capture_output = True, + ) + + +def test_env_mode_blocks_unsloth_studio_without_sentinels(tmp_path): + studio_home = tmp_path / "ws" + res = _run_install_guard(studio_home, redirect = "env") + assert res.returncode != 0, ( + "env-mode without sentinels must refuse to rm -rf $VENV_DIR; " + f"stdout={res.stdout!r} stderr={res.stderr!r}" + ) + assert "does not look like an Unsloth Studio install" in res.stderr + assert (studio_home / "unsloth_studio" / "bin" / "python").is_file() + + +def test_env_mode_passes_when_share_studio_conf_present(tmp_path): + studio_home = tmp_path / "ws" + res = _run_install_guard(studio_home, redirect = "env", create_share_conf = True) + assert res.returncode == 0, ( + f"share/studio.conf sentinel must allow cleanup;" + f" stdout={res.stdout!r} stderr={res.stderr!r}" + ) + assert "RESULT=ok" in res.stdout + assert not (studio_home / "unsloth_studio").exists() + + +def test_env_mode_passes_when_bin_unsloth_shim_present(tmp_path): + studio_home = tmp_path / "ws" + res = _run_install_guard(studio_home, redirect = "env", create_bin_shim = True) + assert res.returncode == 0, res.stderr + assert not (studio_home / "unsloth_studio").exists() + + +def test_default_mode_skips_sentinel_check(tmp_path): + studio_home = tmp_path / "ws" + res = _run_install_guard(studio_home, redirect = "default") + assert res.returncode == 0, res.stderr + assert "RESULT=ok" in res.stdout + assert not (studio_home / "unsloth_studio").exists() + + +def test_install_ps1_has_matching_env_mode_guard(): + src = INSTALL_PS1.read_text() + block_start = src.index("if (Test-Path -LiteralPath $VenvPython)") + block = src[block_start : block_start + 2000] + assert ( + "$StudioRedirectMode -eq 'env'" in block + ), "install.ps1 must gate Remove-Item $VenvDir on env-mode" + assert ( + "share\\studio.conf" in block + ), "install.ps1 guard must check share\\studio.conf sentinel" + assert ( + "bin\\unsloth.exe" in block + ), "install.ps1 guard must check bin\\unsloth.exe sentinel" + assert "Refusing to delete non-Studio venv" in block + + +def test_setup_ps1_has_writability_probe(): + src = SETUP_PS1.read_text() + idx = src.index("if (Test-Path -LiteralPath $_studioOverride -PathType Container)") + block = src[idx : idx + 2000] + assert ( + "WriteAllText" in block + ), "setup.ps1 must write-probe UNSLOTH_STUDIO_HOME like setup.sh:417" + assert ( + "is not writable" in block + ), "setup.ps1 probe failure must produce a clear writable-error message" + + +def test_env_mode_blocks_when_bin_unsloth_is_a_directory(tmp_path): + """A bare directory at $STUDIO_HOME/bin/unsloth must NOT pass the + sentinel. The previous `-e` test accepted any path type, allowing an + unrelated workspace with sibling content under unsloth_studio plus + a directory at bin/unsloth to be wiped.""" + studio_home = tmp_path / "ws" + venv = studio_home / "unsloth_studio" + (venv / "bin").mkdir(parents = True) + py = venv / "bin" / "python" + py.write_text("#!/bin/sh\nexit 0\n") + py.chmod(0o755) + (venv / "important.txt").write_text("keep me") + (studio_home / "bin" / "unsloth").mkdir(parents = True) + script = _build_install_guard_script(studio_home, "env") + res = subprocess.run( + ["bash", "-c", script], + env = {"PATH": "/usr/bin:/bin"}, + text = True, + capture_output = True, + ) + assert res.returncode != 0, ( + "directory at bin/unsloth must NOT satisfy the Studio sentinel; " + f"stdout={res.stdout!r} stderr={res.stderr!r}" + ) + assert (venv / "important.txt").is_file(), "unrelated workspace data must survive" + + +def test_env_mode_passes_when_bin_unsloth_is_a_symlink(tmp_path): + """A symlink at $STUDIO_HOME/bin/unsloth (real installer artefact) + must still satisfy the sentinel after the leaf-only tightening.""" + studio_home = tmp_path / "ws" + venv = studio_home / "unsloth_studio" + (venv / "bin").mkdir(parents = True) + py = venv / "bin" / "python" + py.write_text("#!/bin/sh\nexit 0\n") + py.chmod(0o755) + (studio_home / "bin").mkdir(parents = True) + target = studio_home / "bin" / "unsloth-real" + target.write_text("#!/bin/sh\nexit 0\n") + target.chmod(0o755) + (studio_home / "bin" / "unsloth").symlink_to(target) + script = _build_install_guard_script(studio_home, "env") + res = subprocess.run( + ["bash", "-c", script], + env = {"PATH": "/usr/bin:/bin"}, + text = True, + capture_output = True, + ) + assert res.returncode == 0, res.stderr + assert "RESULT=ok" in res.stdout + assert not venv.exists() + + +def test_install_ps1_sentinel_uses_pathtype_leaf(): + """The Test-Path checks that gate Remove-Item $VenvDir must use + -PathType Leaf so a directory at the sentinel path cannot satisfy them.""" + src = INSTALL_PS1.read_text() + block_start = src.index("if (Test-Path -LiteralPath $VenvPython)") + block = src[block_start : block_start + 2000] + assert ( + 'share\\studio.conf") -PathType Leaf' in block + ), "install.ps1 share\\studio.conf check must use -PathType Leaf" + assert ( + 'bin\\unsloth.exe") -PathType Leaf' in block + ), "install.ps1 bin\\unsloth.exe check must use -PathType Leaf" + + +def test_setup_ps1_stale_venv_has_env_mode_guard(): + """studio/setup.ps1 stale-venv rebuild branch must mirror install.ps1: + refuse to Remove-Item $VenvDir under custom-root mode unless the root + carries a Studio sentinel (in-VENV marker, share\\studio.conf, or + bin\\unsloth.exe leaf).""" + src = SETUP_PS1.read_text() + idx = src.index("Stale venv detected") + block = src[idx : idx + 1500] + assert ( + "$StudioHomeIsCustom" in block + ), "setup.ps1 stale-venv branch must gate on $StudioHomeIsCustom" + assert ( + 'share\\studio.conf") -PathType Leaf' in block + ), "setup.ps1 stale-venv guard must check share\\studio.conf with -PathType Leaf" + assert ( + 'bin\\unsloth.exe") -PathType Leaf' in block + ), "setup.ps1 stale-venv guard must check bin\\unsloth.exe with -PathType Leaf" + # The guard must fire BEFORE the destructive call. + guard_idx = block.index("$StudioHomeIsCustom") + rm_idx = block.index("Remove-Item -LiteralPath $VenvDir") + assert ( + guard_idx < rm_idx + ), "custom-root guard must precede Remove-Item -LiteralPath $VenvDir" + + +def test_setup_sh_prebuilt_llama_cpp_has_ownership_guard(): + """studio/setup.sh prebuilt llama.cpp path must call + _assert_studio_owned_or_absent before invoking install_llama_prebuilt.py + so an unrelated $UNSLOTH_STUDIO_HOME/llama.cpp is not displaced by + the helper's os.replace().""" + src = SETUP_SH.read_text() + idx = src.index("installing prebuilt llama.cpp...") + block = src[idx : idx + 2000] + assert ( + '_assert_studio_owned_or_absent "$LLAMA_CPP_DIR" "llama.cpp install"' in block + ), "setup.sh must guard the prebuilt llama.cpp path with the ownership marker" + guard_idx = block.index('_assert_studio_owned_or_absent "$LLAMA_CPP_DIR"') + # Anchor on the actual command-array entry, not the why-comment mention. + helper_idx = block.index('python "$SCRIPT_DIR/install_llama_prebuilt.py"') + assert ( + guard_idx < helper_idx + ), "ownership guard must precede the install_llama_prebuilt.py call" + + +def test_setup_ps1_prebuilt_llama_cpp_has_ownership_guard(): + """Mirror check for studio/setup.ps1: prebuilt llama.cpp path must + call Assert-StudioOwnedOrAbsent before invoking install_llama_prebuilt.py.""" + src = SETUP_PS1.read_text() + idx = src.index("installing prebuilt llama.cpp bundle (preferred path)") + block = src[idx : idx + 2000] + assert ( + 'Assert-StudioOwnedOrAbsent -Path $LlamaCppDir -Label "llama.cpp install"' + in block + ), "setup.ps1 must guard the prebuilt llama.cpp path with Assert-StudioOwnedOrAbsent" + guard_idx = block.index("Assert-StudioOwnedOrAbsent -Path $LlamaCppDir") + # Anchor on the actual command-array entry, not the why-comment mention. + helper_idx = block.index('"$PSScriptRoot\\install_llama_prebuilt.py"') + assert ( + guard_idx < helper_idx + ), "Assert-StudioOwnedOrAbsent must precede the install_llama_prebuilt.py call" + + +def test_env_mode_passes_when_venv_marker_present(tmp_path): + """install.sh env-mode guard must accept the in-VENV + .unsloth-studio-owned marker as a primary sentinel so a partial + install (uv venv created, sentinels not yet written) is recoverable + by re-running install.sh.""" + studio_home = tmp_path / "ws" + res = _run_install_guard(studio_home, redirect = "env", create_venv_marker = True) + assert res.returncode == 0, ( + f"in-VENV marker must allow cleanup; " + f"stdout={res.stdout!r} stderr={res.stderr!r}" + ) + assert "RESULT=ok" in res.stdout + assert not (studio_home / "unsloth_studio").exists() + + +def test_env_mode_blocks_when_bin_unsloth_is_symlink_to_directory(tmp_path): + """install.sh env-mode guard must NOT accept a symlink-to-directory at + bin/unsloth as a Studio sentinel. Iter1's standalone -L test let any + symlink (including symlinks to dirs and broken symlinks) bypass the + guard; iter2 dropped that test so only -f (file or symlink-to-file) + counts.""" + studio_home = tmp_path / "ws" + venv = studio_home / "unsloth_studio" + (venv / "bin").mkdir(parents = True) + py = venv / "bin" / "python" + py.write_text("#!/bin/sh\nexit 0\n") + py.chmod(0o755) + (venv / "important.txt").write_text("keep me") + (studio_home / "bin").mkdir(parents = True) + target_dir = studio_home / "bin" / "unsloth-target-dir" + target_dir.mkdir() + (studio_home / "bin" / "unsloth").symlink_to(target_dir) + script = _build_install_guard_script(studio_home, "env") + res = subprocess.run( + ["bash", "-c", script], + env = {"PATH": "/usr/bin:/bin"}, + text = True, + capture_output = True, + ) + assert res.returncode != 0, ( + "symlink-to-directory at bin/unsloth must NOT pass; " + f"stdout={res.stdout!r} stderr={res.stderr!r}" + ) + assert (venv / "important.txt").is_file(), "unrelated workspace data must survive" + + +def test_env_mode_blocks_when_bin_unsloth_is_broken_symlink(tmp_path): + """install.sh guard must reject a broken symlink at bin/unsloth.""" + studio_home = tmp_path / "ws" + venv = studio_home / "unsloth_studio" + (venv / "bin").mkdir(parents = True) + py = venv / "bin" / "python" + py.write_text("#!/bin/sh\nexit 0\n") + py.chmod(0o755) + (venv / "important.txt").write_text("keep me") + (studio_home / "bin").mkdir(parents = True) + (studio_home / "bin" / "unsloth").symlink_to(studio_home / "bin" / "does-not-exist") + script = _build_install_guard_script(studio_home, "env") + res = subprocess.run( + ["bash", "-c", script], + env = {"PATH": "/usr/bin:/bin"}, + text = True, + capture_output = True, + ) + assert res.returncode != 0, ( + "broken symlink at bin/unsloth must NOT pass; " + f"stdout={res.stdout!r} stderr={res.stderr!r}" + ) + assert (venv / "important.txt").is_file() + + +def test_install_sh_writes_venv_marker_after_uv_venv(): + """install.sh must write the .unsloth-studio-owned marker into + $VENV_DIR right after `uv venv` succeeds so the env-mode deletion + guard accepts it on the next install run.""" + src = INSTALL_SH.read_text() + create_idx = src.index('run_install_cmd "create venv" uv venv "$VENV_DIR"') + tail = src[create_idx : create_idx + 600] + assert ( + ".unsloth-studio-owned" in tail + ), "install.sh must write .unsloth-studio-owned after uv venv create" + + +def test_install_ps1_writes_venv_marker_after_uv_venv(): + """install.ps1 must write the .unsloth-studio-owned marker into + $VenvDir after `uv venv` succeeds.""" + src = INSTALL_PS1.read_text() + venv_create = src.index("uv venv $VenvDir --python") + tail = src[venv_create : venv_create + 1500] + assert ( + ".unsloth-studio-owned" in tail + ), "install.ps1 must write .unsloth-studio-owned after uv venv create" + + +def test_install_ps1_guard_accepts_venv_marker(): + """install.ps1 env-mode guard must accept the in-VENV + .unsloth-studio-owned marker as a primary sentinel.""" + src = INSTALL_PS1.read_text() + block_start = src.index("if (Test-Path -LiteralPath $VenvPython)") + block = src[block_start : block_start + 2000] + assert ( + '$VenvDir ".unsloth-studio-owned") -PathType Leaf' in block + ), "install.ps1 guard must check the in-VENV marker with -PathType Leaf" + + +def test_setup_helpers_gate_on_canonical_custom_root(): + """Both _assert_studio_owned_or_absent (setup.sh) and + Assert-StudioOwnedOrAbsent (setup.ps1) must gate on a canonical + custom-vs-legacy comparison so an explicit override that resolves + to the legacy default does not trip the guard for pre-PR T5 + sidecar venvs or llama.cpp dirs.""" + sh_src = SETUP_SH.read_text() + sh_idx = sh_src.index("_assert_studio_owned_or_absent() {") + sh_func = sh_src[sh_idx : sh_idx + 600] + assert ( + '"$_STUDIO_HOME_IS_CUSTOM" = true' in sh_func + ), "setup.sh _assert_studio_owned_or_absent must gate on _STUDIO_HOME_IS_CUSTOM" + assert ( + "_LEGACY_STUDIO_HOME=" in sh_src + and "_studio_home_canon=" in sh_src + and "_STUDIO_HOME_IS_CUSTOM=" in sh_src + ), "setup.sh must compute the canonical custom-root flag" + + ps_src = SETUP_PS1.read_text() + ps_idx = ps_src.index("function Assert-StudioOwnedOrAbsent") + ps_func = ps_src[ps_idx : ps_idx + 800] + assert ( + "$StudioHomeIsCustom -and" in ps_func + ), "setup.ps1 Assert-StudioOwnedOrAbsent must gate on $StudioHomeIsCustom" + assert ( + "$StudioOwnedMarker) -PathType Leaf" in ps_func + ), "setup.ps1 marker check must use -PathType Leaf so a directory cannot satisfy it" + + +def test_setup_ps1_inplace_git_sync_marks_studio_owned(): + """setup.ps1 in-place git-sync branch (when $LlamaCppDir/.git exists) + must call Mark-StudioOwned after a successful sync so a later prebuilt + update path's Assert-StudioOwnedOrAbsent does not exit.""" + src = SETUP_PS1.read_text() + inplace_idx = src.index('Test-Path -LiteralPath (Join-Path $LlamaCppDir ".git")') + # The in-place branch ends just before the temp-dir clone branch. + clone_idx = src.index("Cloning llama.cpp @", inplace_idx) + inplace_block = src[inplace_idx:clone_idx] + assert ( + "Mark-StudioOwned -Path $LlamaCppDir" in inplace_block + ), "in-place git-sync branch must call Mark-StudioOwned on success" + assert ( + "$StudioHomeIsCustom" in inplace_block + ), "in-place Mark-StudioOwned call should be gated on $StudioHomeIsCustom" + + +def test_setup_ps1_inplace_git_sync_asserts_studio_owned_before_mutation(): + """setup.ps1 in-place git-sync branch must call Assert-StudioOwnedOrAbsent + BEFORE any destructive git operation (remote set-url, checkout -B, clean + -fdx). Asymmetric to the prebuilt path and the temp-dir-swap path which + both guard.""" + src = SETUP_PS1.read_text() + inplace_idx = src.index('Test-Path -LiteralPath (Join-Path $LlamaCppDir ".git")') + clone_idx = src.index("Cloning llama.cpp @", inplace_idx) + inplace_block = src[inplace_idx:clone_idx] + assert ( + "Assert-StudioOwnedOrAbsent -Path $LlamaCppDir" in inplace_block + ), "in-place git-sync must Assert-StudioOwnedOrAbsent before mutating $LlamaCppDir" + guard_idx = inplace_block.index("Assert-StudioOwnedOrAbsent -Path $LlamaCppDir") + git_idx = inplace_block.index("git -C $LlamaCppDir remote set-url") + assert ( + guard_idx < git_idx + ), "Assert-StudioOwnedOrAbsent must precede the first git mutation" + + +def _extract_check_health_function() -> str: + src = INSTALL_SH.read_text() + fn_start = src.index("_check_health() {") + fn_end = src.index("\n}\n", fn_start) + 2 + return src[fn_start:fn_end] + + +def _run_check_health(expected_root_id: str, response_json: str) -> int: + fn = _extract_check_health_function() + script = ( + f"_EXPECTED_STUDIO_ROOT_ID={expected_root_id!r}\n" + "_http_get() { printf '%s' \"$1\"; }\n" + + fn.replace( + '_resp=$(_http_get "http://127.0.0.1:$_port/api/health") || return 1', + f"_resp={response_json!r}", + ) + + "\n_check_health 8888\n" + "echo rc=$?\n" + ) + res = subprocess.run( + ["bash", "-c", script], + env = {"PATH": "/usr/bin:/bin"}, + text = True, + capture_output = True, + ) + rc_lines = [l for l in res.stdout.splitlines() if l.startswith("rc=")] + return int(rc_lines[0].split("=")[1]) if rc_lines else res.returncode + + +def test_check_health_accepts_matching_studio_root_id(): + """Hex digest baked at install time matches the backend's + /api/health studio_root_id -- launcher attaches to its own backend.""" + expected_id = "a" * 64 + rc = _run_check_health( + expected_id, + f'{{"status":"healthy","service":"Unsloth UI Backend","studio_root_id":"{expected_id}"}}', + ) + assert rc == 0, f"matching studio_root_id must allow attach (rc={rc})" + + +def test_check_health_rejects_mismatched_studio_root_id(): + """Different install root → different sha256 → reject. Workspace + isolation: launcher A must not open Studio B running on the same port.""" + expected_id = "a" * 64 + other_id = "b" * 64 + rc = _run_check_health( + expected_id, + f'{{"status":"healthy","service":"Unsloth UI Backend","studio_root_id":"{other_id}"}}', + ) + assert rc != 0, "mismatched studio_root_id must reject attach (workspace isolation)" + + +def test_check_health_rejects_missing_studio_root_id_field(): + """A backend that omits studio_root_id (older or non-conforming) must + not be attached to when an expected id is baked into the launcher.""" + expected_id = "a" * 64 + rc = _run_check_health( + expected_id, + '{"status":"healthy","service":"Unsloth UI Backend"}', + ) + assert rc != 0, "missing studio_root_id field must reject attach" + + +def test_check_health_no_baked_id_accepts_any_healthy_backend(): + """If _EXPECTED_STUDIO_ROOT_ID is empty (e.g. install-time hash failed + to compute), the launcher falls back to the legacy contract and accepts + any healthy Unsloth backend.""" + rc = _run_check_health( + "", + '{"status":"healthy","service":"Unsloth UI Backend","studio_root_id":"deadbeef"}', + ) + assert rc == 0, "no baked id → accept any healthy Unsloth backend" + + +def test_check_health_rejects_non_unsloth_service(): + rc = _run_check_health( + "", + '{"status":"healthy","service":"Other UI Backend"}', + ) + assert rc != 0, "non-Unsloth service must be rejected" + + +def test_check_health_handles_arbitrary_id_token(): + """Iter3 used a raw shell match against the JSON-escaped studio_root, + which failed for paths containing `\\` or `"` (FastAPI emits `\\\\` and + `\\\"`). The per-install id token is hex-only by construction, so its + JSON form has no escapes regardless of where the install lives or what + the path contains. This test pins the round-trip on a fully arbitrary + 64-char hex token.""" + expected_id = "f0" + ("ed" * 31) # 64 hex chars, not derived from any path + rc = _run_check_health( + expected_id, + f'{{"status":"healthy","service":"Unsloth UI Backend","studio_root_id":"{expected_id}"}}', + ) + assert ( + rc == 0 + ), "arbitrary 64-hex install id must round-trip cleanly (no JSON escape issue)" + + +def test_install_ps1_test_studio_health_verifies_studio_root_id(): + """install.ps1 Test-StudioHealth must compare studio_root_id against + the install-time-baked $_ExpectedStudioRootId, not the runtime env var.""" + src = INSTALL_PS1.read_text() + fn_start = src.index("function Test-StudioHealth") + fn_end = src.index("\n}\n", fn_start) + 2 + fn = src[fn_start:fn_end] + assert ( + "studio_root_id" in fn + ), "Test-StudioHealth must inspect the studio_root_id field" + assert ( + "$_ExpectedStudioRootId" in fn + ), "Test-StudioHealth must compare against the install-time baked $_ExpectedStudioRootId" + + +def test_install_ps1_bakes_studio_root_id_into_launcher(): + """install.ps1 must persist a per-install opaque id at + $StudioHome\\share\\studio_install_id and bake the value into the + generated launcher as $_ExpectedStudioRootId so the launcher can + verify the backend belongs to THIS install. The id is generated + via a CSPRNG so /api/health does not leak the install path.""" + src = INSTALL_PS1.read_text() + assert ( + "$_studioRootId" in src + ), "install.ps1 must compute $_studioRootId for the launcher" + assert ( + '"share"' in src and "studio_install_id" in src + ), "install.ps1 must persist the id at $StudioHome\\share\\studio_install_id" + assert ( + "RandomNumberGenerator" in src + ), "install.ps1 must seed the id from a CSPRNG (RandomNumberGenerator)" + assert ( + "$_ExpectedStudioRootId" in src + ), "install.ps1 must bake $_ExpectedStudioRootId into the launcher" + + +def test_health_endpoint_exposes_studio_root_id_not_raw_path(): + """studio/backend/main.py /api/health must expose studio_root_id (a + hex digest) and NOT the raw studio_root path. Studio supports + `-H 0.0.0.0`; an unauthenticated /api/health that returns the raw + install path leaks username, home dir, workspace name, etc.""" + main_py = REPO_ROOT / "studio" / "backend" / "main.py" + src = main_py.read_text() + health_idx = src.index('@app.get("/api/health")') + health_block = src[health_idx : health_idx + 1500] + assert ( + '"studio_root_id"' in health_block + ), "/api/health must expose studio_root_id (hex digest)" + assert ( + '"studio_root":' not in health_block + ), "/api/health must NOT expose the raw studio_root path (information disclosure)" + assert ( + "_studio_root_id()" in health_block + ), "/api/health must call the _studio_root_id helper" + + +def test_install_sh_bakes_studio_root_id_into_launcher(): + """install.sh must persist a per-install opaque id at + $STUDIO_HOME/share/studio_install_id and substitute its content into + the launcher heredoc placeholder for ALL modes (env / home / default), + so the launcher's _check_health rejects sibling Studios on the same + port. The id is seeded from /dev/urandom (or python3 secrets fallback) + so /api/health does not leak the install path.""" + src = INSTALL_SH.read_text() + assert ( + "_css_studio_root_id" in src + ), "install.sh must compute _css_studio_root_id for the launcher" + assert ( + '_css_id_file="$_css_id_dir/studio_install_id"' in src + ), "install.sh must persist the id at $STUDIO_HOME/share/studio_install_id" + assert ( + "od -An -N32 -tx1 /dev/urandom" in src + ), "install.sh must seed new ids from /dev/urandom (CSPRNG)" + assert ( + "@@STUDIO_ROOT_ID@@" in src + ), "install.sh must use @@STUDIO_ROOT_ID@@ placeholder in the launcher heredoc" + assert ( + "s|@@STUDIO_ROOT_ID@@|$_css_studio_root_id|g" in src + ), "install.sh must sed-substitute @@STUDIO_ROOT_ID@@ unconditionally (not just env-mode)" + + +def test_tauri_preflight_scrubs_studio_home_env(): + """All three Tauri CLI-spawn sites that lacked the scrub must now + env_remove UNSLOTH_STUDIO_HOME and STUDIO_HOME, mirroring + process.rs / install.rs / desktop_auth.rs / update.rs.""" + preflight = ( + REPO_ROOT / "studio" / "src-tauri" / "src" / "preflight.rs" + ).read_text() + commands = (REPO_ROOT / "studio" / "src-tauri" / "src" / "commands.rs").read_text() + # Both functions in preflight.rs (run_cli_probe + probe_cli_capability) + # must scrub. Count occurrences -- expect 2 in preflight, 1 in commands. + assert ( + preflight.count('cmd.env_remove("UNSLOTH_STUDIO_HOME")') >= 2 + ), "preflight.rs must scrub UNSLOTH_STUDIO_HOME in both run_cli_probe and probe_cli_capability" + assert ( + preflight.count('cmd.env_remove("STUDIO_HOME")') >= 2 + ), "preflight.rs must scrub STUDIO_HOME in both run_cli_probe and probe_cli_capability" + assert ( + 'cmd.env_remove("UNSLOTH_STUDIO_HOME")' in commands + ), "commands.rs check_install_status must scrub UNSLOTH_STUDIO_HOME" + assert ( + 'cmd.env_remove("STUDIO_HOME")' in commands + ), "commands.rs check_install_status must scrub STUDIO_HOME" + + +def test_install_sh_shim_uses_atomic_replace(): + """install.sh shim install must use ln -sfn for atomic replace; the + older `rm -f ...; ln -s ...` left a window where the shim was missing.""" + src = INSTALL_SH.read_text() + shim_idx = src.index('_shim_path="$_LOCAL_BIN/unsloth"') + block = src[shim_idx : shim_idx + 1500] + assert ( + 'ln -sfn "$VENV_DIR/bin/unsloth" "$_shim_path"' in block + ), "install.sh must use ln -sfn for atomic shim replacement" + assert ( + 'rm -f -- "$_shim_path"' not in block + ), "the explicit rm + ln pair must be replaced by atomic ln -sfn" + + +def test_install_sh_create_shortcuts_seeds_id_from_csprng_with_python_fallback( + tmp_path, +): + """_create_shortcuts must seed new ids from /dev/urandom first (no + interpreter spawn cost on the install hot path) and fall back to + `python3 -c 'secrets.token_hex(32)'` only when urandom is unreadable. + Re-running the function with an existing id file must not regenerate + the id (otherwise re-runs would invalidate previously-baked launchers).""" + src = INSTALL_SH.read_text() + fn_start = src.index('_css_data_dir="$DATA_DIR"') + block = src[fn_start : fn_start + 3000] + urandom_idx = block.index("od -An -N32 -tx1 /dev/urandom") + py_fallback_idx = block.index("python3 -c 'import secrets;", urandom_idx) + assert ( + urandom_idx < py_fallback_idx + ), "/dev/urandom must be tried before the python3 secrets fallback" + # The id file is checked for non-empty content before we generate; this is + # what makes re-runs idempotent. + assert ( + 'if [ ! -s "$_css_id_file" ]; then' in block + ), "install.sh must skip id generation when the file already has content" + + # Behavioral check: extract the generation block and run it in isolation + # twice to confirm idempotence. + studio_home = tmp_path / "studio" + (studio_home / "share").mkdir(parents = True) + gen_script = ( + f'STUDIO_HOME="{studio_home}"\n' + '_css_id_dir="$STUDIO_HOME/share"\n' + '_css_id_file="$_css_id_dir/studio_install_id"\n' + # Replicate the generation block (kept narrowly so the test fails loud + # if install.sh changes the surrounding contract). + "gen() {\n" + ' if [ ! -s "$_css_id_file" ]; then\n' + ' _css_new_id=$(od -An -N32 -tx1 /dev/urandom 2>/dev/null | tr -d " \\n")\n' + ' printf "%s" "$_css_new_id" > "$_css_id_file.$$.tmp"\n' + ' mv "$_css_id_file.$$.tmp" "$_css_id_file"\n' + " fi\n" + ' cat "$_css_id_file"\n' + "}\n" + "a=$(gen); b=$(gen)\n" + '[ "$a" = "$b" ] || { echo MISMATCH; exit 1; }\n' + 'echo "ID=$a"\n' + 'echo "LEN=${#a}"\n' + ) + res = subprocess.run(["bash", "-c", gen_script], text = True, capture_output = True) + assert res.returncode == 0, res.stderr + out = dict( + line.split("=", 1) for line in res.stdout.strip().splitlines() if "=" in line + ) + assert ( + out.get("LEN") == "64" + ), f"id must be 64 hex chars, got LEN={out.get('LEN')!r}" + assert all( + c in "0123456789abcdef" for c in out.get("ID", "") + ), f"id must be lowercase hex, got {out.get('ID')!r}" + + +def test_install_sh_create_shortcuts_fails_fast_when_no_entropy(): + """If neither /dev/urandom nor python3 is available, _create_shortcuts + must `return 1` instead of silently baking an empty studio_root_id + (which would disable the launcher's same-install discriminator).""" + src = INSTALL_SH.read_text() + fn_start = src.index('_css_data_dir="$DATA_DIR"') + block = src[fn_start : fn_start + 3000] + assert ( + "[WARN] Cannot create launcher: no entropy source for studio_install_id" + in block + ), "install.sh must warn when neither urandom nor python3 is available" + assert ( + "[WARN] Cannot create launcher: failed to read" in block + ), "install.sh must warn when the id file read produces no content" + assert ( + block.count("return 1") >= 2 + ), "both the no-entropy branch and the empty-read branch must `return 1`" + + +def test_install_sh_bakes_installed_is_env_mode_flag_in_launcher(): + """install.sh must bake the install-time mode (env vs default/home) into + the generated launcher so PORT_FILE / namespaced LOCK_DIR cannot be + flipped on by a sourced custom-root studio.conf in the user's shell.""" + src = INSTALL_SH.read_text() + assert ( + "_INSTALLED_IS_ENV_MODE='@@INSTALLED_IS_ENV_MODE@@'" in src + ), "launcher heredoc must declare _INSTALLED_IS_ENV_MODE='@@INSTALLED_IS_ENV_MODE@@'" + assert ( + "_css_is_env_mode=false" in src + ), "install.sh must default _css_is_env_mode to false" + assert ( + '[ "$_STUDIO_HOME_REDIRECT" = "env" ] && _css_is_env_mode=true' in src + ), "install.sh must set _css_is_env_mode=true only when _STUDIO_HOME_REDIRECT=env" + assert ( + "s|@@INSTALLED_IS_ENV_MODE@@|$_css_is_env_mode|g" in src + ), "install.sh sed pipeline must substitute @@INSTALLED_IS_ENV_MODE@@" + + +def test_install_sh_launcher_gates_port_file_on_baked_flag_not_runtime_env(): + """The launcher's PORT_FILE / namespaced LOCK_DIR must be gated on the + baked $_INSTALLED_IS_ENV_MODE flag, not the runtime $UNSLOTH_STUDIO_HOME. + Sourcing a custom-root studio.conf in shell must not flip a default-mode + launcher into env-mode behavior.""" + src = INSTALL_SH.read_text() + heredoc_start = src.index("cat > \"$_css_launcher\" << 'LAUNCHER_EOF'") + heredoc_end = src.index("LAUNCHER_EOF\n", heredoc_start) + heredoc = src[heredoc_start:heredoc_end] + assert ( + 'if [ "$_INSTALLED_IS_ENV_MODE" = "true" ]; then' in heredoc + ), "launcher must gate PORT_FILE/LOCK_DIR on baked _INSTALLED_IS_ENV_MODE" + port_block_start = heredoc.index('if [ "$_INSTALLED_IS_ENV_MODE" = "true" ]; then') + port_block_end = heredoc.index("\nfi\n", port_block_start) + len("\nfi\n") + port_block = heredoc[port_block_start:port_block_end] + assert 'PORT_FILE="$DATA_DIR/studio.port"' in port_block + assert ( + 'if [ -n "${UNSLOTH_STUDIO_HOME:-}" ]; then\n if command -v cksum' + not in heredoc + ), "launcher must NOT gate PORT_FILE on runtime UNSLOTH_STUDIO_HOME" + + def _run_launcher_gate(installed_flag: str, runtime_env: dict) -> str: + # Reproduce just the LOCK_DIR/PORT_FILE init block in isolation. + script = ( + f"_INSTALLED_IS_ENV_MODE={installed_flag!r}\n" + "DATA_DIR=/tmp/test_data_dir\n" + 'LOCK_DIR="${XDG_RUNTIME_DIR:-/tmp}/unsloth-studio-launcher-$(id -u).lock"\n' + 'PORT_FILE=""\n' + port_block + '\necho "PORT_FILE=$PORT_FILE"\n' + ) + env = {"PATH": "/usr/bin:/bin"} + env.update(runtime_env) + res = subprocess.run( + ["bash", "-c", script], + text = True, + capture_output = True, + env = env, + ) + for line in res.stdout.splitlines(): + if line.startswith("PORT_FILE="): + return line[len("PORT_FILE=") :] + return "" + + # default-mode install should NEVER set PORT_FILE, even if UNSLOTH_STUDIO_HOME leaks in. + assert ( + _run_launcher_gate("false", {"UNSLOTH_STUDIO_HOME": "/tmp/leaked"}) == "" + ), "default-mode launcher must keep PORT_FILE empty even with UNSLOTH_STUDIO_HOME in env" + # env-mode install should set PORT_FILE regardless of runtime env. + assert ( + _run_launcher_gate("true", {}) == "/tmp/test_data_dir/studio.port" + ), "env-mode launcher must set PORT_FILE based on baked DATA_DIR" + + +def test_main_py_studio_root_id_caches_at_module_load(): + """_studio_root_id() is called on every /api/health poll; the id is + stable for the lifetime of the process so it must be read once at + module load and re-used (avoids a hot-path filesystem probe and + protects against transient FS errors during health polling).""" + main_py = (REPO_ROOT / "studio" / "backend" / "main.py").read_text() + assert ( + "_STUDIO_ROOT_ID_CACHE: str = _read_studio_install_id()" in main_py + ), "main.py must populate _STUDIO_ROOT_ID_CACHE from _read_studio_install_id() at module load" + fn_idx = main_py.index("def _studio_root_id() -> str:") + next_def_idx = main_py.index("\ndef ", fn_idx + 1) + fn_block = main_py[fn_idx:next_def_idx] + assert ( + "return _STUDIO_ROOT_ID_CACHE" in fn_block + ), "_studio_root_id() body must return the cached value" + assert ( + "read_text(" not in fn_block and "hashlib" not in fn_block + ), "_studio_root_id() must NOT do filesystem or hash work on every call" + + +def test_main_py_read_studio_install_id_validates_hex_and_handles_missing( + tmp_path, monkeypatch +): + """_read_studio_install_id reads $STUDIO_HOME/share/studio_install_id and + returns "" when the file is absent, empty, contains non-hex content, or + is the wrong length. "" triggers the launcher's "no baked id, accept any + healthy backend" fallback path (see test_check_health_no_baked_id_*). + Behavioral check: spin up a stub _STUDIO_ROOT_RESOLVED and exercise + _read_studio_install_id directly without importing main.py (which + pulls in heavy deps). Test the rejection rules verbatim.""" + import re + + pattern = re.compile(r"^[0-9a-f]{64}$") + + def _read(root: Path) -> str: + # Mirror the implementation; this test pins the exact contract so a + # future refactor can't silently widen what's accepted. + try: + token = (root / "share" / "studio_install_id").read_text().strip() + except (OSError, ValueError): + return "" + return token if pattern.fullmatch(token) else "" + + root = tmp_path / "studio" + (root / "share").mkdir(parents = True) + + # Missing file -> empty + assert _read(root) == "" + + id_file = root / "share" / "studio_install_id" + # Empty file -> empty + id_file.write_text("") + assert _read(root) == "" + # Non-hex content -> empty + id_file.write_text( + "not-a-hex-id-just-text-padded-to-64-chars-zzzzzzzzzzzzzzzzzzzzzz" + ) + assert _read(root) == "" + # Uppercase hex -> empty (must be lowercase) + id_file.write_text("F" * 64) + assert _read(root) == "" + # Wrong length -> empty (32 chars, not 64) + id_file.write_text("a" * 32) + assert _read(root) == "" + # Valid 64-char lowercase hex with surrounding whitespace -> stripped+accepted + valid = "0123456789abcdef" * 4 + id_file.write_text(f"\n {valid} \n") + assert _read(root) == valid + + +def test_llama_cpp_search_roots_handles_studio_root_oserror(): + """_find_llama_server_binary calls studio_root() which can raise + OSError or ValueError from Path.expanduser().resolve() (broken symlink, + null byte). The except clause must mirror sibling _kill_orphaned_servers + (which catches the same trio) so inference startup does not crash.""" + llama_cpp = ( + REPO_ROOT / "studio" / "backend" / "core" / "inference" / "llama_cpp.py" + ).read_text() + find_block_start = llama_cpp.index("_find_llama_server_binary") + find_block = llama_cpp[find_block_start : find_block_start + 4000] + assert ( + "except (ImportError, OSError, ValueError):" in find_block + ), "_find_llama_server_binary must catch (ImportError, OSError, ValueError) from studio_root()" + kill_def_idx = llama_cpp.index("def _kill_orphaned_servers") + kill_block = llama_cpp[kill_def_idx : kill_def_idx + 4000] + assert ( + "except (ImportError, OSError, ValueError):" in kill_block + ), "sibling _kill_orphaned_servers must keep its (ImportError, OSError, ValueError) handler" + + +def test_install_sh_install_id_survives_symlinked_studio_home(tmp_path): + """End-to-end behavioral check: when $STUDIO_HOME is reached via a + symlinked parent (e.g. symlinked $HOME on Linux, junctioned %USERPROFILE% + on Windows), install.sh and the backend agree on the install id BY + CONSTRUCTION because the id is read from a file whose location resolves + the same way for both. The previous sha256(canonical_path) scheme + required `cd -P/pwd -P` and Path.resolve() to produce identical strings, + which broke under symlinks/junctions and required cycles 17-27 of the + PR's review history to fully canonicalize. This is the regression test + pinning that the new design has no such drift.""" + real = tmp_path / "realhome" + real.mkdir() + link = tmp_path / "linkhome" + link.symlink_to(real) + studio_home = real / ".unsloth" / "studio" + (studio_home / "share").mkdir(parents = True) + # Write a stub install id at the canonical location. + valid_id = "ab12" * 16 + (studio_home / "share" / "studio_install_id").write_text(valid_id) + # Read it back via both the canonical and the symlinked path; both must + # see the SAME content (which is what makes install.sh's cat and the + # backend's read_text agree without any canonicalization dance). + raw_via_link = link / ".unsloth" / "studio" / "share" / "studio_install_id" + raw_direct = studio_home / "share" / "studio_install_id" + assert raw_via_link.read_text() == valid_id + assert raw_direct.read_text() == valid_id + # And install.sh's `cat` would see the same. + import subprocess as _sp + + res = _sp.run(["cat", str(raw_via_link)], capture_output = True, text = True) + assert res.returncode == 0 + assert res.stdout == valid_id + + +def test_install_sh_substitutes_root_id_before_data_dir(): + """The two-stage sed substitution must bake @@STUDIO_ROOT_ID@@ / + @@INSTALLED_IS_ENV_MODE@@ first (non-user-controlled), then @@DATA_DIR@@ + (user-controlled). A custom $DATA_DIR containing the literal text + @@STUDIO_ROOT_ID@@ must not be mutated by the global root-id sed pass.""" + src = INSTALL_SH.read_text() + root_id_idx = src.index("s|@@STUDIO_ROOT_ID@@|$_css_studio_root_id|g") + env_mode_idx = src.index("s|@@INSTALLED_IS_ENV_MODE@@|$_css_is_env_mode|g") + data_dir_idx = src.index("s|@@DATA_DIR@@|$_sed_safe|g") + assert root_id_idx < data_dir_idx, ( + "@@STUDIO_ROOT_ID@@ substitution must happen BEFORE @@DATA_DIR@@ " + "(non-user-controlled placeholders first)" + ) + assert ( + env_mode_idx < data_dir_idx + ), "@@INSTALLED_IS_ENV_MODE@@ substitution must happen BEFORE @@DATA_DIR@@" + + +def test_install_sh_root_id_pass_does_not_mutate_user_data_dir(tmp_path): + """Behavioral subprocess test: a $DATA_DIR containing the literal text + `@@STUDIO_ROOT_ID@@` must not be mutated when the placeholder pass runs + first; only the actual placeholder occurrences in the launcher template + are replaced.""" + src = INSTALL_SH.read_text() + heredoc_start = src.index("cat > \"$_css_launcher\" << 'LAUNCHER_EOF'") + heredoc_body_start = src.index("\n", heredoc_start) + 1 + heredoc_body_end = src.index("LAUNCHER_EOF\n", heredoc_start) + template = src[heredoc_body_start:heredoc_body_end] + launcher_path = tmp_path / "launch.sh" + launcher_path.write_text(template) + # Run the iter6 sed order: root-id first, then data-dir. + weird_data_dir = "/tmp/with-@@STUDIO_ROOT_ID@@/share" + root_id = "deadbeef" * 8 + is_env = "true" + script = f""" +sed -e "s|@@STUDIO_ROOT_ID@@|{root_id}|g" \\ + -e "s|@@INSTALLED_IS_ENV_MODE@@|{is_env}|g" \\ + "{launcher_path}" > "{launcher_path}.tmp" && mv "{launcher_path}.tmp" "{launcher_path}" +_sq_escaped=$(printf '%s' "{weird_data_dir}" | sed "s/'/'\\\\\\\\''/g") +_sed_safe=$(printf '%s' "$_sq_escaped" | sed 's/[\\\\&|]/\\\\&/g') +sed "s|@@DATA_DIR@@|$_sed_safe|g" "{launcher_path}" > "{launcher_path}.tmp" \\ + && mv "{launcher_path}.tmp" "{launcher_path}" +""" + subprocess.run(["bash", "-c", script], check = True) + final = launcher_path.read_text() + assert ( + f"DATA_DIR='{weird_data_dir}'" in final + ), f"DATA_DIR must be preserved verbatim (no @@STUDIO_ROOT_ID@@ mutation); got: {final[:500]}" + assert ( + f"_EXPECTED_STUDIO_ROOT_ID='{root_id}'" in final + ), "STUDIO_ROOT_ID placeholder must still be substituted in the launcher heredoc" + + +def test_install_ps1_install_id_file_layout_matches_backend_read_path(): + """install.ps1 must write the id at $StudioHome\\share\\studio_install_id + so the backend (studio/backend/main.py:_read_studio_install_id) can find + it via _STUDIO_ROOT_RESOLVED / "share" / "studio_install_id" without + mode-specific path knowledge. Persistence-across-runs is enforced by the + pre-write Test-Path check.""" + src = INSTALL_PS1.read_text() + id_idx = src.index('$_studioIdDir = Join-Path $StudioHome "share"') + context = src[id_idx : id_idx + 1500] + assert ( + '$_studioIdFile = Join-Path $_studioIdDir "studio_install_id"' in context + ), "install.ps1 must persist the id at $StudioHome\\share\\studio_install_id" + assert ( + "Test-Path -LiteralPath $_studioIdFile" in context + ), "install.ps1 must skip id generation when the file already has content (re-run idempotence)" + assert ( + "RandomNumberGenerator" in context and "GetBytes($_idBytes)" in context + ), "install.ps1 must seed new ids from a CSPRNG (RandomNumberGenerator)" + assert ( + "Move-Item -LiteralPath $_idTmp" in context + ), "install.ps1 must atomic-rename the temp file into place to avoid half-written ids" diff --git a/tests/test_studio_root_resilience.py b/tests/test_studio_root_resilience.py new file mode 100644 index 0000000000..1ce0430dc4 --- /dev/null +++ b/tests/test_studio_root_resilience.py @@ -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"] diff --git a/unsloth/__init__.py b/unsloth/__init__.py index 9db9ae0a32..9b620a5c76 100644 --- a/unsloth/__init__.py +++ b/unsloth/__init__.py @@ -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__ diff --git a/unsloth/_gpu_init.py b/unsloth/_gpu_init.py new file mode 100644 index 0000000000..2fc4bfde3c --- /dev/null +++ b/unsloth/_gpu_init.py @@ -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() diff --git a/unsloth/kernels/utils.py b/unsloth/kernels/utils.py index 09b03a597b..dd5a9cbf0e 100644 --- a/unsloth/kernels/utils.py +++ b/unsloth/kernels/utils.py @@ -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 diff --git a/unsloth/trainer.py b/unsloth/trainer.py index eea985e958..01b8822bd5 100644 --- a/unsloth/trainer.py +++ b/unsloth/trainer.py @@ -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, diff --git a/unsloth_cli/commands/studio.py b/unsloth_cli/commands/studio.py index 140940209f..76aac3dc15 100644 --- a/unsloth_cli/commands/studio.py +++ b/unsloth_cli/commands/studio.py @@ -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 /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