Final cleanup

This commit is contained in:
Roland Tannous 2026-03-12 18:28:04 +00:00
commit 985d2e43ee
123 changed files with 7474 additions and 5805 deletions

View file

@ -10,8 +10,8 @@ from cli.commands.ui import ui
from cli.commands.studio import studio_app
app = typer.Typer(
help="Command-line interface for Unsloth training, inference, and export.",
context_settings={"help_option_names": ["-h", "--help"]},
help = "Command-line interface for Unsloth training, inference, and export.",
context_settings = {"help_option_names": ["-h", "--help"]},
)
app.command()(train)
@ -19,4 +19,4 @@ app.command()(inference)
app.command()(export)
app.command("list-checkpoints")(list_checkpoints)
app.command()(ui)
app.add_typer(studio_app, name="studio", help="Unsloth Studio commands.")
app.add_typer(studio_app, name = "studio", help = "Unsloth Studio commands.")

View file

@ -0,0 +1,2 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0

View file

@ -13,14 +13,14 @@ GGUF_QUANTS = ["q4_k_m", "q5_k_m", "q8_0", "f16"]
def list_checkpoints(
outputs_dir: Path = typer.Option(
Path("./outputs"), "--outputs-dir", help="Directory that holds training runs."
Path("./outputs"), "--outputs-dir", help = "Directory that holds training runs."
),
):
"""List checkpoints detected in the outputs directory."""
from studio.backend.core.export import ExportBackend
backend = ExportBackend()
checkpoints = backend.scan_checkpoints(outputs_dir=str(outputs_dir))
checkpoints = backend.scan_checkpoints(outputs_dir = str(outputs_dir))
if not checkpoints:
typer.echo("No checkpoints found.")
raise typer.Exit()
@ -33,96 +33,100 @@ def list_checkpoints(
def export(
checkpoint: Path = typer.Argument(..., help="Path to checkpoint directory."),
output_dir: Path = typer.Argument(..., help="Directory to save exported model."),
checkpoint: Path = typer.Argument(..., help = "Path to checkpoint directory."),
output_dir: Path = typer.Argument(..., help = "Directory to save exported model."),
format: str = typer.Option(
"merged-16bit",
"--format",
"-f",
help=f"Export format: {', '.join(EXPORT_FORMATS)}",
help = f"Export format: {', '.join(EXPORT_FORMATS)}",
),
quantization: str = typer.Option(
"q4_k_m",
"--quantization",
"-q",
help=f"GGUF quantization method: {', '.join(GGUF_QUANTS)}",
help = f"GGUF quantization method: {', '.join(GGUF_QUANTS)}",
),
push_to_hub: bool = typer.Option(
False, "--push-to-hub", help="Push exported model to HuggingFace Hub."
False, "--push-to-hub", help = "Push exported model to HuggingFace Hub."
),
repo_id: Optional[str] = typer.Option(
None, "--repo-id", help="HuggingFace repo ID (username/model-name)."
None, "--repo-id", help = "HuggingFace repo ID (username/model-name)."
),
hf_token: Optional[str] = typer.Option(
None, "--hf-token", envvar="HF_TOKEN", help="HuggingFace token."
None, "--hf-token", envvar = "HF_TOKEN", help = "HuggingFace token."
),
private: bool = typer.Option(
False, "--private", help="Make the HuggingFace repo private."
False, "--private", help = "Make the HuggingFace repo private."
),
max_seq_length: int = typer.Option(2048, "--max-seq-length"),
load_in_4bit: bool = typer.Option(True, "--load-in-4bit/--no-load-in-4bit"),
):
"""Export a checkpoint to various formats (merged, GGUF, LoRA adapter)."""
if format not in EXPORT_FORMATS:
typer.echo(f"Error: Invalid format '{format}'. Choose from: {', '.join(EXPORT_FORMATS)}", err=True)
raise typer.Exit(code=2)
typer.echo(
f"Error: Invalid format '{format}'. Choose from: {', '.join(EXPORT_FORMATS)}",
err = True,
)
raise typer.Exit(code = 2)
if push_to_hub and not repo_id:
typer.echo("Error: --repo-id required when using --push-to-hub", err=True)
raise typer.Exit(code=2)
typer.echo("Error: --repo-id required when using --push-to-hub", err = True)
raise typer.Exit(code = 2)
from studio.backend.core.export import ExportBackend
backend = ExportBackend()
typer.echo(f"Loading checkpoint: {checkpoint}")
success, message = backend.load_checkpoint(
checkpoint_path=str(checkpoint),
max_seq_length=max_seq_length,
load_in_4bit=load_in_4bit,
checkpoint_path = str(checkpoint),
max_seq_length = max_seq_length,
load_in_4bit = load_in_4bit,
)
if not success:
typer.echo(f"Error: {message}", err=True)
raise typer.Exit(code=1)
typer.echo(f"Error: {message}", err = True)
raise typer.Exit(code = 1)
typer.echo(message)
typer.echo(f"Exporting as {format}...")
if format == "merged-16bit":
success, message = backend.export_merged_model(
save_directory=str(output_dir),
format_type="16-bit (FP16)",
push_to_hub=push_to_hub,
repo_id=repo_id,
hf_token=hf_token,
private=private,
save_directory = str(output_dir),
format_type = "16-bit (FP16)",
push_to_hub = push_to_hub,
repo_id = repo_id,
hf_token = hf_token,
private = private,
)
elif format == "merged-4bit":
success, message = backend.export_merged_model(
save_directory=str(output_dir),
format_type="4-bit (FP4)",
push_to_hub=push_to_hub,
repo_id=repo_id,
hf_token=hf_token,
private=private,
save_directory = str(output_dir),
format_type = "4-bit (FP4)",
push_to_hub = push_to_hub,
repo_id = repo_id,
hf_token = hf_token,
private = private,
)
elif format == "gguf":
success, message = backend.export_gguf(
save_directory=str(output_dir),
quantization_method=quantization.upper(),
push_to_hub=push_to_hub,
repo_id=repo_id,
hf_token=hf_token,
save_directory = str(output_dir),
quantization_method = quantization.upper(),
push_to_hub = push_to_hub,
repo_id = repo_id,
hf_token = hf_token,
)
elif format == "lora":
success, message = backend.export_lora_adapter(
save_directory=str(output_dir),
push_to_hub=push_to_hub,
repo_id=repo_id,
hf_token=hf_token,
private=private,
save_directory = str(output_dir),
push_to_hub = push_to_hub,
repo_id = repo_id,
hf_token = hf_token,
private = private,
)
if not success:
typer.echo(f"Error: {message}", err=True)
raise typer.Exit(code=1)
typer.echo(f"Error: {message}", err = True)
raise typer.Exit(code = 1)
typer.echo(message)

View file

@ -8,10 +8,10 @@ import typer
def inference(
model: str = typer.Argument(..., help="HF model id or local path."),
prompt: str = typer.Argument(..., help="Prompt to send to the model."),
model: str = typer.Argument(..., help = "HF model id or local path."),
prompt: str = typer.Argument(..., help = "Prompt to send to the model."),
hf_token: Optional[str] = typer.Option(
None, "--hf-token", envvar="HF_TOKEN", help="Hugging Face token if needed."
None, "--hf-token", envvar = "HF_TOKEN", help = "Hugging Face token if needed."
),
temperature: float = typer.Option(0.7, "--temperature"),
top_p: float = typer.Option(0.9, "--top-p"),
@ -21,7 +21,7 @@ def inference(
system_prompt: str = typer.Option(
"",
"--system-prompt",
help="Optional system prompt to prepend.",
help = "Optional system prompt to prepend.",
),
max_seq_length: int = typer.Option(2048, "--max-seq-length"),
load_in_4bit: bool = typer.Option(True, "--load-in-4bit/--no-load-in-4bit"),
@ -31,36 +31,36 @@ def inference(
inference_backend = get_inference_backend()
model_config = ModelConfig.from_ui_selection(
dropdown_value=model, search_value=None, hf_token=hf_token, is_lora=False
dropdown_value = model, search_value = None, hf_token = hf_token, is_lora = False
)
if not model_config:
typer.echo("Could not resolve model config", err=True)
raise typer.Exit(code=1)
typer.echo("Could not resolve model config", err = True)
raise typer.Exit(code = 1)
if not inference_backend.load_model(
config=model_config,
max_seq_length=max_seq_length,
load_in_4bit=load_in_4bit,
hf_token=hf_token,
config = model_config,
max_seq_length = max_seq_length,
load_in_4bit = load_in_4bit,
hf_token = hf_token,
):
typer.echo("Model load failed", err=True)
raise typer.Exit(code=1)
typer.echo("Model load failed", err = True)
raise typer.Exit(code = 1)
messages = [{"role": "user", "content": prompt}]
stream = inference_backend.generate_chat_response(
messages=messages,
system_prompt=system_prompt,
temperature=temperature,
top_p=top_p,
top_k=top_k,
max_new_tokens=max_new_tokens,
repetition_penalty=repetition_penalty,
messages = messages,
system_prompt = system_prompt,
temperature = temperature,
top_p = top_p,
top_k = top_k,
max_new_tokens = max_new_tokens,
repetition_penalty = repetition_penalty,
)
typer.echo("Assistant:", nl=True)
typer.echo("Assistant:", nl = True)
previous = ""
for chunk in stream:
delta = chunk[len(previous):]
delta = chunk[len(previous) :]
if delta:
sys.stdout.write(delta)
sys.stdout.flush()

View file

@ -10,28 +10,44 @@ from pathlib import Path
from typing import Optional
import typer
studio_app = typer.Typer(help="Unsloth Studio commands.")
studio_app = typer.Typer(help = "Unsloth Studio commands.")
STUDIO_HOME = Path.home() / ".unsloth" / "studio"
# __file__ is cli/commands/studio.py — two parents up is the package root
# (either site-packages or the repo root for editable installs).
_PACKAGE_ROOT = Path(__file__).resolve().parent.parent.parent
def _is_repo_root(path: Path) -> bool:
"""Check if a directory looks like the repo root (actual git clone, not site-packages)."""
return (
(path / ".git").exists()
and (path / "pyproject.toml").is_file()
and (
(path / "studio" / "setup.sh").is_file()
or (path / "studio" / "setup.ps1").is_file()
)
)
def _get_repo_root() -> Optional[Path]:
"""Find the git clone repo root, or None if pure pip install."""
"""Find the git clone repo root.
Used only by setup() checks __file__ first (editable install),
then walks CWD parents (wheel install, user is inside the clone).
"""
# Check 1: __file__ is in the repo (editable install)
candidate = Path(__file__).resolve().parent.parent.parent
if (candidate / "pyproject.toml").is_file() and (candidate / "studio" / "setup.sh").is_file():
return candidate
# Check 2: CWD is the repo (non-editable wheel, running from repo dir)
cwd = Path.cwd()
if (cwd / "pyproject.toml").is_file() and (cwd / "studio" / "setup.sh").is_file():
return cwd
if _is_repo_root(_PACKAGE_ROOT):
return _PACKAGE_ROOT
# Check 2: CWD or any parent is the repo
cwd = Path.cwd().resolve()
for parent in (cwd, *cwd.parents):
if _is_repo_root(parent):
return parent
return None
def _is_git_clone() -> bool:
return _get_repo_root() is not None
def _studio_venv_python() -> Optional[Path]:
"""Return the studio venv Python binary, or None if not set up."""
if platform.system() == "Windows":
@ -42,37 +58,69 @@ def _studio_venv_python() -> Optional[Path]:
def _find_run_py() -> Optional[Path]:
"""Find studio/backend/run.py."""
# 1. Repo root (git clone / editable)
repo = _get_repo_root()
if repo:
run_py = repo / "studio" / "backend" / "run.py"
if run_py.is_file():
return run_py
# 2. Studio venv's site-packages
for match in (STUDIO_HOME / ".venv").glob("lib/python*/site-packages/studio/backend/run.py"):
return match
# 3. Current package's site-packages
run_py = Path(__file__).resolve().parent.parent.parent / "studio" / "backend" / "run.py"
return run_py if run_py.is_file() else None
"""Find studio/backend/run.py.
No CWD dependency works from any directory.
Since studio/ is now a proper package (has __init__.py), it lives in
site-packages after pip install, right next to cli/.
"""
# 1. Relative to __file__ (site-packages or editable repo root)
run_py = _PACKAGE_ROOT / "studio" / "backend" / "run.py"
if run_py.is_file():
return run_py
# 2. Studio venv's site-packages (Linux + Windows layouts)
for pattern in (
"lib/python*/site-packages/studio/backend/run.py",
"Lib/site-packages/studio/backend/run.py",
):
for match in (STUDIO_HOME / ".venv").glob(pattern):
return match
return None
def _find_install_script() -> Optional[Path]:
"""Find studio/install_python_stack.py."""
# 1. Repo root
repo = _get_repo_root()
if repo:
s = repo / "studio" / "install_python_stack.py"
if s.is_file():
return s
# 2. Relative to __file__ (in site-packages)
s = Path(__file__).resolve().parent.parent.parent / "studio" / "install_python_stack.py"
return s if s.is_file() else None
"""Find studio/install_python_stack.py.
No CWD dependency works from any directory.
"""
# 1. Relative to __file__ (site-packages or editable repo root)
s = _PACKAGE_ROOT / "studio" / "install_python_stack.py"
if s.is_file():
return s
# 2. Studio venv's site-packages
for pattern in (
"lib/python*/site-packages/studio/install_python_stack.py",
"Lib/site-packages/studio/install_python_stack.py",
):
for match in (STUDIO_HOME / ".venv").glob(pattern):
return match
return None
def _find_setup_script() -> Optional[Path]:
"""Find studio/setup.sh or studio/setup.ps1.
No CWD dependency works from any directory.
"""
name = "setup.ps1" if platform.system() == "Windows" else "setup.sh"
# 1. Relative to __file__ (site-packages or editable repo root)
s = _PACKAGE_ROOT / "studio" / name
if s.is_file():
return s
# 2. Studio venv's site-packages
for pattern in (
f"lib/python*/site-packages/studio/{name}",
f"Lib/site-packages/studio/{name}",
):
for match in (STUDIO_HOME / ".venv").glob(pattern):
return match
return None
# ── unsloth studio (server) ──────────────────────────────────────────
@studio_app.callback(invoke_without_command=True)
@studio_app.callback(invoke_without_command = True)
def studio_default(
ctx: typer.Context,
port: int = typer.Option(8000, "--port", "-p"),
@ -94,7 +142,14 @@ def studio_default(
if studio_python and run_py:
if not silent:
typer.echo("Launching with studio venv...")
args = [str(studio_python), str(run_py), "--host", host, "--port", str(port)]
args = [
str(studio_python),
str(run_py),
"--host",
host,
"--port",
str(port),
]
if frontend:
args.extend(["--frontend", str(frontend)])
if silent:
@ -108,14 +163,15 @@ def studio_default(
if not silent:
from studio.backend.run import _resolve_external_ip
display_host = _resolve_external_ip() if host == "0.0.0.0" else host
typer.echo(f"Starting Unsloth Studio on http://{display_host}:{port}")
run_server(
host=host,
port=port,
frontend_path=frontend,
silent=silent,
host = host,
port = port,
frontend_path = frontend,
silent = silent,
)
try:
@ -127,25 +183,30 @@ def studio_default(
# ── unsloth studio setup ─────────────────────────────────────────────
@studio_app.command()
def setup():
"""Run one-time Studio environment setup."""
if _is_git_clone():
_dev_setup()
# If we're inside a git clone, use the full setup script (builds frontend, etc.)
repo = _get_repo_root()
if repo:
_dev_setup(repo)
else:
_pip_setup()
def _dev_setup():
def _dev_setup(repo_root: Path):
"""Git-clone: run setup.sh / setup.ps1."""
repo_root = _get_repo_root()
studio_dir = repo_root / "studio"
if platform.system() == "Windows":
script = studio_dir / "setup.ps1"
subprocess.run(["powershell", "-ExecutionPolicy", "Bypass", "-File", str(script)], check=True)
subprocess.run(
["powershell", "-ExecutionPolicy", "Bypass", "-File", str(script)],
check = True,
)
else:
script = studio_dir / "setup.sh"
subprocess.run(["bash", str(script)], check=True)
subprocess.run(["bash", str(script)], check = True)
def _pip_setup():
@ -167,14 +228,14 @@ def _pip_setup():
# 1. Create venv
if not venv_python.is_file():
typer.echo(f" Creating venv at {venv_dir}...")
STUDIO_HOME.mkdir(parents=True, exist_ok=True)
_venv.create(str(venv_dir), with_pip=True)
STUDIO_HOME.mkdir(parents = True, exist_ok = True)
_venv.create(str(venv_dir), with_pip = True)
# 2. Install all Python deps via install_python_stack.py
install_script = _find_install_script()
if install_script:
typer.echo(" Installing Python dependencies...")
subprocess.run([str(venv_python), str(install_script)], check=True)
subprocess.run([str(venv_python), str(install_script)], check = True)
else:
typer.echo("Error: Could not find install_python_stack.py")
raise typer.Exit(1)
@ -184,16 +245,28 @@ def _pip_setup():
typer.echo(f" Transformers 5.x overlay already at {venv_t5_dir}")
else:
typer.echo(" Installing transformers 5.x overlay...")
venv_t5_dir.mkdir(parents=True, exist_ok=True)
venv_t5_dir.mkdir(parents = True, exist_ok = True)
subprocess.run(
[str(venv_pip), "install",
"--target", str(venv_t5_dir), "--no-deps", "transformers==5.2.0"],
check=True,
[
str(venv_pip),
"install",
"--target",
str(venv_t5_dir),
"--no-deps",
"transformers==5.2.0",
],
check = True,
)
subprocess.run(
[str(venv_pip), "install",
"--target", str(venv_t5_dir), "--no-deps", "huggingface_hub==1.3.0"],
check=True,
[
str(venv_pip),
"install",
"--target",
str(venv_t5_dir),
"--no-deps",
"huggingface_hub==1.3.0",
],
check = True,
)
typer.echo(f" Installed to {venv_t5_dir}")
@ -221,12 +294,26 @@ def _build_llama_cpp():
typer.echo(" Building llama.cpp for GGUF inference...")
if llama_dir.exists():
shutil.rmtree(llama_dir)
unsloth_home.mkdir(parents=True, exist_ok=True)
# necessary because shutil.rmtree fails on Windows because .git pack files are read-only
def _force_remove_readonly(func, path, exc_info):
"""Clear read-only flag and retry — needed on Windows for .git pack files."""
import stat
os.chmod(path, stat.S_IWRITE)
func(path)
shutil.rmtree(llama_dir, onerror = _force_remove_readonly)
unsloth_home.mkdir(parents = True, exist_ok = True)
result = subprocess.run(
["git", "clone", "--depth", "1", "https://github.com/ggml-org/llama.cpp.git", str(llama_dir)],
stdout=subprocess.PIPE, stderr=subprocess.STDOUT,
[
"git",
"clone",
"--depth",
"1",
"https://github.com/ggml-org/llama.cpp.git",
str(llama_dir),
],
stdout = subprocess.PIPE,
stderr = subprocess.STDOUT,
)
if result.returncode != 0:
typer.echo(" Failed to clone llama.cpp")
@ -245,7 +332,8 @@ def _build_llama_cpp():
build_dir = llama_dir / "build"
result = subprocess.run(
["cmake", "-S", str(llama_dir), "-B", str(build_dir)] + cmake_args,
stdout=subprocess.PIPE, stderr=subprocess.STDOUT,
stdout = subprocess.PIPE,
stderr = subprocess.STDOUT,
)
if result.returncode != 0:
typer.echo(" cmake configure failed")
@ -253,21 +341,42 @@ def _build_llama_cpp():
ncpu = str(os.cpu_count() or 4)
result = subprocess.run(
["cmake", "--build", str(build_dir), "--config", "Release",
"--target", "llama-server", f"-j{ncpu}"],
stdout=subprocess.PIPE, stderr=subprocess.STDOUT,
[
"cmake",
"--build",
str(build_dir),
"--config",
"Release",
"--target",
"llama-server",
f"-j{ncpu}",
],
stdout = subprocess.PIPE,
stderr = subprocess.STDOUT,
)
if result.returncode != 0:
typer.echo(" llama-server build failed")
return
subprocess.run(
["cmake", "--build", str(build_dir), "--config", "Release",
"--target", "llama-quantize", f"-j{ncpu}"],
stdout=subprocess.PIPE, stderr=subprocess.STDOUT,
[
"cmake",
"--build",
str(build_dir),
"--config",
"Release",
"--target",
"llama-quantize",
f"-j{ncpu}",
],
stdout = subprocess.PIPE,
stderr = subprocess.STDOUT,
)
server_bin = build_dir / "bin" / "llama-server"
if sys.platform == "win32":
server_bin = build_dir / "bin" / "Release" / "llama-server.exe"
else:
server_bin = build_dir / "bin" / "llama-server"
if server_bin.is_file():
typer.echo(f" llama-server built at {server_bin}")
else:

View file

@ -17,18 +17,18 @@ def train(
None,
"--config",
"-c",
help="Path to YAML/JSON config file. CLI flags override config values.",
help = "Path to YAML/JSON config file. CLI flags override config values.",
),
hf_token: Optional[str] = typer.Option(
None, "--hf-token", envvar="HF_TOKEN", help="Hugging Face token if needed."
None, "--hf-token", envvar = "HF_TOKEN", help = "Hugging Face token if needed."
),
wandb_token: Optional[str] = typer.Option(
None, "--wandb-token", envvar="WANDB_API_KEY", help="Weights & Biases API key."
None, "--wandb-token", envvar = "WANDB_API_KEY", help = "Weights & Biases API key."
),
dry_run: bool = typer.Option(
False,
"--dry-run",
help="Show resolved config and exit without training.",
help = "Show resolved config and exit without training.",
),
config_overrides: dict = None,
):
@ -36,14 +36,15 @@ def train(
try:
cfg = load_config(config)
except FileNotFoundError as e:
typer.echo(f"Error: {e}", err=True)
raise typer.Exit(code=2)
typer.echo(f"Error: {e}", err = True)
raise typer.Exit(code = 2)
cfg.apply_overrides(**config_overrides)
# CLI/env tokens take precedence over config
# Handle case where typer.Option isn't resolved (decorator interaction)
from typer.models import OptionInfo
if isinstance(hf_token, OptionInfo):
hf_token = None
if isinstance(wandb_token, OptionInfo):
@ -53,33 +54,38 @@ def train(
if dry_run:
import yaml
data = cfg.model_dump()
data["training"]["output_dir"] = str(data["training"]["output_dir"])
typer.echo(yaml.dump(data, default_flow_style=False, sort_keys=False))
raise typer.Exit(code=0)
typer.echo(yaml.dump(data, default_flow_style = False, sort_keys = False))
raise typer.Exit(code = 0)
if not cfg.model:
typer.echo("Error: provide --model or set model in --config", err=True)
raise typer.Exit(code=2)
typer.echo("Error: provide --model or set model in --config", err = True)
raise typer.Exit(code = 2)
if not cfg.data.dataset and not cfg.data.local_dataset:
typer.echo(
"Error: provide --dataset or --local-dataset (or via --config)", err=True
"Error: provide --dataset or --local-dataset (or via --config)", err = True
)
raise typer.Exit(code=2)
raise typer.Exit(code = 2)
# Check if the model path is a LoRA adapter (has adapter_config.json)
model_path = Path(cfg.model) if cfg.model else None
model_is_lora = model_path and model_path.is_dir() and (model_path / "adapter_config.json").exists()
model_is_lora = (
model_path
and model_path.is_dir()
and (model_path / "adapter_config.json").exists()
)
use_lora = cfg.training.training_type.lower() == "lora"
if model_is_lora and not use_lora:
typer.echo(
"Error: Cannot do full finetuning on a LoRA adapter. "
"Use --training-type lora or provide a base model.",
err=True,
err = True,
)
raise typer.Exit(code=2)
raise typer.Exit(code = 2)
from studio.backend.core.training.trainer import UnslothTrainer
@ -87,38 +93,40 @@ def train(
# Load model (trainer.is_vlm is set after this)
if not trainer.load_model(
model_name=cfg.model,
max_seq_length=cfg.training.max_seq_length,
load_in_4bit=cfg.training.load_in_4bit if use_lora else False,
hf_token=hf_token,
model_name = cfg.model,
max_seq_length = cfg.training.max_seq_length,
load_in_4bit = cfg.training.load_in_4bit if use_lora else False,
hf_token = hf_token,
):
typer.echo("Model load failed", err=True)
raise typer.Exit(code=1)
typer.echo("Model load failed", err = True)
raise typer.Exit(code = 1)
is_vision = trainer.is_vlm
if not trainer.prepare_model_for_training(**cfg.model_kwargs(use_lora, is_vision)):
typer.echo("Model preparation failed", err=True)
raise typer.Exit(code=1)
typer.echo("Model preparation failed", err = True)
raise typer.Exit(code = 1)
result = trainer.load_and_format_dataset(
dataset_source=cfg.data.dataset or "",
format_type=cfg.data.format_type,
local_datasets=cfg.data.local_dataset,
dataset_source = cfg.data.dataset or "",
format_type = cfg.data.format_type,
local_datasets = cfg.data.local_dataset,
)
if result is None:
typer.echo("Dataset load failed", err=True)
raise typer.Exit(code=1)
typer.echo("Dataset load failed", err = True)
raise typer.Exit(code = 1)
ds, eval_ds = result
training_kwargs = cfg.training_kwargs()
training_kwargs["wandb_token"] = wandb_token # CLI/env takes precedence
started = trainer.start_training(dataset=ds, eval_dataset=eval_ds, **training_kwargs)
started = trainer.start_training(
dataset = ds, eval_dataset = eval_ds, **training_kwargs
)
if not started:
typer.echo("Training failed to start", err=True)
raise typer.Exit(code=1)
typer.echo("Training failed to start", err = True)
raise typer.Exit(code = 1)
try:
while trainer.training_thread and trainer.training_thread.is_alive():
@ -132,5 +140,5 @@ def train(
final = trainer.get_training_progress()
if getattr(final, "error", None):
typer.echo(f"Training error: {final.error}", err=True)
raise typer.Exit(code=1)
typer.echo(f"Training error: {final.error}", err = True)
raise typer.Exit(code = 1)

View file

@ -1,6 +1,8 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
import os
import sys
import time
from pathlib import Path
from typing import Optional
@ -9,27 +11,64 @@ import typer
def ui(
port: int = typer.Option(8000, "--port", "-p", help="Port to run the UI server on."),
host: str = typer.Option("0.0.0.0", "--host", "-H", help="Host address to bind to."),
frontend: Optional[Path] = typer.Option(None, "--frontend", "-f", help="Path to frontend build directory."),
silent: bool = typer.Option(False, "--silent", "-q", help="Suppress startup messages."),
port: int = typer.Option(
8000, "--port", "-p", help = "Port to run the UI server on."
),
host: str = typer.Option(
"0.0.0.0", "--host", "-H", help = "Host address to bind to."
),
frontend: Optional[Path] = typer.Option(
None, "--frontend", "-f", help = "Path to frontend build directory."
),
silent: bool = typer.Option(
False, "--silent", "-q", help = "Suppress startup messages."
),
):
"""Launch the Unsloth web UI backend server."""
"""Launch the Unsloth web UI backend server (alias for 'unsloth studio')."""
from cli.commands.studio import _studio_venv_python, _find_run_py, STUDIO_HOME
# Re-execute in studio venv if available and not already inside it
studio_venv_dir = STUDIO_HOME / ".venv"
in_studio_venv = sys.prefix.startswith(str(studio_venv_dir))
if not in_studio_venv:
studio_python = _studio_venv_python()
run_py = _find_run_py()
if studio_python and run_py:
if not silent:
typer.echo("Launching with studio venv...")
args = [
str(studio_python),
str(run_py),
"--host",
host,
"--port",
str(port),
]
if frontend:
args.extend(["--frontend", str(frontend)])
if silent:
args.append("--silent")
os.execvp(str(studio_python), args)
else:
typer.echo("Studio not set up. Run 'unsloth studio setup' first.")
raise typer.Exit(1)
from studio.backend.run import run_server
if not silent:
from studio.backend.run import _resolve_external_ip
display_host = _resolve_external_ip() if host == "0.0.0.0" else host
typer.echo(f"Starting Unsloth Studio on http://{display_host}:{port}")
run_server(
host=host,
port=port,
frontend_path=frontend,
silent=silent,
host = host,
port = port,
frontend_path = frontend,
silent = silent,
)
# Keep running until interrupted
try:
while True:
time.sleep(1)

View file

@ -58,10 +58,10 @@ class LoggingConfig(BaseModel):
class Config(BaseModel):
model: Optional[str] = None
data: DataConfig = Field(default_factory=DataConfig)
training: TrainingConfig = Field(default_factory=TrainingConfig)
lora: LoraConfig = Field(default_factory=LoraConfig)
logging: LoggingConfig = Field(default_factory=LoggingConfig)
data: DataConfig = Field(default_factory = DataConfig)
training: TrainingConfig = Field(default_factory = TrainingConfig)
lora: LoraConfig = Field(default_factory = LoraConfig)
logging: LoggingConfig = Field(default_factory = LoggingConfig)
def apply_overrides(self, **kwargs):
"""Apply CLI overrides by matching arg names to config fields."""
@ -83,7 +83,11 @@ class Config(BaseModel):
# Vision models expect a string (e.g., "all-linear"); fall back to None to use trainer defaults
target_modules = "all-linear" if self.lora.vision_all_linear else None
else:
parsed = [m.strip() for m in str(self.lora.target_modules).split(",") if m and m.strip()]
parsed = [
m.strip()
for m in str(self.lora.target_modules).split(",")
if m and m.strip()
]
target_modules = parsed or None
return {
@ -134,11 +138,12 @@ def load_config(path: Optional[Path]) -> Config:
if not path.exists():
raise FileNotFoundError(f"Config file not found: {path}")
text = path.read_text(encoding="utf-8")
text = path.read_text(encoding = "utf-8")
if path.suffix.lower() in {".yaml", ".yml"}:
data = yaml.safe_load(text) or {}
else:
import json
data = json.loads(text or "{}")
return Config(**data)

View file

@ -80,7 +80,9 @@ def add_options_from_config(config_class: type[BaseModel]) -> Callable:
which will receive a dict of all CLI-provided config values.
"""
fields = _collect_config_fields(config_class)
field_names = {name for name, field_info in fields if not _is_list_type(field_info.annotation)}
field_names = {
name for name, field_info in fields if not _is_list_type(field_info.annotation)
}
def decorator(func: Callable) -> Callable:
sig = inspect.signature(func)
@ -105,22 +107,22 @@ def add_options_from_config(config_class: type[BaseModel]) -> Callable:
default = typer.Option(
None,
f"{flag_name}/--no-{field_name.replace('_', '-')}",
help=help_text,
help = help_text,
)
param = inspect.Parameter(
field_name,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
default=default,
annotation=Optional[bool],
default = default,
annotation = Optional[bool],
)
else:
py_type = _get_python_type(annotation)
default = typer.Option(None, flag_name, help=help_text)
default = typer.Option(None, flag_name, help = help_text)
param = inspect.Parameter(
field_name,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
default=default,
annotation=Optional[py_type],
default = default,
annotation = Optional[py_type],
)
new_params.append(param)
@ -129,7 +131,7 @@ def add_options_from_config(config_class: type[BaseModel]) -> Callable:
if param.name != "config_overrides":
new_params.append(param)
new_sig = sig.replace(parameters=new_params)
new_sig = sig.replace(parameters = new_params)
@functools.wraps(func)
def wrapper(*args, **kwargs):