Studio (#4237)
* Rebuild Studio branch on top of main * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Fix security and code quality issues for Studio PR #4237 - Validate models_dir query param against allowed directory roots to prevent path traversal in /api/models/local endpoint - Replace string startswith() with Path.is_relative_to() for frontend path traversal check in serve_frontend - Sanitize SSE error messages to not leak exception details to clients (4 locations in inference.py) - Bind port-discovery socket to 127.0.0.1 instead of all interfaces in llama_cpp backend - Import datasets_root and resolve_output_dir in embedding training function to fix NameError and use managed output directory - Remove stale .gitignore entries for package-lock.json and test directories so tests can be tracked in version control - Add venv-reexecution logic to ui CLI command matching the studio command behavior * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Move models_dir path validation before try/except block The HTTPException(403) was inside the try/except Exception handler, so it would be caught and re-raised as a 500. Moving the validation before the try block ensures the 403 is returned directly and also makes the control flow clearer for static analysis (path is validated before any filesystem operations). * Use os.path.realpath + startswith for models_dir validation CodeQL py/path-injection does not recognize Path.is_relative_to() as a sanitizer. Switched to os.path.realpath + str.startswith which is a recognized sanitizer pattern in CodeQL's taint analysis. The startswith check uses root_str + os.sep to prevent prefix collisions (e.g. /app/models_evil matching /app/models). * Never pass user input to Path constructor in models_dir validation CodeQL traces taint through Path(resolved) even after a startswith barrier guard. Fix: the user-supplied models_dir is only used as a string for comparison against allowed roots. The Path object passed to _scan_models_dir comes from the trusted allowed_roots list, not from user input. This fully breaks the taint chain. --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
parent
bced78373f
commit
f08aef1804
664 changed files with 103567 additions and 4 deletions
22
cli/__init__.py
Normal file
22
cli/__init__.py
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
# 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 typer
|
||||
|
||||
from cli.commands.train import train
|
||||
from cli.commands.inference import inference
|
||||
from cli.commands.export import export, list_checkpoints
|
||||
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"]},
|
||||
)
|
||||
|
||||
app.command()(train)
|
||||
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.")
|
||||
2
cli/commands/__init__.py
Normal file
2
cli/commands/__init__.py
Normal 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
|
||||
132
cli/commands/export.py
Normal file
132
cli/commands/export.py
Normal file
|
|
@ -0,0 +1,132 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import typer
|
||||
|
||||
|
||||
EXPORT_FORMATS = ["merged-16bit", "merged-4bit", "gguf", "lora"]
|
||||
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."
|
||||
),
|
||||
):
|
||||
"""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))
|
||||
if not checkpoints:
|
||||
typer.echo("No checkpoints found.")
|
||||
raise typer.Exit()
|
||||
|
||||
for model_name, ckpt_list, metadata in checkpoints:
|
||||
typer.echo(f"\n{model_name}:")
|
||||
for display, path, loss in ckpt_list:
|
||||
loss_str = f" (loss: {loss:.4f})" if loss is not None else ""
|
||||
typer.echo(f" {display}{loss_str}: {path}")
|
||||
|
||||
|
||||
def export(
|
||||
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)}",
|
||||
),
|
||||
quantization: str = typer.Option(
|
||||
"q4_k_m",
|
||||
"--quantization",
|
||||
"-q",
|
||||
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."
|
||||
),
|
||||
repo_id: Optional[str] = typer.Option(
|
||||
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."
|
||||
),
|
||||
private: bool = typer.Option(
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
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,
|
||||
)
|
||||
if not success:
|
||||
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,
|
||||
)
|
||||
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,
|
||||
)
|
||||
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,
|
||||
)
|
||||
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,
|
||||
)
|
||||
|
||||
if not success:
|
||||
typer.echo(f"Error: {message}", err = True)
|
||||
raise typer.Exit(code = 1)
|
||||
|
||||
typer.echo(message)
|
||||
69
cli/commands/inference.py
Normal file
69
cli/commands/inference.py
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
# 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 sys
|
||||
from typing import Optional
|
||||
|
||||
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."),
|
||||
hf_token: Optional[str] = typer.Option(
|
||||
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"),
|
||||
top_k: int = typer.Option(40, "--top-k"),
|
||||
max_new_tokens: int = typer.Option(256, "--max-new-tokens"),
|
||||
repetition_penalty: float = typer.Option(1.1, "--repetition-penalty"),
|
||||
system_prompt: str = typer.Option(
|
||||
"",
|
||||
"--system-prompt",
|
||||
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"),
|
||||
):
|
||||
"""Run a single inference using the specified model."""
|
||||
from studio.backend.core import ModelConfig, get_inference_backend
|
||||
|
||||
inference_backend = get_inference_backend()
|
||||
model_config = ModelConfig.from_ui_selection(
|
||||
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)
|
||||
|
||||
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,
|
||||
):
|
||||
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,
|
||||
)
|
||||
|
||||
typer.echo("Assistant:", nl = True)
|
||||
previous = ""
|
||||
for chunk in stream:
|
||||
delta = chunk[len(previous) :]
|
||||
if delta:
|
||||
sys.stdout.write(delta)
|
||||
sys.stdout.flush()
|
||||
previous = chunk
|
||||
sys.stdout.write("\n")
|
||||
sys.stdout.flush()
|
||||
374
cli/commands/studio.py
Normal file
374
cli/commands/studio.py
Normal file
|
|
@ -0,0 +1,374 @@
|
|||
# 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 platform
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
import typer
|
||||
|
||||
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.
|
||||
|
||||
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)
|
||||
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 _studio_venv_python() -> Optional[Path]:
|
||||
"""Return the studio venv Python binary, or None if not set up."""
|
||||
if platform.system() == "Windows":
|
||||
p = STUDIO_HOME / ".venv" / "Scripts" / "python.exe"
|
||||
else:
|
||||
p = STUDIO_HOME / ".venv" / "bin" / "python"
|
||||
return p if p.is_file() else None
|
||||
|
||||
|
||||
def _find_run_py() -> Optional[Path]:
|
||||
"""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.
|
||||
|
||||
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)
|
||||
def studio_default(
|
||||
ctx: typer.Context,
|
||||
port: int = typer.Option(8000, "--port", "-p"),
|
||||
host: str = typer.Option("0.0.0.0", "--host", "-H"),
|
||||
frontend: Optional[Path] = typer.Option(None, "--frontend", "-f"),
|
||||
silent: bool = typer.Option(False, "--silent", "-q"),
|
||||
):
|
||||
"""Launch the Unsloth Studio server."""
|
||||
if ctx.invoked_subcommand is not None:
|
||||
return
|
||||
|
||||
# Always use the studio venv if it exists and we're not already in 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,
|
||||
)
|
||||
|
||||
try:
|
||||
while True:
|
||||
time.sleep(1)
|
||||
except KeyboardInterrupt:
|
||||
typer.echo("\nShutting down...")
|
||||
|
||||
|
||||
# ── unsloth studio setup ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@studio_app.command()
|
||||
def setup():
|
||||
"""Run one-time Studio environment 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(repo_root: Path):
|
||||
"""Git-clone: run setup.sh / setup.ps1."""
|
||||
studio_dir = repo_root / "studio"
|
||||
if platform.system() == "Windows":
|
||||
script = studio_dir / "setup.ps1"
|
||||
subprocess.run(
|
||||
["powershell", "-ExecutionPolicy", "Bypass", "-File", str(script)],
|
||||
check = True,
|
||||
)
|
||||
else:
|
||||
script = studio_dir / "setup.sh"
|
||||
subprocess.run(["bash", str(script)], check = True)
|
||||
|
||||
|
||||
def _pip_setup():
|
||||
"""Pip-install: create studio venv, install all deps, build extras."""
|
||||
import venv as _venv
|
||||
|
||||
venv_dir = STUDIO_HOME / ".venv"
|
||||
venv_t5_dir = STUDIO_HOME / ".venv_t5"
|
||||
|
||||
if platform.system() == "Windows":
|
||||
venv_python = venv_dir / "Scripts" / "python.exe"
|
||||
venv_pip = venv_dir / "Scripts" / "pip.exe"
|
||||
else:
|
||||
venv_python = venv_dir / "bin" / "python"
|
||||
venv_pip = venv_dir / "bin" / "pip"
|
||||
|
||||
typer.echo("Setting up Unsloth Studio...")
|
||||
|
||||
# 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)
|
||||
|
||||
# 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)
|
||||
else:
|
||||
typer.echo("Error: Could not find install_python_stack.py")
|
||||
raise typer.Exit(1)
|
||||
|
||||
# 3. Pre-install transformers 5.x overlay
|
||||
if venv_t5_dir.is_dir() and any(venv_t5_dir.iterdir()):
|
||||
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)
|
||||
subprocess.run(
|
||||
[
|
||||
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,
|
||||
)
|
||||
typer.echo(f" Installed to {venv_t5_dir}")
|
||||
|
||||
# 4. Build llama.cpp
|
||||
_build_llama_cpp()
|
||||
|
||||
typer.echo("")
|
||||
typer.echo("Setup complete! Run 'unsloth studio' to start.")
|
||||
|
||||
|
||||
def _build_llama_cpp():
|
||||
"""Build llama.cpp at ~/.unsloth/llama.cpp/."""
|
||||
import shutil
|
||||
|
||||
unsloth_home = Path.home() / ".unsloth"
|
||||
llama_dir = unsloth_home / "llama.cpp"
|
||||
|
||||
if not shutil.which("cmake"):
|
||||
typer.echo(" cmake not found — skipping llama.cpp build")
|
||||
return
|
||||
if not shutil.which("git"):
|
||||
typer.echo(" git not found — skipping llama.cpp build")
|
||||
return
|
||||
|
||||
typer.echo(" Building llama.cpp for GGUF inference...")
|
||||
|
||||
if llama_dir.exists():
|
||||
shutil.rmtree(llama_dir)
|
||||
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,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
typer.echo(" Failed to clone llama.cpp")
|
||||
return
|
||||
|
||||
cmake_args = []
|
||||
nvcc_path = shutil.which("nvcc")
|
||||
if not nvcc_path and Path("/usr/local/cuda/bin/nvcc").is_file():
|
||||
nvcc_path = "/usr/local/cuda/bin/nvcc"
|
||||
if nvcc_path:
|
||||
typer.echo(f" Building with CUDA (nvcc: {nvcc_path})...")
|
||||
cmake_args.append("-DGGML_CUDA=ON")
|
||||
else:
|
||||
typer.echo(" Building CPU-only...")
|
||||
|
||||
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,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
typer.echo(" cmake configure failed")
|
||||
return
|
||||
|
||||
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,
|
||||
)
|
||||
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,
|
||||
)
|
||||
|
||||
server_bin = build_dir / "bin" / "llama-server"
|
||||
if server_bin.is_file():
|
||||
typer.echo(f" llama-server built at {server_bin}")
|
||||
else:
|
||||
typer.echo(" llama-server binary not found after build")
|
||||
144
cli/commands/train.py
Normal file
144
cli/commands/train.py
Normal file
|
|
@ -0,0 +1,144 @@
|
|||
# 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 time
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import typer
|
||||
|
||||
from cli.config import Config, load_config
|
||||
from cli.options import add_options_from_config
|
||||
|
||||
|
||||
@add_options_from_config(Config)
|
||||
def train(
|
||||
config: Optional[Path] = typer.Option(
|
||||
None,
|
||||
"--config",
|
||||
"-c",
|
||||
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."
|
||||
),
|
||||
wandb_token: Optional[str] = typer.Option(
|
||||
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.",
|
||||
),
|
||||
config_overrides: dict = None,
|
||||
):
|
||||
"""Launch training using the existing Unsloth training backend."""
|
||||
try:
|
||||
cfg = load_config(config)
|
||||
except FileNotFoundError as e:
|
||||
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):
|
||||
wandb_token = None
|
||||
hf_token = hf_token or cfg.logging.hf_token
|
||||
wandb_token = wandb_token or cfg.logging.wandb_token
|
||||
|
||||
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)
|
||||
|
||||
if not cfg.model:
|
||||
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
|
||||
)
|
||||
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()
|
||||
)
|
||||
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,
|
||||
)
|
||||
raise typer.Exit(code = 2)
|
||||
|
||||
from studio.backend.core.training.trainer import UnslothTrainer
|
||||
|
||||
trainer = UnslothTrainer()
|
||||
|
||||
# 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,
|
||||
):
|
||||
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)
|
||||
|
||||
result = trainer.load_and_format_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)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
if not started:
|
||||
typer.echo("Training failed to start", err = True)
|
||||
raise typer.Exit(code = 1)
|
||||
|
||||
try:
|
||||
while trainer.training_thread and trainer.training_thread.is_alive():
|
||||
time.sleep(1)
|
||||
except KeyboardInterrupt:
|
||||
typer.echo("Stopping training (Ctrl+C detected)...")
|
||||
trainer.stop_training()
|
||||
finally:
|
||||
if trainer.training_thread:
|
||||
trainer.training_thread.join()
|
||||
|
||||
final = trainer.get_training_progress()
|
||||
if getattr(final, "error", None):
|
||||
typer.echo(f"Training error: {final.error}", err = True)
|
||||
raise typer.Exit(code = 1)
|
||||
76
cli/commands/ui.py
Normal file
76
cli/commands/ui.py
Normal file
|
|
@ -0,0 +1,76 @@
|
|||
# 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
|
||||
|
||||
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."
|
||||
),
|
||||
):
|
||||
"""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,
|
||||
)
|
||||
|
||||
try:
|
||||
while True:
|
||||
time.sleep(1)
|
||||
except KeyboardInterrupt:
|
||||
typer.echo("\nShutting down...")
|
||||
149
cli/config.py
Normal file
149
cli/config.py
Normal file
|
|
@ -0,0 +1,149 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Literal, Optional, List
|
||||
|
||||
import yaml
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class DataConfig(BaseModel):
|
||||
dataset: Optional[str] = None
|
||||
local_dataset: Optional[List[str]] = None
|
||||
format_type: Literal["auto", "alpaca", "chatml", "sharegpt"] = "auto"
|
||||
|
||||
|
||||
class TrainingConfig(BaseModel):
|
||||
training_type: Literal["lora", "full"] = "lora"
|
||||
max_seq_length: int = 2048
|
||||
load_in_4bit: bool = True
|
||||
output_dir: Path = Path("./outputs")
|
||||
num_epochs: int = 3
|
||||
learning_rate: float = 2e-4
|
||||
batch_size: int = 2
|
||||
gradient_accumulation_steps: int = 4
|
||||
warmup_steps: int = 5
|
||||
max_steps: int = 0
|
||||
save_steps: int = 0
|
||||
weight_decay: float = 0.01
|
||||
random_seed: int = 3407
|
||||
packing: bool = False
|
||||
train_on_completions: bool = False
|
||||
gradient_checkpointing: Literal["unsloth", "true", "none"] = "unsloth"
|
||||
|
||||
|
||||
class LoraConfig(BaseModel):
|
||||
lora_r: int = 64
|
||||
lora_alpha: int = 16
|
||||
lora_dropout: float = 0.0
|
||||
target_modules: str = "q_proj,k_proj,v_proj,o_proj,gate_proj,up_proj,down_proj"
|
||||
vision_all_linear: bool = False
|
||||
use_rslora: bool = False
|
||||
use_loftq: bool = False
|
||||
finetune_vision_layers: bool = True
|
||||
finetune_language_layers: bool = True
|
||||
finetune_attention_modules: bool = True
|
||||
finetune_mlp_modules: bool = True
|
||||
|
||||
|
||||
class LoggingConfig(BaseModel):
|
||||
enable_wandb: bool = False
|
||||
wandb_project: str = "unsloth-training"
|
||||
wandb_token: Optional[str] = None
|
||||
enable_tensorboard: bool = False
|
||||
tensorboard_dir: str = "runs"
|
||||
hf_token: Optional[str] = None
|
||||
|
||||
|
||||
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)
|
||||
|
||||
def apply_overrides(self, **kwargs):
|
||||
"""Apply CLI overrides by matching arg names to config fields."""
|
||||
for key, value in kwargs.items():
|
||||
if value is None:
|
||||
continue
|
||||
if hasattr(self, key):
|
||||
setattr(self, key, value)
|
||||
else:
|
||||
for section in (self.data, self.training, self.lora, self.logging):
|
||||
if hasattr(section, key):
|
||||
setattr(section, key, value)
|
||||
break
|
||||
|
||||
def model_kwargs(self, use_lora: bool, is_vision: bool) -> dict:
|
||||
"""Return kwargs for trainer.prepare_model_for_training()."""
|
||||
# Determine target modules based on model type
|
||||
if use_lora and is_vision:
|
||||
# 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()
|
||||
]
|
||||
target_modules = parsed or None
|
||||
|
||||
return {
|
||||
"use_lora": use_lora,
|
||||
"finetune_vision_layers": self.lora.finetune_vision_layers,
|
||||
"finetune_language_layers": self.lora.finetune_language_layers,
|
||||
"finetune_attention_modules": self.lora.finetune_attention_modules,
|
||||
"finetune_mlp_modules": self.lora.finetune_mlp_modules,
|
||||
"target_modules": target_modules,
|
||||
"lora_r": self.lora.lora_r,
|
||||
"lora_alpha": self.lora.lora_alpha,
|
||||
"lora_dropout": self.lora.lora_dropout,
|
||||
"use_gradient_checkpointing": self.training.gradient_checkpointing,
|
||||
"use_rslora": self.lora.use_rslora,
|
||||
"use_loftq": self.lora.use_loftq,
|
||||
}
|
||||
|
||||
def training_kwargs(self) -> dict:
|
||||
"""Return kwargs for trainer.start_training()."""
|
||||
return {
|
||||
"output_dir": str(self.training.output_dir),
|
||||
"num_epochs": self.training.num_epochs,
|
||||
"learning_rate": self.training.learning_rate,
|
||||
"batch_size": self.training.batch_size,
|
||||
"gradient_accumulation_steps": self.training.gradient_accumulation_steps,
|
||||
"warmup_steps": self.training.warmup_steps,
|
||||
"max_steps": self.training.max_steps,
|
||||
"save_steps": self.training.save_steps,
|
||||
"weight_decay": self.training.weight_decay,
|
||||
"random_seed": self.training.random_seed,
|
||||
"packing": self.training.packing,
|
||||
"train_on_completions": self.training.train_on_completions,
|
||||
"max_seq_length": self.training.max_seq_length,
|
||||
"enable_wandb": self.logging.enable_wandb,
|
||||
"wandb_project": self.logging.wandb_project,
|
||||
"wandb_token": self.logging.wandb_token,
|
||||
"enable_tensorboard": self.logging.enable_tensorboard,
|
||||
"tensorboard_dir": self.logging.tensorboard_dir,
|
||||
}
|
||||
|
||||
|
||||
def load_config(path: Optional[Path]) -> Config:
|
||||
"""Load config from YAML/JSON file, or return defaults if no path given."""
|
||||
if not path:
|
||||
return Config()
|
||||
|
||||
path = Path(path)
|
||||
if not path.exists():
|
||||
raise FileNotFoundError(f"Config file not found: {path}")
|
||||
|
||||
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)
|
||||
153
cli/options.py
Normal file
153
cli/options.py
Normal file
|
|
@ -0,0 +1,153 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
"""Generate Typer CLI options from Pydantic models."""
|
||||
|
||||
import functools
|
||||
import inspect
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable, Optional, get_args, get_origin
|
||||
|
||||
import typer
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
def _python_name_to_cli_flag(name: str) -> str:
|
||||
"""Convert python_name to --cli-flag."""
|
||||
return "--" + name.replace("_", "-")
|
||||
|
||||
|
||||
def _unwrap_optional(annotation: Any) -> Any:
|
||||
"""Unwrap Optional[X] to X."""
|
||||
origin = get_origin(annotation)
|
||||
if origin is not None:
|
||||
args = get_args(annotation)
|
||||
if type(None) in args:
|
||||
non_none = [a for a in args if a is not type(None)]
|
||||
if non_none:
|
||||
return non_none[0]
|
||||
return annotation
|
||||
|
||||
|
||||
def _is_bool_field(annotation: Any) -> bool:
|
||||
"""Check if field is a boolean (including Optional[bool])."""
|
||||
return _unwrap_optional(annotation) is bool
|
||||
|
||||
|
||||
def _is_list_type(annotation: Any) -> bool:
|
||||
"""Check if type is a List."""
|
||||
return get_origin(annotation) is list
|
||||
|
||||
|
||||
def _get_python_type(annotation: Any) -> type:
|
||||
"""Get the Python type for annotation."""
|
||||
unwrapped = _unwrap_optional(annotation)
|
||||
if unwrapped in (str, int, float, bool, Path):
|
||||
return unwrapped
|
||||
return str
|
||||
|
||||
|
||||
def _collect_config_fields(config_class: type[BaseModel]) -> list[tuple[str, Any]]:
|
||||
"""
|
||||
Collect all fields from a config class, flattening nested models. Returns list of
|
||||
(name, field_info) tuples. Raises ValueError on duplicate field names.
|
||||
"""
|
||||
fields = []
|
||||
seen_names: set[str] = set()
|
||||
|
||||
for name, field_info in config_class.model_fields.items():
|
||||
annotation = field_info.annotation
|
||||
# Skip nested models - recurse into them
|
||||
if isinstance(annotation, type) and issubclass(annotation, BaseModel):
|
||||
for nested_name, nested_field in annotation.model_fields.items():
|
||||
if nested_name in seen_names:
|
||||
raise ValueError(f"Duplicate field name '{nested_name}' in config")
|
||||
seen_names.add(nested_name)
|
||||
fields.append((nested_name, nested_field))
|
||||
else:
|
||||
if name in seen_names:
|
||||
raise ValueError(f"Duplicate field name '{name}' in config")
|
||||
seen_names.add(name)
|
||||
fields.append((name, field_info))
|
||||
return fields
|
||||
|
||||
|
||||
def add_options_from_config(config_class: type[BaseModel]) -> Callable:
|
||||
"""
|
||||
Decorator that adds CLI options for all fields in a Pydantic config model.
|
||||
|
||||
The decorated function should declare a `config_overrides: dict = None` parameter
|
||||
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)
|
||||
}
|
||||
|
||||
def decorator(func: Callable) -> Callable:
|
||||
sig = inspect.signature(func)
|
||||
original_params = list(sig.parameters.values())
|
||||
original_param_names = {p.name for p in original_params}
|
||||
|
||||
# Build new parameters: config fields first, then original params
|
||||
new_params = []
|
||||
|
||||
for field_name, field_info in fields:
|
||||
# Skip fields already defined in function signature (e.g., with envvar)
|
||||
if field_name in original_param_names:
|
||||
continue
|
||||
annotation = field_info.annotation
|
||||
if _is_list_type(annotation):
|
||||
continue
|
||||
|
||||
flag_name = _python_name_to_cli_flag(field_name)
|
||||
help_text = field_info.description or ""
|
||||
|
||||
if _is_bool_field(annotation):
|
||||
default = typer.Option(
|
||||
None,
|
||||
f"{flag_name}/--no-{field_name.replace('_', '-')}",
|
||||
help = help_text,
|
||||
)
|
||||
param = inspect.Parameter(
|
||||
field_name,
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
default = default,
|
||||
annotation = Optional[bool],
|
||||
)
|
||||
else:
|
||||
py_type = _get_python_type(annotation)
|
||||
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],
|
||||
)
|
||||
new_params.append(param)
|
||||
|
||||
# Add original params, excluding config_overrides (will be injected)
|
||||
for param in original_params:
|
||||
if param.name != "config_overrides":
|
||||
new_params.append(param)
|
||||
|
||||
new_sig = sig.replace(parameters = new_params)
|
||||
|
||||
@functools.wraps(func)
|
||||
def wrapper(*args, **kwargs):
|
||||
config_overrides = {}
|
||||
for key in list(kwargs.keys()):
|
||||
if key in field_names:
|
||||
if kwargs[key] is not None:
|
||||
config_overrides[key] = kwargs[key]
|
||||
# Only delete if not an explicitly declared parameter
|
||||
if key not in original_param_names:
|
||||
del kwargs[key]
|
||||
|
||||
kwargs["config_overrides"] = config_overrides
|
||||
return func(*args, **kwargs)
|
||||
|
||||
wrapper.__signature__ = new_sig
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
Loading…
Add table
Add a link
Reference in a new issue