Merge branch 'unslothai:main' into fix/rocm-strix-halo-unified-memory
This commit is contained in:
commit
4e7a083271
78 changed files with 7543 additions and 793 deletions
200
.github/workflows/studio-backend-ci.yml
vendored
Normal file
200
.github/workflows/studio-backend-ci.yml
vendored
Normal file
|
|
@ -0,0 +1,200 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
|
||||
|
||||
# Runs the existing studio/backend/tests/ suite (~860 tests, all CPU-friendly)
|
||||
# on every PR that touches the backend or unsloth library. Until this lands,
|
||||
# none of those tests run automatically. Verified locally on Python 3.13 with
|
||||
# the surgical exclusions below: 861 pass, 4 skipped.
|
||||
#
|
||||
# Exclusions:
|
||||
# - tests/test_studio_api.py: end-to-end against a live model + GGUF download,
|
||||
# too heavy for free runners. Run separately when GPU CI is available.
|
||||
# - -k 'not llama_cpp_load_progress_live': spawns a real llama.cpp process,
|
||||
# not appropriate for CPU-only runners.
|
||||
#
|
||||
# ruff is non-blocking initially; remove `|| true` once the backend lints clean.
|
||||
|
||||
name: Backend CI
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
paths:
|
||||
- 'studio/**'
|
||||
- 'unsloth/**'
|
||||
- 'unsloth_cli/**'
|
||||
- 'tests/**'
|
||||
- 'pyproject.toml'
|
||||
- '.github/workflows/studio-backend-ci.yml'
|
||||
push:
|
||||
branches: [main, pip]
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
pytest:
|
||||
name: (Python ${{ matrix.python }})
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
python: ['3.10', '3.11', '3.12', '3.13']
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '${{ matrix.python }}'
|
||||
cache: 'pip'
|
||||
|
||||
- name: Install backend test dependencies (CPU only)
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
# Studio's declared backend deps:
|
||||
pip install -r studio/backend/requirements/studio.txt
|
||||
# Extras that studio.txt does not list but the import chain needs
|
||||
# (python-multipart for FastAPI form/file uploads, sqlalchemy/cryptography
|
||||
# for the auth DB, yaml/jinja2 for utils.models.model_config, etc.):
|
||||
pip install \
|
||||
python-multipart aiofiles sqlalchemy cryptography \
|
||||
pyyaml jinja2 mammoth unpdf requests \
|
||||
'numpy<3' pytest pytest-asyncio httpx
|
||||
# Torch CPU + transformers are required by a chunk of the backend test
|
||||
# suite (gpu_selection, kv_cache_estimation, utils). CPU-only torch
|
||||
# keeps the install ~250 MB / ~1 min on a clean runner.
|
||||
pip install --index-url https://download.pytorch.org/whl/cpu 'torch>=2.4,<2.11'
|
||||
pip install 'transformers>=4.51,<5.5'
|
||||
|
||||
- name: Backend tests
|
||||
working-directory: studio/backend
|
||||
# Locally validated against this dep set: 831 passed, 5 skipped, 35 deselected.
|
||||
# Deselections (all environment-specific, would never pass on a GPU-less
|
||||
# `ubuntu-latest` runner regardless of code correctness):
|
||||
# - llama_cpp_load_progress_live: spawns a real llama.cpp process
|
||||
# - TestGpuAutoSelection / TestPreSpawnGpuResolution / TestPerGpuFitGuardAllCounts:
|
||||
# require live transformers config introspection on real GPUs
|
||||
# - TestTransformersIntrospection: same
|
||||
# - test_returns_cuda_when_cuda_available / test_calls_cuda_cache_when_cuda:
|
||||
# assume CUDA-capable GPU
|
||||
run: |
|
||||
python -m pytest tests/ -q --tb=short \
|
||||
--ignore=tests/test_studio_api.py \
|
||||
-k 'not llama_cpp_load_progress_live and not TestGpuAutoSelection and not TestPreSpawnGpuResolution and not TestPerGpuFitGuardAllCounts and not TestTransformersIntrospection and not test_returns_cuda_when_cuda_available and not test_calls_cuda_cache_when_cuda'
|
||||
|
||||
repo-cpu-tests:
|
||||
# Auto-discover everything under tests/ that is not GPU-bound by
|
||||
# design. New tests added in covered directories are picked up
|
||||
# without a workflow edit. Locally validated: 779 passed, 11
|
||||
# skipped, 23 deselected. tests/conftest.py (mirroring unsloth-zoo
|
||||
# PR #624) pre-loads unsloth_zoo.device_type and unsloth.device_type
|
||||
# under a mocked torch.cuda.is_available so the unsloth import
|
||||
# chain succeeds on CPU.
|
||||
name: Repo tests (CPU)
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.12'
|
||||
cache: 'pip'
|
||||
|
||||
- name: Install deps (shared shape with backend pytest job)
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r studio/backend/requirements/studio.txt
|
||||
pip install \
|
||||
python-multipart aiofiles sqlalchemy cryptography \
|
||||
pyyaml jinja2 mammoth unpdf requests typer \
|
||||
'numpy<3' pytest pytest-asyncio httpx
|
||||
# torchvision is needed because unsloth_zoo.vision_utils imports
|
||||
# it at module scope and is reached via unsloth.models._utils.
|
||||
pip install --index-url https://download.pytorch.org/whl/cpu \
|
||||
'torch>=2.4,<2.11' 'torchvision<0.26'
|
||||
pip install 'transformers>=4.51,<5.5'
|
||||
# bitsandbytes is a hard import in unsloth/models/_utils.py.
|
||||
# Recent versions ship a CPU build so it installs on a free
|
||||
# Linux runner; the kernels still raise on use, but import
|
||||
# succeeds and the package collects.
|
||||
pip install 'bitsandbytes>=0.45'
|
||||
# unsloth.device_type imports unsloth_zoo.utils.Version at module
|
||||
# scope, so the conftest harness needs unsloth_zoo on the path
|
||||
# even though it is an optional dep of unsloth.
|
||||
pip install 'unsloth_zoo>=2026.5.1'
|
||||
pip install -e . --no-deps
|
||||
|
||||
- name: Repo tests (CPU, auto-discovered)
|
||||
env:
|
||||
# tests/python/* import install_python_stack from studio/.
|
||||
PYTHONPATH: ${{ github.workspace }}/studio
|
||||
# Skip lazy compilation work the unsloth import chain wants to
|
||||
# do at import time on a real GPU.
|
||||
UNSLOTH_COMPILE_DISABLE: '1'
|
||||
# --ignore: GPU-bound directories (qlora and saving need real
|
||||
# weights / GPU; tests/sh is a shell suite the next step
|
||||
# handles; tests/utils is a helpers folder, not tests).
|
||||
# State-sensitive hardware-spoofing files are pulled out and run
|
||||
# in isolation in the next step because they mutate
|
||||
# hardware.py module globals (IS_ROCM / DEVICE) and pollute
|
||||
# downstream tests.
|
||||
# -m: honour markers already declared in tests/python/conftest.py
|
||||
# (`server` = needs studio venv, `e2e` = needs network).
|
||||
# --deselect: two registry tests that hit huggingface_hub for
|
||||
# live model existence checks; they belong on a network job.
|
||||
run: |
|
||||
python -m pytest tests/ -q --tb=short \
|
||||
--ignore=tests/qlora \
|
||||
--ignore=tests/saving \
|
||||
--ignore=tests/utils \
|
||||
--ignore=tests/sh \
|
||||
--ignore=tests/studio/test_hardware_dispatch_matrix.py \
|
||||
--ignore=tests/studio/test_is_mlx_dispatch_gate.py \
|
||||
-m 'not server and not e2e' \
|
||||
--deselect tests/test_model_registry.py::test_model_registration \
|
||||
--deselect tests/test_model_registry.py::test_all_model_registration
|
||||
|
||||
- name: Hardware-spoof tests (state-sensitive, run in isolation)
|
||||
env:
|
||||
PYTHONPATH: ${{ github.workspace }}/studio
|
||||
UNSLOTH_COMPILE_DISABLE: '1'
|
||||
# These two files mutate hardware.py module globals at runtime
|
||||
# via the spoof fixtures, which leaks state into any other test
|
||||
# that imports hardware. Run them in their own pytest invocation
|
||||
# so the leak does not cross file boundaries.
|
||||
run: |
|
||||
python -m pytest -q --tb=short \
|
||||
tests/studio/test_hardware_dispatch_matrix.py \
|
||||
tests/studio/test_is_mlx_dispatch_gate.py
|
||||
|
||||
- name: Shell installer tests
|
||||
# Subset that does not depend on a writable / pristine install.sh
|
||||
# tree; test_install_host_defaults.sh checks install.ps1 layout
|
||||
# which has drifted (separate followup).
|
||||
run: |
|
||||
set -e
|
||||
for s in \
|
||||
tests/sh/test_get_torch_index_url.sh \
|
||||
tests/sh/test_mac_intel_compat.sh \
|
||||
tests/sh/test_tauri_install_exit_order.sh \
|
||||
tests/sh/test_torch_constraint.sh; do
|
||||
echo "::group::$s"
|
||||
bash "$s"
|
||||
echo "::endgroup::"
|
||||
done
|
||||
|
||||
ruff:
|
||||
name: Backend ruff lint (non-blocking)
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.12'
|
||||
cache: 'pip'
|
||||
- run: pip install ruff
|
||||
- name: ruff check (non-blocking until accumulated drift is cleared)
|
||||
run: ruff check studio/backend || true
|
||||
108
.github/workflows/studio-frontend-ci.yml
vendored
Normal file
108
.github/workflows/studio-frontend-ci.yml
vendored
Normal file
|
|
@ -0,0 +1,108 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
|
||||
|
||||
# Frontend PR gate: lockfile freshness, typecheck, build, and a bundle grep
|
||||
# that catches the 2026.5.1 chat-history regression at the JS level.
|
||||
#
|
||||
# biome runs as non-blocking for now: the codebase currently has accumulated
|
||||
# ~470 errors and ~1650 warnings against the existing biome config. Surfacing
|
||||
# the count in CI lets us drive it down without forcing a fleet-wide cleanup
|
||||
# in the same PR. Drop `continue-on-error` once that number is zero.
|
||||
|
||||
name: Frontend CI
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
paths:
|
||||
- 'studio/frontend/**'
|
||||
- '.github/workflows/studio-frontend-ci.yml'
|
||||
push:
|
||||
branches: [main, pip]
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
build:
|
||||
name: Frontend build + bundle sanity
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
defaults:
|
||||
run:
|
||||
working-directory: studio/frontend
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
# FIXME: drop this step once @assistant-ui/* and assistant-stream
|
||||
# leave 0.x -- on 1.x, caret ranges are conventional. Until then,
|
||||
# every 0.minor on this surface is a SemVer-major (this is exactly
|
||||
# how 2026.5.1 shipped a broken chat runtime: ^0.12.19 quietly
|
||||
# resolved to 0.12.28).
|
||||
- name: '@assistant-ui must be pinned exactly (no caret/tilde)'
|
||||
working-directory: ${{ github.workspace }}
|
||||
run: |
|
||||
set -e
|
||||
if grep -nE '"(@assistant-ui/[a-z-]+|assistant-stream)":[[:space:]]*"[\^~]' studio/frontend/package.json; then
|
||||
echo "::error file=studio/frontend/package.json::These packages must be pinned to exact versions until they leave 0.x. Drop the leading ^ or ~."
|
||||
exit 1
|
||||
fi
|
||||
echo "All assistant-ui packages are pinned exactly."
|
||||
|
||||
- uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
cache: 'npm'
|
||||
cache-dependency-path: studio/frontend/package-lock.json
|
||||
|
||||
- name: Lockfile must agree with package.json (npm ci is strict)
|
||||
run: npm ci --no-fund --no-audit
|
||||
|
||||
- name: npm ci must not have modified the working tree
|
||||
working-directory: ${{ github.workspace }}
|
||||
run: |
|
||||
if ! git diff --quiet -- studio/frontend; then
|
||||
echo "::error::npm ci modified files; commit the updated lockfile"
|
||||
git status -- studio/frontend
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Typecheck
|
||||
run: npm run typecheck
|
||||
|
||||
- name: Build
|
||||
run: npm run build
|
||||
|
||||
- name: Built bundle must not contain Studio's unstable_Provider call site
|
||||
run: |
|
||||
set -e
|
||||
JS=$(ls dist/assets/index-*.js | head -1)
|
||||
HITS=$(grep -c 'unstable_Provider:' "$JS" || echo 0)
|
||||
echo "main bundle: $JS"
|
||||
echo "unstable_Provider: hits=$HITS (assistant-ui internals contribute up to 3)"
|
||||
if [ "$HITS" -gt 3 ]; then
|
||||
echo "::error file=studio/frontend/src/features/chat/runtime-provider.tsx::Studio bundle still passes unstable_Provider through useRemoteThreadListRuntime; this is the 2026.5.1 chat-history regression. Pass adapters directly into useLocalRuntime instead."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Bundle size budget (75 MB)
|
||||
run: |
|
||||
SIZE=$(du -sb dist | cut -f1)
|
||||
BUDGET=$((75 * 1024 * 1024))
|
||||
echo "dist size: $SIZE bytes ($((SIZE/1024/1024)) MB), budget: $BUDGET bytes (75 MB)"
|
||||
if [ "$SIZE" -gt "$BUDGET" ]; then
|
||||
echo "::error::studio/frontend/dist/ exceeded the 75 MB budget. Drop dead deps (e.g. the unused next dep) or split chunks."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Biome (non-blocking until accumulated drift is cleared)
|
||||
continue-on-error: true
|
||||
run: npm run biome:check
|
||||
|
||||
- name: Upload built dist on failure
|
||||
if: failure()
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: studio-frontend-dist
|
||||
path: studio/frontend/dist
|
||||
retention-days: 3
|
||||
185
.github/workflows/studio-inference-smoke.yml
vendored
Normal file
185
.github/workflows/studio-inference-smoke.yml
vendored
Normal file
|
|
@ -0,0 +1,185 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
|
||||
|
||||
# End-to-end smoke: install Studio via install.sh --local --no-torch, download
|
||||
# a tiny GGUF, boot Studio, log in, change password, load the model, send a
|
||||
# chat completion, assert a non-empty response. Only workflow that tests "the
|
||||
# app actually works".
|
||||
#
|
||||
# Model: Qwen3.5-2B UD-IQ3_XXS (~890 MiB) -- small enough that the cache miss
|
||||
# is cheap and inference fits in the 25 min CPU-runner budget. GGUF is cached
|
||||
# across runs via actions/cache.
|
||||
|
||||
name: Studio GGUF CI
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
paths:
|
||||
- 'studio/**'
|
||||
- 'unsloth/**'
|
||||
- 'unsloth_cli/**'
|
||||
- 'install.sh'
|
||||
- 'pyproject.toml'
|
||||
- '.github/workflows/studio-inference-smoke.yml'
|
||||
push:
|
||||
branches: [main, pip]
|
||||
# Manual trigger for pre-warming the GGUF cache on main, or re-running
|
||||
# against an arbitrary branch without pushing a no-op commit.
|
||||
workflow_dispatch:
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
env:
|
||||
GGUF_REPO: unsloth/Qwen3.5-2B-GGUF
|
||||
GGUF_FILE: Qwen3.5-2B-UD-IQ3_XXS.gguf
|
||||
STUDIO_PORT: '18888'
|
||||
|
||||
jobs:
|
||||
inference:
|
||||
name: Studio boots, loads a GGUF, answers a chat completion
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 25
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Linux dependencies for llama.cpp prebuilt
|
||||
run: |
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y --no-install-recommends \
|
||||
libcurl4-openssl-dev libssl-dev jq
|
||||
|
||||
- uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
cache: 'npm'
|
||||
cache-dependency-path: studio/frontend/package-lock.json
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.12'
|
||||
cache: 'pip'
|
||||
|
||||
- name: Cache GGUF model file
|
||||
id: cache-gguf
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: gguf-cache
|
||||
key: ${{ runner.os }}-gguf-${{ env.GGUF_REPO }}-${{ env.GGUF_FILE }}-v1
|
||||
|
||||
- name: Download GGUF if cache miss
|
||||
if: steps.cache-gguf.outputs.cache-hit != 'true'
|
||||
run: |
|
||||
# huggingface-cli was deprecated in huggingface_hub 1.13; the new CLI is `hf`.
|
||||
python -m pip install --upgrade huggingface_hub hf_transfer
|
||||
mkdir -p gguf-cache
|
||||
HF_HUB_ENABLE_HF_TRANSFER=1 \
|
||||
hf download "$GGUF_REPO" "$GGUF_FILE" --local-dir gguf-cache
|
||||
|
||||
- name: Install Studio (--local, --no-torch keeps the install lean)
|
||||
run: |
|
||||
mkdir -p logs
|
||||
set -o pipefail
|
||||
bash install.sh --local --no-torch 2>&1 | tee logs/install.log
|
||||
|
||||
- name: Assert llama.cpp prebuilt was installed (no source-build fallback)
|
||||
# ubuntu-latest is CPU-only x86_64, so studio/setup.sh should route
|
||||
# to ggml-org/llama.cpp and grab bin-ubuntu-x64.tar.gz. A source
|
||||
# build here means the routing regressed.
|
||||
run: |
|
||||
if grep -q "falling back to source build" logs/install.log; then
|
||||
echo "::error::llama.cpp prebuilt path failed on ubuntu-latest. studio/setup.sh routing regressed; CPU-only Linux x86_64 should hit ggml-org/llama.cpp's bin-ubuntu-x64.tar.gz."
|
||||
grep -E "llama-prebuilt|llama.cpp" logs/install.log | tail -60
|
||||
exit 1
|
||||
fi
|
||||
if ! grep -qE "prebuilt installed and validated|prebuilt up to date and validated" logs/install.log; then
|
||||
echo "::error::install.log does not contain the success marker for the llama.cpp prebuilt path. Did setup.sh skip the prebuilt install?"
|
||||
grep -E "llama-prebuilt|llama.cpp" logs/install.log | tail -60
|
||||
exit 1
|
||||
fi
|
||||
echo "llama.cpp prebuilt path used successfully"
|
||||
|
||||
- name: Reset auth + start Studio in the background
|
||||
run: |
|
||||
unsloth studio reset-password
|
||||
mkdir -p logs
|
||||
UNSLOTH_API_ONLY=1 unsloth studio -H 127.0.0.1 -p "$STUDIO_PORT" \
|
||||
> logs/studio.log 2>&1 &
|
||||
echo "STUDIO_PID=$!" >> "$GITHUB_ENV"
|
||||
|
||||
- name: Wait for /api/health
|
||||
run: |
|
||||
for i in $(seq 1 60); do
|
||||
if curl -fs "http://127.0.0.1:${STUDIO_PORT}/api/health" > /tmp/health.json; then
|
||||
echo "ready after ${i}s"
|
||||
cat /tmp/health.json
|
||||
jq -e '.status == "healthy"' /tmp/health.json
|
||||
exit 0
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
echo "Studio did not become healthy in 60s"
|
||||
tail -200 logs/studio.log
|
||||
exit 1
|
||||
|
||||
- name: Login + change bootstrap password
|
||||
run: |
|
||||
PW=$(cat ~/.unsloth/studio/auth/.bootstrap_password)
|
||||
NEW="CIPasswordSmoke12345!"
|
||||
TOKEN=$(curl -fs -X POST "http://127.0.0.1:${STUDIO_PORT}/api/auth/login" \
|
||||
-H 'content-type: application/json' \
|
||||
-d "{\"username\":\"unsloth\",\"password\":\"$PW\"}" | jq -r .access_token)
|
||||
curl -fs -X POST "http://127.0.0.1:${STUDIO_PORT}/api/auth/change-password" \
|
||||
-H "Authorization: Bearer $TOKEN" -H 'content-type: application/json' \
|
||||
-d "{\"current_password\":\"$PW\",\"new_password\":\"$NEW\"}" > /dev/null
|
||||
# Re-login to clear must_change_password flag.
|
||||
NEW_TOKEN=$(curl -fs -X POST "http://127.0.0.1:${STUDIO_PORT}/api/auth/login" \
|
||||
-H 'content-type: application/json' \
|
||||
-d "{\"username\":\"unsloth\",\"password\":\"$NEW\"}" | jq -r .access_token)
|
||||
echo "TOKEN=$NEW_TOKEN" >> "$GITHUB_ENV"
|
||||
|
||||
- name: Load the GGUF into Studio
|
||||
run: |
|
||||
GGUF_PATH="$GITHUB_WORKSPACE/gguf-cache/${GGUF_FILE}"
|
||||
ls -lh "$GGUF_PATH"
|
||||
curl -fs -X POST "http://127.0.0.1:${STUDIO_PORT}/api/inference/load" \
|
||||
-H "Authorization: Bearer $TOKEN" -H 'content-type: application/json' \
|
||||
--max-time 600 \
|
||||
-d "{\"model_path\":\"$GGUF_PATH\",\"is_lora\":false,\"max_seq_length\":2048}" \
|
||||
| jq '{status, display_name, is_gguf, context_length}'
|
||||
|
||||
- name: Send a chat completion + assert non-empty response
|
||||
run: |
|
||||
RESP=$(curl -fs -X POST "http://127.0.0.1:${STUDIO_PORT}/api/inference/chat/completions" \
|
||||
-H "Authorization: Bearer $TOKEN" -H 'content-type: application/json' \
|
||||
--max-time 900 \
|
||||
-d '{
|
||||
"messages":[{"role":"user","content":"Say hello in one short sentence."}],
|
||||
"max_tokens":40,
|
||||
"stream":false
|
||||
}')
|
||||
echo "raw response: $RESP"
|
||||
CONTENT=$(echo "$RESP" | jq -r '.choices[0].message.content // empty')
|
||||
echo "model response: $CONTENT"
|
||||
if [ -z "$CONTENT" ]; then
|
||||
echo "::error::Empty assistant response from Studio"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Stop Studio
|
||||
if: always()
|
||||
run: |
|
||||
kill "${STUDIO_PID}" || true
|
||||
sleep 2
|
||||
ss -tln | grep ":${STUDIO_PORT}" || true
|
||||
|
||||
- name: Upload Studio + install logs on failure
|
||||
if: failure()
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: studio-inference-log
|
||||
path: |
|
||||
logs/studio.log
|
||||
logs/install.log
|
||||
retention-days: 7
|
||||
105
.github/workflows/studio-tauri-smoke.yml
vendored
Normal file
105
.github/workflows/studio-tauri-smoke.yml
vendored
Normal file
|
|
@ -0,0 +1,105 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
|
||||
|
||||
# PR-time smoke for the Tauri desktop wrapper. Builds the frontend and the
|
||||
# Tauri Linux debug binary, with no codesigning. Catches:
|
||||
# - tauri.conf.json drift
|
||||
# - src-tauri Cargo.toml or rust source breakage
|
||||
# - Tauri CLI version drift (we pin 2.10.1, matching release-desktop.yml)
|
||||
# - frontend output not picked up by Tauri's distDir
|
||||
#
|
||||
# Linux-only on a free `ubuntu-latest` runner. Mac and Windows desktop builds
|
||||
# stay in release-desktop.yml (manual `workflow_dispatch`) because they need
|
||||
# code-signing secrets and ~30 min of runner time each.
|
||||
|
||||
name: Studio Tauri CI
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
paths:
|
||||
- 'studio/frontend/**'
|
||||
- 'studio/src-tauri/**'
|
||||
- '.github/workflows/studio-tauri-smoke.yml'
|
||||
push:
|
||||
branches: [main, pip]
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
linux-debug-build:
|
||||
name: Tauri Linux debug build (no codesign)
|
||||
runs-on: ubuntu-22.04
|
||||
timeout-minutes: 25
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Linux native deps for Tauri / WebKit2GTK
|
||||
run: |
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y \
|
||||
libwebkit2gtk-4.1-dev libayatana-appindicator3-dev \
|
||||
librsvg2-dev libxdo-dev libssl-dev patchelf
|
||||
|
||||
- uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '24'
|
||||
cache: 'npm'
|
||||
cache-dependency-path: studio/frontend/package-lock.json
|
||||
|
||||
- uses: dtolnay/rust-toolchain@stable
|
||||
|
||||
- uses: swatinem/rust-cache@v2
|
||||
with:
|
||||
workspaces: studio/src-tauri -> target
|
||||
|
||||
- name: Install pinned Tauri CLI (matches release-desktop.yml)
|
||||
run: npm install --save-dev --prefix studio @tauri-apps/cli@2.10.1
|
||||
|
||||
- name: Verify pinned Tauri CLI version
|
||||
run: |
|
||||
out="$(npx --prefix studio tauri --version)"
|
||||
echo "$out"
|
||||
[ "$out" = "tauri-cli 2.10.1" ] || { echo "::error::expected tauri-cli 2.10.1, got $out"; exit 1; }
|
||||
|
||||
- name: Frontend build (npm ci, vite)
|
||||
working-directory: studio/frontend
|
||||
run: |
|
||||
npm ci --no-fund --no-audit
|
||||
npm run build
|
||||
test -f dist/index.html
|
||||
|
||||
- name: Tauri debug build (Linux, no bundle, no codesign)
|
||||
# `--debug` + `--no-bundle` keeps this lean: compiles the Rust crate,
|
||||
# confirms the frontend dist is wired into Tauri, but skips the AppImage
|
||||
# / .deb production. Code signing is irrelevant because we never produce
|
||||
# a distributable artifact.
|
||||
env:
|
||||
TAURI_SIGNING_PRIVATE_KEY: ''
|
||||
TAURI_SIGNING_PRIVATE_KEY_PASSWORD: ''
|
||||
run: npx --prefix studio tauri build --debug --no-bundle
|
||||
|
||||
- name: Inspect produced binary
|
||||
run: |
|
||||
BIN=$(find studio/src-tauri/target/debug -maxdepth 1 -type f -executable 2>/dev/null \
|
||||
| grep -Ev '\.(d|so|dylib|dll)$' \
|
||||
| grep -Ev '/(deps|build|examples)$' \
|
||||
| head -1)
|
||||
echo "binary: $BIN"
|
||||
if [ -z "$BIN" ]; then
|
||||
echo "::error::Tauri debug binary not produced"
|
||||
ls -la studio/src-tauri/target/debug/ || true
|
||||
exit 1
|
||||
fi
|
||||
file "$BIN"
|
||||
du -h "$BIN"
|
||||
|
||||
- uses: actions/upload-artifact@v4
|
||||
if: failure()
|
||||
with:
|
||||
name: tauri-debug-build
|
||||
path: |
|
||||
studio/src-tauri/target/debug
|
||||
studio/frontend/dist
|
||||
retention-days: 3
|
||||
124
.github/workflows/wheel-smoke.yml
vendored
Normal file
124
.github/workflows/wheel-smoke.yml
vendored
Normal file
|
|
@ -0,0 +1,124 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
|
||||
|
||||
# Builds the PyPI wheel from the PR branch, then verifies the built wheel
|
||||
# actually contains what we expect to ship and does NOT contain the broken
|
||||
# Studio bundle that 2026.5.1 published. This is the single workflow that
|
||||
# would have blocked the 2026.5.1 release before twine upload.
|
||||
#
|
||||
# Verified locally end-to-end against this branch:
|
||||
# - python -m build produces unsloth-<version>-py3-none-any.whl in 13s
|
||||
# - wheel content sanity passes:
|
||||
# lockfile shipped, frontend dist shipped,
|
||||
# no node_modules in wheel, no bun.lock in wheel,
|
||||
# main bundle has unstable_Provider hits=1 (assistant-ui internals only).
|
||||
# - Studio backend imports cleanly from the installed wheel with the
|
||||
# lightweight dep set below.
|
||||
|
||||
name: Wheel CI
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
paths:
|
||||
- 'pyproject.toml'
|
||||
- 'studio/**'
|
||||
- 'unsloth/**'
|
||||
- 'unsloth_cli/**'
|
||||
- '.github/workflows/wheel-smoke.yml'
|
||||
push:
|
||||
branches: [main, pip]
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
wheel:
|
||||
name: Wheel build + content sanity + import smoke
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
cache: 'npm'
|
||||
cache-dependency-path: studio/frontend/package-lock.json
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.12'
|
||||
|
||||
- name: Build frontend
|
||||
run: |
|
||||
cd studio/frontend
|
||||
npm ci --no-fund --no-audit
|
||||
npm run build
|
||||
|
||||
- name: Build wheel + sdist
|
||||
run: |
|
||||
python -m pip install --upgrade pip build
|
||||
rm -rf dist build ./*.egg-info
|
||||
python -m build
|
||||
|
||||
- name: Wheel content sanity
|
||||
run: |
|
||||
python - <<'PY'
|
||||
import zipfile, glob, sys
|
||||
w = glob.glob("dist/unsloth-*.whl")
|
||||
if not w:
|
||||
print("FAIL: no wheel produced"); sys.exit(2)
|
||||
w = w[0]
|
||||
print(f"wheel: {w}")
|
||||
with zipfile.ZipFile(w) as z:
|
||||
n = z.namelist()
|
||||
checks = {
|
||||
"lockfile shipped": any(s.endswith("studio/frontend/package-lock.json") for s in n),
|
||||
"frontend dist shipped": any(s.endswith("studio/frontend/dist/index.html") for s in n),
|
||||
"no node_modules": not any("studio/frontend/node_modules/" in s for s in n),
|
||||
"no bun.lock": not any(s.endswith("studio/frontend/bun.lock") for s in n),
|
||||
}
|
||||
js = [s for s in n
|
||||
if "studio/frontend/dist/assets/" in s
|
||||
and s.endswith(".js")
|
||||
and "/index-" in s]
|
||||
if not js:
|
||||
print("FAIL: no main bundle index-*.js in wheel"); sys.exit(2)
|
||||
data = z.read(js[0]).decode("utf-8", "replace")
|
||||
hits = data.count("unstable_Provider:")
|
||||
print(f"main bundle: {js[0]}")
|
||||
print(f"unstable_Provider hits: {hits} (>=4 indicates 2026.5.1 regression)")
|
||||
checks["bundle has no Studio unstable_Provider call site"] = (hits < 4)
|
||||
|
||||
print()
|
||||
for k, v in checks.items():
|
||||
print(f" [{'PASS' if v else 'FAIL'}] {k}")
|
||||
sys.exit(0 if all(checks.values()) else 1)
|
||||
PY
|
||||
|
||||
- name: Studio backend import smoke
|
||||
# Imports `studio.backend.main:app` from the freshly-installed wheel in
|
||||
# a clean venv. This catches the class of bug that 2026.5.1 shipped with:
|
||||
# frontend dist missing, package-lock.json missing, or the wheel's Python
|
||||
# source tree broken in a way that surfaces only at app construction time.
|
||||
run: |
|
||||
python -m venv /tmp/v
|
||||
/tmp/v/bin/pip install --upgrade pip
|
||||
/tmp/v/bin/pip install -r studio/backend/requirements/studio.txt
|
||||
/tmp/v/bin/pip install \
|
||||
python-multipart aiofiles sqlalchemy cryptography \
|
||||
pyyaml jinja2 mammoth unpdf requests \
|
||||
'numpy<3'
|
||||
/tmp/v/bin/pip install --no-deps dist/unsloth-*.whl
|
||||
# Run from /tmp so Python imports the installed package, not the source tree.
|
||||
cd /tmp
|
||||
/tmp/v/bin/python -c "from studio.backend.main import app; print('Studio backend OK:', app.title)"
|
||||
|
||||
- name: Upload wheel on failure
|
||||
if: failure()
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: unsloth-wheel
|
||||
path: dist/
|
||||
retention-days: 7
|
||||
5
.gitignore
vendored
5
.gitignore
vendored
|
|
@ -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/
|
||||
|
|
|
|||
390
install.ps1
390
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 ""
|
||||
}
|
||||
|
|
|
|||
421
install.sh
421
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>
|
||||
PLIST_EOF
|
||||
|
||||
# Executable stub
|
||||
cat > "$_css_macos_dir/launch-studio" << STUB_EOF
|
||||
# Executable stub: same single-quoted-heredoc + sed-substitute
|
||||
# pattern as launch-studio.sh so $-vars in $_css_data_dir don't
|
||||
# expand at .app launch time.
|
||||
_css_sq_dir=$(printf '%s' "$_css_data_dir" | sed "s/'/'\\\\''/g")
|
||||
_css_sed_dir=$(printf '%s' "$_css_sq_dir" | sed 's/[\\&|]/\\&/g')
|
||||
cat > "$_css_macos_dir/launch-studio" << 'STUB_EOF'
|
||||
#!/bin/sh
|
||||
exec "$HOME/.local/share/unsloth/launch-studio.sh" "\$@"
|
||||
exec '@@DATA_DIR@@/launch-studio.sh' "$@"
|
||||
STUB_EOF
|
||||
sed "s|@@DATA_DIR@@|$_css_sed_dir|g" "$_css_macos_dir/launch-studio" \
|
||||
> "$_css_macos_dir/launch-studio.tmp" \
|
||||
&& mv "$_css_macos_dir/launch-studio.tmp" "$_css_macos_dir/launch-studio"
|
||||
chmod +x "$_css_macos_dir/launch-studio"
|
||||
|
||||
# Build AppIcon.icns from unsloth-gem.png (2240x2240)
|
||||
|
|
@ -1079,11 +1348,28 @@ mkdir -p "$STUDIO_HOME"
|
|||
_MIGRATED=false
|
||||
|
||||
if [ -x "$VENV_DIR/bin/python" ]; then
|
||||
# why: matching guard to the .venv branch below -- in env-mode
|
||||
# $STUDIO_HOME is a user-chosen workspace, so refuse to nuke an
|
||||
# existing $STUDIO_HOME/unsloth_studio that lacks Studio sentinels.
|
||||
# Accept the in-VENV ownership marker so partial-install retries are
|
||||
# not blocked. Sentinels must be regular files: -f follows symlinks
|
||||
# to files (the legitimate ln -s shim shape) but rejects directories
|
||||
# and broken/dir-targeted symlinks.
|
||||
if [ "$_STUDIO_HOME_REDIRECT" = "env" ] \
|
||||
&& [ ! -f "$VENV_DIR/.unsloth-studio-owned" ] \
|
||||
&& [ ! -f "$STUDIO_HOME/share/studio.conf" ] \
|
||||
&& [ ! -f "$STUDIO_HOME/bin/unsloth" ]; then
|
||||
echo "ERROR: $VENV_DIR already exists but does not look like an Unsloth Studio install." >&2
|
||||
echo " Move it aside or choose an empty UNSLOTH_STUDIO_HOME." >&2
|
||||
exit 1
|
||||
fi
|
||||
# New layout already exists — replace only after preserving rollback copy.
|
||||
substep "preserving existing environment for rollback..."
|
||||
_start_studio_venv_replacement "$VENV_DIR"
|
||||
elif [ -x "$STUDIO_HOME/.venv/bin/python" ]; then
|
||||
elif [ "$_STUDIO_HOME_REDIRECT" != "env" ] && [ -x "$STUDIO_HOME/.venv/bin/python" ]; then
|
||||
# Old layout exists — validate before migrating.
|
||||
# Skip in env-mode so we don't rm -rf an unrelated .venv at the
|
||||
# workspace root (e.g. user's existing project Python venv).
|
||||
# In no-torch mode, a missing torch package is expected; validate Python only.
|
||||
substep "found legacy Studio environment, validating..."
|
||||
_legacy_ok=false
|
||||
|
|
@ -1132,6 +1418,13 @@ if [ ! -x "$VENV_DIR/bin/python" ]; then
|
|||
run_install_cmd "create venv" uv venv "$VENV_DIR" --python "$PYTHON_VERSION"
|
||||
fi
|
||||
|
||||
# Mark the freshly-created venv as Studio-owned so a partial install can be
|
||||
# repaired by re-running install.sh; the env-mode deletion guard above accepts
|
||||
# this marker as the primary sentinel.
|
||||
if [ -x "$VENV_DIR/bin/python" ]; then
|
||||
: > "$VENV_DIR/.unsloth-studio-owned" 2>/dev/null || true
|
||||
fi
|
||||
|
||||
# Guard against Python 3.13.8 torch import bug on Apple Silicon
|
||||
# (skip when the user explicitly chose a version via --python)
|
||||
if [ -z "$_USER_PYTHON" ] && [ "$OS" = "macos" ] && [ "$_ARCH" = "arm64" ]; then
|
||||
|
|
@ -1143,6 +1436,9 @@ if [ -z "$_USER_PYTHON" ] && [ "$OS" = "macos" ] && [ "$_ARCH" = "arm64" ]; then
|
|||
rm -rf "$VENV_DIR"
|
||||
PYTHON_VERSION="3.12"
|
||||
run_install_cmd "recreate venv" uv venv "$VENV_DIR" --python "$PYTHON_VERSION"
|
||||
if [ -x "$VENV_DIR/bin/python" ]; then
|
||||
: > "$VENV_DIR/.unsloth-studio-owned" 2>/dev/null || true
|
||||
fi
|
||||
fi
|
||||
fi
|
||||
|
||||
|
|
@ -1721,6 +2017,12 @@ else
|
|||
fi
|
||||
fi
|
||||
|
||||
# ── Install mlx-vlm on Apple Silicon (optional, for VLM training) ──
|
||||
if [ "$OS" = "macos" ] && [ "$_ARCH" = "arm64" ]; then
|
||||
substep "installing mlx-vlm (VLM training support)..."
|
||||
run_install_cmd "install mlx-vlm" uv pip install --python "$_VENV_PY" mlx-vlm
|
||||
fi
|
||||
|
||||
# ── Run studio setup ──
|
||||
tauri_log "STEP" "Running Studio setup"
|
||||
# When --local, use the repo's own setup.sh directly.
|
||||
|
|
@ -1768,7 +2070,17 @@ _SKIP_FRONTEND=0
|
|||
if [ "$TAURI_MODE" = true ]; then
|
||||
_SKIP_FRONTEND=1
|
||||
fi
|
||||
# Prepend UNSLOTH_STUDIO_HOME=$STUDIO_HOME to "$@" for env-override installs
|
||||
# without word-splitting on whitespace paths.
|
||||
_run_setup_with_studio_home() {
|
||||
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
|
||||
UNSLOTH_STUDIO_HOME="$STUDIO_HOME" "$@"
|
||||
else
|
||||
"$@"
|
||||
fi
|
||||
}
|
||||
if [ "$STUDIO_LOCAL_INSTALL" = true ]; then
|
||||
_run_setup_with_studio_home env \
|
||||
SKIP_STUDIO_BASE="$_SKIP_BASE" \
|
||||
SKIP_STUDIO_FRONTEND="$_SKIP_FRONTEND" \
|
||||
STUDIO_PACKAGE_NAME="$PACKAGE_NAME" \
|
||||
|
|
@ -1782,6 +2094,7 @@ else
|
|||
# the same session) does not silently flip a normal install onto the
|
||||
# local-dev path in setup.sh and install_python_stack.py. Mirrors the
|
||||
# reset already done in install.ps1 for PowerShell.
|
||||
_run_setup_with_studio_home env \
|
||||
SKIP_STUDIO_BASE="$_SKIP_BASE" \
|
||||
SKIP_STUDIO_FRONTEND="$_SKIP_FRONTEND" \
|
||||
STUDIO_PACKAGE_NAME="$PACKAGE_NAME" \
|
||||
|
|
@ -1791,36 +2104,53 @@ else
|
|||
bash "$SETUP_SH" </dev/null || _SETUP_EXIT=$?
|
||||
fi
|
||||
|
||||
# ── Make 'unsloth' available globally via ~/.local/bin ──
|
||||
mkdir -p "$HOME/.local/bin"
|
||||
ln -sf "$VENV_DIR/bin/unsloth" "$HOME/.local/bin/unsloth"
|
||||
# ── Make 'unsloth' available via $_LOCAL_BIN (resolved earlier) ──
|
||||
# Env-mode: $_LOCAL_BIN is $STUDIO_HOME/bin; skip shell-rc PATH append so we
|
||||
# don't pollute the user's profile with a workspace-scoped path.
|
||||
mkdir -p "$_LOCAL_BIN"
|
||||
# ln -sf into an existing dir creates link inside it. Refuse to delete a
|
||||
# real directory at the shim path -- that could destroy unrelated user data.
|
||||
_shim_path="$_LOCAL_BIN/unsloth"
|
||||
if [ -d "$_shim_path" ] && [ ! -L "$_shim_path" ]; then
|
||||
echo "ERROR: $_shim_path is a directory; refusing to delete it." >&2
|
||||
echo " Move or remove it manually, then re-run the installer." >&2
|
||||
exit 1
|
||||
fi
|
||||
# why: -sfn is atomic and -n prevents descent into a symlink-to-directory at
|
||||
# the shim path (the directory guard above already rejects a real directory).
|
||||
ln -sfn "$VENV_DIR/bin/unsloth" "$_shim_path"
|
||||
|
||||
_LOCAL_BIN="$HOME/.local/bin"
|
||||
case ":$PATH:" in
|
||||
*":$_LOCAL_BIN:"*) ;; # already on PATH
|
||||
*)
|
||||
_SHELL_PROFILE=""
|
||||
if [ -n "${ZSH_VERSION:-}" ] || [ "$(basename "${SHELL:-}")" = "zsh" ]; then
|
||||
_SHELL_PROFILE="$HOME/.zshrc"
|
||||
elif [ -f "$HOME/.bashrc" ]; then
|
||||
_SHELL_PROFILE="$HOME/.bashrc"
|
||||
elif [ -f "$HOME/.profile" ]; then
|
||||
_SHELL_PROFILE="$HOME/.profile"
|
||||
fi
|
||||
|
||||
if [ -n "$_SHELL_PROFILE" ]; then
|
||||
if ! grep -q '\.local/bin' "$_SHELL_PROFILE" 2>/dev/null; then
|
||||
echo '' >> "$_SHELL_PROFILE"
|
||||
echo '# Added by Unsloth installer' >> "$_SHELL_PROFILE"
|
||||
echo 'export PATH="$HOME/.local/bin:$PATH"' >> "$_SHELL_PROFILE"
|
||||
step "path" "added ~/.local/bin to PATH in $_SHELL_PROFILE"
|
||||
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
|
||||
export PATH="$_LOCAL_BIN:$PATH"
|
||||
step "path" "exported $_LOCAL_BIN for this session (no rc-file append in env-override mode)"
|
||||
else
|
||||
_SHELL_PROFILE=""
|
||||
if [ -n "${ZSH_VERSION:-}" ] || [ "$(basename "${SHELL:-}")" = "zsh" ]; then
|
||||
_SHELL_PROFILE="$HOME/.zshrc"
|
||||
elif [ -f "$HOME/.bashrc" ]; then
|
||||
_SHELL_PROFILE="$HOME/.bashrc"
|
||||
elif [ -f "$HOME/.profile" ]; then
|
||||
_SHELL_PROFILE="$HOME/.profile"
|
||||
fi
|
||||
if [ -n "$_SHELL_PROFILE" ]; then
|
||||
if ! grep -q '\.local/bin' "$_SHELL_PROFILE" 2>/dev/null; then
|
||||
echo '' >> "$_SHELL_PROFILE"
|
||||
echo '# Added by Unsloth installer' >> "$_SHELL_PROFILE"
|
||||
echo 'export PATH="$HOME/.local/bin:$PATH"' >> "$_SHELL_PROFILE"
|
||||
step "path" "added ~/.local/bin to PATH in $_SHELL_PROFILE"
|
||||
fi
|
||||
fi
|
||||
export PATH="$_LOCAL_BIN:$PATH"
|
||||
fi
|
||||
export PATH="$_LOCAL_BIN:$PATH"
|
||||
;;
|
||||
esac
|
||||
|
||||
# Non-Tauri installs keep shortcuts even if setup reports failure.
|
||||
# create_studio_shortcuts gates persistent menu shortcuts on env-mode;
|
||||
# launcher + studio.conf + icon are always written.
|
||||
if [ "$TAURI_MODE" != true ]; then
|
||||
create_studio_shortcuts "$VENV_ABS_BIN/unsloth" "$OS"
|
||||
fi
|
||||
|
|
@ -1883,10 +2213,21 @@ if [ -t 1 ]; then
|
|||
esac
|
||||
else
|
||||
step "launch" "manual commands:"
|
||||
substep "unsloth studio -p 8888"
|
||||
substep "or activate env first:"
|
||||
substep "source ${VENV_DIR}/bin/activate"
|
||||
substep "unsloth studio -p 8888"
|
||||
# Single-quote-escape so paths with spaces / apostrophes copy-paste cleanly.
|
||||
_li_shim_q="'$(printf '%s' "${_LOCAL_BIN}/unsloth" | sed "s/'/'\\\\''/g")'"
|
||||
_li_act_q="'$(printf '%s' "${VENV_DIR}/bin/activate" | sed "s/'/'\\\\''/g")'"
|
||||
if [ "$_STUDIO_HOME_REDIRECT" = "env" ]; then
|
||||
# Env-mode skips the rc PATH append, so print the absolute shim path.
|
||||
substep "$_li_shim_q studio -p 8888"
|
||||
substep "or activate env first:"
|
||||
substep "source $_li_act_q"
|
||||
substep "unsloth studio -p 8888"
|
||||
else
|
||||
substep "unsloth studio -p 8888"
|
||||
substep "or activate env first:"
|
||||
substep "source $_li_act_q"
|
||||
substep "unsloth studio -p 8888"
|
||||
fi
|
||||
substep "(add -H 0.0.0.0 to allow network / cloud access)"
|
||||
echo ""
|
||||
fi
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
395
studio/backend/core/inference/mlx_inference.py
Normal file
395
studio/backend/core/inference/mlx_inference.py
Normal file
|
|
@ -0,0 +1,395 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
"""MLX inference backend for Apple Silicon.
|
||||
|
||||
Drop-in replacement for InferenceBackend — same interface, uses mlx-lm/mlx-vlm
|
||||
instead of torch/transformers for model loading and generation.
|
||||
"""
|
||||
|
||||
import threading
|
||||
from typing import Optional, Generator
|
||||
from loggers import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class MLXInferenceBackend:
|
||||
def __init__(self):
|
||||
self.models = {}
|
||||
self.active_model_name = None
|
||||
self.loading_models = set()
|
||||
self.loaded_local_models = []
|
||||
self.device = "mlx"
|
||||
self._generation_lock = threading.Lock()
|
||||
|
||||
# MLX state
|
||||
self._model = None
|
||||
self._tokenizer = None
|
||||
self._processor = None
|
||||
self._is_vlm = False
|
||||
self._config = {}
|
||||
|
||||
# Recorded for unload to release pinned memory back to the OS.
|
||||
self._memory_limits_applied = {}
|
||||
|
||||
def _configure_memory_limits(self):
|
||||
"""Apply Metal memory caps before loading a model.
|
||||
|
||||
Mirrors MLXTrainer._configure_memory_limits's defaults:
|
||||
memory_limit = 85% of recommended working-set,
|
||||
wired_limit = min(recommended, memory_limit). Recorded so unload
|
||||
can lower wired_limit back to release pinned RAM.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
if not mx.metal.is_available():
|
||||
return
|
||||
info = mx.device_info()
|
||||
rec_bytes = info.get("max_recommended_working_set_size")
|
||||
if not rec_bytes or rec_bytes <= 0:
|
||||
return
|
||||
rec_gb = rec_bytes / 1e9
|
||||
memory_limit_gb = rec_gb * 0.85
|
||||
wired_limit_gb = min(rec_gb, memory_limit_gb)
|
||||
mx.set_memory_limit(int(memory_limit_gb * 1e9))
|
||||
mx.set_wired_limit(int(wired_limit_gb * 1e9))
|
||||
self._memory_limits_applied = {
|
||||
"memory_limit_gb": memory_limit_gb,
|
||||
"wired_limit_gb": wired_limit_gb,
|
||||
"recommended_gb": rec_gb,
|
||||
}
|
||||
logger.info(
|
||||
"MLX memory caps: memory_limit=%.2f GB, wired_limit=%.2f GB",
|
||||
memory_limit_gb,
|
||||
wired_limit_gb,
|
||||
)
|
||||
|
||||
def load_model(
|
||||
self,
|
||||
config,
|
||||
max_seq_length = 2048,
|
||||
load_in_4bit = True,
|
||||
hf_token = None,
|
||||
trust_remote_code = False,
|
||||
gpu_ids = None,
|
||||
dtype = None,
|
||||
) -> bool:
|
||||
import mlx.core as mx
|
||||
|
||||
model_name = config.identifier if hasattr(config, "identifier") else str(config)
|
||||
is_vision = getattr(config, "is_vision", False)
|
||||
|
||||
if hf_token:
|
||||
import os
|
||||
|
||||
os.environ["HF_TOKEN"] = hf_token
|
||||
self._configure_memory_limits()
|
||||
|
||||
is_lora = getattr(config, "is_lora", False)
|
||||
|
||||
logger.info(
|
||||
"Loading %s via %s (is_lora=%s)",
|
||||
model_name,
|
||||
"mlx-vlm" if is_vision else "mlx-lm",
|
||||
is_lora,
|
||||
)
|
||||
|
||||
try:
|
||||
from unsloth_zoo.mlx_loader import FastMLXModel
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"Unsloth: MLX inference requires unsloth-zoo with the MLX modules "
|
||||
"(unsloth_zoo.mlx_loader). Reinstall via install.sh on Apple Silicon."
|
||||
) from e
|
||||
|
||||
model, tokenizer_or_processor = FastMLXModel.from_pretrained(
|
||||
model_name,
|
||||
max_seq_length = max_seq_length,
|
||||
dtype = dtype,
|
||||
load_in_4bit = load_in_4bit,
|
||||
token = hf_token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
text_only = False if is_vision else True,
|
||||
)
|
||||
|
||||
if is_vision:
|
||||
processor = tokenizer_or_processor
|
||||
self._model = model
|
||||
self._processor = processor
|
||||
self._tokenizer = getattr(processor, "tokenizer", processor)
|
||||
self._is_vlm = True
|
||||
else:
|
||||
tokenizer = tokenizer_or_processor
|
||||
self._model = model
|
||||
self._tokenizer = tokenizer
|
||||
self._processor = None
|
||||
self._is_vlm = False
|
||||
|
||||
self.active_model_name = model_name
|
||||
self.models[model_name] = {
|
||||
"model": self._model,
|
||||
"tokenizer": self._tokenizer,
|
||||
"processor": self._processor,
|
||||
"is_vision": is_vision,
|
||||
"is_lora": getattr(config, "is_lora", False),
|
||||
"is_audio": False,
|
||||
"audio_type": None,
|
||||
"has_audio_input": False,
|
||||
}
|
||||
|
||||
logger.info("Model %s loaded successfully", model_name)
|
||||
return True
|
||||
|
||||
def unload_model(self, model_name: str) -> bool:
|
||||
import mlx.core as mx
|
||||
import gc
|
||||
|
||||
if model_name in self.models:
|
||||
del self.models[model_name]
|
||||
self._model = None
|
||||
self._tokenizer = None
|
||||
self._processor = None
|
||||
if self.active_model_name == model_name:
|
||||
self.active_model_name = None
|
||||
gc.collect()
|
||||
mx.clear_cache()
|
||||
|
||||
if mx.metal.is_available() and self._memory_limits_applied and not self.models:
|
||||
try:
|
||||
mx.set_wired_limit(0)
|
||||
logger.info("MLX wired_limit released back to OS on unload")
|
||||
except Exception as e:
|
||||
logger.warning("Failed to release wired_limit: %s", e)
|
||||
self._memory_limits_applied = {}
|
||||
logger.info("Model %s unloaded", model_name)
|
||||
return True
|
||||
|
||||
def generate_chat_response(
|
||||
self,
|
||||
messages,
|
||||
system_prompt = "",
|
||||
image = None,
|
||||
temperature = 0.7,
|
||||
top_p = 0.9,
|
||||
top_k = 40,
|
||||
min_p = 0.0,
|
||||
max_new_tokens = 256,
|
||||
repetition_penalty = 1.0,
|
||||
cancel_event = None,
|
||||
) -> Generator[str, None, None]:
|
||||
if self._model is None:
|
||||
raise RuntimeError("No model loaded")
|
||||
|
||||
# Build messages with system prompt
|
||||
full_messages = []
|
||||
if system_prompt:
|
||||
full_messages.append({"role": "system", "content": system_prompt})
|
||||
full_messages.extend(messages)
|
||||
|
||||
# Inject image into the last user message for VLM
|
||||
if self._is_vlm and image is not None:
|
||||
for msg in reversed(full_messages):
|
||||
if msg.get("role") == "user":
|
||||
content = msg.get("content", "")
|
||||
if isinstance(content, str):
|
||||
msg["content"] = [
|
||||
{"type": "image"},
|
||||
{"type": "text", "text": content},
|
||||
]
|
||||
elif isinstance(content, list):
|
||||
# Prepend image if not already there
|
||||
has_image = any(
|
||||
p.get("type") == "image"
|
||||
for p in content
|
||||
if isinstance(p, dict)
|
||||
)
|
||||
if not has_image:
|
||||
content.insert(0, {"type": "image"})
|
||||
break
|
||||
|
||||
if self._is_vlm:
|
||||
yield from self._generate_vlm(
|
||||
full_messages,
|
||||
image,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
min_p,
|
||||
max_new_tokens,
|
||||
repetition_penalty,
|
||||
cancel_event,
|
||||
)
|
||||
else:
|
||||
yield from self._generate_text(
|
||||
full_messages,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
min_p,
|
||||
max_new_tokens,
|
||||
repetition_penalty,
|
||||
cancel_event,
|
||||
)
|
||||
|
||||
def _generate_text(
|
||||
self,
|
||||
messages,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
min_p,
|
||||
max_new_tokens,
|
||||
repetition_penalty,
|
||||
cancel_event,
|
||||
):
|
||||
from mlx_lm import stream_generate
|
||||
from mlx_lm.sample_utils import make_sampler, make_logits_processors
|
||||
|
||||
prompt = self._tokenizer.apply_chat_template(
|
||||
messages,
|
||||
tokenize = False,
|
||||
add_generation_prompt = True,
|
||||
)
|
||||
if prompt is None:
|
||||
raise RuntimeError(
|
||||
"apply_chat_template returned None — tokenizer may be incompatible"
|
||||
)
|
||||
|
||||
sampler = make_sampler(
|
||||
temp = temperature,
|
||||
top_p = top_p,
|
||||
top_k = int(top_k or 0),
|
||||
min_p = float(min_p or 0.0),
|
||||
min_tokens_to_keep = 1,
|
||||
)
|
||||
# Only build a logits processor when we actually have a non-trivial
|
||||
# repetition penalty (1.0 is the no-op value).
|
||||
logits_processors = None
|
||||
if repetition_penalty is not None and float(repetition_penalty) not in (
|
||||
0.0,
|
||||
1.0,
|
||||
):
|
||||
logits_processors = make_logits_processors(
|
||||
repetition_penalty = float(repetition_penalty),
|
||||
)
|
||||
|
||||
token_ids = []
|
||||
logger.info(
|
||||
"Generating: prompt_len=%d, max_tokens=%d, model=%s, tokenizer=%s",
|
||||
len(prompt),
|
||||
max_new_tokens,
|
||||
type(self._model).__name__,
|
||||
type(self._tokenizer).__name__,
|
||||
)
|
||||
with self._generation_lock:
|
||||
try:
|
||||
gen_kwargs = dict(
|
||||
prompt = prompt,
|
||||
max_tokens = max_new_tokens,
|
||||
sampler = sampler,
|
||||
)
|
||||
if logits_processors is not None:
|
||||
gen_kwargs["logits_processors"] = logits_processors
|
||||
for response in stream_generate(
|
||||
self._model,
|
||||
self._tokenizer,
|
||||
**gen_kwargs,
|
||||
):
|
||||
token_ids.append(response.token)
|
||||
# Decode full sequence with skip_special_tokens — same as GPU
|
||||
cumulative = self._tokenizer.decode(
|
||||
token_ids,
|
||||
skip_special_tokens = True,
|
||||
)
|
||||
yield cumulative
|
||||
|
||||
if cancel_event and cancel_event.is_set():
|
||||
break
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
logger.error("stream_generate failed:\n%s", traceback.format_exc())
|
||||
raise
|
||||
|
||||
def _generate_vlm(
|
||||
self,
|
||||
messages,
|
||||
image,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
min_p,
|
||||
max_new_tokens,
|
||||
repetition_penalty,
|
||||
cancel_event,
|
||||
):
|
||||
from mlx_vlm import stream_generate as vlm_stream
|
||||
|
||||
# Apply chat template
|
||||
chat_fn = getattr(self._processor, "apply_chat_template", None)
|
||||
if (
|
||||
chat_fn is None
|
||||
or not hasattr(self._processor, "chat_template")
|
||||
or self._processor.chat_template is None
|
||||
):
|
||||
tok = getattr(self._processor, "tokenizer", self._processor)
|
||||
chat_fn = tok.apply_chat_template
|
||||
|
||||
prompt = chat_fn(messages, tokenize = False, add_generation_prompt = True)
|
||||
|
||||
# For VLM: always use mlx_vlm's stream_generate which handles
|
||||
# pixel_values properly (passes None for text-only, image for VLM)
|
||||
images = [image] if image is not None else None
|
||||
|
||||
cumulative = ""
|
||||
logger.info(
|
||||
"VLM generating: prompt_len=%d, has_image=%s",
|
||||
len(prompt),
|
||||
image is not None,
|
||||
)
|
||||
# mlx_vlm.stream_generate forwards **kwargs into generate_step, which
|
||||
# accepts temp/top_p/top_k/repetition_penalty (and builds the sampler
|
||||
# + logits_processors internally). Pass them through.
|
||||
# NOTE: mlx_vlm.generate_step expects ``temperature=`` (long form) —
|
||||
# passing ``temp=`` silently falls into **kwargs and is ignored,
|
||||
# leaving generation stuck at the default 0.0 (greedy).
|
||||
vlm_kwargs = dict(
|
||||
max_tokens = max_new_tokens,
|
||||
temperature = temperature,
|
||||
top_p = top_p,
|
||||
top_k = int(top_k or 0),
|
||||
min_p = float(min_p or 0.0),
|
||||
)
|
||||
if repetition_penalty is not None and float(repetition_penalty) not in (
|
||||
0.0,
|
||||
1.0,
|
||||
):
|
||||
vlm_kwargs["repetition_penalty"] = float(repetition_penalty)
|
||||
|
||||
with self._generation_lock:
|
||||
for response in vlm_stream(
|
||||
self._model,
|
||||
self._processor,
|
||||
prompt,
|
||||
images,
|
||||
**vlm_kwargs,
|
||||
):
|
||||
token_text = (
|
||||
response.text if hasattr(response, "text") else str(response)
|
||||
)
|
||||
cumulative += token_text
|
||||
yield cumulative
|
||||
if cancel_event and cancel_event.is_set():
|
||||
break
|
||||
|
||||
def generate_with_adapter_control(
|
||||
self, use_adapter = None, cancel_event = None, **gen_kwargs
|
||||
) -> Generator[str, None, None]:
|
||||
# MLX LoRA adapter toggling not yet supported — generate normally
|
||||
yield from self.generate_chat_response(cancel_event = cancel_event, **gen_kwargs)
|
||||
|
||||
def reset_generation_state(self):
|
||||
import mlx.core as mx
|
||||
import gc
|
||||
|
||||
gc.collect()
|
||||
mx.clear_cache()
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
157
studio/backend/tests/test_mlx_inference_backend.py
Normal file
157
studio/backend/tests/test_mlx_inference_backend.py
Normal file
|
|
@ -0,0 +1,157 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
|
||||
import sys
|
||||
import types
|
||||
from types import SimpleNamespace
|
||||
|
||||
|
||||
class _DummyMetal:
|
||||
@staticmethod
|
||||
def is_available():
|
||||
return False
|
||||
|
||||
|
||||
class _DummyMX:
|
||||
metal = _DummyMetal()
|
||||
|
||||
@staticmethod
|
||||
def set_wired_limit(_limit):
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def device_info():
|
||||
return {"max_recommended_working_set_size": 1024}
|
||||
|
||||
|
||||
class _DummyTokenizer:
|
||||
pass
|
||||
|
||||
|
||||
class _DummyProcessor:
|
||||
tokenizer = _DummyTokenizer()
|
||||
|
||||
|
||||
class _DummyModel:
|
||||
pass
|
||||
|
||||
|
||||
def _install_fake_mlx(monkeypatch):
|
||||
mlx_pkg = types.ModuleType("mlx")
|
||||
mlx_core = types.ModuleType("mlx.core")
|
||||
mlx_core.metal = _DummyMetal()
|
||||
mlx_core.set_wired_limit = _DummyMX.set_wired_limit
|
||||
mlx_core.device_info = _DummyMX.device_info
|
||||
mlx_pkg.core = mlx_core
|
||||
monkeypatch.setitem(sys.modules, "mlx", mlx_pkg)
|
||||
monkeypatch.setitem(sys.modules, "mlx.core", mlx_core)
|
||||
|
||||
|
||||
def _install_fake_fast_mlx(monkeypatch, calls):
|
||||
class _FastMLXModel:
|
||||
@staticmethod
|
||||
def from_pretrained(*args, **kwargs):
|
||||
calls.append((args, kwargs))
|
||||
if kwargs["text_only"] is False:
|
||||
return _DummyModel(), _DummyProcessor()
|
||||
return _DummyModel(), _DummyTokenizer()
|
||||
|
||||
unsloth_zoo_pkg = types.ModuleType("unsloth_zoo")
|
||||
mlx_loader = types.ModuleType("unsloth_zoo.mlx_loader")
|
||||
mlx_loader.FastMLXModel = _FastMLXModel
|
||||
unsloth_zoo_pkg.mlx_loader = mlx_loader
|
||||
monkeypatch.setitem(sys.modules, "unsloth_zoo", unsloth_zoo_pkg)
|
||||
monkeypatch.setitem(sys.modules, "unsloth_zoo.mlx_loader", mlx_loader)
|
||||
|
||||
|
||||
def test_mlx_inference_text_load_forwards_studio_settings(monkeypatch):
|
||||
_install_fake_mlx(monkeypatch)
|
||||
calls = []
|
||||
_install_fake_fast_mlx(monkeypatch, calls)
|
||||
|
||||
from core.inference.mlx_inference import MLXInferenceBackend
|
||||
|
||||
backend = MLXInferenceBackend()
|
||||
config = SimpleNamespace(identifier = "fake/text", is_vision = False, is_lora = False)
|
||||
|
||||
assert backend.load_model(
|
||||
config,
|
||||
max_seq_length = 4096,
|
||||
load_in_4bit = False,
|
||||
hf_token = "hf-token",
|
||||
trust_remote_code = True,
|
||||
dtype = "float16",
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
(
|
||||
("fake/text",),
|
||||
{
|
||||
"max_seq_length": 4096,
|
||||
"dtype": "float16",
|
||||
"load_in_4bit": False,
|
||||
"token": "hf-token",
|
||||
"trust_remote_code": True,
|
||||
"text_only": True,
|
||||
},
|
||||
)
|
||||
]
|
||||
assert backend._is_vlm is False
|
||||
assert isinstance(backend._tokenizer, _DummyTokenizer)
|
||||
|
||||
|
||||
def test_mlx_inference_vlm_lora_uses_unsloth_loader_without_native_adapter_rewrite(
|
||||
monkeypatch,
|
||||
tmp_path,
|
||||
):
|
||||
_install_fake_mlx(monkeypatch)
|
||||
calls = []
|
||||
_install_fake_fast_mlx(monkeypatch, calls)
|
||||
|
||||
def _native_vlm_load(*_args, **_kwargs):
|
||||
raise AssertionError("Studio MLX VLM inference must use FastMLXModel")
|
||||
|
||||
mlx_vlm = types.ModuleType("mlx_vlm")
|
||||
mlx_vlm.load = _native_vlm_load
|
||||
monkeypatch.setitem(sys.modules, "mlx_vlm", mlx_vlm)
|
||||
|
||||
adapter_dir = tmp_path / "adapter"
|
||||
adapter_dir.mkdir()
|
||||
cfg_path = adapter_dir / "adapter_config.json"
|
||||
original_cfg = '{"base_model_name_or_path": "fake/base", "rank": 8}\n'
|
||||
cfg_path.write_text(original_cfg)
|
||||
|
||||
from core.inference.mlx_inference import MLXInferenceBackend
|
||||
|
||||
backend = MLXInferenceBackend()
|
||||
config = SimpleNamespace(
|
||||
identifier = str(adapter_dir),
|
||||
is_vision = True,
|
||||
is_lora = True,
|
||||
base_model = "fake/base",
|
||||
)
|
||||
|
||||
assert backend.load_model(
|
||||
config,
|
||||
max_seq_length = 8192,
|
||||
load_in_4bit = True,
|
||||
hf_token = "hf-token",
|
||||
trust_remote_code = True,
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
(
|
||||
(str(adapter_dir),),
|
||||
{
|
||||
"max_seq_length": 8192,
|
||||
"dtype": None,
|
||||
"load_in_4bit": True,
|
||||
"token": "hf-token",
|
||||
"trust_remote_code": True,
|
||||
"text_only": False,
|
||||
},
|
||||
)
|
||||
]
|
||||
assert cfg_path.read_text() == original_cfg
|
||||
assert backend._is_vlm is True
|
||||
assert isinstance(backend._processor, _DummyProcessor)
|
||||
assert isinstance(backend._tokenizer, _DummyTokenizer)
|
||||
83
studio/backend/tests/test_mlx_training_worker_config.py
Normal file
83
studio/backend/tests/test_mlx_training_worker_config.py
Normal file
|
|
@ -0,0 +1,83 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _load_worker_module():
|
||||
stub_names = (
|
||||
"structlog",
|
||||
"loggers",
|
||||
"utils",
|
||||
"utils.hardware",
|
||||
"utils.wheel_utils",
|
||||
)
|
||||
previous_modules = {name: sys.modules.get(name) for name in stub_names}
|
||||
|
||||
try:
|
||||
sys.modules["structlog"] = types.ModuleType("structlog")
|
||||
|
||||
loggers = types.ModuleType("loggers")
|
||||
loggers.get_logger = lambda *_args, **_kwargs: None
|
||||
sys.modules["loggers"] = loggers
|
||||
|
||||
utils = types.ModuleType("utils")
|
||||
utils.__path__ = []
|
||||
sys.modules["utils"] = utils
|
||||
|
||||
hardware = types.ModuleType("utils.hardware")
|
||||
hardware.apply_gpu_ids = lambda *_args, **_kwargs: None
|
||||
sys.modules["utils.hardware"] = hardware
|
||||
|
||||
wheel_utils = types.ModuleType("utils.wheel_utils")
|
||||
for name in (
|
||||
"direct_wheel_url",
|
||||
"flash_attn_wheel_url",
|
||||
"install_wheel",
|
||||
"probe_torch_wheel_env",
|
||||
"url_exists",
|
||||
):
|
||||
setattr(wheel_utils, name, lambda *_args, **_kwargs: None)
|
||||
sys.modules["utils.wheel_utils"] = wheel_utils
|
||||
|
||||
worker_path = (
|
||||
Path(__file__).resolve().parents[1] / "core" / "training" / "worker.py"
|
||||
)
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
"mlx_training_worker_under_test", worker_path
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
finally:
|
||||
for name, module in previous_modules.items():
|
||||
if module is None:
|
||||
sys.modules.pop(name, None)
|
||||
else:
|
||||
sys.modules[name] = module
|
||||
|
||||
|
||||
_worker = _load_worker_module()
|
||||
_normalize_mlx_studio_optimizer = _worker._normalize_mlx_studio_optimizer
|
||||
_normalize_mlx_studio_scheduler = _worker._normalize_mlx_studio_scheduler
|
||||
|
||||
|
||||
def test_mlx_studio_optimizer_aliases_are_explicit():
|
||||
assert _normalize_mlx_studio_optimizer("adamw_8bit") == "adamw"
|
||||
assert _normalize_mlx_studio_optimizer("paged_adamw_8bit") == "adamw"
|
||||
assert _normalize_mlx_studio_optimizer("adafactor") == "adafactor"
|
||||
|
||||
|
||||
def test_mlx_studio_rejects_unknown_optimizer():
|
||||
with pytest.raises(ValueError, match = "Unsupported optimizer for MLX training"):
|
||||
_normalize_mlx_studio_optimizer("adamw_typo")
|
||||
|
||||
|
||||
def test_mlx_studio_rejects_unknown_scheduler():
|
||||
with pytest.raises(ValueError, match = "Unsupported LR scheduler for MLX training"):
|
||||
_normalize_mlx_studio_scheduler("linear_typo")
|
||||
183
studio/backend/tests/test_training_raw_support.py
Normal file
183
studio/backend/tests/test_training_raw_support.py
Normal file
|
|
@ -0,0 +1,183 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
import asyncio
|
||||
import importlib.util
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from datasets import Dataset
|
||||
|
||||
from core.training.training import TrainingBackend
|
||||
from models.training import TrainingStartRequest
|
||||
from utils.datasets import format_dataset, format_and_template_dataset
|
||||
from utils.datasets.raw_text import prepare_raw_text_dataset
|
||||
|
||||
_BACKEND_ROOT = Path(__file__).resolve().parent.parent
|
||||
|
||||
|
||||
def _load_route_module(name: str, relative_path: str):
|
||||
spec = importlib.util.spec_from_file_location(name, _BACKEND_ROOT / relative_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
class TestTrainingRawSupport(unittest.TestCase):
|
||||
def test_training_backend_preserves_cpt_4bit_and_embedding_lr(self):
|
||||
backend = TrainingBackend()
|
||||
|
||||
class DummyProcess:
|
||||
pid = 12345
|
||||
|
||||
def start(self):
|
||||
return None
|
||||
|
||||
class DummyThread:
|
||||
def start(self):
|
||||
return None
|
||||
|
||||
dummy_queue = object()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"core.training.training.prepare_gpu_selection",
|
||||
return_value = ([0], {"selection_mode": "auto"}),
|
||||
),
|
||||
patch(
|
||||
"core.training.training._CTX.Queue",
|
||||
side_effect = [dummy_queue, dummy_queue],
|
||||
),
|
||||
patch(
|
||||
"core.training.training._CTX.Process", return_value = DummyProcess()
|
||||
) as mock_process,
|
||||
patch(
|
||||
"core.training.training.threading.Thread",
|
||||
return_value = DummyThread(),
|
||||
),
|
||||
):
|
||||
backend.start_training(
|
||||
job_id = "test-cpt-raw",
|
||||
model_name = "unsloth/test-bnb-4bit",
|
||||
training_type = "Continued Pretraining",
|
||||
format_type = "raw",
|
||||
load_in_4bit = True,
|
||||
embedding_learning_rate = 1e-5,
|
||||
)
|
||||
|
||||
config = mock_process.call_args.kwargs["kwargs"]["config"]
|
||||
self.assertTrue(config["load_in_4bit"])
|
||||
self.assertEqual(config["embedding_learning_rate"], 1e-5)
|
||||
|
||||
def test_training_route_forwards_embedding_learning_rate(self):
|
||||
training_route = _load_route_module(
|
||||
"training_route_module_raw_support",
|
||||
"routes/training.py",
|
||||
)
|
||||
captured: dict = {}
|
||||
|
||||
class DummyBackend:
|
||||
current_job_id = None
|
||||
|
||||
def is_training_active(self):
|
||||
return False
|
||||
|
||||
def start_training(self, **kwargs):
|
||||
captured.update(kwargs)
|
||||
return True
|
||||
|
||||
request = TrainingStartRequest(
|
||||
model_name = "unsloth/test-bnb-4bit",
|
||||
training_type = "Continued Pretraining",
|
||||
format_type = "raw",
|
||||
load_in_4bit = True,
|
||||
embedding_learning_rate = 1e-5,
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
training_route,
|
||||
"get_training_backend",
|
||||
return_value = DummyBackend(),
|
||||
),
|
||||
patch.object(training_route, "load_model_defaults", return_value = {}),
|
||||
patch(
|
||||
"core.inference.get_inference_backend",
|
||||
return_value = type(
|
||||
"InferenceBackend",
|
||||
(),
|
||||
{"active_model_name": None},
|
||||
)(),
|
||||
),
|
||||
patch(
|
||||
"core.export.get_export_backend",
|
||||
return_value = type(
|
||||
"ExportBackend",
|
||||
(),
|
||||
{"current_checkpoint": None},
|
||||
)(),
|
||||
),
|
||||
):
|
||||
response = asyncio.run(
|
||||
training_route.start_training(request, current_subject = "test-user")
|
||||
)
|
||||
|
||||
self.assertEqual(response.status, "queued")
|
||||
self.assertEqual(captured["embedding_learning_rate"], 1e-5)
|
||||
self.assertTrue(captured["load_in_4bit"])
|
||||
|
||||
def test_format_dataset_supports_raw_text(self):
|
||||
dataset = Dataset.from_dict(
|
||||
{
|
||||
"body": ["hello", "world"],
|
||||
"title": ["a", "b"],
|
||||
"id": [1, 2],
|
||||
}
|
||||
)
|
||||
|
||||
result = format_dataset(dataset, format_type = "raw")
|
||||
|
||||
self.assertEqual(result["final_format"], "raw_text")
|
||||
self.assertIn("text", result["dataset"].column_names)
|
||||
self.assertEqual(result["dataset"][0]["text"], "hello")
|
||||
self.assertFalse(result["requires_manual_mapping"])
|
||||
|
||||
def test_format_and_template_dataset_supports_raw_text_without_template(self):
|
||||
dataset = Dataset.from_dict({"body": ["hello raw world"]})
|
||||
|
||||
result = format_and_template_dataset(
|
||||
dataset,
|
||||
model_name = "unsloth/test",
|
||||
tokenizer = None,
|
||||
format_type = "raw",
|
||||
)
|
||||
|
||||
self.assertTrue(result["success"])
|
||||
self.assertEqual(result["final_format"], "raw_text")
|
||||
self.assertEqual(result["dataset"][0]["text"], "hello raw world")
|
||||
|
||||
def test_prepare_raw_text_dataset_drops_null_rows_before_appending_eos(self):
|
||||
dataset = Dataset.from_dict({"text": ["hello", None, "world"]})
|
||||
|
||||
result = prepare_raw_text_dataset(
|
||||
dataset,
|
||||
mode_label = "CPT",
|
||||
split_name = "train",
|
||||
eos_token = "<eos>",
|
||||
append_eos = True,
|
||||
)
|
||||
|
||||
self.assertEqual(len(result.dataset), 2)
|
||||
self.assertEqual(result.dataset[0]["text"], "hello<eos>")
|
||||
self.assertEqual(result.dataset[1]["text"], "world<eos>")
|
||||
self.assertTrue(
|
||||
any(
|
||||
"null or non-string 'text' values" in notice.message
|
||||
for notice in result.notices
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
142
studio/backend/utils/datasets/raw_text.py
Normal file
142
studio/backend/utils/datasets/raw_text.py
Normal file
|
|
@ -0,0 +1,142 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
"""
|
||||
Shared helpers for raw-text dataset preparation.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Literal
|
||||
|
||||
from datasets import Dataset
|
||||
|
||||
|
||||
@dataclass(frozen = True)
|
||||
class RawTextNotice:
|
||||
message: str
|
||||
level: Literal["info", "warning"]
|
||||
update_status: bool = False
|
||||
|
||||
|
||||
@dataclass(frozen = True)
|
||||
class RawTextPreparationResult:
|
||||
dataset: Dataset
|
||||
notices: list[RawTextNotice]
|
||||
|
||||
|
||||
def _string_columns(dataset: Dataset) -> list[str]:
|
||||
feature_map = getattr(dataset, "features", {}) or {}
|
||||
string_cols: list[str] = []
|
||||
for col in dataset.column_names:
|
||||
feature = feature_map.get(col)
|
||||
dtype = str(getattr(feature, "dtype", ""))
|
||||
if dtype in {"string", "large_string"}:
|
||||
string_cols.append(col)
|
||||
return string_cols
|
||||
|
||||
|
||||
def _split_scope(split_name: str | None) -> str:
|
||||
return f"the {split_name} split" if split_name else "this dataset"
|
||||
|
||||
|
||||
def _drop_invalid_text_rows(
|
||||
dataset: Dataset,
|
||||
*,
|
||||
mode_title: str,
|
||||
split_scope: str,
|
||||
) -> tuple[Dataset, list[RawTextNotice]]:
|
||||
filtered_dataset = dataset.filter(lambda ex: isinstance(ex["text"], str))
|
||||
dropped_rows = len(dataset) - len(filtered_dataset)
|
||||
if not dropped_rows:
|
||||
return filtered_dataset, []
|
||||
|
||||
if len(filtered_dataset) == 0:
|
||||
raise ValueError(
|
||||
f"{mode_title} training requires at least one string 'text' value "
|
||||
f"in {split_scope}; all {dropped_rows} rows were null or non-string."
|
||||
)
|
||||
|
||||
return filtered_dataset, [
|
||||
RawTextNotice(
|
||||
message = (
|
||||
f"{mode_title}: dropped {dropped_rows:,} row(s) with null or "
|
||||
f"non-string 'text' values from {split_scope}"
|
||||
),
|
||||
level = "warning",
|
||||
update_status = True,
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
def prepare_raw_text_dataset(
|
||||
dataset: Dataset,
|
||||
*,
|
||||
mode_label: str = "raw text",
|
||||
split_name: str | None = None,
|
||||
eos_token: str | None = None,
|
||||
append_eos: bool = False,
|
||||
) -> RawTextPreparationResult:
|
||||
notices: list[RawTextNotice] = []
|
||||
mode_title = mode_label.capitalize()
|
||||
split_scope = _split_scope(split_name)
|
||||
|
||||
if "text" not in dataset.column_names:
|
||||
string_cols = _string_columns(dataset)
|
||||
if not string_cols:
|
||||
raise ValueError(
|
||||
f"{mode_title} training requires a string 'text' column but none "
|
||||
f"was found in {split_scope} (columns: {dataset.column_names})."
|
||||
)
|
||||
|
||||
renamed_col = string_cols[0]
|
||||
if len(string_cols) > 1:
|
||||
notices.append(
|
||||
RawTextNotice(
|
||||
message = (
|
||||
f"{mode_title}: dataset has {len(string_cols)} string "
|
||||
f"columns ({string_cols}); auto-selecting '{renamed_col}' "
|
||||
"as the training text. Rename the intended column to "
|
||||
"'text' to override."
|
||||
),
|
||||
level = "warning",
|
||||
update_status = True,
|
||||
)
|
||||
)
|
||||
notices.append(
|
||||
RawTextNotice(
|
||||
message = (
|
||||
f"{mode_title}: renaming column '{renamed_col}' -> 'text' "
|
||||
f"for {split_scope}"
|
||||
),
|
||||
level = "info",
|
||||
)
|
||||
)
|
||||
dataset = dataset.rename_column(renamed_col, "text")
|
||||
|
||||
dataset, invalid_row_notices = _drop_invalid_text_rows(
|
||||
dataset,
|
||||
mode_title = mode_title,
|
||||
split_scope = split_scope,
|
||||
)
|
||||
notices.extend(invalid_row_notices)
|
||||
|
||||
if append_eos:
|
||||
if not eos_token:
|
||||
notices.append(
|
||||
RawTextNotice(
|
||||
message = (
|
||||
f"{mode_title}: tokenizer has no eos_token; skipping EOS "
|
||||
"append. Model will not learn document boundaries."
|
||||
),
|
||||
level = "warning",
|
||||
)
|
||||
)
|
||||
else:
|
||||
|
||||
def _append_eos(ex, _eos = eos_token):
|
||||
text = ex["text"]
|
||||
return {"text": text if text.endswith(_eos) else text + _eos}
|
||||
|
||||
dataset = dataset.map(_append_eos)
|
||||
|
||||
return RawTextPreparationResult(dataset = dataset, notices = notices)
|
||||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -38,11 +38,11 @@ import {
|
|||
Download03Icon,
|
||||
GemIcon,
|
||||
Globe02Icon,
|
||||
HelpCircleIcon,
|
||||
Search01Icon,
|
||||
PowerIcon,
|
||||
PencilEdit02Icon,
|
||||
LayoutAlignLeftIcon,
|
||||
HelpCircleIcon,
|
||||
Settings02Icon,
|
||||
ZapIcon,
|
||||
} from "@hugeicons/core-free-icons";
|
||||
|
|
|
|||
|
|
@ -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 }) => {
|
||||
|
|
|
|||
|
|
@ -50,7 +50,7 @@ export async function fetchDeviceType(): Promise<DeviceType> {
|
|||
if (res.ok) {
|
||||
const data = (await res.json()) as { device_type?: string; chat_only?: boolean };
|
||||
const deviceType = data.device_type ?? detectLocalPlatform();
|
||||
const chatOnly = data.chat_only ?? deviceType === "mac";
|
||||
const chatOnly = data.chat_only ?? false;
|
||||
usePlatformStore.setState({ deviceType, chatOnly, fetched: true });
|
||||
return deviceType;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -74,6 +74,7 @@ export const METHOD_LABELS: Record<TrainingMethod, string> = {
|
|||
qlora: "QLoRA",
|
||||
lora: "LoRA",
|
||||
full: "Full Fine-tune",
|
||||
cpt: "Continued Pretraining",
|
||||
};
|
||||
|
||||
export const GUIDE_STEPS = [
|
||||
|
|
|
|||
|
|
@ -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() {
|
||||
|
|
|
|||
|
|
@ -366,6 +366,7 @@ export function ModelSelectionStep() {
|
|||
<SelectItem value="qlora">QLoRA (4-bit)</SelectItem>
|
||||
<SelectItem value="lora">LoRA (16-bit)</SelectItem>
|
||||
<SelectItem value="full">Full Fine-tune</SelectItem>
|
||||
<SelectItem value="cpt">Continued Pretraining</SelectItem>
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ import { Badge } from "@/components/ui/badge";
|
|||
import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
|
||||
import { Separator } from "@/components/ui/separator";
|
||||
import { useTrainingConfigStore } from "@/features/training";
|
||||
import { getTrainingMethodLabel } from "@/features/training/lib/training-methods";
|
||||
import { useHardwareInfo } from "@/hooks";
|
||||
import { isAdapterMethod } from "@/types/training";
|
||||
import { ChipIcon, Database02Icon, GpuIcon, Settings04Icon } from "@hugeicons/core-free-icons";
|
||||
|
|
@ -102,6 +103,7 @@ export function SummaryStep() {
|
|||
|
||||
const showLoraParams = isAdapterMethod(trainingMethod);
|
||||
const datasetName = datasetSource === "upload" ? uploadedFile : dataset;
|
||||
const trainingMethodLabel = getTrainingMethodLabel(trainingMethod);
|
||||
|
||||
return (
|
||||
<div className="grid grid-cols-2 gap-3">
|
||||
|
|
@ -150,7 +152,7 @@ export function SummaryStep() {
|
|||
<Separator className="my-2" />
|
||||
<div className="space-y-1 text-sm">
|
||||
<Row label="Type" value={modelType} capitalize />
|
||||
<Row label="Method" value={trainingMethod === "qlora" ? "QLoRA" : trainingMethod === "lora" ? "LoRA" : "Full"} />
|
||||
<Row label="Method" value={trainingMethodLabel} />
|
||||
</div>
|
||||
</CardContent>
|
||||
</Card>
|
||||
|
|
@ -199,7 +201,7 @@ export function SummaryStep() {
|
|||
<div className="flex flex-1 flex-col">
|
||||
<span className="text-xs text-muted-foreground">Training</span>
|
||||
<span className="text-sm font-medium">
|
||||
{trainingMethod === "qlora" ? "QLoRA" : trainingMethod === "lora" ? "LoRA" : "Full"}
|
||||
{trainingMethodLabel}
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@
|
|||
import type { TrainingViewData } from "@/features/training";
|
||||
import { getTrainingRun } from "@/features/training";
|
||||
import type { TrainingRunDetailResponse } from "@/features/training";
|
||||
import { parseBackendTrainingMethod } from "@/features/training/lib/training-methods";
|
||||
import { type ReactElement, useEffect, useState } from "react";
|
||||
import { ChartsSection } from "./sections/charts-section";
|
||||
import { ProgressSection } from "./sections/progress-section";
|
||||
|
|
@ -12,15 +13,6 @@ interface HistoricalTrainingViewProps {
|
|||
runId: string;
|
||||
}
|
||||
|
||||
function normalizeTrainingMethod(config: Record<string, unknown>): string {
|
||||
const type = config?.training_type as string | undefined;
|
||||
if (!type || type === "Full Finetuning") return "full";
|
||||
if (type === "LoRA/QLoRA") {
|
||||
return config?.load_in_4bit ? "qlora" : "lora";
|
||||
}
|
||||
return "full";
|
||||
}
|
||||
|
||||
function mapToViewData(detail: TrainingRunDetailResponse): TrainingViewData {
|
||||
const { run, metrics } = detail;
|
||||
|
||||
|
|
@ -79,7 +71,10 @@ function mapToViewData(detail: TrainingRunDetailResponse): TrainingViewData {
|
|||
error: run.status === "error" ? run.error_message : null,
|
||||
isTrainingRunning: false,
|
||||
modelName: run.model_name,
|
||||
trainingMethod: normalizeTrainingMethod(detail.config),
|
||||
trainingMethod: parseBackendTrainingMethod(
|
||||
detail.config?.training_type,
|
||||
detail.config?.load_in_4bit,
|
||||
),
|
||||
lossHistory,
|
||||
lrHistory,
|
||||
gradNormHistory,
|
||||
|
|
@ -143,7 +138,11 @@ export function HistoricalTrainingView({
|
|||
loraRank: detail.config.lora_r as number | undefined,
|
||||
loraAlpha: detail.config.lora_alpha as number | undefined,
|
||||
loraDropout: detail.config.lora_dropout as number | undefined,
|
||||
loraVariant: detail.config.use_rslora ? "rsLoRA" : undefined,
|
||||
loraVariant: detail.config.use_rslora
|
||||
? "rslora"
|
||||
: detail.config.use_loftq
|
||||
? "loftq"
|
||||
: "lora",
|
||||
}
|
||||
: undefined;
|
||||
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import { Badge } from "@/components/ui/badge";
|
|||
import { Spinner } from "@/components/ui/spinner";
|
||||
import { useTrainingActions, useTrainingConfigStore } from "@/features/training";
|
||||
import { checkDatasetFormat } from "@/features/training/api/datasets-api";
|
||||
import { isRawTextDatasetFormat } from "@/features/training/lib/training-methods";
|
||||
import type { CheckFormatResponse } from "@/features/training/types/datasets";
|
||||
import { Database02Icon, AlertCircleIcon } from "@hugeicons/core-free-icons";
|
||||
import { HugeiconsIcon } from "@hugeicons/react";
|
||||
|
|
@ -90,10 +91,11 @@ export function DatasetPreviewDialog({
|
|||
const effectiveIsAudio = !!data?.is_audio;
|
||||
const effectiveIsVlm = isVlm || !!data?.is_image;
|
||||
|
||||
const isRawFormat = isRawTextDatasetFormat(datasetFormat);
|
||||
const hasHeuristicMapping = !data?.requires_manual_mapping && !!data?.suggested_mapping;
|
||||
const mappingEnabled = !!data?.requires_manual_mapping || hasHeuristicMapping;
|
||||
const mappingEnabled = !isRawFormat && (!!data?.requires_manual_mapping || hasHeuristicMapping);
|
||||
const showMappingFooter = mode === "mapping" && mappingEnabled;
|
||||
const mappingOk = isMappingComplete(manualMapping, effectiveIsVlm, datasetFormat, effectiveIsAudio);
|
||||
const mappingOk = isRawFormat || isMappingComplete(manualMapping, effectiveIsVlm, datasetFormat, effectiveIsAudio);
|
||||
const availableRoles = getAvailableRoles(effectiveIsVlm, datasetFormat, effectiveIsAudio);
|
||||
const isHfDataset = datasetSource === "huggingface";
|
||||
|
||||
|
|
@ -413,7 +415,7 @@ export function DatasetPreviewDialog({
|
|||
<MetaRow label="Source" value={sourceLabel} />
|
||||
<MetaRow
|
||||
label="Format"
|
||||
value={data.detected_format || "--"}
|
||||
value={isRawFormat ? "Raw Text" : (data.detected_format || "--")}
|
||||
/>
|
||||
<MetaRow
|
||||
label="Total Rows"
|
||||
|
|
@ -441,7 +443,7 @@ export function DatasetPreviewDialog({
|
|||
/>
|
||||
</div>
|
||||
|
||||
{data.warning && (
|
||||
{data.warning && !isRawFormat && (
|
||||
<div className="rounded-lg border border-amber-200 bg-amber-50 px-4 py-3 text-xs text-amber-700 dark:border-amber-800 dark:bg-amber-950 dark:text-amber-400 mb-4 flex items-start gap-2.5">
|
||||
<HugeiconsIcon icon={AlertCircleIcon} className="size-4 shrink-0 mt-0.5" />
|
||||
<span>{data.warning}</span>
|
||||
|
|
|
|||
|
|
@ -913,6 +913,7 @@ export function DatasetSection() {
|
|||
<SelectItem value="alpaca">Alpaca</SelectItem>
|
||||
<SelectItem value="chatml">ChatML</SelectItem>
|
||||
<SelectItem value="sharegpt">ShareGPT</SelectItem>
|
||||
<SelectItem value="raw">Raw Text</SelectItem>
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -67,6 +67,7 @@ const METHOD_DOTS: Record<string, string> = {
|
|||
qlora: "bg-emerald-400",
|
||||
lora: "bg-blue-400",
|
||||
full: "bg-amber-400",
|
||||
cpt: "bg-purple-400",
|
||||
};
|
||||
|
||||
const DARK_TRIGGER =
|
||||
|
|
@ -570,7 +571,9 @@ export function ModelSection() {
|
|||
</TooltipTrigger>
|
||||
<TooltipContent className="max-w-xs">
|
||||
QLoRA uses 4-bit quantization for lowest VRAM. LoRA uses
|
||||
16-bit. Full updates all weights.{" "}
|
||||
16-bit. Full updates all weights. CPT (Continued Pretraining)
|
||||
trains on raw text to adapt the model to a new domain without
|
||||
chat formatting.{" "}
|
||||
<a
|
||||
href="https://unsloth.ai/docs/get-started/fine-tuning-llms-guide/lora-hyperparameters-guide"
|
||||
target="_blank"
|
||||
|
|
@ -617,6 +620,14 @@ export function ModelSection() {
|
|||
Full Fine-tune
|
||||
</span>
|
||||
</SelectItem>
|
||||
<SelectItem value="cpt">
|
||||
<span className="flex items-center gap-2">
|
||||
<span
|
||||
className={`size-2 shrink-0 rounded-full ${METHOD_DOTS.cpt}`}
|
||||
/>
|
||||
Continued Pretraining
|
||||
</span>
|
||||
</SelectItem>
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
// SPDX-License-Identifier: AGPL-3.0-only
|
||||
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
import { usePlatformStore } from "@/config/env";
|
||||
import { SectionCard } from "@/components/section-card";
|
||||
import { Checkbox } from "@/components/ui/checkbox";
|
||||
import {
|
||||
|
|
@ -33,11 +34,14 @@ import {
|
|||
} from "@/components/ui/tooltip";
|
||||
import {
|
||||
CONTEXT_LENGTHS,
|
||||
CPT_TARGET_MODULES,
|
||||
LR_SCHEDULER_OPTIONS,
|
||||
OPTIMIZER_OPTIONS,
|
||||
TARGET_MODULES,
|
||||
} from "@/config/training";
|
||||
import { useMaxStepsEpochsToggle, useTrainingConfigStore } from "@/features/training";
|
||||
import { isRawTextDatasetFormat } from "@/features/training/lib/training-methods";
|
||||
import { isAdapterMethod } from "@/types/training";
|
||||
import type { GradientCheckpointing } from "@/types/training";
|
||||
import {
|
||||
ArrowDown01Icon,
|
||||
|
|
@ -124,10 +128,14 @@ function SliderRow({
|
|||
|
||||
export function ParamsSection(): ReactElement {
|
||||
const store = useTrainingConfigStore();
|
||||
const isLora = store.trainingMethod !== "full";
|
||||
const platformDeviceType = usePlatformStore((s) => s.deviceType);
|
||||
const isLora = isAdapterMethod(store.trainingMethod);
|
||||
const isCpt = store.trainingMethod === "cpt";
|
||||
const isRawText = isRawTextDatasetFormat(store.datasetFormat);
|
||||
const showVisionLora = store.isVisionModel && store.isDatasetImage === true;
|
||||
const [loraOpen, setLoraOpen] = useState(false);
|
||||
const [hyperOpen, setHyperOpen] = useState(false);
|
||||
const needsExpandedHeight = isCpt || (isLora && loraOpen) || hyperOpen;
|
||||
const [ctxInput, setCtxInput] = useState(String(store.contextLength));
|
||||
const ctxAnchorRef = useRef<HTMLDivElement>(null);
|
||||
const ctxItems = CONTEXT_LENGTHS.map(String);
|
||||
|
|
@ -166,7 +174,7 @@ export function ParamsSection(): ReactElement {
|
|||
title="Parameters"
|
||||
description="Configure training hyperparameters"
|
||||
accent="orange"
|
||||
className={`${(isLora && loraOpen) || hyperOpen
|
||||
className={`${needsExpandedHeight
|
||||
? "min-h-studio-config-column"
|
||||
: "h-studio-config-column"} duration-150`}
|
||||
>
|
||||
|
|
@ -376,10 +384,62 @@ export function ParamsSection(): ReactElement {
|
|||
className="w-full font-mono"
|
||||
/>
|
||||
<p className="text-[10px] text-muted-foreground">
|
||||
Recommended: 2e-4 for LoRA, 2e-5 for full fine-tune
|
||||
Recommended: 2e-4 for LoRA, 5e-5 for CPT, 2e-5 for full fine-tune
|
||||
</p>
|
||||
</div>
|
||||
|
||||
{/* Embedding Learning Rate (CPT only) */}
|
||||
{isCpt && (
|
||||
<div className="flex flex-col gap-2">
|
||||
<span className="flex items-center gap-1.5 text-xs font-medium text-muted-foreground">
|
||||
Embedding Learning Rate
|
||||
<Tooltip>
|
||||
<TooltipTrigger asChild={true}>
|
||||
<button
|
||||
type="button"
|
||||
className="text-foreground/70 hover:text-foreground"
|
||||
>
|
||||
<HugeiconsIcon
|
||||
icon={InformationCircleIcon}
|
||||
className="size-3"
|
||||
/>
|
||||
</button>
|
||||
</TooltipTrigger>
|
||||
<TooltipContent>
|
||||
Only used when CPT is training <code>embed_tokens</code>.
|
||||
Embeddings are easier to destabilize than LoRA weights, so
|
||||
they usually need a smaller LR. Leave blank to use
|
||||
<code>lr/10</code>; typical working range is 2x-10x smaller
|
||||
than the main LR. Increase it only if vocabulary or
|
||||
domain-token adaptation is too slow.
|
||||
</TooltipContent>
|
||||
</Tooltip>
|
||||
</span>
|
||||
<Input
|
||||
type="number"
|
||||
step="0.00001"
|
||||
min="0"
|
||||
max="1"
|
||||
placeholder={`auto (${(store.learningRate / 10).toExponential(1)})`}
|
||||
value={store.embeddingLearningRate ?? ""}
|
||||
onChange={(e) => {
|
||||
const raw = e.target.value;
|
||||
if (raw === "") {
|
||||
store.setEmbeddingLearningRate(null);
|
||||
return;
|
||||
}
|
||||
const n = Number(raw);
|
||||
store.setEmbeddingLearningRate(Number.isFinite(n) ? n : null);
|
||||
}}
|
||||
className="w-full font-mono"
|
||||
/>
|
||||
<p className="text-[10px] text-muted-foreground">
|
||||
Leave blank to use lr/10 (recommended). Typical range is
|
||||
2x-10x smaller than the main learning rate.
|
||||
</p>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* LoRA Settings */}
|
||||
{isLora && (
|
||||
<Collapsible open={loraOpen} onOpenChange={setLoraOpen}>
|
||||
|
|
@ -514,7 +574,7 @@ export function ParamsSection(): ReactElement {
|
|||
Target Modules
|
||||
</span>
|
||||
<div className="flex flex-wrap gap-1.5">
|
||||
{TARGET_MODULES.map((mod) => {
|
||||
{(isCpt ? CPT_TARGET_MODULES : TARGET_MODULES).map((mod) => {
|
||||
const active = store.targetModules.includes(mod);
|
||||
return (
|
||||
<button
|
||||
|
|
@ -883,7 +943,11 @@ export function ParamsSection(): ReactElement {
|
|||
<SelectContent>
|
||||
<SelectItem value="none">None</SelectItem>
|
||||
<SelectItem value="true">Standard</SelectItem>
|
||||
<SelectItem value="unsloth">Unsloth</SelectItem>
|
||||
{platformDeviceType === "mac" ? (
|
||||
<SelectItem value="mlx">MLX</SelectItem>
|
||||
) : (
|
||||
<SelectItem value="unsloth">Unsloth</SelectItem>
|
||||
)}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</Row>
|
||||
|
|
@ -902,7 +966,7 @@ export function ParamsSection(): ReactElement {
|
|||
</label>
|
||||
</div>
|
||||
)}
|
||||
{!store.isEmbeddingModel && (
|
||||
{!store.isEmbeddingModel && !isCpt && !isRawText && (
|
||||
<div className="flex items-center gap-2">
|
||||
<Checkbox
|
||||
id="trainOnCompletions"
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ import {
|
|||
useTrainingConfigStore,
|
||||
useTrainingRuntimeStore,
|
||||
} from "@/features/training";
|
||||
import { getTrainingMethodLabel } from "@/features/training/lib/training-methods";
|
||||
import type { TrainingViewData } from "@/features/training";
|
||||
import { useGpuUtilization } from "@/hooks";
|
||||
import { cn } from "@/lib/utils";
|
||||
|
|
@ -86,6 +87,7 @@ export function ProgressSection({
|
|||
configOverride,
|
||||
}: ProgressSectionProps): ReactElement {
|
||||
const navigate = useNavigate();
|
||||
const trainingMethodLabel = getTrainingMethodLabel(data.trainingMethod);
|
||||
|
||||
const config = useTrainingConfigStore(
|
||||
useShallow((state) => ({
|
||||
|
|
@ -272,7 +274,7 @@ export function ProgressSection({
|
|||
{data.modelName || "--"}
|
||||
</MetricStat>
|
||||
<MetricStat label="Method">
|
||||
{data.trainingMethod === "qlora" ? "QLoRA" : data.trainingMethod === "lora" ? "LoRA" : "Full"}
|
||||
{trainingMethodLabel}
|
||||
</MetricStat>
|
||||
</div>
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ import { checkDatasetFormat } from "../api/datasets-api";
|
|||
import { getTrainingRun } from "../api/history-api";
|
||||
import { buildTrainingStartPayload } from "../api/mappers";
|
||||
import { resetTraining, startTraining, stopTraining } from "../api/train-api";
|
||||
import { isRawTextDatasetFormat } from "../lib/training-methods";
|
||||
import { syncTrainingRuntimeFromBackend } from "../lib/sync-runtime";
|
||||
import { validateTrainingConfig } from "../lib/validation";
|
||||
import { useDatasetPreviewDialogStore } from "../stores/dataset-preview-dialog-store";
|
||||
|
|
@ -88,7 +89,10 @@ export function useTrainingActions() {
|
|||
});
|
||||
}
|
||||
|
||||
const needsReview = check.requires_manual_mapping || check.detected_format === "custom_heuristic";
|
||||
const isRawFormat = isRawTextDatasetFormat(config.datasetFormat);
|
||||
const needsReview =
|
||||
!isRawFormat &&
|
||||
(check.requires_manual_mapping || check.detected_format === "custom_heuristic");
|
||||
if (needsReview && !hasManualMapping(config, isVlm, isAudio)) {
|
||||
// Pre-fill from suggested_mapping or VLM detected columns
|
||||
const hint: Record<string, string> = {};
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,48 @@
|
|||
// SPDX-License-Identifier: AGPL-3.0-only
|
||||
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
import type { DatasetFormat, TrainingMethod } from "@/types/training";
|
||||
|
||||
const BACKEND_TRAINING_TYPE: Record<TrainingMethod, string> = {
|
||||
qlora: "LoRA/QLoRA",
|
||||
lora: "LoRA/QLoRA",
|
||||
full: "Full Finetuning",
|
||||
cpt: "Continued Pretraining",
|
||||
};
|
||||
|
||||
const TRAINING_METHOD_LABELS: Record<TrainingMethod, string> = {
|
||||
qlora: "QLoRA",
|
||||
lora: "LoRA",
|
||||
full: "Full",
|
||||
cpt: "CPT",
|
||||
};
|
||||
|
||||
export function toBackendTrainingType(trainingMethod: TrainingMethod): string {
|
||||
return BACKEND_TRAINING_TYPE[trainingMethod];
|
||||
}
|
||||
|
||||
export function getTrainingMethodLabel(
|
||||
trainingMethod: TrainingMethod | string,
|
||||
): string {
|
||||
if (Object.prototype.hasOwnProperty.call(TRAINING_METHOD_LABELS, trainingMethod)) {
|
||||
return TRAINING_METHOD_LABELS[trainingMethod as TrainingMethod];
|
||||
}
|
||||
return TRAINING_METHOD_LABELS.full;
|
||||
}
|
||||
|
||||
export function parseBackendTrainingMethod(
|
||||
trainingType: unknown,
|
||||
loadIn4Bit: unknown,
|
||||
): TrainingMethod {
|
||||
if (trainingType === "Continued Pretraining") return "cpt";
|
||||
if (trainingType === "LoRA/QLoRA") {
|
||||
return loadIn4Bit ? "qlora" : "lora";
|
||||
}
|
||||
return "full";
|
||||
}
|
||||
|
||||
export function isRawTextDatasetFormat(
|
||||
datasetFormat: DatasetFormat,
|
||||
): boolean {
|
||||
return datasetFormat === "raw";
|
||||
}
|
||||
|
|
@ -1,15 +1,17 @@
|
|||
// SPDX-License-Identifier: AGPL-3.0-only
|
||||
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
import { DEFAULT_HYPERPARAMS, LR_DEFAULT_FULL, LR_DEFAULT_LORA, STEPS } from "@/config/training";
|
||||
import { CPT_TARGET_MODULES, DEFAULT_HYPERPARAMS, LR_DEFAULT_CPT, LR_DEFAULT_FULL, LR_DEFAULT_LORA, STEPS, TARGET_MODULES } from "@/config/training";
|
||||
import { authFetch } from "@/features/auth";
|
||||
import { isAdapterMethod } from "@/types/training";
|
||||
import type { DatasetFormat } from "@/types/training";
|
||||
import type { ModelType, StepNumber, TrainingMethod } from "@/types/training";
|
||||
import { create } from "zustand";
|
||||
import { persist } from "zustand/middleware";
|
||||
import { checkDatasetFormat } from "../api/datasets-api";
|
||||
import { checkVisionModel, getModelConfig } from "../api/models-api";
|
||||
import { mapBackendModelConfigToTrainingPatch } from "../lib/model-defaults";
|
||||
import { isRawTextDatasetFormat } from "../lib/training-methods";
|
||||
import type { BackendModelConfig } from "../api/models-api";
|
||||
import type { TrainingConfigState, TrainingConfigStore } from "../types/config";
|
||||
|
||||
|
|
@ -108,6 +110,11 @@ let _learningRateManuallySet = false;
|
|||
// setTrainingMethod can restore it when switching back from full to adapter.
|
||||
let _yamlLearningRate: number | undefined = undefined;
|
||||
|
||||
// Track whether entering CPT auto-forced datasetFormat="raw" so that
|
||||
// leaving CPT can restore the prior user-visible format.
|
||||
let _datasetFormatBeforeCpt: DatasetFormat | null = null;
|
||||
let _datasetFormatAutoForcedByCpt = false;
|
||||
|
||||
const NON_PERSISTED_STATE_KEYS: ReadonlySet<keyof TrainingConfigState> = new Set([
|
||||
"modelType",
|
||||
"isCheckingVision",
|
||||
|
|
@ -156,6 +163,123 @@ function canProceedForStep(state: TrainingConfigState): boolean {
|
|||
}
|
||||
}
|
||||
|
||||
type TrainingMethodStatePatch = Partial<
|
||||
Pick<
|
||||
TrainingConfigState,
|
||||
| "trainingMethod"
|
||||
| "learningRate"
|
||||
| "loraRank"
|
||||
| "loraAlpha"
|
||||
| "loraVariant"
|
||||
| "targetModules"
|
||||
| "datasetFormat"
|
||||
| "trainOnCompletions"
|
||||
>
|
||||
>;
|
||||
|
||||
function getCptTrainingPatch(): TrainingMethodStatePatch {
|
||||
return {
|
||||
loraRank: 128,
|
||||
loraAlpha: 32,
|
||||
loraVariant: "rslora",
|
||||
targetModules: CPT_TARGET_MODULES,
|
||||
datasetFormat: "raw",
|
||||
trainOnCompletions: false,
|
||||
};
|
||||
}
|
||||
|
||||
function getCptModelDefaultsPatch(): TrainingMethodStatePatch {
|
||||
return {
|
||||
...getCptTrainingPatch(),
|
||||
learningRate: LR_DEFAULT_CPT,
|
||||
};
|
||||
}
|
||||
|
||||
function getRestoreFromCptPatch(): TrainingMethodStatePatch {
|
||||
return {
|
||||
loraRank: DEFAULT_HYPERPARAMS.loraRank,
|
||||
loraAlpha: DEFAULT_HYPERPARAMS.loraAlpha,
|
||||
loraVariant: DEFAULT_HYPERPARAMS.loraVariant,
|
||||
targetModules: TARGET_MODULES,
|
||||
};
|
||||
}
|
||||
|
||||
function clearCptDatasetFormatTracking(): void {
|
||||
_datasetFormatBeforeCpt = null;
|
||||
_datasetFormatAutoForcedByCpt = false;
|
||||
}
|
||||
|
||||
function recordCptDatasetFormatOverride(currentDatasetFormat: DatasetFormat): void {
|
||||
if (isRawTextDatasetFormat(currentDatasetFormat)) {
|
||||
clearCptDatasetFormatTracking();
|
||||
return;
|
||||
}
|
||||
_datasetFormatBeforeCpt = currentDatasetFormat;
|
||||
_datasetFormatAutoForcedByCpt = true;
|
||||
}
|
||||
|
||||
function getRestoreDatasetFormatFromCptPatch(): TrainingMethodStatePatch {
|
||||
if (!_datasetFormatAutoForcedByCpt || _datasetFormatBeforeCpt == null) {
|
||||
clearCptDatasetFormatTracking();
|
||||
return {};
|
||||
}
|
||||
|
||||
const previousDatasetFormat = _datasetFormatBeforeCpt;
|
||||
clearCptDatasetFormatTracking();
|
||||
return { datasetFormat: previousDatasetFormat };
|
||||
}
|
||||
|
||||
function resolveTrainingMethodLearningRate(
|
||||
prevMethod: TrainingMethod,
|
||||
nextMethod: TrainingMethod,
|
||||
): number | undefined {
|
||||
if (_learningRateManuallySet) {
|
||||
return undefined;
|
||||
}
|
||||
|
||||
const wasCpt = prevMethod === "cpt";
|
||||
const wasAdapter = isAdapterMethod(prevMethod);
|
||||
const nowAdapter = isAdapterMethod(nextMethod);
|
||||
|
||||
if (nextMethod === "cpt") {
|
||||
return LR_DEFAULT_CPT;
|
||||
}
|
||||
if (wasCpt && nowAdapter) {
|
||||
return _yamlLearningRate ?? LR_DEFAULT_LORA;
|
||||
}
|
||||
if (wasAdapter && nowAdapter) {
|
||||
return undefined;
|
||||
}
|
||||
return nowAdapter ? _yamlLearningRate ?? LR_DEFAULT_LORA : LR_DEFAULT_FULL;
|
||||
}
|
||||
|
||||
function buildTrainingMethodPatch(
|
||||
prevMethod: TrainingMethod,
|
||||
nextMethod: TrainingMethod,
|
||||
currentDatasetFormat: DatasetFormat,
|
||||
): TrainingMethodStatePatch {
|
||||
const patch: TrainingMethodStatePatch = { trainingMethod: nextMethod };
|
||||
|
||||
if (prevMethod !== "cpt" && nextMethod === "cpt") {
|
||||
recordCptDatasetFormatOverride(currentDatasetFormat);
|
||||
Object.assign(patch, getCptTrainingPatch());
|
||||
}
|
||||
if (prevMethod === "cpt" && nextMethod !== "cpt") {
|
||||
Object.assign(
|
||||
patch,
|
||||
getRestoreFromCptPatch(),
|
||||
getRestoreDatasetFormatFromCptPatch(),
|
||||
);
|
||||
}
|
||||
|
||||
const learningRate = resolveTrainingMethodLearningRate(prevMethod, nextMethod);
|
||||
if (learningRate !== undefined) {
|
||||
patch.learningRate = learningRate;
|
||||
}
|
||||
|
||||
return patch;
|
||||
}
|
||||
|
||||
export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
||||
persist(
|
||||
(set, get) => {
|
||||
|
|
@ -216,11 +340,14 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
// Auto-select training method based on model size vs GPU memory.
|
||||
// If model_size * 1.5 * context_scale fits in free VRAM, use LoRA 16-bit.
|
||||
// Otherwise use QLoRA 4-bit.
|
||||
// Auto-select LoRA vs QLoRA based on GPU memory.
|
||||
// Skip if user has manually chosen CPT -- don't override it.
|
||||
const modelSizeBytes = modelDetails.model_size_bytes;
|
||||
if (modelSizeBytes && modelSizeBytes > 0) {
|
||||
if (modelSizeBytes && modelSizeBytes > 0 && get().trainingMethod !== "cpt") {
|
||||
void autoSelectTrainingMethod(modelSizeBytes, patch.contextLength ?? get().contextLength)
|
||||
.then((method) => {
|
||||
if (get().selectedModel !== modelName) return;
|
||||
if (get().trainingMethod === "cpt") return;
|
||||
if (method) {
|
||||
const lrPatch = !_learningRateManuallySet && !modelConfigHasLR
|
||||
? { learningRate: method === "full" ? LR_DEFAULT_FULL : LR_DEFAULT_LORA }
|
||||
|
|
@ -230,8 +357,16 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
});
|
||||
}
|
||||
|
||||
// Preserve CPT hyperparams: YAML adapter defaults (r/alpha/targets/LR)
|
||||
// are tuned for standard LoRA and would otherwise clobber CPT settings.
|
||||
const cptOverrides =
|
||||
get().trainingMethod === "cpt"
|
||||
? getCptModelDefaultsPatch()
|
||||
: {};
|
||||
|
||||
set({
|
||||
...patch,
|
||||
...cptOverrides,
|
||||
modelType: inferredModelType,
|
||||
isVisionModel: modelDetails.is_vision,
|
||||
isEmbeddingModel: isEmbedding,
|
||||
|
|
@ -396,29 +531,14 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
void loadAndApplyModelDefaults(state.selectedModel);
|
||||
},
|
||||
setTrainingMethod: (trainingMethod) => {
|
||||
if (_learningRateManuallySet) {
|
||||
set({ trainingMethod });
|
||||
return;
|
||||
}
|
||||
|
||||
const prev = get().trainingMethod;
|
||||
const wasAdapter = isAdapterMethod(prev);
|
||||
const nowAdapter = isAdapterMethod(trainingMethod);
|
||||
|
||||
// qlora <-> lora: same LR range, don't touch learning rate
|
||||
if (wasAdapter && nowAdapter) {
|
||||
set({ trainingMethod });
|
||||
return;
|
||||
}
|
||||
|
||||
// Category changed (adapter <-> full)
|
||||
if (nowAdapter) {
|
||||
// Switching TO adapter: restore YAML LR if available
|
||||
set({ trainingMethod, learningRate: _yamlLearningRate ?? LR_DEFAULT_LORA });
|
||||
} else {
|
||||
// Switching TO full: no YAML full-LR exists, use constant
|
||||
set({ trainingMethod, learningRate: LR_DEFAULT_FULL });
|
||||
}
|
||||
const state = get();
|
||||
set(
|
||||
buildTrainingMethodPatch(
|
||||
state.trainingMethod,
|
||||
trainingMethod,
|
||||
state.datasetFormat,
|
||||
),
|
||||
);
|
||||
},
|
||||
setHfToken: (hfToken) =>
|
||||
set({ hfToken: hfToken.trim().replace(/^["']+|["']+$/g, "") }),
|
||||
|
|
@ -448,7 +568,26 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
runDatasetCheck(uploadedFile, "train");
|
||||
}
|
||||
},
|
||||
setDatasetFormat: (datasetFormat) => set({ datasetFormat }),
|
||||
setDatasetFormat: (datasetFormat) =>
|
||||
set((state) => {
|
||||
if (state.trainingMethod === "cpt") {
|
||||
if (isRawTextDatasetFormat(datasetFormat)) {
|
||||
clearCptDatasetFormatTracking();
|
||||
}
|
||||
return {
|
||||
datasetFormat: "raw",
|
||||
trainOnCompletions: false,
|
||||
};
|
||||
}
|
||||
|
||||
return {
|
||||
datasetFormat,
|
||||
trainOnCompletions:
|
||||
isRawTextDatasetFormat(datasetFormat)
|
||||
? false
|
||||
: state.trainOnCompletions,
|
||||
};
|
||||
}),
|
||||
setDataset: (dataset) => {
|
||||
_datasetCheckController?.abort();
|
||||
_datasetCheckController = null;
|
||||
|
|
@ -566,6 +705,8 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
_learningRateManuallySet = true;
|
||||
set({ learningRate });
|
||||
},
|
||||
setEmbeddingLearningRate: (embeddingLearningRate) =>
|
||||
set({ embeddingLearningRate }),
|
||||
setOptimizerType: (optimizerType) => set({ optimizerType }),
|
||||
setLrSchedulerType: (lrSchedulerType) => set({ lrSchedulerType }),
|
||||
setLoraRank: (loraRank) => set({ loraRank }),
|
||||
|
|
@ -608,6 +749,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
_trainOnCompletionsManuallySet = false;
|
||||
_learningRateManuallySet = false;
|
||||
_yamlLearningRate = undefined;
|
||||
clearCptDatasetFormatTracking();
|
||||
set(initialState);
|
||||
},
|
||||
resetToModelDefaults: () => {
|
||||
|
|
@ -629,7 +771,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
},
|
||||
{
|
||||
name: "unsloth_training_config_v1",
|
||||
version: 9,
|
||||
version: 10,
|
||||
migrate: (persisted, version) => {
|
||||
const s = persisted as Record<string, unknown>;
|
||||
if (version < 2 && s.datasetSubset == null && s.datasetConfig != null) {
|
||||
|
|
@ -665,6 +807,17 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
s.weightDecay = DEFAULT_HYPERPARAMS.weightDecay;
|
||||
}
|
||||
}
|
||||
if (version < 10 && s.trainingMethod === "cpt") {
|
||||
// Backfill CPT defaults for state persisted before they existed.
|
||||
s.loraRank = 128;
|
||||
s.loraAlpha = 32;
|
||||
s.loraVariant = "rslora";
|
||||
s.targetModules = CPT_TARGET_MODULES;
|
||||
s.datasetFormat = "raw";
|
||||
if (s.learningRate == null || s.learningRate === LR_DEFAULT_LORA) {
|
||||
s.learningRate = LR_DEFAULT_CPT;
|
||||
}
|
||||
}
|
||||
return s as unknown as TrainingConfigStore;
|
||||
},
|
||||
partialize: partializePersistedState,
|
||||
|
|
|
|||
|
|
@ -21,6 +21,8 @@ export interface TrainingStartRequest {
|
|||
custom_format_mapping?: Record<string, unknown> | null;
|
||||
num_epochs: number;
|
||||
learning_rate: string;
|
||||
/** Optional CPT embedding LR. If omitted, backend uses lr/10; typical range is 2x-10x smaller than main LR. */
|
||||
embedding_learning_rate?: number | null;
|
||||
batch_size: number;
|
||||
gradient_accumulation_steps: number;
|
||||
warmup_steps: number | null;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import { listModels } from "@huggingface/hub";
|
|||
import { type CachedResult, cachedModelInfo, primeCacheFromListing } from "@/lib/hf-cache";
|
||||
import { useCallback, useMemo } from "react";
|
||||
import { useHfPaginatedSearch } from "./use-hf-paginated-search";
|
||||
import { usePlatformStore } from "@/config/env";
|
||||
|
||||
export interface HfModelResult {
|
||||
id: string;
|
||||
|
|
@ -16,7 +17,8 @@ export interface HfModelResult {
|
|||
isGguf: boolean;
|
||||
}
|
||||
|
||||
const EXCLUDED_TAGS = new Set([
|
||||
/** Tags to exclude on GPU (CUDA/ROCm) — MLX models won't load on GPU. */
|
||||
const EXCLUDED_TAGS_GPU = new Set([
|
||||
"gptq",
|
||||
"awq",
|
||||
"exl2",
|
||||
|
|
@ -28,6 +30,18 @@ const EXCLUDED_TAGS = new Set([
|
|||
"ctranslate2",
|
||||
]);
|
||||
|
||||
/** Tags to exclude on MLX (Mac) — GPU-only quant formats won't load on MLX. */
|
||||
const EXCLUDED_TAGS_MLX = new Set([
|
||||
"gptq",
|
||||
"awq",
|
||||
"exl2",
|
||||
"onnx",
|
||||
"openvino",
|
||||
"coreml",
|
||||
"tflite",
|
||||
"ctranslate2",
|
||||
]);
|
||||
|
||||
// Embedding / sentence-transformer models ship with onnx/openvino as additional
|
||||
// export formats — they should not be excluded by the tag check above.
|
||||
const EMBEDDING_TAGS = new Set([
|
||||
|
|
@ -77,7 +91,7 @@ function estimateSizeFromDtypes(
|
|||
return total > 0 ? total : undefined;
|
||||
}
|
||||
|
||||
function makeMapModel(excludeGguf: boolean) {
|
||||
function makeMapModel(excludeGguf: boolean, excludedTags: Set<string>) {
|
||||
return (raw: unknown): HfModelResult | null => {
|
||||
const m = raw as {
|
||||
name: string;
|
||||
|
|
@ -87,7 +101,7 @@ function makeMapModel(excludeGguf: boolean) {
|
|||
tags?: string[];
|
||||
};
|
||||
const isEmbedding = m.tags?.some((t) => EMBEDDING_TAGS.has(t));
|
||||
if (!isEmbedding && m.tags?.some((t) => EXCLUDED_TAGS.has(t))) {
|
||||
if (!isEmbedding && m.tags?.some((t) => excludedTags.has(t))) {
|
||||
return null;
|
||||
}
|
||||
const isGguf =
|
||||
|
|
@ -314,7 +328,9 @@ export function useHfModelSearch(
|
|||
[trimmed, searchQuery, pinnedId, task, accessToken, priorityIds],
|
||||
);
|
||||
|
||||
const mapModel = useMemo(() => makeMapModel(excludeGguf), [excludeGguf]);
|
||||
const deviceType = usePlatformStore((s) => s.deviceType);
|
||||
const excludedTags = deviceType === "mac" ? EXCLUDED_TAGS_MLX : EXCLUDED_TAGS_GPU;
|
||||
const mapModel = useMemo(() => makeMapModel(excludeGguf, excludedTags), [excludeGguf, excludedTags]);
|
||||
const search = useHfPaginatedSearch(createIter, mapModel);
|
||||
|
||||
// Secondary sort guarantee: unsloth models always float to the top.
|
||||
|
|
|
|||
|
|
@ -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 });
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
211
studio/setup.ps1
211
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 {
|
||||
|
|
|
|||
103
studio/setup.sh
103
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
|
||||
|
|
|
|||
|
|
@ -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) => {
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -102,6 +102,11 @@ async fn run_cli_probe(bin: &std::path::Path, args: &[&str]) -> bool {
|
|||
cmd.env_remove("PYTHONPATH");
|
||||
}
|
||||
|
||||
// Tauri uses the legacy root regardless of UNSLOTH_STUDIO_HOME / STUDIO_HOME;
|
||||
// probe subprocesses must follow the same isolation as process.rs.
|
||||
cmd.env_remove("UNSLOTH_STUDIO_HOME");
|
||||
cmd.env_remove("STUDIO_HOME");
|
||||
|
||||
#[cfg(windows)]
|
||||
{
|
||||
use std::os::windows::process::CommandExt;
|
||||
|
|
@ -135,6 +140,11 @@ async fn probe_cli_capability(bin: &std::path::Path) -> Option<DesktopCapability
|
|||
cmd.env_remove("PYTHONPATH");
|
||||
}
|
||||
|
||||
// Tauri uses the legacy root regardless of UNSLOTH_STUDIO_HOME / STUDIO_HOME;
|
||||
// probe subprocesses must follow the same isolation as process.rs.
|
||||
cmd.env_remove("UNSLOTH_STUDIO_HOME");
|
||||
cmd.env_remove("STUDIO_HOME");
|
||||
|
||||
#[cfg(windows)]
|
||||
{
|
||||
use std::os::windows::process::CommandExt;
|
||||
|
|
|
|||
|
|
@ -316,6 +316,12 @@ pub fn start_backend(
|
|||
cmd.env_remove("PYTHONPATH");
|
||||
}
|
||||
|
||||
// Tauri uses the legacy root regardless of UNSLOTH_STUDIO_HOME / STUDIO_HOME;
|
||||
// scrub so the spawned Python backend can't diverge. UNSLOTH_LLAMA_CPP_PATH
|
||||
// is a pre-existing user-controlled llama.cpp dir override; keep it.
|
||||
cmd.env_remove("UNSLOTH_STUDIO_HOME");
|
||||
cmd.env_remove("STUDIO_HOME");
|
||||
|
||||
// On Windows, launch the backend directly with hidden-window flags.
|
||||
// The app process is assigned to a KILL_ON_JOB_CLOSE job in main.rs, so
|
||||
// children inherit crash-safe cleanup without the buggy per-child JobObject wrapper.
|
||||
|
|
|
|||
|
|
@ -61,6 +61,11 @@ fn spawn_update(
|
|||
cmd.env_remove("PYTHONPATH");
|
||||
}
|
||||
|
||||
// Tauri manages the legacy root; scrub so 'unsloth studio update' targets
|
||||
// the same install the desktop app uses, not an inherited custom root.
|
||||
cmd.env_remove("UNSLOTH_STUDIO_HOME");
|
||||
cmd.env_remove("STUDIO_HOME");
|
||||
|
||||
#[cfg(windows)]
|
||||
let mut child: Box<dyn ChildWrapper + Send> = {
|
||||
use std::os::windows::process::CommandExt;
|
||||
|
|
|
|||
141
tests/conftest.py
Normal file
141
tests/conftest.py
Normal file
|
|
@ -0,0 +1,141 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
|
||||
|
||||
"""GPU-free test harness.
|
||||
|
||||
unsloth's import chain hits unsloth_zoo.device_type, which calls
|
||||
get_device_type() at import time and raises NotImplementedError on CI
|
||||
runners with no CUDA / XPU / HIP visible. Pre-load the real
|
||||
unsloth_zoo.device_type under a temporarily-mocked
|
||||
torch.cuda.is_available() so its @cache permanently captures "cuda".
|
||||
On a real accelerator the pre-load is skipped and detection runs
|
||||
normally.
|
||||
|
||||
Mirrors the conftest harness in unslothai/unsloth-zoo PR #624.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
|
||||
|
||||
def _has_real_accelerator() -> bool:
|
||||
try:
|
||||
import torch
|
||||
except Exception:
|
||||
return False
|
||||
for probe in (
|
||||
lambda: hasattr(torch, "cuda") and torch.cuda.is_available(),
|
||||
lambda: hasattr(torch, "xpu") and torch.xpu.is_available(),
|
||||
lambda: hasattr(torch, "accelerator") and torch.accelerator.is_available(),
|
||||
):
|
||||
try:
|
||||
if probe():
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
return False
|
||||
|
||||
|
||||
def _preload_device_type(package: str, prereqs: tuple[str, ...] = ()) -> bool:
|
||||
"""Pre-load <package>.device_type under a mocked
|
||||
torch.cuda.is_available() == True so its @cache permanently
|
||||
captures "cuda". prereqs lists submodule names of <package> that
|
||||
must be loaded first (e.g. 'utils' for unsloth_zoo). Returns False
|
||||
if the package or any prerequisite cannot be imported, in which
|
||||
case the caller falls back to a stub."""
|
||||
target = f"{package}.device_type"
|
||||
if target in sys.modules:
|
||||
return True
|
||||
pkg_spec = importlib.util.find_spec(package)
|
||||
if pkg_spec is None or not pkg_spec.submodule_search_locations:
|
||||
return False
|
||||
pkg_path = pkg_spec.submodule_search_locations[0]
|
||||
|
||||
skeleton_already = package in sys.modules
|
||||
if not skeleton_already:
|
||||
skel = types.ModuleType(package)
|
||||
skel.__path__ = [pkg_path]
|
||||
skel.__spec__ = pkg_spec
|
||||
skel.__package__ = package
|
||||
sys.modules[package] = skel
|
||||
|
||||
try:
|
||||
for prereq in prereqs:
|
||||
full = f"{package}.{prereq}"
|
||||
if full in sys.modules:
|
||||
continue
|
||||
prereq_path = os.path.join(pkg_path, f"{prereq}.py")
|
||||
prereq_spec = importlib.util.spec_from_file_location(full, prereq_path)
|
||||
prereq_mod = importlib.util.module_from_spec(prereq_spec)
|
||||
sys.modules[full] = prereq_mod
|
||||
prereq_spec.loader.exec_module(prereq_mod)
|
||||
|
||||
device_type_path = os.path.join(pkg_path, "device_type.py")
|
||||
dt_spec = importlib.util.spec_from_file_location(target, device_type_path)
|
||||
dt_mod = importlib.util.module_from_spec(dt_spec)
|
||||
sys.modules[target] = dt_mod
|
||||
|
||||
import torch
|
||||
|
||||
_orig_is_avail = torch.cuda.is_available
|
||||
torch.cuda.is_available = lambda: True # type: ignore[assignment]
|
||||
try:
|
||||
dt_spec.loader.exec_module(dt_mod)
|
||||
finally:
|
||||
torch.cuda.is_available = _orig_is_avail
|
||||
except Exception:
|
||||
sys.modules.pop(target, None)
|
||||
return False
|
||||
finally:
|
||||
if not skeleton_already:
|
||||
sys.modules.pop(package, None)
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def _patch_torch_cuda_for_import() -> None:
|
||||
"""Stub torch.cuda.* probes that fire at IMPORT time of unsloth /
|
||||
unsloth_zoo when DEVICE_TYPE was forced to "cuda" above. These are
|
||||
queries, not real GPU work, so returning plausible Ampere values
|
||||
lets the import chain finish; tests that touch real tensors run on
|
||||
CPU like normal."""
|
||||
try:
|
||||
import torch.cuda.memory as _cuda_memory # type: ignore
|
||||
|
||||
_cuda_memory.mem_get_info = lambda *a, **k: (0, 80 * 1024**3)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
import torch
|
||||
|
||||
torch.cuda.get_device_capability = lambda *a, **k: (8, 0)
|
||||
torch.cuda.is_bf16_supported = lambda *a, **k: True
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _install_device_type_stub(name: str) -> None:
|
||||
stub = types.ModuleType(name)
|
||||
stub.DEVICE_TYPE = "cuda"
|
||||
stub.DEVICE_TYPE_TORCH = "cuda"
|
||||
stub.DEVICE_COUNT = 1
|
||||
stub.ALLOW_PREQUANTIZED_MODELS = False
|
||||
stub.is_hip = lambda: False
|
||||
stub.get_device_type = lambda: "cuda"
|
||||
stub.get_device_count = lambda: 1
|
||||
stub.device_synchronize = lambda *a, **k: None
|
||||
stub.device_empty_cache = lambda *a, **k: None
|
||||
stub.device_is_bf16_supported = lambda *a, **k: False
|
||||
sys.modules[name] = stub
|
||||
|
||||
|
||||
if not _has_real_accelerator():
|
||||
if not _preload_device_type("unsloth_zoo", prereqs = ("utils",)):
|
||||
_install_device_type_stub("unsloth_zoo.device_type")
|
||||
if not _preload_device_type("unsloth"):
|
||||
_install_device_type_stub("unsloth.device_type")
|
||||
_patch_torch_cuda_for_import()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
46
tests/python/test_gpu_init_ldconfig_guard.py
Normal file
46
tests/python/test_gpu_init_ldconfig_guard.py
Normal file
|
|
@ -0,0 +1,46 @@
|
|||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
GPU_INIT = REPO_ROOT / "unsloth" / "_gpu_init.py"
|
||||
|
||||
|
||||
def _find_geteuid_guard(tree: ast.AST):
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.If):
|
||||
continue
|
||||
for sub in ast.walk(node.test):
|
||||
if isinstance(sub, ast.Call) and isinstance(sub.func, ast.Attribute):
|
||||
if sub.func.attr == "geteuid":
|
||||
return node
|
||||
return None
|
||||
|
||||
|
||||
def test_gpu_init_has_geteuid_guard():
|
||||
tree = ast.parse(GPU_INIT.read_text())
|
||||
guard = _find_geteuid_guard(tree)
|
||||
assert (
|
||||
guard is not None
|
||||
), "_gpu_init.py must guard ldconfig recovery on os.geteuid()"
|
||||
|
||||
|
||||
def test_ldconfig_calls_only_inside_geteuid_guard():
|
||||
src = GPU_INIT.read_text()
|
||||
tree = ast.parse(src)
|
||||
guard = _find_geteuid_guard(tree)
|
||||
assert guard is not None
|
||||
guard_src = ast.get_source_segment(src, guard) or ""
|
||||
ldconfig_lines = [
|
||||
line for line in src.splitlines() if "ldconfig" in line and "os.system" in line
|
||||
]
|
||||
for line in ldconfig_lines:
|
||||
assert line.strip() in guard_src, (
|
||||
"os.system('ldconfig ...') must live inside the geteuid guard, "
|
||||
f"but found unguarded: {line!r}"
|
||||
)
|
||||
|
||||
|
||||
def test_non_root_branch_warns_when_bnb_present():
|
||||
src = GPU_INIT.read_text()
|
||||
assert "elif bnb is not None" in src
|
||||
assert "sudo ldconfig" in src
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
# =============================================================================
|
||||
|
|
|
|||
122
tests/studio/test_export_output_path_contract.py
Normal file
122
tests/studio/test_export_output_path_contract.py
Normal file
|
|
@ -0,0 +1,122 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
|
||||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
EXPORT = REPO_ROOT / "studio" / "backend" / "core" / "export" / "export.py"
|
||||
|
||||
EXPORT_FNS = (
|
||||
"export_merged_model",
|
||||
"export_base_model",
|
||||
"export_gguf",
|
||||
"export_lora_adapter",
|
||||
)
|
||||
|
||||
|
||||
def _find_method(tree, cls_name, method_name):
|
||||
for cls in ast.walk(tree):
|
||||
if isinstance(cls, ast.ClassDef) and cls.name == cls_name:
|
||||
for item in cls.body:
|
||||
if isinstance(item, ast.FunctionDef) and item.name == method_name:
|
||||
return item
|
||||
return None
|
||||
|
||||
|
||||
def _return_tuple_arity(fn):
|
||||
arities = []
|
||||
for node in ast.walk(fn):
|
||||
if isinstance(node, ast.Return) and isinstance(node.value, ast.Tuple):
|
||||
arities.append(len(node.value.elts))
|
||||
return arities
|
||||
|
||||
|
||||
def test_export_methods_return_three_tuple_annotation():
|
||||
tree = ast.parse(EXPORT.read_text())
|
||||
for fn_name in EXPORT_FNS:
|
||||
fn = _find_method(tree, "ExportBackend", fn_name)
|
||||
assert fn is not None, f"missing ExportBackend.{fn_name}"
|
||||
ret = fn.returns
|
||||
assert isinstance(ret, ast.Subscript), f"{fn_name} return must be Tuple[...]"
|
||||
slc = ret.slice
|
||||
elts = slc.elts if isinstance(slc, ast.Tuple) else None
|
||||
assert (
|
||||
elts is not None and len(elts) == 3
|
||||
), f"{fn_name} return annotation must be a 3-tuple, got {ast.dump(ret)}"
|
||||
|
||||
|
||||
def test_export_methods_return_three_element_tuples():
|
||||
tree = ast.parse(EXPORT.read_text())
|
||||
for fn_name in EXPORT_FNS:
|
||||
fn = _find_method(tree, "ExportBackend", fn_name)
|
||||
assert fn is not None
|
||||
arities = _return_tuple_arity(fn)
|
||||
assert arities, f"{fn_name} has no tuple-return statements"
|
||||
for arity in arities:
|
||||
assert arity == 3, f"{fn_name} return tuple arity {arity}, expected 3"
|
||||
|
||||
|
||||
def test_local_save_assigns_output_path():
|
||||
tree = ast.parse(EXPORT.read_text())
|
||||
for fn_name in EXPORT_FNS:
|
||||
fn = _find_method(tree, "ExportBackend", fn_name)
|
||||
assert fn is not None
|
||||
assigns = []
|
||||
for node in ast.walk(fn):
|
||||
if isinstance(node, ast.Assign):
|
||||
for tgt in node.targets:
|
||||
if isinstance(tgt, ast.Name) and tgt.id == "output_path":
|
||||
assigns.append(node)
|
||||
non_none = [
|
||||
a
|
||||
for a in assigns
|
||||
if not (isinstance(a.value, ast.Constant) and a.value.value is None)
|
||||
]
|
||||
assert non_none, f"{fn_name} never assigns a non-None output_path"
|
||||
|
||||
|
||||
def test_gpu_save_method_bound_for_hub_only():
|
||||
tree = ast.parse(EXPORT.read_text())
|
||||
fn = _find_method(tree, "ExportBackend", "export_merged_model")
|
||||
assert fn is not None
|
||||
found_pre_save_method = False
|
||||
for node in ast.walk(fn):
|
||||
if isinstance(node, ast.Try):
|
||||
for stmt in node.body:
|
||||
if isinstance(stmt, ast.If):
|
||||
test = stmt.test
|
||||
if isinstance(test, ast.Name) and test.id == "_IS_MLX":
|
||||
for sub in ast.walk(
|
||||
ast.Module(body = stmt.orelse, type_ignores = [])
|
||||
):
|
||||
if isinstance(sub, ast.Assign) and any(
|
||||
isinstance(t, ast.Name) and t.id == "save_method"
|
||||
for t in sub.targets
|
||||
):
|
||||
found_pre_save_method = True
|
||||
break
|
||||
if found_pre_save_method:
|
||||
break
|
||||
if found_pre_save_method:
|
||||
break
|
||||
assert found_pre_save_method, (
|
||||
"GPU save_method must be assigned at the top of the try block, "
|
||||
"before the `if save_directory:` guard, so Hub-only export does not "
|
||||
"raise UnboundLocalError."
|
||||
)
|
||||
|
||||
|
||||
def test_mlx_hub_only_uses_temp_directory():
|
||||
src = EXPORT.read_text()
|
||||
assert (
|
||||
src.count("tempfile.TemporaryDirectory") >= 3
|
||||
), "expected TemporaryDirectory in merged, base, and lora hub-push paths"
|
||||
assert "import tempfile" in src.split("class ExportBackend")[0]
|
||||
|
||||
|
||||
def test_is_mlx_imported_from_unsloth():
|
||||
src = EXPORT.read_text()
|
||||
assert "from unsloth import" in src
|
||||
head = src.split("class ExportBackend")[0]
|
||||
assert "_IS_MLX" in head
|
||||
assert "_IS_MLX = platform.system()" not in src
|
||||
395
tests/studio/test_hardware_dispatch_matrix.py
Normal file
395
tests/studio/test_hardware_dispatch_matrix.py
Normal file
|
|
@ -0,0 +1,395 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
"""
|
||||
Comprehensive hardware dispatch matrix for Studio.
|
||||
|
||||
Drives every supported hardware profile from a single test host by
|
||||
spoofing platform / torch.cuda / torch.xpu / sys.modules['mlx'] so we
|
||||
can exercise the CUDA, ROCm, XPU, MLX, and CPU dispatch paths
|
||||
deterministically without real hardware.
|
||||
|
||||
Profiles checked:
|
||||
|
||||
nvidia_cuda Linux x86_64 + torch.cuda.is_available()=True,
|
||||
torch.version.hip=None
|
||||
amd_rocm Linux x86_64 + torch.cuda.is_available()=True,
|
||||
torch.version.hip="6.1" (PyTorch ROCm aliases
|
||||
torch.cuda.* over HIP)
|
||||
intel_xpu Linux x86_64 + torch.cuda off, torch.xpu.is_available()=True
|
||||
apple_silicon_mlx Darwin arm64 + cuda off + xpu off + mlx importable
|
||||
apple_silicon_no_mlx Darwin arm64 + everything off (no mlx pkg)
|
||||
linux_arm64_with_mlx Linux arm64 + mlx importable -- gate must NOT activate
|
||||
(canary against accidental Linux-arm64 hijack)
|
||||
cpu_only Linux x86_64 + nothing -- pure CPU fallback
|
||||
|
||||
For each profile we assert three contracts:
|
||||
|
||||
1. ``unsloth._IS_MLX`` (re-evaluated under the spoof).
|
||||
2. ``utils.hardware.detect_hardware()`` ``DeviceType`` and ``IS_ROCM``.
|
||||
3. ``utils.hardware.is_apple_silicon()``.
|
||||
|
||||
Add a row to ``PROFILES`` to extend coverage; tests parametrize over it
|
||||
automatically. No real hardware required.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import importlib.machinery
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
STUDIO_BACKEND = REPO_ROOT / "studio" / "backend"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Profile definition
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class HardwareProfile:
|
||||
name: str
|
||||
system: str # platform.system() value
|
||||
machine: str # platform.machine() value
|
||||
cuda_available: bool # torch.cuda.is_available() value
|
||||
hip_version: Optional[
|
||||
str
|
||||
] # torch.version.hip; None for NVIDIA, "6.1" etc. for ROCm
|
||||
xpu_available: bool # torch.xpu.is_available() value
|
||||
has_mlx: bool # whether to inject a fake mlx into sys.modules
|
||||
mps_available: bool # torch.backends.mps.is_available() value
|
||||
|
||||
expect_is_mlx: bool # unsloth._IS_MLX
|
||||
expect_device_type: (
|
||||
str # Studio DeviceType (uppercased name: "CUDA"/"XPU"/"MLX"/"CPU")
|
||||
)
|
||||
expect_is_rocm: bool # Studio IS_ROCM
|
||||
expect_apple_silicon: bool # Studio is_apple_silicon()
|
||||
extra_notes: str = ""
|
||||
|
||||
|
||||
PROFILES = [
|
||||
HardwareProfile(
|
||||
name = "nvidia_cuda",
|
||||
system = "Linux",
|
||||
machine = "x86_64",
|
||||
cuda_available = True,
|
||||
hip_version = None,
|
||||
xpu_available = False,
|
||||
has_mlx = False,
|
||||
mps_available = False,
|
||||
expect_is_mlx = False,
|
||||
expect_device_type = "CUDA",
|
||||
expect_is_rocm = False,
|
||||
expect_apple_silicon = False,
|
||||
),
|
||||
HardwareProfile(
|
||||
name = "amd_rocm",
|
||||
system = "Linux",
|
||||
machine = "x86_64",
|
||||
cuda_available = True,
|
||||
hip_version = "6.1",
|
||||
xpu_available = False,
|
||||
has_mlx = False,
|
||||
mps_available = False,
|
||||
expect_is_mlx = False,
|
||||
expect_device_type = "CUDA",
|
||||
expect_is_rocm = True,
|
||||
expect_apple_silicon = False,
|
||||
extra_notes = "PyTorch ROCm reuses torch.cuda.* over HIP; "
|
||||
"Studio still uses DeviceType.CUDA but flips IS_ROCM=True.",
|
||||
),
|
||||
HardwareProfile(
|
||||
name = "intel_xpu",
|
||||
system = "Linux",
|
||||
machine = "x86_64",
|
||||
cuda_available = False,
|
||||
hip_version = None,
|
||||
xpu_available = True,
|
||||
has_mlx = False,
|
||||
mps_available = False,
|
||||
expect_is_mlx = False,
|
||||
expect_device_type = "XPU",
|
||||
expect_is_rocm = False,
|
||||
expect_apple_silicon = False,
|
||||
),
|
||||
HardwareProfile(
|
||||
name = "apple_silicon_mlx",
|
||||
system = "Darwin",
|
||||
machine = "arm64",
|
||||
cuda_available = False,
|
||||
hip_version = None,
|
||||
xpu_available = False,
|
||||
has_mlx = True,
|
||||
mps_available = True,
|
||||
expect_is_mlx = True,
|
||||
expect_device_type = "MLX",
|
||||
expect_is_rocm = False,
|
||||
expect_apple_silicon = True,
|
||||
),
|
||||
HardwareProfile(
|
||||
name = "apple_silicon_no_mlx",
|
||||
system = "Darwin",
|
||||
machine = "arm64",
|
||||
cuda_available = False,
|
||||
hip_version = None,
|
||||
xpu_available = False,
|
||||
has_mlx = False,
|
||||
mps_available = True,
|
||||
expect_is_mlx = False,
|
||||
expect_device_type = "CPU",
|
||||
expect_is_rocm = False,
|
||||
expect_apple_silicon = True,
|
||||
extra_notes = "Mac without mlx falls through to CPU (chat-only).",
|
||||
),
|
||||
HardwareProfile(
|
||||
name = "linux_arm64_with_mlx",
|
||||
system = "Linux",
|
||||
machine = "arm64",
|
||||
cuda_available = False,
|
||||
hip_version = None,
|
||||
xpu_available = False,
|
||||
has_mlx = True,
|
||||
mps_available = False,
|
||||
expect_is_mlx = False,
|
||||
expect_device_type = "CPU",
|
||||
expect_is_rocm = False,
|
||||
expect_apple_silicon = False,
|
||||
extra_notes = "Canary: Linux ARM64 with mlx package installed must NOT "
|
||||
"trigger MLX dispatch; the system check is what guards it.",
|
||||
),
|
||||
HardwareProfile(
|
||||
name = "cpu_only",
|
||||
system = "Linux",
|
||||
machine = "x86_64",
|
||||
cuda_available = False,
|
||||
hip_version = None,
|
||||
xpu_available = False,
|
||||
has_mlx = False,
|
||||
mps_available = False,
|
||||
expect_is_mlx = False,
|
||||
expect_device_type = "CPU",
|
||||
expect_is_rocm = False,
|
||||
expect_apple_silicon = False,
|
||||
),
|
||||
]
|
||||
|
||||
PROFILE_IDS = [p.name for p in PROFILES]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Spoofing helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def spoof_hardware(monkeypatch):
|
||||
"""Return a function that applies a HardwareProfile to the live process.
|
||||
|
||||
Idempotent: each call re-applies the profile. Cleanup happens
|
||||
automatically when the test exits via monkeypatch.
|
||||
"""
|
||||
|
||||
def _apply(profile: HardwareProfile) -> None:
|
||||
import platform
|
||||
import torch
|
||||
|
||||
# platform spoof (used by both the unsloth gate and Studio's helpers)
|
||||
monkeypatch.setattr(platform, "system", lambda: profile.system)
|
||||
monkeypatch.setattr(platform, "machine", lambda: profile.machine)
|
||||
|
||||
# torch.cuda.is_available
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: profile.cuda_available)
|
||||
# detect_hardware reads torch.cuda.get_device_properties(0).name when
|
||||
# cuda_available is True. On a CPU CI runner that triggers _cuda_init
|
||||
# and crashes with "No CUDA GPUs are available". Stub it so the
|
||||
# dispatch path under test runs end-to-end.
|
||||
if profile.cuda_available:
|
||||
stub_props = types.SimpleNamespace(
|
||||
name = "Stub GPU" if not profile.hip_version else "Stub AMD GPU",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
torch.cuda,
|
||||
"get_device_properties",
|
||||
lambda i = 0: stub_props,
|
||||
raising = False,
|
||||
)
|
||||
|
||||
# torch.version.hip — None on NVIDIA, "6.1" etc. on ROCm
|
||||
torch_version = torch.version
|
||||
monkeypatch.setattr(torch_version, "hip", profile.hip_version, raising = False)
|
||||
|
||||
# torch.xpu.is_available + get_device_name -- detect_hardware reads both.
|
||||
# Real torch.xpu.get_device_name requires the XPU-compiled torch build,
|
||||
# so always stub it under the spoof to keep tests hardware-agnostic.
|
||||
if hasattr(torch, "xpu"):
|
||||
monkeypatch.setattr(
|
||||
torch.xpu, "is_available", lambda: profile.xpu_available
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
torch.xpu,
|
||||
"get_device_name",
|
||||
lambda i = 0: "Intel XPU (stub)",
|
||||
raising = False,
|
||||
)
|
||||
elif profile.xpu_available:
|
||||
xpu_stub = types.SimpleNamespace(
|
||||
is_available = lambda: True,
|
||||
get_device_name = lambda i = 0: "Intel XPU (stub)",
|
||||
)
|
||||
monkeypatch.setattr(torch, "xpu", xpu_stub, raising = False)
|
||||
|
||||
# torch.backends.mps.is_available
|
||||
if hasattr(torch.backends, "mps"):
|
||||
monkeypatch.setattr(
|
||||
torch.backends.mps, "is_available", lambda: profile.mps_available
|
||||
)
|
||||
|
||||
# mlx + mlx.core in sys.modules
|
||||
if profile.has_mlx:
|
||||
fake_mlx = types.ModuleType("mlx")
|
||||
fake_mlx.__spec__ = importlib.machinery.ModuleSpec("mlx", loader = None)
|
||||
fake_mlx.__path__ = []
|
||||
fake_mlx_core = types.ModuleType("mlx.core")
|
||||
fake_mlx.core = fake_mlx_core
|
||||
monkeypatch.setitem(sys.modules, "mlx", fake_mlx)
|
||||
monkeypatch.setitem(sys.modules, "mlx.core", fake_mlx_core)
|
||||
else:
|
||||
monkeypatch.delitem(sys.modules, "mlx", raising = False)
|
||||
monkeypatch.delitem(sys.modules, "mlx.core", raising = False)
|
||||
real_find_spec = importlib.util.find_spec
|
||||
|
||||
def _no_mlx(name, *args, **kwargs):
|
||||
if name == "mlx":
|
||||
return None
|
||||
return real_find_spec(name, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(importlib.util, "find_spec", _no_mlx)
|
||||
|
||||
return _apply
|
||||
|
||||
|
||||
def _evaluate_unsloth_is_mlx_gate() -> bool:
|
||||
"""Re-evaluate the exact expression from unsloth/__init__.py:20-24."""
|
||||
import importlib.util
|
||||
import platform
|
||||
|
||||
return (
|
||||
platform.system() == "Darwin"
|
||||
and platform.machine() == "arm64"
|
||||
and importlib.util.find_spec("mlx") is not None
|
||||
)
|
||||
|
||||
|
||||
def _import_studio_hardware_module():
|
||||
"""Lazy-load Studio's hardware module under the bare-imports layout."""
|
||||
if str(STUDIO_BACKEND) not in sys.path:
|
||||
sys.path.insert(0, str(STUDIO_BACKEND))
|
||||
# Force a fresh import so detect_hardware re-runs under the current spoofs.
|
||||
sys.modules.pop("utils.hardware.hardware", None)
|
||||
sys.modules.pop("utils.hardware", None)
|
||||
from utils.hardware import hardware as hw # type: ignore
|
||||
|
||||
return hw
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize("profile", PROFILES, ids = PROFILE_IDS)
|
||||
def test_unsloth_is_mlx_gate_matches_profile(profile, spoof_hardware):
|
||||
"""The _IS_MLX expression in unsloth/__init__.py flips correctly per profile."""
|
||||
spoof_hardware(profile)
|
||||
actual = _evaluate_unsloth_is_mlx_gate()
|
||||
assert actual is profile.expect_is_mlx, (
|
||||
f"profile {profile.name}: expected _IS_MLX={profile.expect_is_mlx}, "
|
||||
f"got {actual}. {profile.extra_notes}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("profile", PROFILES, ids = PROFILE_IDS)
|
||||
def test_studio_detect_hardware_matches_profile(profile, spoof_hardware):
|
||||
"""Studio's detect_hardware() routes to the right DeviceType per profile."""
|
||||
spoof_hardware(profile)
|
||||
hw = _import_studio_hardware_module()
|
||||
detected = hw.detect_hardware()
|
||||
expected = getattr(hw.DeviceType, profile.expect_device_type)
|
||||
assert detected == expected, (
|
||||
f"profile {profile.name}: expected {profile.expect_device_type}, "
|
||||
f"got {detected!r}. {profile.extra_notes}"
|
||||
)
|
||||
assert hw.IS_ROCM is profile.expect_is_rocm, (
|
||||
f"profile {profile.name}: expected IS_ROCM={profile.expect_is_rocm}, "
|
||||
f"got {hw.IS_ROCM}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("profile", PROFILES, ids = PROFILE_IDS)
|
||||
def test_studio_is_apple_silicon_matches_profile(profile, spoof_hardware):
|
||||
"""Studio's is_apple_silicon() helper agrees with platform spoof."""
|
||||
spoof_hardware(profile)
|
||||
hw = _import_studio_hardware_module()
|
||||
assert hw.is_apple_silicon() is profile.expect_apple_silicon, (
|
||||
f"profile {profile.name}: expected is_apple_silicon={profile.expect_apple_silicon}, "
|
||||
f"got {hw.is_apple_silicon()}"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Negative-space tests: catch regressions where the dispatch order changes.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_cuda_takes_priority_over_mlx_when_both_available(spoof_hardware):
|
||||
"""If both CUDA and MLX are available, Studio MUST pick CUDA. This is the
|
||||
canary that protects every existing GPU user from being silently routed
|
||||
to MLX after future refactors.
|
||||
"""
|
||||
profile = HardwareProfile(
|
||||
name = "cuda_plus_mlx",
|
||||
system = "Darwin",
|
||||
machine = "arm64",
|
||||
cuda_available = True,
|
||||
hip_version = None,
|
||||
xpu_available = False,
|
||||
has_mlx = True,
|
||||
mps_available = True,
|
||||
expect_is_mlx = True,
|
||||
expect_device_type = "CUDA",
|
||||
expect_is_rocm = False,
|
||||
expect_apple_silicon = True,
|
||||
)
|
||||
spoof_hardware(profile)
|
||||
hw = _import_studio_hardware_module()
|
||||
assert hw.detect_hardware() == hw.DeviceType.CUDA
|
||||
|
||||
|
||||
def test_xpu_takes_priority_over_mlx_when_both_available(spoof_hardware):
|
||||
"""XPU is selected over MLX in the dispatch order."""
|
||||
profile = HardwareProfile(
|
||||
name = "xpu_plus_mlx",
|
||||
system = "Darwin",
|
||||
machine = "arm64",
|
||||
cuda_available = False,
|
||||
hip_version = None,
|
||||
xpu_available = True,
|
||||
has_mlx = True,
|
||||
mps_available = True,
|
||||
expect_is_mlx = True,
|
||||
expect_device_type = "XPU",
|
||||
expect_is_rocm = False,
|
||||
expect_apple_silicon = True,
|
||||
)
|
||||
spoof_hardware(profile)
|
||||
hw = _import_studio_hardware_module()
|
||||
assert hw.detect_hardware() == hw.DeviceType.XPU
|
||||
213
tests/studio/test_is_mlx_dispatch_gate.py
Normal file
213
tests/studio/test_is_mlx_dispatch_gate.py
Normal file
|
|
@ -0,0 +1,213 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
|
||||
"""
|
||||
Regression tests for the CUDA-vs-MLX dispatch gates Studio relies on.
|
||||
|
||||
Two gates drive every dispatch decision in Studio's MLX path:
|
||||
|
||||
1. ``unsloth._IS_MLX`` at the top of ``unsloth/__init__.py`` -- evaluated
|
||||
once at import time and read by Studio worker code to choose between
|
||||
the GPU and MLX trainer / inference / export paths. Defined as
|
||||
``Darwin AND arm64 AND find_spec("mlx") is not None``.
|
||||
|
||||
2. ``utils.hardware.detect_hardware()`` -- runtime probe in the Studio
|
||||
backend. Priority order: CUDA -> XPU -> MLX -> CPU. The MLX branch is
|
||||
reached only when both CUDA and XPU are unavailable AND the host is
|
||||
Apple Silicon AND mlx is importable.
|
||||
|
||||
These gates are the canaries for "MLX support accidentally hijacks
|
||||
CUDA/AMD/Intel users". The tests here:
|
||||
|
||||
* verify the source-level structure of the ``_IS_MLX`` expression so an
|
||||
accidental rewrite (e.g. dropping the ``arm64`` check) is caught,
|
||||
* exercise the runtime gate logic under a spoofed Darwin+arm64 platform
|
||||
with a fake ``mlx`` module in ``sys.modules`` to confirm both gates
|
||||
flip True together,
|
||||
* confirm that on the actual Linux+CUDA test host both gates remain in
|
||||
their CUDA-side state.
|
||||
|
||||
No real MLX install is required; uses the same ``monkeypatch.setitem``
|
||||
fake-mlx pattern as ``test_mlx_inference_backend.py``.
|
||||
"""
|
||||
|
||||
import ast
|
||||
import importlib
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
UNSLOTH_INIT = REPO_ROOT / "unsloth" / "__init__.py"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1. Source-level structure check on _IS_MLX (no platform dependencies).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_is_mlx_gate_uses_three_required_predicates():
|
||||
"""The _IS_MLX assignment must AND together exactly the three checks
|
||||
that Studio depends on: Darwin OS, arm64 machine, and an importable
|
||||
mlx package. Dropping any one of them silently breaks dispatch.
|
||||
"""
|
||||
tree = ast.parse(UNSLOTH_INIT.read_text())
|
||||
|
||||
target = None
|
||||
for node in ast.walk(tree):
|
||||
if (
|
||||
isinstance(node, ast.Assign)
|
||||
and len(node.targets) == 1
|
||||
and isinstance(node.targets[0], ast.Name)
|
||||
and node.targets[0].id == "_IS_MLX"
|
||||
):
|
||||
target = node.value
|
||||
break
|
||||
assert target is not None, "_IS_MLX assignment not found in unsloth/__init__.py"
|
||||
assert isinstance(target, ast.BoolOp) and isinstance(
|
||||
target.op, ast.And
|
||||
), "_IS_MLX must be a BoolOp(And) of platform + mlx checks"
|
||||
|
||||
expr_src = ast.unparse(target)
|
||||
assert (
|
||||
"platform.system()" in expr_src and "Darwin" in expr_src
|
||||
), "_IS_MLX must check platform.system() == 'Darwin'"
|
||||
assert (
|
||||
"platform.machine()" in expr_src and "arm64" in expr_src
|
||||
), "_IS_MLX must check platform.machine() == 'arm64'"
|
||||
assert (
|
||||
"find_spec" in expr_src and "'mlx'" in expr_src
|
||||
), "_IS_MLX must check importlib.util.find_spec('mlx')"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 2. Runtime gate behavior with the platform spoofed to Apple Silicon and a
|
||||
# fake mlx module in sys.modules. Re-evaluates the same expression
|
||||
# rather than reloading unsloth (which would cascade-reload torch).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _evaluate_is_mlx_gate(platform_module, importlib_util):
|
||||
"""Re-evaluate the _IS_MLX expression using injected dependencies.
|
||||
|
||||
Mirrors the assignment in unsloth/__init__.py exactly.
|
||||
"""
|
||||
return (
|
||||
platform_module.system() == "Darwin"
|
||||
and platform_module.machine() == "arm64"
|
||||
and importlib_util.find_spec("mlx") is not None
|
||||
)
|
||||
|
||||
|
||||
def test_is_mlx_gate_true_on_apple_silicon_with_mlx_present(monkeypatch):
|
||||
import platform
|
||||
import importlib.util
|
||||
|
||||
# Inject a fake mlx package so find_spec returns a non-None ModuleSpec.
|
||||
fake_mlx = types.ModuleType("mlx")
|
||||
fake_mlx.__spec__ = importlib.machinery.ModuleSpec("mlx", loader = None)
|
||||
fake_mlx.__path__ = []
|
||||
monkeypatch.setitem(sys.modules, "mlx", fake_mlx)
|
||||
|
||||
monkeypatch.setattr(platform, "system", lambda: "Darwin")
|
||||
monkeypatch.setattr(platform, "machine", lambda: "arm64")
|
||||
|
||||
assert _evaluate_is_mlx_gate(platform, importlib.util) is True
|
||||
|
||||
|
||||
def test_is_mlx_gate_false_when_mlx_missing(monkeypatch):
|
||||
import platform
|
||||
import importlib.util
|
||||
|
||||
# Apple Silicon platform but no mlx package -> gate must be False.
|
||||
monkeypatch.delitem(sys.modules, "mlx", raising = False)
|
||||
monkeypatch.setattr(platform, "system", lambda: "Darwin")
|
||||
monkeypatch.setattr(platform, "machine", lambda: "arm64")
|
||||
|
||||
real_find_spec = importlib.util.find_spec
|
||||
|
||||
def _no_mlx(name, *args, **kwargs):
|
||||
if name == "mlx":
|
||||
return None
|
||||
return real_find_spec(name, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(importlib.util, "find_spec", _no_mlx)
|
||||
|
||||
assert _evaluate_is_mlx_gate(platform, importlib.util) is False
|
||||
|
||||
|
||||
def test_is_mlx_gate_false_on_non_apple_silicon():
|
||||
"""On the real Linux+CUDA / AMD / Intel test host, the gate stays False."""
|
||||
import platform
|
||||
import importlib.util
|
||||
|
||||
if platform.system() == "Darwin" and platform.machine() == "arm64":
|
||||
# On a Mac CI runner this assertion would not apply; skip there.
|
||||
import pytest
|
||||
|
||||
pytest.skip("Test host is Apple Silicon; CUDA-side canary doesn't apply.")
|
||||
|
||||
assert _evaluate_is_mlx_gate(platform, importlib.util) is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3. Studio's runtime detect_hardware() picks MLX only when CUDA + XPU are
|
||||
# both unavailable AND the host is Apple Silicon AND mlx is importable.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _import_studio_hardware():
|
||||
"""Lazy import for the Studio hardware module, with the bare-imports
|
||||
convention that Studio uses (studio/backend on sys.path).
|
||||
"""
|
||||
studio_backend = REPO_ROOT / "studio" / "backend"
|
||||
if str(studio_backend) not in sys.path:
|
||||
sys.path.insert(0, str(studio_backend))
|
||||
from utils.hardware import hardware as hw # type: ignore
|
||||
|
||||
return hw
|
||||
|
||||
|
||||
def test_detect_hardware_picks_mlx_when_only_apple_silicon_available(monkeypatch):
|
||||
hw = _import_studio_hardware()
|
||||
|
||||
# Force CUDA + XPU paths off so detect_hardware falls through to MLX.
|
||||
import torch
|
||||
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: False)
|
||||
if hasattr(torch, "xpu"):
|
||||
monkeypatch.setattr(torch.xpu, "is_available", lambda: False)
|
||||
|
||||
# Spoof Apple Silicon and provide an importable mlx.core for _has_mlx().
|
||||
import platform
|
||||
|
||||
monkeypatch.setattr(platform, "system", lambda: "Darwin")
|
||||
monkeypatch.setattr(platform, "machine", lambda: "arm64")
|
||||
|
||||
fake_mlx = types.ModuleType("mlx")
|
||||
fake_mlx_core = types.ModuleType("mlx.core")
|
||||
fake_mlx.core = fake_mlx_core
|
||||
monkeypatch.setitem(sys.modules, "mlx", fake_mlx)
|
||||
monkeypatch.setitem(sys.modules, "mlx.core", fake_mlx_core)
|
||||
|
||||
detected = hw.detect_hardware()
|
||||
assert detected == hw.DeviceType.MLX, f"expected MLX, got {detected!r}"
|
||||
|
||||
|
||||
def test_detect_hardware_picks_cuda_on_real_host():
|
||||
"""Canary: on a real CUDA host the MLX branch must NOT be taken even
|
||||
if mlx happens to be importable. Protects CUDA/AMD/Intel users from
|
||||
accidental MLX dispatch when MLX support is added.
|
||||
"""
|
||||
import torch
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
import pytest
|
||||
|
||||
pytest.skip("No CUDA available on this host; canary not applicable.")
|
||||
|
||||
hw = _import_studio_hardware()
|
||||
detected = hw.detect_hardware()
|
||||
assert (
|
||||
detected == hw.DeviceType.CUDA
|
||||
), f"CUDA host must dispatch to CUDA, got {detected!r}"
|
||||
90
tests/studio/test_mlx_training_worker_behaviors.py
Normal file
90
tests/studio/test_mlx_training_worker_behaviors.py
Normal file
|
|
@ -0,0 +1,90 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
|
||||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
WORKER = REPO_ROOT / "studio" / "backend" / "core" / "training" / "worker.py"
|
||||
|
||||
|
||||
def _find_func(tree, name):
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.FunctionDef) and node.name == name:
|
||||
return node
|
||||
return None
|
||||
|
||||
|
||||
def test_run_mlx_training_passes_token_to_from_pretrained():
|
||||
tree = ast.parse(WORKER.read_text())
|
||||
fn = _find_func(tree, "_run_mlx_training")
|
||||
assert fn is not None
|
||||
found = False
|
||||
for node in ast.walk(fn):
|
||||
if (
|
||||
isinstance(node, ast.Call)
|
||||
and isinstance(node.func, ast.Attribute)
|
||||
and node.func.attr == "from_pretrained"
|
||||
and isinstance(node.func.value, ast.Name)
|
||||
and node.func.value.id == "FastMLXModel"
|
||||
):
|
||||
kwarg_names = {kw.arg for kw in node.keywords if kw.arg}
|
||||
assert (
|
||||
"token" in kwarg_names
|
||||
), f"FastMLXModel.from_pretrained must forward token=hf_token; got {kwarg_names!r}"
|
||||
found = True
|
||||
assert found, "FastMLXModel.from_pretrained call not found in _run_mlx_training"
|
||||
|
||||
|
||||
def test_wandb_init_strips_secret_keys():
|
||||
src = WORKER.read_text()
|
||||
assert "_wandb_sensitive" in src, "expected a sensitive-key set near wandb.init"
|
||||
assert '"hf_token"' in src and '"wandb_token"' in src
|
||||
assert (
|
||||
"config = dict(config)" not in src
|
||||
), "wandb.init received raw config dict; secrets would leak"
|
||||
|
||||
|
||||
def test_local_dataset_loader_uses_load_dataset_path():
|
||||
src = WORKER.read_text()
|
||||
assert "_resolve_local_files" in src
|
||||
assert "_loader_for_files" in src
|
||||
assert "data_files = all_files" in src or "data_files=all_files" in src
|
||||
|
||||
|
||||
def test_send_aliases_status_message_to_message():
|
||||
src = WORKER.read_text()
|
||||
assert 'kwargs["message"] = sm' in src or 'kwargs["message"]=sm' in src
|
||||
|
||||
|
||||
def test_slice_uses_inclusive_end_and_handles_zero():
|
||||
src = WORKER.read_text()
|
||||
assert "min(end + 1, len(ds))" in src or "min(end+1, len(ds))" in src
|
||||
assert "slice_start if slice_start is not None else 0" in src
|
||||
assert "slice_end if slice_end is not None else len(ds) - 1" in src
|
||||
|
||||
|
||||
def test_poll_stop_returns_on_broken_pipe():
|
||||
src = WORKER.read_text()
|
||||
assert "except (EOFError, OSError)" in src
|
||||
lines = src.splitlines()
|
||||
for i, line in enumerate(lines):
|
||||
if "except (EOFError, OSError)" in line:
|
||||
for j in range(i + 1, min(i + 6, len(lines))):
|
||||
stripped = lines[j].strip()
|
||||
if not stripped or stripped.startswith("#"):
|
||||
continue
|
||||
assert stripped.startswith(
|
||||
"return"
|
||||
), f"expected return after EOFError/OSError, got {stripped!r}"
|
||||
break
|
||||
break
|
||||
else:
|
||||
raise AssertionError("EOFError/OSError handler not found in worker.py")
|
||||
|
||||
|
||||
def test_unsloth_zoo_mlx_imports_have_friendly_error():
|
||||
src = WORKER.read_text()
|
||||
assert "from unsloth_zoo.mlx_loader import FastMLXModel" in src
|
||||
assert "from unsloth_zoo.mlx_trainer import" in src
|
||||
assert "raise ImportError" in src
|
||||
assert "install.sh" in src
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
1021
tests/test_studio_install_workspace_guard.py
Normal file
1021
tests/test_studio_install_workspace_guard.py
Normal file
File diff suppressed because it is too large
Load diff
154
tests/test_studio_root_resilience.py
Normal file
154
tests/test_studio_root_resilience.py
Normal file
|
|
@ -0,0 +1,154 @@
|
|||
"""Resilience checks for Studio install-root inference under hostile
|
||||
filesystem conditions:
|
||||
- _infer_studio_home_from_venv must NOT propagate PermissionError /
|
||||
OSError out through studio_root() (it would crash module import in
|
||||
run.py / main.py / transformers_version.py / model_config.py).
|
||||
- _kill_orphaned_servers must catch (ImportError, OSError, ValueError)
|
||||
on the studio_root() probe so a transient resolve / sentinel failure
|
||||
cannot crash server startup.
|
||||
- _find_llama_server_binary must keep the custom-root in search_roots
|
||||
when the inner resolve() comparison itself fails."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import re
|
||||
import sys
|
||||
import textwrap
|
||||
from pathlib import Path
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||
STORAGE_ROOTS = (
|
||||
REPO_ROOT / "studio" / "backend" / "utils" / "paths" / "storage_roots.py"
|
||||
)
|
||||
LLAMA_CPP = REPO_ROOT / "studio" / "backend" / "core" / "inference" / "llama_cpp.py"
|
||||
|
||||
|
||||
def _load(name: str, path: Path):
|
||||
spec = importlib.util.spec_from_file_location(name, path)
|
||||
assert spec is not None and spec.loader is not None
|
||||
mod = importlib.util.module_from_spec(spec)
|
||||
sys.modules[name] = mod
|
||||
spec.loader.exec_module(mod)
|
||||
return mod
|
||||
|
||||
|
||||
def test_infer_studio_home_swallows_permission_error(tmp_path, monkeypatch):
|
||||
candidate = tmp_path / "fake_root"
|
||||
venv = candidate / "unsloth_studio"
|
||||
venv.mkdir(parents = True)
|
||||
monkeypatch.setattr(sys, "prefix", str(venv))
|
||||
sys.modules.pop("sr_perm", None)
|
||||
mod = _load("sr_perm", STORAGE_ROOTS)
|
||||
with mock.patch.object(Path, "is_file", side_effect = PermissionError("denied")):
|
||||
# Must NOT raise.
|
||||
assert mod._infer_studio_home_from_venv() is None
|
||||
|
||||
|
||||
def test_studio_root_does_not_crash_on_permission_error(tmp_path, monkeypatch):
|
||||
"""studio_root() must remain callable even when the venv inference
|
||||
encounters a restricted filesystem; it should fall through to the
|
||||
legacy default."""
|
||||
candidate = tmp_path / "fake_root"
|
||||
venv = candidate / "unsloth_studio"
|
||||
venv.mkdir(parents = True)
|
||||
monkeypatch.setattr(sys, "prefix", str(venv))
|
||||
monkeypatch.delenv("UNSLOTH_STUDIO_HOME", raising = False)
|
||||
monkeypatch.delenv("STUDIO_HOME", raising = False)
|
||||
sys.modules.pop("sr_studio_perm", None)
|
||||
mod = _load("sr_studio_perm", STORAGE_ROOTS)
|
||||
with mock.patch.object(Path, "is_file", side_effect = OSError("ebusy")):
|
||||
result = mod.studio_root()
|
||||
assert result == Path.home() / ".unsloth" / "studio"
|
||||
|
||||
|
||||
def test_kill_orphan_catches_oserror_from_studio_root():
|
||||
"""_kill_orphaned_servers must catch (ImportError, OSError, ValueError)
|
||||
on the studio_root() probe specifically; the sister function
|
||||
_find_llama_server_binary uses the same broader catch on its own probe."""
|
||||
src = LLAMA_CPP.read_text()
|
||||
fn_start = src.index("def _kill_orphaned_servers")
|
||||
fn_body = src[fn_start : fn_start + 4000]
|
||||
# The studio_root() probe in this fn is the one that imports as `_sr`
|
||||
# and assigns `_resolved_sr = _sr()`. Find the except that closes it.
|
||||
probe_idx = fn_body.index("storage_roots import studio_root as _sr")
|
||||
# The matching except is the next `except ...:` after the inner
|
||||
# OSError/ValueError block that wraps resolve().
|
||||
after = fn_body[probe_idx:]
|
||||
# Skip over the inner `except (OSError, ValueError):` that wraps resolve().
|
||||
inner_idx = after.index("except (OSError, ValueError):")
|
||||
after_inner = after[inner_idx + len("except (OSError, ValueError):") :]
|
||||
outer_match = re.search(r"except\s*\(?[^)]*?\)?:", after_inner)
|
||||
assert outer_match, "outer except for studio_root probe missing"
|
||||
clause = outer_match.group(0)
|
||||
assert (
|
||||
"OSError" in clause and "ValueError" in clause
|
||||
), f"_kill_orphaned_servers studio_root probe catch too narrow: {clause!r}"
|
||||
|
||||
|
||||
def _exec_search_roots_block(
|
||||
home: Path, studio_root_value: Path, resolve_raises: bool
|
||||
) -> list[Path]:
|
||||
"""Extract _find_llama_server_binary's env-mode search_roots block
|
||||
and execute it with controlled inputs."""
|
||||
src = LLAMA_CPP.read_text()
|
||||
block_start = src.index('legacy_llama = Path.home() / ".unsloth" / "llama.cpp"')
|
||||
block_end = src.index("_seen_roots: set[str]", block_start)
|
||||
raw = src[block_start:block_end]
|
||||
indent = " " * 8
|
||||
block = textwrap.dedent(indent + raw)
|
||||
fake_module = type(sys)("fake_storage_roots")
|
||||
fake_module.studio_root = lambda: studio_root_value
|
||||
sys.modules["utils.paths.storage_roots"] = fake_module
|
||||
try:
|
||||
original_resolve = Path.resolve
|
||||
|
||||
def _resolve(self, *a, **k):
|
||||
if resolve_raises:
|
||||
raise OSError("ebusy")
|
||||
return original_resolve(self, *a, **k)
|
||||
|
||||
with (
|
||||
mock.patch.object(Path, "home", classmethod(lambda cls: home)),
|
||||
mock.patch.object(Path, "resolve", _resolve),
|
||||
):
|
||||
ns: dict = {"Path": Path}
|
||||
exec(block, ns) # noqa: S102
|
||||
return ns["search_roots"]
|
||||
finally:
|
||||
sys.modules.pop("utils.paths.storage_roots", None)
|
||||
|
||||
|
||||
def test_search_roots_keeps_custom_when_resolve_fails(tmp_path):
|
||||
home = tmp_path / "home"
|
||||
home.mkdir()
|
||||
custom = tmp_path / "custom_studio"
|
||||
custom.mkdir()
|
||||
roots = _exec_search_roots_block(
|
||||
home = home, studio_root_value = custom, resolve_raises = True
|
||||
)
|
||||
# On resolve() failure, the inner except falls back to direct equality;
|
||||
# custom != legacy_studio so the custom root must remain in search_roots.
|
||||
assert (
|
||||
custom / "llama.cpp" in roots
|
||||
), f"custom root dropped on resolve() failure: {roots}"
|
||||
# custom-mode discovery excludes the legacy tree to match _kill_orphaned_servers.
|
||||
assert (
|
||||
(home / ".unsloth" / "llama.cpp") not in roots
|
||||
), f"legacy llama path must not appear in custom-mode search_roots: {roots}"
|
||||
|
||||
|
||||
def test_search_roots_default_mode_uses_legacy_only(tmp_path):
|
||||
home = tmp_path / "home"
|
||||
home.mkdir()
|
||||
legacy = home / ".unsloth" / "studio"
|
||||
legacy.mkdir(parents = True)
|
||||
roots = _exec_search_roots_block(
|
||||
home = home, studio_root_value = legacy, resolve_raises = False
|
||||
)
|
||||
# Default mode: only legacy_llama.
|
||||
assert roots == [home / ".unsloth" / "llama.cpp"]
|
||||
|
|
@ -12,348 +12,117 @@
|
|||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import warnings, importlib, sys
|
||||
from packaging.version import Version
|
||||
import os, re, subprocess, inspect, functools
|
||||
import numpy as np
|
||||
import os, platform, importlib.util
|
||||
|
||||
# Log Unsloth is being used
|
||||
os.environ["UNSLOTH_IS_PRESENT"] = "1"
|
||||
|
||||
# Check if modules that need patching are already imported
|
||||
critical_modules = ["trl", "transformers", "peft"]
|
||||
already_imported = [mod for mod in critical_modules if mod in sys.modules]
|
||||
|
||||
# Fix some issues before importing other packages
|
||||
from .import_fixes import (
|
||||
fix_message_factory_issue,
|
||||
check_fbgemm_gpu_version,
|
||||
disable_broken_causal_conv1d,
|
||||
disable_broken_vllm,
|
||||
configure_amdgpu_asic_id_table_path,
|
||||
torchvision_compatibility_check,
|
||||
fix_diffusers_warnings,
|
||||
fix_huggingface_hub,
|
||||
# Detect Apple Silicon + MLX before any torch/numpy imports
|
||||
_IS_MLX = (
|
||||
platform.system() == "Darwin"
|
||||
and platform.machine() == "arm64"
|
||||
and importlib.util.find_spec("mlx") is not None
|
||||
)
|
||||
|
||||
# Configure libdrm ids table path early so ROCm can resolve AMD GPU names.
|
||||
configure_amdgpu_asic_id_table_path()
|
||||
disable_broken_causal_conv1d()
|
||||
disable_broken_vllm()
|
||||
fix_message_factory_issue()
|
||||
check_fbgemm_gpu_version()
|
||||
torchvision_compatibility_check()
|
||||
fix_diffusers_warnings()
|
||||
fix_huggingface_hub()
|
||||
del configure_amdgpu_asic_id_table_path
|
||||
del disable_broken_causal_conv1d
|
||||
del disable_broken_vllm
|
||||
del fix_message_factory_issue
|
||||
del check_fbgemm_gpu_version
|
||||
del torchvision_compatibility_check
|
||||
del fix_diffusers_warnings
|
||||
del fix_huggingface_hub
|
||||
|
||||
# This check is critical because Unsloth optimizes these libraries by modifying
|
||||
# their code at import time. If they're imported first, the original (slower,
|
||||
# more memory-intensive) implementations will be used instead of Unsloth's
|
||||
# optimized versions, potentially causing OOM errors or slower training.
|
||||
if already_imported:
|
||||
# stacklevel=2 makes warning point to user's import line rather than this library code,
|
||||
# showing them exactly where to fix the import order in their script
|
||||
warnings.warn(
|
||||
f"WARNING: Unsloth should be imported before [{', '.join(already_imported)}] "
|
||||
f"to ensure all optimizations are applied. Your code may run slower or encounter "
|
||||
f"memory issues without these optimizations.\n\n"
|
||||
f"Please restructure your imports with 'import unsloth' at the top of your file.",
|
||||
stacklevel = 2,
|
||||
)
|
||||
del already_imported, critical_modules
|
||||
|
||||
# Unsloth currently does not work on multi GPU setups - sadly we are a 2 brother team so
|
||||
# enabling it will require much more work, so we have to prioritize. Please understand!
|
||||
# We do have a beta version, which you can contact us about!
|
||||
# Thank you for your understanding and we appreciate it immensely!
|
||||
|
||||
# Fixes https://github.com/unslothai/unsloth/issues/1266
|
||||
os.environ["PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION"] = "python"
|
||||
|
||||
# [TODO] Check why some GPUs don't work
|
||||
# "pinned_use_cuda_host_register:True,"\
|
||||
# "pinned_num_register_threads:8"
|
||||
|
||||
|
||||
from importlib.metadata import version as importlib_version
|
||||
from importlib.metadata import PackageNotFoundError
|
||||
|
||||
# Check for unsloth_zoo
|
||||
try:
|
||||
unsloth_zoo_version = importlib_version("unsloth_zoo")
|
||||
if Version(unsloth_zoo_version) < Version("2026.3.4"):
|
||||
print(
|
||||
"Unsloth: Please update Unsloth and Unsloth-Zoo to the latest version!\n"
|
||||
"Do this via `pip install --upgrade --force-reinstall --no-cache-dir --no-deps unsloth unsloth_zoo`"
|
||||
)
|
||||
# if os.environ.get("UNSLOTH_DISABLE_AUTO_UPDATES", "0") == "0":
|
||||
# try:
|
||||
# os.system("pip install --upgrade --no-cache-dir --no-deps unsloth_zoo")
|
||||
# except:
|
||||
# try:
|
||||
# os.system("pip install --upgrade --no-cache-dir --no-deps --user unsloth_zoo")
|
||||
# except:
|
||||
# raise ImportError("Unsloth: Please update unsloth_zoo via `pip install --upgrade --no-cache-dir --no-deps unsloth_zoo`")
|
||||
import unsloth_zoo
|
||||
except PackageNotFoundError:
|
||||
raise ImportError(
|
||||
f"Unsloth: Please install unsloth_zoo via `pip install unsloth_zoo` then retry!"
|
||||
)
|
||||
except:
|
||||
raise
|
||||
del PackageNotFoundError, importlib_version
|
||||
|
||||
# Try importing PyTorch and check version
|
||||
try:
|
||||
import torch
|
||||
except ModuleNotFoundError:
|
||||
raise ImportError(
|
||||
"Unsloth: Pytorch is not installed. Go to https://pytorch.org/.\n"
|
||||
"We have some installation instructions on our Github page."
|
||||
)
|
||||
except:
|
||||
raise
|
||||
|
||||
from unsloth_zoo.device_type import (
|
||||
is_hip,
|
||||
get_device_type,
|
||||
DEVICE_TYPE,
|
||||
DEVICE_TYPE_TORCH,
|
||||
DEVICE_COUNT,
|
||||
ALLOW_PREQUANTIZED_MODELS,
|
||||
)
|
||||
|
||||
# Fix other issues
|
||||
from .import_fixes import (
|
||||
fix_xformers_performance_issue,
|
||||
fix_vllm_aimv2_issue,
|
||||
check_vllm_torch_sm100_compatibility,
|
||||
fix_vllm_guided_decoding_params,
|
||||
fix_trl_vllm_ascend,
|
||||
fix_vllm_pdl_blackwell,
|
||||
fix_triton_compiled_kernel_missing_attrs,
|
||||
patch_trunc_normal_precision_issue,
|
||||
ignore_logger_messages,
|
||||
patch_ipykernel_hf_xet,
|
||||
patch_trackio,
|
||||
patch_datasets,
|
||||
patch_enable_input_require_grads,
|
||||
fix_openenv_no_vllm,
|
||||
patch_openspiel_env_async,
|
||||
fix_executorch,
|
||||
patch_vllm_for_notebooks,
|
||||
patch_torchcodec_audio_decoder,
|
||||
disable_torchcodec_if_broken,
|
||||
disable_broken_wandb,
|
||||
patch_peft_weight_converter_compatibility,
|
||||
)
|
||||
|
||||
fix_xformers_performance_issue()
|
||||
fix_vllm_aimv2_issue()
|
||||
# Check vLLM + torch < 2.9.0 + SM100 compatibility BEFORE importing vLLM
|
||||
check_vllm_torch_sm100_compatibility()
|
||||
fix_vllm_guided_decoding_params()
|
||||
fix_trl_vllm_ascend()
|
||||
fix_vllm_pdl_blackwell()
|
||||
fix_triton_compiled_kernel_missing_attrs()
|
||||
patch_trunc_normal_precision_issue()
|
||||
ignore_logger_messages()
|
||||
patch_ipykernel_hf_xet()
|
||||
patch_trackio()
|
||||
patch_datasets()
|
||||
patch_enable_input_require_grads()
|
||||
fix_openenv_no_vllm()
|
||||
patch_openspiel_env_async()
|
||||
fix_executorch()
|
||||
patch_vllm_for_notebooks()
|
||||
patch_torchcodec_audio_decoder()
|
||||
disable_torchcodec_if_broken()
|
||||
disable_broken_wandb()
|
||||
patch_peft_weight_converter_compatibility()
|
||||
|
||||
del fix_xformers_performance_issue
|
||||
del fix_vllm_aimv2_issue
|
||||
del check_vllm_torch_sm100_compatibility
|
||||
del fix_vllm_guided_decoding_params
|
||||
del fix_trl_vllm_ascend
|
||||
del fix_vllm_pdl_blackwell
|
||||
del fix_triton_compiled_kernel_missing_attrs
|
||||
del patch_trunc_normal_precision_issue
|
||||
del ignore_logger_messages
|
||||
del patch_ipykernel_hf_xet
|
||||
del patch_trackio
|
||||
del patch_datasets
|
||||
del patch_enable_input_require_grads
|
||||
del fix_openenv_no_vllm
|
||||
del patch_openspiel_env_async
|
||||
del fix_executorch
|
||||
del patch_vllm_for_notebooks
|
||||
del patch_torchcodec_audio_decoder
|
||||
del disable_torchcodec_if_broken
|
||||
del disable_broken_wandb
|
||||
del patch_peft_weight_converter_compatibility
|
||||
|
||||
# Torch 2.4 has including_emulation
|
||||
if DEVICE_TYPE == "cuda":
|
||||
major_version, minor_version = torch.cuda.get_device_capability()
|
||||
SUPPORTS_BFLOAT16 = major_version >= 8
|
||||
|
||||
old_is_bf16_supported = torch.cuda.is_bf16_supported
|
||||
if "including_emulation" in str(inspect.signature(old_is_bf16_supported)):
|
||||
|
||||
def is_bf16_supported(including_emulation = False):
|
||||
return old_is_bf16_supported(including_emulation)
|
||||
|
||||
torch.cuda.is_bf16_supported = is_bf16_supported
|
||||
else:
|
||||
|
||||
def is_bf16_supported():
|
||||
return SUPPORTS_BFLOAT16
|
||||
|
||||
torch.cuda.is_bf16_supported = is_bf16_supported
|
||||
del major_version, minor_version
|
||||
elif DEVICE_TYPE == "hip":
|
||||
SUPPORTS_BFLOAT16 = torch.cuda.is_bf16_supported()
|
||||
elif DEVICE_TYPE == "xpu":
|
||||
# torch.xpu.is_bf16_supported() does not have including_emulation
|
||||
# set SUPPORTS_BFLOAT16 as torch.xpu.is_bf16_supported()
|
||||
SUPPORTS_BFLOAT16 = torch.xpu.is_bf16_supported()
|
||||
|
||||
# For Gradio HF Spaces?
|
||||
# if "SPACE_AUTHOR_NAME" not in os.environ and "SPACE_REPO_NAME" not in os.environ:
|
||||
import triton
|
||||
|
||||
if DEVICE_TYPE == "cuda":
|
||||
libcuda_dirs = lambda: None
|
||||
if Version(triton.__version__) >= Version("3.0.0"):
|
||||
try:
|
||||
from triton.backends.nvidia.driver import libcuda_dirs
|
||||
except:
|
||||
pass
|
||||
else:
|
||||
from triton.common.build import libcuda_dirs
|
||||
|
||||
# Try loading bitsandbytes and triton
|
||||
if _IS_MLX:
|
||||
try:
|
||||
import bitsandbytes as bnb
|
||||
except:
|
||||
print(
|
||||
"Unsloth: `bitsandbytes` is not installed - 4bit QLoRA unallowed, but 16bit and full finetuning works!"
|
||||
)
|
||||
bnb = None
|
||||
import unsloth_zoo
|
||||
except ImportError as _e:
|
||||
raise ImportError(
|
||||
"Unsloth: MLX support requires `unsloth-zoo` with MLX modules. "
|
||||
"Reinstall with `pip install unsloth-zoo` or rerun install.sh."
|
||||
) from _e
|
||||
# The mlx_trainer / mlx_loader submodules ship with unsloth-zoo's MLX
|
||||
# support. An older installed unsloth-zoo (e.g. from PyPI before the
|
||||
# MLX release lands) will satisfy `import unsloth_zoo` but be missing
|
||||
# these submodules. Surface the same friendly install hint instead of
|
||||
# a raw ImportError on the submodule path.
|
||||
try:
|
||||
cdequantize_blockwise_fp32 = bnb.functional.lib.cdequantize_blockwise_fp32
|
||||
libcuda_dirs()
|
||||
except:
|
||||
# Only run the ldconfig recovery when we can actually run
|
||||
# ldconfig (root). On non-root environments (shared HPC,
|
||||
# locked-down containers, CI runners, etc.) the recovery would
|
||||
# shell out to `ldconfig` and fail with "Permission denied",
|
||||
# which is especially noisy for users who don't even have
|
||||
# bitsandbytes installed and are just doing 16bit/full
|
||||
# finetuning. libcuda_dirs() is used by both triton and bnb,
|
||||
# so we still run the recovery whenever we're root, regardless
|
||||
# of whether bnb is installed.
|
||||
if hasattr(os, "geteuid") and os.geteuid() == 0:
|
||||
warnings.warn("Unsloth: Running `ldconfig /usr/lib64-nvidia` to link CUDA.")
|
||||
from unsloth_zoo.mlx_trainer import MLXTrainer, MLXTrainingConfig
|
||||
from unsloth_zoo.mlx_loader import FastMLXModel
|
||||
except ImportError as _e:
|
||||
raise ImportError(
|
||||
"Unsloth: MLX support requires an unsloth-zoo build that includes "
|
||||
"`unsloth_zoo.mlx_trainer` and `unsloth_zoo.mlx_loader`. Upgrade with "
|
||||
"`pip install -U unsloth-zoo` or rerun install.sh."
|
||||
) from _e
|
||||
|
||||
if os.path.exists("/usr/lib64-nvidia"):
|
||||
os.system("ldconfig /usr/lib64-nvidia")
|
||||
elif os.path.exists("/usr/local"):
|
||||
# Sometimes bitsandbytes cannot be linked properly in Runpod for example
|
||||
possible_cudas = (
|
||||
subprocess.check_output(["ls", "-al", "/usr/local"])
|
||||
.decode("utf-8")
|
||||
.split("\n")
|
||||
)
|
||||
find_cuda = re.compile(r"[\s](cuda\-[\d\.]{2,})$")
|
||||
possible_cudas = [find_cuda.search(x) for x in possible_cudas]
|
||||
possible_cudas = [x.group(1) for x in possible_cudas if x is not None]
|
||||
# Load raw_text helpers without executing dataprep/__init__.py, which
|
||||
# imports synthetic.py -> torch and would defeat the torch-free MLX path.
|
||||
from pathlib import Path as _Path
|
||||
|
||||
# Try linking cuda folder, or everything in local
|
||||
if len(possible_cudas) == 0:
|
||||
os.system("ldconfig /usr/local/")
|
||||
else:
|
||||
find_number = re.compile(r"([\d\.]{2,})")
|
||||
latest_cuda = np.argsort(
|
||||
[float(find_number.search(x).group(1)) for x in possible_cudas]
|
||||
)[::-1][0]
|
||||
latest_cuda = possible_cudas[latest_cuda]
|
||||
os.system(f"ldconfig /usr/local/{latest_cuda}")
|
||||
del find_number, latest_cuda
|
||||
del possible_cudas, find_cuda
|
||||
_raw_text_path = _Path(__file__).resolve().parent / "dataprep" / "raw_text.py"
|
||||
_raw_text_spec = importlib.util.spec_from_file_location(
|
||||
"unsloth._mlx_raw_text", _raw_text_path
|
||||
)
|
||||
if _raw_text_spec is None or _raw_text_spec.loader is None:
|
||||
raise ImportError("Unsloth: could not load MLX raw_text dataprep helpers.")
|
||||
_raw_text = importlib.util.module_from_spec(_raw_text_spec)
|
||||
_raw_text_spec.loader.exec_module(_raw_text)
|
||||
RawTextDataLoader = _raw_text.RawTextDataLoader
|
||||
TextPreprocessor = _raw_text.TextPreprocessor
|
||||
del _raw_text, _raw_text_spec, _raw_text_path, _Path
|
||||
|
||||
if bnb is not None:
|
||||
importlib.reload(bnb)
|
||||
importlib.reload(triton)
|
||||
try:
|
||||
libcuda_dirs = lambda: None
|
||||
if Version(triton.__version__) >= Version("3.0.0"):
|
||||
try:
|
||||
from triton.backends.nvidia.driver import libcuda_dirs
|
||||
except:
|
||||
pass
|
||||
else:
|
||||
from triton.common.build import libcuda_dirs
|
||||
cdequantize_blockwise_fp32 = (
|
||||
bnb.functional.lib.cdequantize_blockwise_fp32
|
||||
)
|
||||
libcuda_dirs()
|
||||
except:
|
||||
warnings.warn(
|
||||
"Unsloth: CUDA is not linked properly.\n"
|
||||
"Try running `python -m bitsandbytes` then `python -m xformers.info`\n"
|
||||
"We tried running `ldconfig /usr/lib64-nvidia` ourselves, but it didn't work.\n"
|
||||
"You need to run in your terminal `sudo ldconfig /usr/lib64-nvidia` yourself, then import Unsloth.\n"
|
||||
"Also try `sudo ldconfig /usr/local/cuda-xx.x` - find the latest cuda version.\n"
|
||||
"Unsloth will still run for now, but maybe it might crash - let's hope it works!"
|
||||
)
|
||||
elif bnb is not None:
|
||||
# Non-root + bnb installed: we can't run ldconfig ourselves,
|
||||
# but bnb is going to crash later when the user actually uses
|
||||
# 4bit quantization - tell them how to fix it manually so
|
||||
# they're not surprised by an opaque error down the road.
|
||||
warnings.warn(
|
||||
"Unsloth: CUDA is not linked properly.\n"
|
||||
"You need to run in your terminal `sudo ldconfig /usr/lib64-nvidia` yourself, then import Unsloth.\n"
|
||||
"Also try `sudo ldconfig /usr/local/cuda-xx.x` - find the latest cuda version.\n"
|
||||
"Unsloth will still run for now, but maybe it might crash - let's hope it works!"
|
||||
__version__ = unsloth_zoo.__version__
|
||||
DEVICE_TYPE = "mlx"
|
||||
|
||||
class FastLanguageModel:
|
||||
@staticmethod
|
||||
def from_pretrained(*args, **kwargs):
|
||||
return FastMLXModel.from_pretrained(*args, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
def get_peft_model(*args, **kwargs):
|
||||
return FastMLXModel.get_peft_model(*args, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
def for_inference(*args, **kwargs):
|
||||
return args[0] if args else None
|
||||
|
||||
class FastVisionModel(FastLanguageModel):
|
||||
@staticmethod
|
||||
def from_pretrained(*args, **kwargs):
|
||||
kwargs.setdefault("text_only", False)
|
||||
return FastMLXModel.from_pretrained(*args, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
def for_training(*args, **kwargs):
|
||||
return args[0] if args else None
|
||||
|
||||
FastTextModel = FastLanguageModel
|
||||
FastModel = FastLanguageModel
|
||||
|
||||
class FastSentenceTransformer:
|
||||
@staticmethod
|
||||
def from_pretrained(*args, **kwargs):
|
||||
raise NotImplementedError(
|
||||
"Unsloth: FastSentenceTransformer is not yet supported on MLX."
|
||||
)
|
||||
del libcuda_dirs
|
||||
elif DEVICE_TYPE == "hip":
|
||||
# NO-OP for rocm device
|
||||
pass
|
||||
elif DEVICE_TYPE == "xpu":
|
||||
import bitsandbytes as bnb
|
||||
|
||||
# TODO: check triton for intel installed properly.
|
||||
pass
|
||||
@staticmethod
|
||||
def get_peft_model(*args, **kwargs):
|
||||
raise NotImplementedError(
|
||||
"Unsloth: FastSentenceTransformer is not yet supported on MLX."
|
||||
)
|
||||
|
||||
from .models import *
|
||||
from .models import __version__
|
||||
from .save import *
|
||||
from .chat_templates import *
|
||||
from .tokenizer_utils import *
|
||||
from .trainer import *
|
||||
def is_bfloat16_supported():
|
||||
try:
|
||||
import mlx.core as mx
|
||||
|
||||
# Export dataprep utilities for CLI and downstream users
|
||||
from .dataprep.raw_text import RawTextDataLoader, TextPreprocessor
|
||||
from unsloth_zoo.rl_environments import (
|
||||
check_python_modules,
|
||||
create_locked_down_function,
|
||||
execute_with_time_limit,
|
||||
Benchmarker,
|
||||
is_port_open,
|
||||
launch_openenv,
|
||||
)
|
||||
name = mx.device_info().get("device_name", "") or ""
|
||||
return not name.startswith(("Apple M1", "Apple M2"))
|
||||
except Exception:
|
||||
return True
|
||||
|
||||
# Patch TRL trainers for backwards compatibility
|
||||
_patch_trl_trainer()
|
||||
is_bf16_supported = is_bfloat16_supported
|
||||
|
||||
class UnslothVisionDataCollator:
|
||||
def __init__(self, *args, **kwargs):
|
||||
raise NotImplementedError(
|
||||
"Unsloth: UnslothVisionDataCollator is not used on MLX. "
|
||||
"Use the MLX trainer/data path instead."
|
||||
)
|
||||
|
||||
else:
|
||||
# GPU path: load everything from _gpu_init
|
||||
from ._gpu_init import *
|
||||
from ._gpu_init import __version__
|
||||
|
|
|
|||
346
unsloth/_gpu_init.py
Normal file
346
unsloth/_gpu_init.py
Normal file
|
|
@ -0,0 +1,346 @@
|
|||
# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import warnings, importlib, sys
|
||||
from packaging.version import Version
|
||||
import os, re, subprocess, inspect, functools
|
||||
import numpy as np
|
||||
|
||||
# Log Unsloth is being used
|
||||
os.environ["UNSLOTH_IS_PRESENT"] = "1"
|
||||
|
||||
# Check if modules that need patching are already imported
|
||||
critical_modules = ["trl", "transformers", "peft"]
|
||||
already_imported = [mod for mod in critical_modules if mod in sys.modules]
|
||||
|
||||
# Fix some issues before importing other packages
|
||||
from .import_fixes import (
|
||||
fix_message_factory_issue,
|
||||
check_fbgemm_gpu_version,
|
||||
disable_broken_causal_conv1d,
|
||||
disable_broken_vllm,
|
||||
configure_amdgpu_asic_id_table_path,
|
||||
torchvision_compatibility_check,
|
||||
fix_diffusers_warnings,
|
||||
fix_huggingface_hub,
|
||||
)
|
||||
|
||||
# Configure libdrm ids table path early so ROCm can resolve AMD GPU names.
|
||||
configure_amdgpu_asic_id_table_path()
|
||||
disable_broken_causal_conv1d()
|
||||
disable_broken_vllm()
|
||||
fix_message_factory_issue()
|
||||
check_fbgemm_gpu_version()
|
||||
torchvision_compatibility_check()
|
||||
fix_diffusers_warnings()
|
||||
fix_huggingface_hub()
|
||||
del configure_amdgpu_asic_id_table_path
|
||||
del disable_broken_causal_conv1d
|
||||
del disable_broken_vllm
|
||||
del fix_message_factory_issue
|
||||
del check_fbgemm_gpu_version
|
||||
del torchvision_compatibility_check
|
||||
del fix_diffusers_warnings
|
||||
del fix_huggingface_hub
|
||||
|
||||
# This check is critical because Unsloth optimizes these libraries by modifying
|
||||
# their code at import time. If they're imported first, the original (slower,
|
||||
# more memory-intensive) implementations will be used instead of Unsloth's
|
||||
# optimized versions, potentially causing OOM errors or slower training.
|
||||
if already_imported:
|
||||
# stacklevel=2 makes warning point to user's import line rather than this library code,
|
||||
# showing them exactly where to fix the import order in their script
|
||||
warnings.warn(
|
||||
f"WARNING: Unsloth should be imported before [{', '.join(already_imported)}] "
|
||||
f"to ensure all optimizations are applied. Your code may run slower or encounter "
|
||||
f"memory issues without these optimizations.\n\n"
|
||||
f"Please restructure your imports with 'import unsloth' at the top of your file.",
|
||||
stacklevel = 2,
|
||||
)
|
||||
del already_imported, critical_modules
|
||||
|
||||
# Unsloth currently does not work on multi GPU setups - sadly we are a 2 brother team so
|
||||
# enabling it will require much more work, so we have to prioritize. Please understand!
|
||||
# We do have a beta version, which you can contact us about!
|
||||
# Thank you for your understanding and we appreciate it immensely!
|
||||
|
||||
# Fixes https://github.com/unslothai/unsloth/issues/1266
|
||||
os.environ["PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION"] = "python"
|
||||
|
||||
# [TODO] Check why some GPUs don't work
|
||||
# "pinned_use_cuda_host_register:True,"\
|
||||
# "pinned_num_register_threads:8"
|
||||
|
||||
|
||||
from importlib.metadata import version as importlib_version
|
||||
from importlib.metadata import PackageNotFoundError
|
||||
|
||||
# Check for unsloth_zoo
|
||||
try:
|
||||
unsloth_zoo_version = importlib_version("unsloth_zoo")
|
||||
if Version(unsloth_zoo_version) < Version("2026.3.4"):
|
||||
print(
|
||||
"Unsloth: Please update Unsloth and Unsloth-Zoo to the latest version!\n"
|
||||
"Do this via `pip install --upgrade --force-reinstall --no-cache-dir --no-deps unsloth unsloth_zoo`"
|
||||
)
|
||||
# if os.environ.get("UNSLOTH_DISABLE_AUTO_UPDATES", "0") == "0":
|
||||
# try:
|
||||
# os.system("pip install --upgrade --no-cache-dir --no-deps unsloth_zoo")
|
||||
# except:
|
||||
# try:
|
||||
# os.system("pip install --upgrade --no-cache-dir --no-deps --user unsloth_zoo")
|
||||
# except:
|
||||
# raise ImportError("Unsloth: Please update unsloth_zoo via `pip install --upgrade --no-cache-dir --no-deps unsloth_zoo`")
|
||||
import unsloth_zoo
|
||||
except PackageNotFoundError:
|
||||
raise ImportError(
|
||||
f"Unsloth: Please install unsloth_zoo via `pip install unsloth_zoo` then retry!"
|
||||
)
|
||||
except:
|
||||
raise
|
||||
del PackageNotFoundError, importlib_version
|
||||
|
||||
# Try importing PyTorch and check version
|
||||
try:
|
||||
import torch
|
||||
except ModuleNotFoundError:
|
||||
raise ImportError(
|
||||
"Unsloth: Pytorch is not installed. Go to https://pytorch.org/.\n"
|
||||
"We have some installation instructions on our Github page."
|
||||
)
|
||||
except:
|
||||
raise
|
||||
|
||||
from unsloth_zoo.device_type import (
|
||||
is_hip,
|
||||
get_device_type,
|
||||
DEVICE_TYPE,
|
||||
DEVICE_TYPE_TORCH,
|
||||
DEVICE_COUNT,
|
||||
ALLOW_PREQUANTIZED_MODELS,
|
||||
)
|
||||
|
||||
# Fix other issues
|
||||
from .import_fixes import (
|
||||
fix_xformers_performance_issue,
|
||||
fix_vllm_aimv2_issue,
|
||||
check_vllm_torch_sm100_compatibility,
|
||||
fix_vllm_guided_decoding_params,
|
||||
fix_vllm_pdl_blackwell,
|
||||
fix_triton_compiled_kernel_missing_attrs,
|
||||
patch_trunc_normal_precision_issue,
|
||||
ignore_logger_messages,
|
||||
patch_ipykernel_hf_xet,
|
||||
patch_trackio,
|
||||
patch_datasets,
|
||||
patch_enable_input_require_grads,
|
||||
fix_openenv_no_vllm,
|
||||
patch_openspiel_env_async,
|
||||
fix_executorch,
|
||||
patch_vllm_for_notebooks,
|
||||
patch_torchcodec_audio_decoder,
|
||||
disable_torchcodec_if_broken,
|
||||
disable_broken_wandb,
|
||||
fix_trl_vllm_ascend,
|
||||
patch_peft_weight_converter_compatibility,
|
||||
)
|
||||
|
||||
fix_xformers_performance_issue()
|
||||
fix_vllm_aimv2_issue()
|
||||
# Check vLLM + torch < 2.9.0 + SM100 compatibility BEFORE importing vLLM
|
||||
check_vllm_torch_sm100_compatibility()
|
||||
fix_vllm_guided_decoding_params()
|
||||
fix_trl_vllm_ascend()
|
||||
fix_vllm_pdl_blackwell()
|
||||
fix_triton_compiled_kernel_missing_attrs()
|
||||
patch_trunc_normal_precision_issue()
|
||||
ignore_logger_messages()
|
||||
patch_ipykernel_hf_xet()
|
||||
patch_trackio()
|
||||
patch_datasets()
|
||||
patch_enable_input_require_grads()
|
||||
fix_openenv_no_vllm()
|
||||
patch_openspiel_env_async()
|
||||
fix_executorch()
|
||||
patch_vllm_for_notebooks()
|
||||
patch_torchcodec_audio_decoder()
|
||||
disable_torchcodec_if_broken()
|
||||
disable_broken_wandb()
|
||||
patch_peft_weight_converter_compatibility()
|
||||
|
||||
del fix_xformers_performance_issue
|
||||
del fix_vllm_aimv2_issue
|
||||
del check_vllm_torch_sm100_compatibility
|
||||
del fix_vllm_guided_decoding_params
|
||||
del fix_trl_vllm_ascend
|
||||
del fix_vllm_pdl_blackwell
|
||||
del fix_triton_compiled_kernel_missing_attrs
|
||||
del patch_trunc_normal_precision_issue
|
||||
del ignore_logger_messages
|
||||
del patch_ipykernel_hf_xet
|
||||
del patch_trackio
|
||||
del patch_datasets
|
||||
del patch_enable_input_require_grads
|
||||
del fix_openenv_no_vllm
|
||||
del patch_openspiel_env_async
|
||||
del fix_executorch
|
||||
del patch_vllm_for_notebooks
|
||||
del patch_torchcodec_audio_decoder
|
||||
del disable_torchcodec_if_broken
|
||||
del disable_broken_wandb
|
||||
del patch_peft_weight_converter_compatibility
|
||||
|
||||
# Torch 2.4 has including_emulation
|
||||
if DEVICE_TYPE == "cuda":
|
||||
major_version, minor_version = torch.cuda.get_device_capability()
|
||||
SUPPORTS_BFLOAT16 = major_version >= 8
|
||||
|
||||
old_is_bf16_supported = torch.cuda.is_bf16_supported
|
||||
if "including_emulation" in str(inspect.signature(old_is_bf16_supported)):
|
||||
|
||||
def is_bf16_supported(including_emulation = False):
|
||||
return old_is_bf16_supported(including_emulation)
|
||||
|
||||
torch.cuda.is_bf16_supported = is_bf16_supported
|
||||
else:
|
||||
|
||||
def is_bf16_supported():
|
||||
return SUPPORTS_BFLOAT16
|
||||
|
||||
torch.cuda.is_bf16_supported = is_bf16_supported
|
||||
del major_version, minor_version
|
||||
elif DEVICE_TYPE == "hip":
|
||||
SUPPORTS_BFLOAT16 = torch.cuda.is_bf16_supported()
|
||||
elif DEVICE_TYPE == "xpu":
|
||||
# torch.xpu.is_bf16_supported() does not have including_emulation
|
||||
# set SUPPORTS_BFLOAT16 as torch.xpu.is_bf16_supported()
|
||||
SUPPORTS_BFLOAT16 = torch.xpu.is_bf16_supported()
|
||||
|
||||
# For Gradio HF Spaces?
|
||||
# if "SPACE_AUTHOR_NAME" not in os.environ and "SPACE_REPO_NAME" not in os.environ:
|
||||
import triton
|
||||
|
||||
if DEVICE_TYPE == "cuda":
|
||||
libcuda_dirs = lambda: None
|
||||
if Version(triton.__version__) >= Version("3.0.0"):
|
||||
try:
|
||||
from triton.backends.nvidia.driver import libcuda_dirs
|
||||
except:
|
||||
pass
|
||||
else:
|
||||
from triton.common.build import libcuda_dirs
|
||||
|
||||
# Try loading bitsandbytes and triton
|
||||
try:
|
||||
import bitsandbytes as bnb
|
||||
except:
|
||||
print(
|
||||
"Unsloth: `bitsandbytes` is not installed - 4bit QLoRA unallowed, but 16bit and full finetuning works!"
|
||||
)
|
||||
bnb = None
|
||||
try:
|
||||
cdequantize_blockwise_fp32 = bnb.functional.lib.cdequantize_blockwise_fp32
|
||||
libcuda_dirs()
|
||||
except:
|
||||
if hasattr(os, "geteuid") and os.geteuid() == 0:
|
||||
warnings.warn("Unsloth: Running `ldconfig /usr/lib64-nvidia` to link CUDA.")
|
||||
|
||||
if os.path.exists("/usr/lib64-nvidia"):
|
||||
os.system("ldconfig /usr/lib64-nvidia")
|
||||
elif os.path.exists("/usr/local"):
|
||||
# Sometimes bitsandbytes cannot be linked properly in Runpod for example
|
||||
possible_cudas = (
|
||||
subprocess.check_output(["ls", "-al", "/usr/local"])
|
||||
.decode("utf-8")
|
||||
.split("\n")
|
||||
)
|
||||
find_cuda = re.compile(r"[\s](cuda\-[\d\.]{2,})$")
|
||||
possible_cudas = [find_cuda.search(x) for x in possible_cudas]
|
||||
possible_cudas = [x.group(1) for x in possible_cudas if x is not None]
|
||||
|
||||
# Try linking cuda folder, or everything in local
|
||||
if len(possible_cudas) == 0:
|
||||
os.system("ldconfig /usr/local/")
|
||||
else:
|
||||
find_number = re.compile(r"([\d\.]{2,})")
|
||||
latest_cuda = np.argsort(
|
||||
[float(find_number.search(x).group(1)) for x in possible_cudas]
|
||||
)[::-1][0]
|
||||
latest_cuda = possible_cudas[latest_cuda]
|
||||
os.system(f"ldconfig /usr/local/{latest_cuda}")
|
||||
del find_number, latest_cuda
|
||||
del possible_cudas, find_cuda
|
||||
|
||||
if bnb is not None:
|
||||
importlib.reload(bnb)
|
||||
importlib.reload(triton)
|
||||
try:
|
||||
libcuda_dirs = lambda: None
|
||||
if Version(triton.__version__) >= Version("3.0.0"):
|
||||
try:
|
||||
from triton.backends.nvidia.driver import libcuda_dirs
|
||||
except:
|
||||
pass
|
||||
else:
|
||||
from triton.common.build import libcuda_dirs
|
||||
cdequantize_blockwise_fp32 = (
|
||||
bnb.functional.lib.cdequantize_blockwise_fp32
|
||||
)
|
||||
libcuda_dirs()
|
||||
except:
|
||||
warnings.warn(
|
||||
"Unsloth: CUDA is not linked properly.\n"
|
||||
"Try running `python -m bitsandbytes` then `python -m xformers.info`\n"
|
||||
"We tried running `ldconfig /usr/lib64-nvidia` ourselves, but it didn't work.\n"
|
||||
"You need to run in your terminal `sudo ldconfig /usr/lib64-nvidia` yourself, then import Unsloth.\n"
|
||||
"Also try `sudo ldconfig /usr/local/cuda-xx.x` - find the latest cuda version.\n"
|
||||
"Unsloth will still run for now, but maybe it might crash - let's hope it works!"
|
||||
)
|
||||
elif bnb is not None:
|
||||
warnings.warn(
|
||||
"Unsloth: CUDA is not linked properly.\n"
|
||||
"You need to run in your terminal `sudo ldconfig /usr/lib64-nvidia` yourself, then import Unsloth.\n"
|
||||
"Also try `sudo ldconfig /usr/local/cuda-xx.x` - find the latest cuda version.\n"
|
||||
"Unsloth will still run for now, but maybe it might crash - let's hope it works!"
|
||||
)
|
||||
del libcuda_dirs
|
||||
elif DEVICE_TYPE == "hip":
|
||||
# NO-OP for rocm device
|
||||
pass
|
||||
elif DEVICE_TYPE == "xpu":
|
||||
import bitsandbytes as bnb
|
||||
|
||||
# TODO: check triton for intel installed properly.
|
||||
pass
|
||||
|
||||
from .models import *
|
||||
from .models import __version__
|
||||
from .save import *
|
||||
from .chat_templates import *
|
||||
from .tokenizer_utils import *
|
||||
from .trainer import *
|
||||
|
||||
# Export dataprep utilities for CLI and downstream users
|
||||
from .dataprep.raw_text import RawTextDataLoader, TextPreprocessor
|
||||
from unsloth_zoo.rl_environments import (
|
||||
check_python_modules,
|
||||
create_locked_down_function,
|
||||
execute_with_time_limit,
|
||||
Benchmarker,
|
||||
is_port_open,
|
||||
launch_openenv,
|
||||
)
|
||||
|
||||
# Patch TRL trainers for backwards compatibility
|
||||
_patch_trl_trainer()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -20,7 +20,73 @@ import typer
|
|||
|
||||
studio_app = typer.Typer(help = "Unsloth Studio commands.")
|
||||
|
||||
STUDIO_HOME = Path.home() / ".unsloth" / "studio"
|
||||
|
||||
# Resolve install root: UNSLOTH_STUDIO_HOME, then STUDIO_HOME alias, then
|
||||
# sys.prefix inference (so a direct call to <root>/bin/unsloth resolves after
|
||||
# the installer's env var has expired), then legacy ~/.unsloth/studio.
|
||||
# UNSLOTH_STUDIO_HOME wins when both env vars are set.
|
||||
def _looks_like_installer_managed_studio_home(candidate: Path) -> bool:
|
||||
"""Sentinel check (studio.conf or bin shim) so a dev venv named
|
||||
unsloth_studio is not misidentified as a custom Studio root.
|
||||
"""
|
||||
shim_name = "unsloth.exe" if platform.system() == "Windows" else "unsloth"
|
||||
return (candidate / "share" / "studio.conf").is_file() or (
|
||||
candidate / "bin" / shim_name
|
||||
).is_file()
|
||||
|
||||
|
||||
def _resolve_studio_home() -> tuple[Path, bool]:
|
||||
override = (os.environ.get("UNSLOTH_STUDIO_HOME") or "").strip()
|
||||
if not override:
|
||||
override = (os.environ.get("STUDIO_HOME") or "").strip()
|
||||
if override:
|
||||
try:
|
||||
return Path(override).expanduser().resolve(), True
|
||||
except (OSError, ValueError):
|
||||
return Path(override).expanduser(), True
|
||||
try:
|
||||
prefix = Path(sys.prefix).resolve()
|
||||
if prefix.name == "unsloth_studio":
|
||||
inferred = prefix.parent
|
||||
legacy = (Path.home() / ".unsloth" / "studio").resolve()
|
||||
if inferred != legacy and _looks_like_installer_managed_studio_home(
|
||||
inferred
|
||||
):
|
||||
return inferred, True
|
||||
except (OSError, ValueError):
|
||||
pass
|
||||
return Path.home() / ".unsloth" / "studio", False
|
||||
|
||||
|
||||
STUDIO_HOME, _STUDIO_HOME_IS_CUSTOM = _resolve_studio_home()
|
||||
|
||||
|
||||
def _ensure_studio_env_exported() -> None:
|
||||
"""Re-export UNSLOTH_STUDIO_HOME / UNSLOTH_LLAMA_CPP_PATH only for real
|
||||
custom roots so subprocesses inherit the right install. Called from each
|
||||
studio subcommand entry rather than at import time, to avoid leaking env
|
||||
state into unrelated importers (tests, --help, CLI introspection).
|
||||
"""
|
||||
if not _STUDIO_HOME_IS_CUSTOM:
|
||||
return
|
||||
# Truthy-check (not setdefault) so a blank UNSLOTH_STUDIO_HOME= does not
|
||||
# suppress the inferred custom root.
|
||||
if not os.environ.get("UNSLOTH_STUDIO_HOME"):
|
||||
os.environ["UNSLOTH_STUDIO_HOME"] = str(STUDIO_HOME)
|
||||
# When override == legacy default, llama.cpp stays at ~/.unsloth/llama.cpp.
|
||||
try:
|
||||
_legacy_studio = (Path.home() / ".unsloth" / "studio").resolve()
|
||||
_is_legacy = STUDIO_HOME.resolve() == _legacy_studio
|
||||
except (OSError, ValueError):
|
||||
_is_legacy = STUDIO_HOME == (Path.home() / ".unsloth" / "studio")
|
||||
if _is_legacy:
|
||||
_llama_dir = Path.home() / ".unsloth" / "llama.cpp"
|
||||
else:
|
||||
_llama_dir = STUDIO_HOME / "llama.cpp"
|
||||
if not os.environ.get("UNSLOTH_LLAMA_CPP_PATH"):
|
||||
os.environ["UNSLOTH_LLAMA_CPP_PATH"] = str(_llama_dir)
|
||||
|
||||
|
||||
BOOTSTRAP_PASSWORD_FILE = ".bootstrap_password"
|
||||
DESKTOP_SECRET_FILE = ".desktop_secret"
|
||||
DEFAULT_ADMIN_USERNAME = "unsloth"
|
||||
|
|
@ -427,6 +493,8 @@ def studio_default(
|
|||
),
|
||||
):
|
||||
"""Launch the Unsloth Studio server."""
|
||||
# Runs before any subcommand; covers run/setup/update/etc in one place.
|
||||
_ensure_studio_env_exported()
|
||||
if ctx.invoked_subcommand is not None:
|
||||
return
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue