Studio: Expose GPU memory mode in unsloth run and unsloth start (#7421)

* Add CLI GPU memory mode selection

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Preserve manual GPU layer overrides

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Lee Jackson <130007945+Imagineer99@users.noreply.github.com>
This commit is contained in:
oobabooga 2026-07-27 09:54:50 -03:00 committed by GitHub
commit 0b34377778
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 277 additions and 18 deletions

View file

@ -20,7 +20,7 @@ import time
import urllib.error
import urllib.request
from pathlib import Path
from typing import NamedTuple, NoReturn, Optional
from typing import Literal, NamedTuple, NoReturn, Optional
from urllib.parse import urlencode, urlparse
import click
@ -183,6 +183,17 @@ _TENSOR_PARALLEL_OPTION = typer.Option(
rich_help_panel = _PANEL_MODEL,
help = "Split a GGUF across GPUs by tensor instead of by layer (multi-GPU only).",
)
_GPU_MEMORY_MODE_OPTION = typer.Option(
None,
"--gpu-memory-mode",
rich_help_panel = _PANEL_MODEL,
help = (
"GPU memory strategy for GGUF models loaded by this command. Auto lets "
"Unsloth manage placement. Manual with default layers and context delegates "
"placement and sizing to llama.cpp --fit. Omit when attaching to preserve "
"the running model's mode."
),
)
# Server knobs. Only used when `unsloth start` auto-starts the server (--serve);
# they have no effect when attaching to a server someone else already started.
@ -459,6 +470,7 @@ class LoadOptions(NamedTuple):
max_seq_length: int = 0
load_in_4bit: bool = True
tensor_parallel: bool = False
gpu_memory_mode: Optional[Literal["auto", "manual"]] = None
class ServerOptions(NamedTuple):
@ -993,6 +1005,8 @@ def _start_studio_server(
command += ["--no-load-in-4bit"]
if load.tensor_parallel:
command += ["--tensor-parallel"]
if load.gpu_memory_mode is not None:
command += ["--gpu-memory-mode", load.gpu_memory_mode]
log_path = Path(tempfile.gettempdir()) / f"unsloth-start-server-{os.getpid()}.log"
typer.echo("Starting Unsloth server")
@ -1463,7 +1477,11 @@ def _resolve_model(
# without reloading when the variant AND settings match, so a second session running
# the same command still attaches without evicting the first.
load_has_overrides = bool(
load.gguf_variant or load.max_seq_length or not load.load_in_4bit or load.tensor_parallel
load.gguf_variant
or load.max_seq_length
or not load.load_in_4bit
or load.tensor_parallel
or load.gpu_memory_mode is not None
)
# /v1/models also lists cached-but-unloaded catalog entries (loaded == False);
# matching one would skip /api/inference/load and leave the agent pointed at a
@ -1517,6 +1535,10 @@ def _resolve_model(
payload["load_in_4bit"] = False
if load.tensor_parallel:
payload["tensor_parallel"] = True
if load.gpu_memory_mode is not None:
payload["gpu_memory_mode"] = load.gpu_memory_mode
if load.gpu_memory_mode == "manual":
payload["gpu_layers"] = -1
loaded = _load_model_with_progress(base, key, requested, load, payload)
if loaded.get("status") == "already_loaded":
typer.echo(f"Reusing loaded model: {_display_model_spec(requested, load.gguf_variant)}")
@ -3022,6 +3044,7 @@ def claude(
max_seq_length: int = _CONTEXT_OPTION,
load_in_4bit: bool = _LOAD_4BIT_OPTION,
tensor_parallel: bool = _TENSOR_PARALLEL_OPTION,
gpu_memory_mode: Optional[Literal["auto", "manual"]] = _GPU_MEMORY_MODE_OPTION,
enable_tools: bool = _ENABLE_TOOLS_OPTION,
tool_call_healing: Optional[bool] = _TOOL_CALL_HEALING_OPTION,
tool_call_nudging: Optional[bool] = _TOOL_CALL_NUDGING_OPTION,
@ -3042,7 +3065,7 @@ def claude(
base, key, entry = _connect(
api_key,
model,
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel),
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel, gpu_memory_mode),
serve = serve,
launch = launch,
server_options = ServerOptions(
@ -3139,6 +3162,7 @@ def codex(
max_seq_length: int = _CONTEXT_OPTION,
load_in_4bit: bool = _LOAD_4BIT_OPTION,
tensor_parallel: bool = _TENSOR_PARALLEL_OPTION,
gpu_memory_mode: Optional[Literal["auto", "manual"]] = _GPU_MEMORY_MODE_OPTION,
enable_tools: bool = _ENABLE_TOOLS_OPTION,
tool_call_healing: Optional[bool] = _TOOL_CALL_HEALING_OPTION,
tool_call_nudging: Optional[bool] = _TOOL_CALL_NUDGING_OPTION,
@ -3159,7 +3183,7 @@ def codex(
base, key, entry = _connect(
api_key,
model,
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel),
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel, gpu_memory_mode),
serve = serve,
launch = launch,
server_options = ServerOptions(
@ -3237,6 +3261,7 @@ def openclaw(
max_seq_length: int = _CONTEXT_OPTION,
load_in_4bit: bool = _LOAD_4BIT_OPTION,
tensor_parallel: bool = _TENSOR_PARALLEL_OPTION,
gpu_memory_mode: Optional[Literal["auto", "manual"]] = _GPU_MEMORY_MODE_OPTION,
enable_tools: bool = _ENABLE_TOOLS_OPTION,
tool_call_healing: Optional[bool] = _TOOL_CALL_HEALING_OPTION,
tool_call_nudging: Optional[bool] = _TOOL_CALL_NUDGING_OPTION,
@ -3257,7 +3282,7 @@ def openclaw(
base, key, entry = _connect(
api_key,
model,
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel),
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel, gpu_memory_mode),
serve = serve,
launch = launch,
server_options = ServerOptions(
@ -3317,6 +3342,7 @@ def opencode(
max_seq_length: int = _CONTEXT_OPTION,
load_in_4bit: bool = _LOAD_4BIT_OPTION,
tensor_parallel: bool = _TENSOR_PARALLEL_OPTION,
gpu_memory_mode: Optional[Literal["auto", "manual"]] = _GPU_MEMORY_MODE_OPTION,
enable_tools: bool = _ENABLE_TOOLS_OPTION,
tool_call_healing: Optional[bool] = _TOOL_CALL_HEALING_OPTION,
tool_call_nudging: Optional[bool] = _TOOL_CALL_NUDGING_OPTION,
@ -3337,7 +3363,7 @@ def opencode(
base, key, entry = _connect(
api_key,
model,
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel),
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel, gpu_memory_mode),
serve = serve,
launch = launch,
server_options = ServerOptions(
@ -3477,6 +3503,7 @@ def hermes(
max_seq_length: int = _CONTEXT_OPTION,
load_in_4bit: bool = _LOAD_4BIT_OPTION,
tensor_parallel: bool = _TENSOR_PARALLEL_OPTION,
gpu_memory_mode: Optional[Literal["auto", "manual"]] = _GPU_MEMORY_MODE_OPTION,
enable_tools: bool = _ENABLE_TOOLS_OPTION,
tool_call_healing: Optional[bool] = _TOOL_CALL_HEALING_OPTION,
tool_call_nudging: Optional[bool] = _TOOL_CALL_NUDGING_OPTION,
@ -3499,7 +3526,7 @@ def hermes(
base, key, entry = _connect(
api_key,
model,
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel),
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel, gpu_memory_mode),
serve = serve,
launch = launch,
server_options = ServerOptions(
@ -3533,6 +3560,7 @@ def pi(
max_seq_length: int = _CONTEXT_OPTION,
load_in_4bit: bool = _LOAD_4BIT_OPTION,
tensor_parallel: bool = _TENSOR_PARALLEL_OPTION,
gpu_memory_mode: Optional[Literal["auto", "manual"]] = _GPU_MEMORY_MODE_OPTION,
enable_tools: bool = _ENABLE_TOOLS_OPTION,
tool_call_healing: Optional[bool] = _TOOL_CALL_HEALING_OPTION,
tool_call_nudging: Optional[bool] = _TOOL_CALL_NUDGING_OPTION,
@ -3553,7 +3581,7 @@ def pi(
base, key, entry = _connect(
api_key,
model,
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel),
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel, gpu_memory_mode),
serve = serve,
launch = launch,
server_options = ServerOptions(

View file

@ -19,7 +19,7 @@ import urllib.error
import urllib.request
from datetime import datetime, timezone
from pathlib import Path
from typing import List, Optional
from typing import List, Literal, Optional
import typer
from unsloth_cli.commands import _password_prompt
@ -1172,6 +1172,7 @@ def _load_model_via_http(
gguf_variant: Optional[str],
max_seq_length: int,
load_in_4bit: bool,
gpu_memory_mode: Literal["auto", "manual"] = "auto",
tensor_parallel: bool = False,
llama_extra_args: Optional[List[str]] = None,
timeout: int = 600,
@ -1188,6 +1189,9 @@ def _load_model_via_http(
}
if gguf_variant:
payload["gguf_variant"] = gguf_variant
if gpu_memory_mode == "manual":
payload["gpu_memory_mode"] = "manual"
payload["gpu_layers"] = -1
if tensor_parallel:
payload["tensor_parallel"] = True
if llama_extra_args:
@ -1728,6 +1732,16 @@ def run(
rich_help_panel = _RUN_PANEL_MODEL,
help = "Runtime context length in tokens (0 = model default for GGUF; 2048 for hub models)",
),
gpu_memory_mode: Literal["auto", "manual"] = typer.Option(
"auto",
"--gpu-memory-mode",
rich_help_panel = _RUN_PANEL_MODEL,
help = (
"GPU memory strategy for GGUF models. Auto lets Unsloth select GPUs "
"and cap context to fit VRAM. Manual with default layers and context "
"delegates placement and sizing to llama.cpp --fit."
),
),
load_in_4bit: bool = typer.Option(
True, "--load-in-4bit/--no-load-in-4bit", rich_help_panel = _RUN_PANEL_MODEL
),
@ -2127,6 +2141,8 @@ def run(
"--host",
host,
]
if gpu_memory_mode != "auto":
args.extend(["--gpu-memory-mode", gpu_memory_mode])
if gguf_variant:
args.extend(["--gguf-variant", gguf_variant])
# Forward the explicit polarity; a future default flip on one
@ -2244,6 +2260,7 @@ def run(
gguf_variant = gguf_variant,
max_seq_length = max_seq_length,
load_in_4bit = load_in_4bit,
gpu_memory_mode = gpu_memory_mode,
tensor_parallel = tensor_parallel,
llama_extra_args = extra_llama_args,
)

View file

@ -1993,6 +1993,7 @@ def test_start_studio_server_forwards_tool_flags_via_command_and_env(monkeypatch
start._start_studio_server("http://127.0.0.1:8888", "unsloth/M-GGUF", start.LoadOptions())
cmd, env = captured["command"], captured["kwargs"]["env"]
assert "--disable-tools" in cmd and "--enable-tools" not in cmd
assert "--gpu-memory-mode" not in cmd
assert env["UNSLOTH_DISABLE_TOOL_CALL_HEALING"] == "0"
assert env["UNSLOTH_TOOL_CALL_NUDGE"] == "1"
@ -2206,6 +2207,59 @@ def test_connect_load_knobs_reach_server_even_when_id_loaded(fake_studio):
]
@pytest.mark.parametrize(
"command_name", ["claude", "codex", "openclaw", "opencode", "hermes", "pi"]
)
def test_start_agents_expose_gpu_memory_mode_option(command_name):
import inspect
command = getattr(start, command_name)
opt = inspect.signature(command).parameters["gpu_memory_mode"].default
assert set(getattr(opt, "param_decls", None) or []) == {"--gpu-memory-mode"}
assert getattr(opt, "default", None) is None
assert getattr(opt, "rich_help_panel", None) == start._PANEL_MODEL
@pytest.mark.parametrize(
"mode,expected",
[
("auto", {"model_path": MODEL["id"], "gpu_memory_mode": "auto"}),
(
"manual",
{
"model_path": MODEL["id"],
"gpu_memory_mode": "manual",
"gpu_layers": -1,
},
),
],
)
def test_start_gpu_memory_mode_reaches_running_server(fake_studio, mode, expected):
result = CliRunner().invoke(
start.start_app,
[
"claude",
"--no-launch",
"--model",
MODEL["id"],
"--gpu-memory-mode",
mode,
],
)
assert result.exit_code == 0, result.output
loads = [call for call in fake_studio if call[1].endswith("/api/inference/load")]
assert loads == [("POST", f"{BASE}/api/inference/load", expected)]
def test_start_rejects_invalid_gpu_memory_mode(fake_studio):
result = CliRunner().invoke(
start.start_app,
["claude", "--no-launch", "--gpu-memory-mode", "invalid"],
)
assert result.exit_code != 0
assert "Invalid value for '--gpu-memory-mode'" in result.output
def test_connect_model_variant_suffix_loads_split_repo(fake_studio):
# When the model is not already loaded, the `:QUANT` suffix becomes the gguf_variant
# and the load uses the bare (valid) repo id, mirroring `unsloth run repo --gguf-variant`.
@ -2617,7 +2671,11 @@ def test_start_studio_server_builds_command_and_waits(monkeypatch, capsys):
"http://127.0.0.1:8888",
"unsloth/Qwen3-1.7B-GGUF:UD-Q4_K_XL",
start.LoadOptions(
gguf_variant = "UD-Q4_K_XL", max_seq_length = 8192, load_in_4bit = True, tensor_parallel = True
gguf_variant = "UD-Q4_K_XL",
max_seq_length = 8192,
load_in_4bit = True,
tensor_parallel = True,
gpu_memory_mode = "manual",
),
)
cmd = captured["command"]
@ -2627,6 +2685,7 @@ def test_start_studio_server_builds_command_and_waits(monkeypatch, capsys):
assert cmd[cmd.index("--gguf-variant") + 1] == "UD-Q4_K_XL"
assert cmd[cmd.index("--context-length") + 1] == "8192"
assert "--tensor-parallel" in cmd
assert cmd[cmd.index("--gpu-memory-mode") + 1] == "manual"
assert "--start-api-key-marker" not in cmd
assert captured["kwargs"]["env"][start._START_API_KEY_MARKER_ENV] == "1"
assert start.os.environ[start._START_API_KEY_MARKER_ENV] == "parent"

View file

@ -13,7 +13,9 @@ canonicaliser and the legacy `-m` / `-hfr` / `-f` shim.
from __future__ import annotations
import json
import sys
from io import BytesIO
from pathlib import Path
import pytest
@ -63,6 +65,19 @@ def test_context_length_alias_is_registered():
assert "--context-length" in flags
def test_gpu_memory_mode_option_is_registered_with_auto_default():
"""The GPU placement policy is a first-class model option."""
studio_mod = _load_run_command()
import inspect
sig = inspect.signature(studio_mod.run)
opt = sig.parameters["gpu_memory_mode"].default
flags = set(getattr(opt, "param_decls", None) or [])
assert flags == {"--gpu-memory-mode"}
assert getattr(opt, "default", None) == "auto"
assert getattr(opt, "rich_help_panel", None) == "Model"
def test_parallel_default_is_four():
"""Default must stay at 4 so plain `unsloth studio run` is unchanged."""
studio_mod = _load_run_command()
@ -377,6 +392,75 @@ def test_reexec_forwards_context_length_alias(monkeypatch):
assert "--context-length" not in argv, argv
def test_reexec_forwards_manual_gpu_memory_mode(monkeypatch):
"""An explicit manual policy must survive the Studio venv re-exec."""
result, captured = _invoke_run(
monkeypatch,
_BASE + ["--gpu-memory-mode", "manual"],
)
assert len(captured) == 1, result.output
argv = captured[0]["argv"]
assert _value_after(argv, "--gpu-memory-mode") == "manual", argv
def test_reexec_omits_default_gpu_memory_mode(monkeypatch):
"""The default stays compatible with older Studio venv launchers."""
result, captured = _invoke_run(monkeypatch, _BASE)
assert len(captured) == 1, result.output
assert "--gpu-memory-mode" not in captured[0]["argv"]
def test_run_rejects_invalid_gpu_memory_mode(monkeypatch):
result, captured = _invoke_run(
monkeypatch,
_BASE + ["--gpu-memory-mode", "invalid"],
)
assert result.exit_code != 0
assert captured == []
@pytest.mark.parametrize(
"mode,expected",
[
("auto", {"model_path": "owner/model-GGUF", "max_seq_length": 0, "load_in_4bit": True}),
(
"manual",
{
"model_path": "owner/model-GGUF",
"max_seq_length": 0,
"load_in_4bit": True,
"gpu_memory_mode": "manual",
"gpu_layers": -1,
},
),
],
)
def test_load_model_http_payload_for_gpu_memory_mode(monkeypatch, mode, expected):
"""Manual plus untouched layer and context settings matches the UI payload."""
studio_mod = _load_run_command()
captured = {}
def urlopen(request, timeout):
captured["request"] = request
captured["timeout"] = timeout
return BytesIO(b'{"model": "owner/model-GGUF"}')
monkeypatch.setattr(studio_mod.urllib.request, "urlopen", urlopen)
result = studio_mod._load_model_via_http(
port = 8888,
api_key = "sk-test",
model = "owner/model-GGUF",
gguf_variant = None,
max_seq_length = 0,
load_in_4bit = True,
gpu_memory_mode = mode,
)
assert result == {"model": "owner/model-GGUF"}
assert json.loads(captured["request"].data) == expected
assert captured["request"].get_header("Authorization") == "Bearer sk-test"
def test_reexec_mixed_parallel_with_passthrough(monkeypatch):
"""--parallel + llama-server pass-through flags must all reach the child."""
result, captured = _invoke_run(