unsloth/studio/backend/tests/test_llama_server_args.py
oobabooga 72e67ae5a6
Studio: Add Tensor-Parallel llama.cpp support (#6040)
* Studio: Add Tensor-Parallel llama.cpp support

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

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

* Studio: harden Tensor-Parallel fallback and GPU selection

* Studio: reconcile split-mode extras and harden tensor-split planning

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

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

* Studio: reconcile split-mode extras in backend duplicate-load guard

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

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

* Studio: preserve inherited non-tensor split modes on reload

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

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

* Studio: honor cancellation in tensor fallback, preserve tensor mode on rollback, and don't raise an explicit small context

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

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

* Studio: reconcile split-mode in reload check and strip it on tensor downgrade

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

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

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

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

* Strip --tensor-split alongside --split-mode so inherited ratios don't override the tensor planner

An inherited or stale --tensor-split in llama_extra_args was appended after
Studio's computed --tensor-split and won last in llama.cpp, re-introducing the
asymmetric-GPU OOM tensor mode is meant to prevent. Group -ts/--tensor-split
into the split-mode shadow set so it is stripped on inherit and on the layer
fallback; parse_split_mode_override still keys on the mode value only.

* Drop quantized KV for the tensor attempt and report native max context

Tensor mode aborts on a quantized KV cache, so a user with q8_0/q4_1 etc. who
enabled Tensor Parallelism silently fell back to layer split. Clear the cache
type (and strip inherited/explicit --cache-type) for the tensor attempt only;
the layer fallback re-runs with tensor off and keeps the user's choice.

Also report max_available_ctx from the native context, not an explicit small
-c, so the context slider no longer warns too early in tensor mode.

* Reconcile inherited split-mode extras in the already-loaded check

When a same-model load omitted llama_extra_args, the tensor comparison resolved
the raw (None) request and treated an inherited --split-mode tensor server as a
mismatch, forcing a needless reload. Compare using the stored extras stripped
the same way the reload strips them.

* Pass tensor_parallel through compare-mode loads

The generalized compare path loaded each GGUF without tensor_parallel, so
compare ran layer split even with the toggle on and left the settings sheet
stale. Send the toggle and hydrate the loaded state from the response, matching
the main chat and recipe load paths.

* Add --tensor-parallel flag to unsloth studio run

The headless one-liner could only reach tensor mode by passing --split-mode
tensor as a raw llama.cpp extra. Add a first-class --tensor-parallel/
--no-tensor-parallel option that sets the tensor_parallel field on the
/api/inference/load payload, forwarded through the studio-venv re-exec like the
other polarity flags. Matches the web UI toggle and the API field.

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

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

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: danielhanchen <michaelhan2050@gmail.com>
2026-06-12 04:00:52 -07:00

740 lines
25 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Unit tests for the llama-server pass-through args validator.
The validator is the boundary between user CLI/HTTP input and the
llama-server subprocess. These tests pin denylist behaviour so it doesn't
regress when new managed flags are added.
"""
from __future__ import annotations
import importlib.util
import re
from pathlib import Path
import pytest
# Load llama_server_args.py directly to avoid dragging in the full backend
# chain via core/inference/__init__.py. The validator is dependency-free.
_LSA_PATH = Path(__file__).resolve().parent.parent / "core" / "inference" / "llama_server_args.py"
_spec = importlib.util.spec_from_file_location("_lsa_test_only", _LSA_PATH)
_lsa = importlib.util.module_from_spec(_spec)
_spec.loader.exec_module(_lsa)
is_managed_flag = _lsa.is_managed_flag
parse_cache_override = _lsa.parse_cache_override
parse_ctx_override = _lsa.parse_ctx_override
parse_split_mode_override = _lsa.parse_split_mode_override
resolve_cache_type_kv = _lsa.resolve_cache_type_kv
resolve_tensor_parallel = _lsa.resolve_tensor_parallel
strip_shadowing_flags = _lsa.strip_shadowing_flags
strip_split_mode_only = _lsa.strip_split_mode_only
extra_args_disable_mmproj = _lsa.extra_args_disable_mmproj
validate_extra_args = _lsa.validate_extra_args
# ── Pass-through (allowed) ───────────────────────────────────────────
@pytest.mark.parametrize(
"args",
[
# Sampling
["--top-k", "20"],
["--top-p", "0.9", "--min-p", "0.05"],
["--seed", "-1"], # negative value, not a flag
["--temp", "0.0"],
["--repeat-penalty", "1.05"],
["--mirostat", "2", "--mirostat-lr", "0.1"],
["--xtc-probability", "0.05", "--xtc-threshold", "0.1"],
["--dry-multiplier", "0.5"],
# Tier-2 knobs that map to LoadRequest fields
["--cache-type-k", "q8_0"],
["--cache-type-v", "q8_0"],
["--chat-template-file", "/tmp/tpl.jinja"],
["--chat-template-kwargs", '{"reasoning_effort":"high"}'],
["--spec-type", "ngram-mod"],
["--spec-default"],
# MTP path (llama.cpp #22673).
["--spec-type", "draft-mtp"],
["--spec-type", "draft-mtp", "--spec-draft-n-max", "6"],
[
"--spec-type",
"ngram-mod,draft-mtp",
"--spec-draft-n-max",
"3",
"--spec-ngram-mod-n-match",
"24",
"--spec-ngram-mod-n-min",
"48",
"--spec-ngram-mod-n-max",
"64",
],
# Reasoning controls
["--reasoning-format", "deepseek"],
["-rea", "auto"],
# Soft-managed: user flags last-wins over Studio's auto-set version.
# --parallel / -np / --n-parallel are hard-denied (KV-cache + slot
# count would desync); use `unsloth studio run --parallel N` instead.
["-c", "131072"],
["--ctx-size", "8192"],
["--flash-attn", "off"],
["-fa", "on"],
["--no-context-shift"],
["--context-shift"],
["--jinja"],
["--no-jinja"],
["-ngl", "-1"],
["--gpu-layers", "32"],
["-t", "16"],
["--threads", "32"],
["-fit", "off"],
["--fit", "on"],
["--fit-ctx", "8192"],
],
)
def test_pass_through_allowed(args):
assert validate_extra_args(args) == args
def test_none_returns_empty_list():
assert validate_extra_args(None) == []
def test_empty_list_returns_empty_list():
assert validate_extra_args([]) == []
def test_value_with_equals_form_passes_through():
assert validate_extra_args(["--top-k=20"]) == ["--top-k=20"]
def test_non_flag_token_passes_through():
# Bare positionals are passed through; llama-server can reject them.
assert validate_extra_args(["foo"]) == ["foo"]
# ── Denylist (rejected) ──────────────────────────────────────────────
@pytest.mark.parametrize(
"denied",
[
# Parallel slots -- owned by the typer --parallel flag.
"-np",
"--parallel",
"--n-parallel",
# Model identity (every alias; bumping llama.cpp must keep every
# form rejected, not just the long one).
"-m",
"--model",
"-mu",
"--model-url",
"-dr",
"--docker-repo",
"-hf",
"-hfr",
"--hf-repo",
"-hff",
"--hf-file",
"-hfv",
"-hfrv",
"--hf-repo-v",
"-hffv",
"--hf-file-v",
"-hft",
"--hf-token",
"-mm",
"--mmproj",
"-mmu",
"--mmproj-url",
# Networking (Studio binds + proxies)
"--host",
"--port",
"--path",
"--api-prefix",
"--reuse-port",
# Auth / TLS
"--api-key",
"--api-key-file",
"--ssl-key-file",
"--ssl-cert-file",
# Single-model server (legacy --webui + current --ui group)
"--webui",
"--no-webui",
"--ui",
"--no-ui",
"--ui-config",
"--ui-config-file",
"--ui-mcp-proxy",
"--no-ui-mcp-proxy",
"--models-dir",
"--models-preset",
"--models-max",
"--models-autoload",
"--no-models-autoload",
# Server-mode flips: --embedding / --rerank restrict llama-server to
# those endpoints and break Studio's chat hop.
"--embedding",
"--embeddings",
"--rerank",
"--reranking",
# llama-server's own --tools clashes with Studio's tool policy.
"--tools",
],
)
def test_denylist_rejects_all_aliases(denied):
with pytest.raises(ValueError, match = denied):
validate_extra_args([denied, "value"])
@pytest.mark.parametrize(
"args,offending",
[
# Pass-through --parallel would last-wins-override the real slot
# count while Studio's KV-cache fit + llama_parallel_slots stay at
# the typer value -- plan vs. process disagree.
(["--parallel", "8"], "--parallel"),
(["--parallel=8"], "--parallel"),
(["--n-parallel", "16"], "--n-parallel"),
(["--n-parallel=16"], "--n-parallel"),
(["-np", "32"], "-np"),
# Attached short form: Click clusters it CLI-side; HTTP /load with
# `["-np8"]` must still resolve to managed.
(["-np8"], "-np"),
(["-np64"], "-np"),
# Out-of-range values that would bypass the typer 1..64 guard.
(["--parallel", "999"], "--parallel"),
(["-np", "0"], "-np"),
(["-np999"], "-np"),
# Signed attached forms; `-np-1` must not slip past.
(["-np-1"], "-np"),
(["-np+1"], "-np"),
],
)
def test_parallel_flags_are_managed(args, offending):
with pytest.raises(ValueError, match = re.escape(offending)):
validate_extra_args(args)
def test_denylist_rejects_equals_form():
with pytest.raises(ValueError, match = "--port"):
validate_extra_args(["--port=9000"])
@pytest.mark.parametrize(
"padded",
[" --parallel", "--parallel ", "\t--parallel", " -np", "-np \n", "-np\t"],
)
def test_denylist_rejects_whitespace_padded_forms(padded):
# `_flag_name` trims whitespace before lookup; else a trailing space
# could slip a managed flag past the boundary.
with pytest.raises(ValueError, match = "parallel|np"):
validate_extra_args([padded, "8"])
@pytest.mark.parametrize(
"attached",
["-np8x", "-np-1foo", "-np+1bar", "-np9zzz"],
)
def test_denylist_rejects_np_with_digit_prefix_and_junk(attached):
# Backend `_flag_name` must classify the same forms the CLI rewriter
# expands, else HTTP /load could smuggle `-np8x` through.
with pytest.raises(ValueError, match = "np"):
validate_extra_args([attached])
def test_denylist_rejects_short_form_when_long_is_denied():
# `-m` is the short form of --model; rejecting only the long form
# would leave a trivial bypass.
with pytest.raises(ValueError, match = "-m"):
validate_extra_args(["-m", "/some/other/path.gguf"])
def test_denylist_message_names_offending_flag():
with pytest.raises(ValueError) as excinfo:
validate_extra_args(["--top-k", "20", "--api-key", "secret"])
assert "--api-key" in str(excinfo.value)
def test_first_denied_flag_short_circuits():
# Validation stops at the first denied flag; the message names it.
with pytest.raises(ValueError, match = "--port"):
validate_extra_args(["--port", "1", "--host", "x"])
# ── Numeric values that look flag-ish ─────────────────────────────────
@pytest.mark.parametrize("value", ["-1", "-0.5", "-42", "-.5"])
def test_negative_number_value_is_not_flag(value):
# `--seed -1`: the -1 is a value, not a flag.
assert validate_extra_args(["--seed", value]) == ["--seed", value]
# ── is_managed_flag helper ───────────────────────────────────────────
def test_is_managed_flag_true_for_denied():
assert is_managed_flag("--port") is True
assert is_managed_flag("--api-key") is True
assert is_managed_flag("-m") is True
assert is_managed_flag("--model") is True
# Parallel slots owned by the typer --parallel flag.
assert is_managed_flag("--parallel") is True
assert is_managed_flag("--n-parallel") is True
assert is_managed_flag("-np") is True
# Normalised forms must classify like the canonical token so
# is_managed_flag filtering stays in sync with validate_extra_args.
assert is_managed_flag("-np8") is True
assert is_managed_flag("--parallel=8") is True
assert is_managed_flag("--port=9000") is True
def test_is_managed_flag_false_for_pass_through():
assert is_managed_flag("--top-k") is False
assert is_managed_flag("--cache-type-k") is False
assert is_managed_flag("--chat-template-file") is False
# Soft-managed flags pass through (last-wins override)
assert is_managed_flag("-c") is False
assert is_managed_flag("--ctx-size") is False
assert is_managed_flag("--flash-attn") is False
assert is_managed_flag("-ngl") is False
assert is_managed_flag("--threads") is False
# ── strip_shadowing_flags ─────────────────────────────────────────────
def test_strip_shadowing_flags_drops_context_when_requested():
out = strip_shadowing_flags(
["-c", "4096", "--top-k", "20"],
strip_context = True,
strip_cache = False,
strip_spec = False,
strip_template = False,
)
assert out == ["--top-k", "20"]
def test_strip_shadowing_flags_keeps_context_when_not_requested():
out = strip_shadowing_flags(
["-c", "4096", "--top-k", "20"],
strip_context = False,
strip_cache = False,
strip_spec = False,
strip_template = False,
)
assert out == ["-c", "4096", "--top-k", "20"]
def test_strip_shadowing_flags_keeps_chat_template_when_template_disabled():
# No chat_template_override supplied; inherited
# --chat-template-file must survive.
out = strip_shadowing_flags(
["--chat-template-file", "/tmp/custom.jinja", "--top-k", "20"],
strip_context = True,
strip_cache = True,
strip_spec = True,
strip_template = False,
)
assert out == ["--chat-template-file", "/tmp/custom.jinja", "--top-k", "20"]
def test_strip_shadowing_flags_drops_template_when_requested():
out = strip_shadowing_flags(
["--chat-template-file", "/tmp/custom.jinja", "--top-k", "20"],
strip_template = True,
)
assert out == ["--top-k", "20"]
def test_strip_shadowing_flags_keeps_cache_when_cache_disabled():
out = strip_shadowing_flags(
["--cache-type-k", "q8_0", "--cache-type-v", "q8_0", "--top-k", "20"],
strip_cache = False,
)
assert out == ["--cache-type-k", "q8_0", "--cache-type-v", "q8_0", "--top-k", "20"]
def test_strip_shadowing_flags_keeps_spec_when_spec_disabled():
out = strip_shadowing_flags(
["--spec-type", "ngram-mod", "--draft-min", "48", "--top-k", "20"],
strip_spec = False,
)
assert out == ["--spec-type", "ngram-mod", "--draft-min", "48", "--top-k", "20"]
def test_strip_shadowing_flags_drops_mtp_flags_when_requested():
# MTP / draft-mtp flags must drop when speculative_type re-applies.
out = strip_shadowing_flags(
[
"--spec-type",
"draft-mtp",
"--spec-draft-n-max",
"6",
"--spec-ngram-mod-n-match",
"24",
"--spec-ngram-mod-n-min",
"48",
"--spec-ngram-mod-n-max",
"6",
"--top-k",
"20",
],
strip_spec = True,
)
assert out == ["--top-k", "20"]
def test_is_managed_flag_false_for_mtp_pass_through():
assert is_managed_flag("--spec-draft-n-max") is False
assert is_managed_flag("--spec-ngram-mod-n-match") is False
assert is_managed_flag("--spec-ngram-mod-n-min") is False
assert is_managed_flag("--spec-ngram-mod-n-max") is False
# ── parse_ctx_override ───────────────────────────────────────────────
@pytest.mark.parametrize(
"args,expected",
[
(None, None),
([], None),
(["--top-k", "20"], None),
(["--ctx-size", "128000"], 128000),
(["--ctx-size=128000"], 128000),
(["-c", "128000"], 128000),
(["-c=128000"], 128000),
(["-c", "4096", "--ctx-size", "128000"], 128000),
],
)
def test_parse_ctx_override(args, expected):
assert parse_ctx_override(args) == expected
@pytest.mark.parametrize(
"args",
[
["--ctx-size"],
["--ctx-size", "--top-k"],
["--ctx-size", "abc"],
["--ctx-size=abc"],
["-c", "-1"],
],
)
def test_parse_ctx_override_rejects_malformed_values(args):
with pytest.raises(ValueError, match = "ctx-size|'-c'"):
parse_ctx_override(args)
def test_validate_extra_args_rejects_malformed_ctx_override():
with pytest.raises(ValueError, match = "ctx-size"):
validate_extra_args(["--ctx-size", "abc"])
# ── parse_cache_override ─────────────────────────────────────────────
@pytest.mark.parametrize(
"args,expected",
[
(None, None),
([], None),
(["--top-k", "20"], None),
(["--cache-type-k", "q8_0"], "q8_0"),
(["-ctk", "q4_0"], "q4_0"),
(["-ctv", "q4_0"], "q4_0"),
(["--cache-type-k=q4_0"], "q4_0"),
(["-ctk", "f16", "-ctk", "q8_0"], "q8_0"),
],
)
def test_parse_cache_override(args, expected):
assert parse_cache_override(args) == expected
@pytest.mark.parametrize(
"args",
[
["-ctk"],
["-ctk", "-c", "4096"],
],
)
def test_parse_cache_override_rejects_malformed_values(args):
with pytest.raises(ValueError, match = "cache-type|'-ctk'"):
parse_cache_override(args)
def test_resolve_cache_type_kv_uses_override_when_present():
assert resolve_cache_type_kv(["--cache-type-k", "q8_0"], "f16") == "q8_0"
def test_resolve_cache_type_kv_uses_fallback_without_override():
assert resolve_cache_type_kv(["--top-k", "20"], "f16") == "f16"
def test_strip_shadowing_flags_boolean_does_not_consume_next_token():
# `--spec-default` is boolean; drop just the flag, keep the next token.
out = strip_shadowing_flags(["--spec-default", "ngram-mod"], strip_spec = True)
assert out == ["ngram-mod"]
def test_strip_shadowing_flags_jinja_boolean_preserves_positional():
out = strip_shadowing_flags(["--jinja", "trailing-positional"], strip_template = True)
assert out == ["trailing-positional"]
def test_strip_shadowing_flags_no_jinja_boolean_preserves_positional():
out = strip_shadowing_flags(["--no-jinja", "trailing-positional"], strip_template = True)
assert out == ["trailing-positional"]
def test_strip_shadowing_flags_equals_form_drops_only_the_flag():
out = strip_shadowing_flags(["--ctx-size=4096", "--seed", "-1"], strip_context = True)
assert out == ["--seed", "-1"]
def test_strip_shadowing_flags_handles_none_input():
assert strip_shadowing_flags(None) == []
def test_strip_shadowing_flags_handles_empty_input():
assert strip_shadowing_flags([]) == []
def test_strip_shadowing_flags_defaults_strip_everything():
# The route's already-loaded comparator calls with no kwargs to
# detect ANY shadowing flag in stored extras.
out = strip_shadowing_flags(
["-c", "4096", "--cache-type-k", "q8_0", "--spec-default", "--jinja"]
)
assert out == []
# ── --split-mode (Tensor Parallelism toggle) ─────────────────────────
# Soft-shadowed exactly like --cache-type-*: pass-through allowed (keeps
# the row/none/layer modes the boolean toggle doesn't expose), stripped
# on inherit, and reconciled back into the round-tripped tensor_parallel
# state.
@pytest.mark.parametrize(
"args",
[
["--split-mode", "tensor"],
["--split-mode", "row"],
["--split-mode", "none"],
["--split-mode", "layer"],
["-sm", "tensor"],
["--split-mode=row"],
["-sm=tensor"],
],
)
def test_split_mode_passes_through(args):
# Not denylisted -- a user keeps row/none/layer via extras.
assert validate_extra_args(args) == args
def test_split_mode_is_not_managed():
assert is_managed_flag("--split-mode") is False
assert is_managed_flag("-sm") is False
@pytest.mark.parametrize(
"args,expected",
[
(None, None),
([], None),
(["--top-k", "20"], None),
(["--split-mode", "tensor"], "tensor"),
(["--split-mode", "row"], "row"),
(["-sm", "none"], "none"),
(["--split-mode=layer"], "layer"),
(["-sm=tensor"], "tensor"),
# last-wins when supplied twice
(["-sm", "row", "--split-mode", "tensor"], "tensor"),
],
)
def test_parse_split_mode_override(args, expected):
assert parse_split_mode_override(args) == expected
@pytest.mark.parametrize(
"args",
[
["--split-mode"],
["-sm"],
["--split-mode", "-c", "4096"], # next token is a flag, not a value
],
)
def test_parse_split_mode_override_rejects_malformed_values(args):
with pytest.raises(ValueError, match = "split-mode|'-sm'"):
parse_split_mode_override(args)
def test_validate_extra_args_rejects_malformed_split_mode():
# Validation catches a value-less --split-mode at the boundary,
# mirroring the early --ctx-size / --cache-type checks.
with pytest.raises(ValueError, match = "split-mode"):
validate_extra_args(["--split-mode"])
@pytest.mark.parametrize(
"args,fallback,expected",
[
# No override -> fall back to the toggle value, both directions.
(["--top-k", "20"], True, True),
(["--top-k", "20"], False, False),
(None, True, True),
([], False, False),
# Explicit override wins: tensor -> on, anything else -> off,
# regardless of the toggle fallback.
(["--split-mode", "tensor"], False, True),
(["-sm", "tensor"], False, True),
(["--split-mode", "row"], True, False),
(["--split-mode", "none"], True, False),
(["--split-mode", "layer"], True, False),
(["--split-mode=tensor"], False, True),
# Case-insensitive on the mode string.
(["--split-mode", "TENSOR"], False, True),
# last-wins across multiple --split-mode flags.
(["-sm", "tensor", "--split-mode", "row"], True, False),
],
)
def test_resolve_tensor_parallel(args, fallback, expected):
assert resolve_tensor_parallel(args, fallback) is expected
def test_strip_shadowing_flags_drops_split_mode_when_requested():
out = strip_shadowing_flags(
["--split-mode", "row", "--top-k", "20"],
strip_context = False,
strip_cache = False,
strip_spec = False,
strip_template = False,
strip_split_mode = True,
)
assert out == ["--top-k", "20"]
def test_extra_args_disable_mmproj_detects_flag():
assert extra_args_disable_mmproj(["--no-mmproj"]) is True
assert extra_args_disable_mmproj(["--threads", "12", "--no-mmproj"]) is True
assert extra_args_disable_mmproj(["--no-mmproj-auto"]) is True
def test_extra_args_disable_mmproj_false_when_absent():
assert extra_args_disable_mmproj(None) is False
assert extra_args_disable_mmproj(["--threads", "12"]) is False
def test_extra_args_disable_mmproj_last_wins():
assert extra_args_disable_mmproj(["--no-mmproj", "--mmproj-auto"]) is False
assert extra_args_disable_mmproj(["--mmproj-auto", "--no-mmproj-auto"]) is True
def test_strip_shadowing_flags_drops_model_draft_with_spec():
# --model-draft (and aliases) are Studio-managed since the separate
# MTP drafter support: an inherited copy must not last-wins-override
# the auto-detected drafter.
out = strip_shadowing_flags(
["--model-draft", "/old/mtp.gguf", "-md", "/old2.gguf", "--top-k", "20"],
strip_context = False,
strip_cache = False,
strip_spec = True,
strip_template = False,
)
assert out == ["--top-k", "20"]
def test_strip_shadowing_flags_keeps_split_mode_when_not_requested():
# No tensor_parallel field supplied on the Apply -> an inherited
# --split-mode survives (mirrors the chat-template keep behavior).
out = strip_shadowing_flags(
["--split-mode", "row", "--top-k", "20"],
strip_context = True,
strip_cache = True,
strip_spec = True,
strip_template = True,
strip_split_mode = False,
)
assert out == ["--split-mode", "row", "--top-k", "20"]
def test_strip_shadowing_flags_drops_split_mode_short_alias_and_equals():
assert strip_shadowing_flags(["-sm", "tensor", "--top-k", "20"], strip_split_mode = True) == [
"--top-k",
"20",
]
assert strip_shadowing_flags(["--split-mode=row", "--seed", "-1"], strip_split_mode = True) == [
"--seed",
"-1",
]
def test_strip_shadowing_flags_defaults_strip_split_mode_too():
# The route's already-loaded comparator (no kwargs) must see a stored
# --split-mode as a shadowing flag so it forces a reload.
assert strip_shadowing_flags(["--split-mode", "tensor"]) == []
@pytest.mark.parametrize(
"args",
[
["--split-mode", "tensor", "-c", "4096"],
["-sm", "tensor", "-c", "4096"],
["--split-mode=tensor", "-c", "4096"],
["-sm=tensor", "-c", "4096"],
],
)
def test_strip_split_mode_only_keeps_other_shadow_flags(args):
# Every --split-mode form (long/short, space/=) is dropped; -c survives.
assert strip_split_mode_only(args) == ["-c", "4096"]
def test_strip_split_mode_only_preserves_none_and_empty():
# None means "inherit"; [] means "explicit empty" -- both must round-trip.
assert strip_split_mode_only(None) is None
assert strip_split_mode_only([]) == []
def test_strip_shadowing_flags_drops_tensor_split_with_split_mode():
# --tensor-split is coupled to the split mode: stripped together so a stale
# ratio can't override Studio's computed tensor split. Other flags survive.
out = strip_shadowing_flags(
["--split-mode", "row", "--tensor-split", "1,1", "--top-k", "20"],
strip_context = False,
strip_cache = False,
strip_spec = False,
strip_template = False,
strip_split_mode = True,
)
assert out == ["--top-k", "20"]
def test_strip_shadowing_flags_keeps_tensor_split_when_not_requested():
# strip_split_mode=False keeps the whole split group (mode + ratios).
assert strip_shadowing_flags(
["--tensor-split", "1,1", "--top-k", "20"], strip_split_mode = False
) == ["--tensor-split", "1,1", "--top-k", "20"]
def test_strip_split_mode_only_drops_tensor_split_too():
# Downgrade / layer fallback must drop the coupled --tensor-split (all forms).
assert strip_split_mode_only(
["--split-mode", "tensor", "--tensor-split", "1,1", "-c", "4096"]
) == ["-c", "4096"]
assert strip_split_mode_only(["-sm=tensor", "-ts=3,1"]) == []
def test_strip_shadowing_flags_keeps_model_draft_without_spec():
out = strip_shadowing_flags(
["--model-draft", "/custom/mtp.gguf"],
strip_context = True,
strip_cache = False,
strip_spec = False,
strip_template = False,
)
assert out == ["--model-draft", "/custom/mtp.gguf"]