* Studio: override chat template for unsloth/gemma-4-*-GGUF with bundled gemma-4.jinja The chat templates baked into the shipped unsloth/gemma-4-*-GGUF quants predate Google's gemma-4 chat-template PR #118 and lack the preserve_thinking flag, so Studio cannot surface the "Preserve thinking" toggle for Gemma 4. Bundle the updated template and override the embedded one at llama-server launch via --chat-template-file, scoped to the gemma-4 GGUF family, so users do not need to re-download any quant. - Add studio/backend/assets/chat_templates/gemma-4.jinja (PR #118 based; preserve_thinking defaults false, the one deliberate divergence from upstream). - Add core/inference/chat_templates.py: gemma-4 GGUF matcher plus an effective-override resolver (explicit user template still wins). - Wire the resolver into routes/inference.py ahead of the reload-dedup check and both load_model calls so the live backend and the incoming request compare against the same template text (no spurious reloads). - Default preserve_thinking off in the launch-time chat_template_kwargs so direct API callers match the UI default. - Ship the asset via package-data and add unit tests. * Studio: ship E2B/E4B edge variant of the bundled Gemma 4 template Google ships two distinct gemma-4 chat templates: E2B and E4B omit the empty "<|channel>thought<channel|>" block on enable_thinking=false, while the 12b/26B-A4B/31B family emits it (confirmed against google/gemma-4-E2B-it, -E4B-it, -12b-it, -26B-A4B-it, -31B-it; the two families differ only in that one block). The single PR #118 based template followed the larger-model behavior, which is wrong for the E2B/E4B GGUFs this feature most targets. - Add studio/backend/assets/chat_templates/gemma-4-edge.jinja: identical to gemma-4.jinja minus the empty-thought-block, matching E2B/E4B behavior. - Route unsloth/gemma-4-E2B-it-GGUF and -E4B-it-GGUF to the edge template; 12b/26B-A4B/31B keep gemma-4.jinja. - Extend tests for the edge matcher, per-family routing, and the empty-thought block difference (off for edge, on for standard). * Studio: address review feedback on the gemma-4 template override - Normalize owner-less shorthand model ids in the template matcher: a bare "gemma-4-E2B-it-GGUF" is canonicalized to "unsloth/" the same way ModelConfig.from_identifier does, so shorthand loads still get the override (and the preserve_thinking capability) instead of falling back to the embedded template. - Scope the test's module stubs with unittest.mock.patch.dict instead of sys.modules.setdefault, and only stub deps that are missing, so the global module registry is not polluted for tests that run afterwards. - Guard the Jinja render tests with pytest.importorskip("jinja2") so the suite stays runnable in minimal Studio environments where jinja2 is not present. - Add tests for shorthand resolution. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Studio: address 10-reviewer P1 findings on the gemma-4 template override - /status no longer surfaces Studio's auto-applied bundled template as a user-authored chat_template_override. The frontend adopts that field as editable state and would otherwise re-send the gemma-4 template as an explicit override for a later, unrelated model. /status now reports None when the live override equals the model's auto-resolved bundled template. - When a bundled family template is in effect, strip an inherited --chat-template-file from llama_extra_args too (not only when the raw request set chat_template_override). Otherwise a stale inherited template, appended last, shadows the bundled one while Studio reports the bundled template's capabilities. - Write the temp chat-template file as UTF-8 explicitly, and keep the bundled templates ASCII (replaced em dashes), so non-UTF-8 Windows locales cannot raise UnicodeEncodeError or emit a mis-encoded template. Added an ASCII guard test. --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
6844 lines
308 KiB
Python
6844 lines
308 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
|
|
|
|
"""llama-server inference backend for GGUF models.
|
|
|
|
Manages a llama-server subprocess and proxies chat completions through its
|
|
OpenAI-compatible /v1/chat/completions endpoint.
|
|
"""
|
|
|
|
import atexit
|
|
import contextlib
|
|
import json
|
|
import os
|
|
import re
|
|
import struct
|
|
import structlog
|
|
from loggers import get_logger
|
|
import shutil
|
|
import signal
|
|
import socket
|
|
import subprocess
|
|
import sys
|
|
import threading
|
|
import time
|
|
from pathlib import Path
|
|
from typing import Callable, Generator, Iterable, List, Optional
|
|
|
|
import httpx
|
|
|
|
from core.inference.llama_server_args import (
|
|
extra_args_disable_mmproj,
|
|
parse_cache_override,
|
|
parse_ctx_override,
|
|
parse_split_mode_override,
|
|
resolve_cache_type_kv,
|
|
resolve_requested_ctx,
|
|
resolve_tensor_parallel,
|
|
strip_shadowing_flags,
|
|
strip_split_mode_only,
|
|
)
|
|
from core.tool_healing import (
|
|
_TC_END_TAG_RE,
|
|
_TC_FUNC_CLOSE_RE,
|
|
_TC_FUNC_START_RE,
|
|
_TC_JSON_START_RE,
|
|
_TC_PARAM_CLOSE_RE,
|
|
_TC_PARAM_START_RE,
|
|
_TOOL_ALL_PATS,
|
|
_TOOL_CLOSED_PATS,
|
|
parse_tool_calls_from_text,
|
|
strip_tool_call_markup,
|
|
)
|
|
from utils.native_path_leases import child_env_without_native_path_secret
|
|
from utils.subprocess_compat import (
|
|
windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs,
|
|
)
|
|
from core.inference.tool_call_parser import (
|
|
RAG_MAX_SEARCHES_PER_TURN,
|
|
RAG_SEARCH_CAP_NUDGE,
|
|
TOOL_XML_SIGNALS,
|
|
parse_tool_calls_from_text as _shared_parse_tool_calls_from_text,
|
|
)
|
|
from core.inference.tool_loop_controller import (
|
|
ToolLoopController,
|
|
tool_event_provenance,
|
|
)
|
|
from state.tool_approvals import (
|
|
TOOL_REJECTED_MESSAGE,
|
|
abort_tool_decision,
|
|
begin_tool_decision,
|
|
new_approval_id,
|
|
wait_tool_decision,
|
|
)
|
|
|
|
logger = get_logger(__name__)
|
|
|
|
|
|
def _wsl_system_rocm_lib_dirs() -> "list[str]":
|
|
"""System ROCm lib dir(s) to load before a prebuilt's bundled HIP, on WSL.
|
|
|
|
The bundled bare-metal HIP can't drive WSL's /dev/dxg and segfaults on the
|
|
first GPU call; the system ROCm libs (libamdhip64 + librocdxg) can, while
|
|
the bundle still supplies libggml-hip / librocblas (gfx1151 kernels).
|
|
Mirrors install_llama_prebuilt._wsl_system_rocm_lib_dirs so a prebuilt that
|
|
passed install validation runs the same at serve time. No-op off a ROCDXG
|
|
WSL host (needs /dev/dxg, "microsoft" /proc/version, librocdxg in /opt/rocm).
|
|
"""
|
|
try:
|
|
if not os.path.exists("/dev/dxg"):
|
|
return []
|
|
with open("/proc/version", encoding = "utf-8", errors = "replace") as fh:
|
|
if "microsoft" not in fh.read().lower():
|
|
return []
|
|
except OSError:
|
|
return []
|
|
out: "list[str]" = []
|
|
for d in ("/opt/rocm/lib", "/opt/rocm/lib64"):
|
|
if os.path.exists(os.path.join(d, "librocdxg.so")) or os.path.exists(
|
|
os.path.join(d, "librocdxg.so.1")
|
|
):
|
|
out.append(d)
|
|
return out
|
|
|
|
|
|
# ── Pre-compiled patterns for plan-without-action re-prompt ──
|
|
# Forward-looking intent signals: the model is describing what it *will*
|
|
# do rather than giving a final answer.
|
|
_INTENT_SIGNAL = re.compile(
|
|
r"(?i)("
|
|
# Direct intent ("I'll ...", "Let me ...", straight + curly apostrophes).
|
|
# Excludes "I can"/"I should"/"I want to"/"let's" (common in answers).
|
|
# Negative lookahead drops negated forms ("I will not") so a refusal
|
|
# doesn't trigger a re-prompt.
|
|
r"\b(i['\u2019](ll|m going to|m gonna)|i am (going to|gonna)|i will|i shall|let me|allow me)\b(?!\s+(?:not|never)\b)"
|
|
r"|"
|
|
# Step/plan framing: "First ...", "Step 1:", "Here's my plan"
|
|
r"\b(?:first\b|step \d+:?|here['\u2019]?s (?:my |the |a )?(?:plan|approach))"
|
|
r"|"
|
|
# "Now I" / "Next I" patterns
|
|
r"\b(?:now i|next i)\b"
|
|
r")"
|
|
)
|
|
_MAX_REPROMPTS = 1
|
|
|
|
# Without max_tokens, llama-server defaults n_predict = n_ctx (up to 262144 for
|
|
# Qwen3.5), causing many-minute zombie decodes when cancel fails.
|
|
# t_max_predict_ms is a wall-clock backstop but per the llama.cpp README only
|
|
# fires after a newline, so we keep a token cap as the front-line limiter.
|
|
# The cap is the effective context length when known, else this floor. 4096 was
|
|
# too low: Qwen3 / gpt-oss reasoning traces and max_tokens-omitting OpenAI-API
|
|
# callers (langchain, llama-index, curl) got truncated mid-sentence.
|
|
_DEFAULT_MAX_TOKENS_FLOOR = 32768
|
|
_DEFAULT_T_MAX_PREDICT_MS = 600_000 # 10 min
|
|
_REPROMPT_MAX_CHARS = 2000
|
|
_FORCED_REPEAT_PLAN_SIGNAL = re.compile(
|
|
r"\b(?:i\s+will|i'll|let\s+me|going\s+to|need\s+to|call|use|run|search|fetch|render)\b",
|
|
re.I,
|
|
)
|
|
_FINAL_ANSWER_SIGNAL = re.compile(
|
|
r"\b(?:final\s+answer|answer\s*:|here\s+is|here's|in\s+summary|result\s*:)\b",
|
|
re.I,
|
|
)
|
|
|
|
|
|
def _is_short_intent_without_action(text: str) -> bool:
|
|
stripped = text.strip()
|
|
return 0 < len(stripped) < _REPROMPT_MAX_CHARS and _INTENT_SIGNAL.search(stripped) is not None
|
|
|
|
|
|
def _should_suppress_forced_no_tool_output(text: str) -> bool:
|
|
"""Suppress only repeated forced-turn planning text, not final answers."""
|
|
stripped = text.strip()
|
|
if not stripped or len(stripped) >= _REPROMPT_MAX_CHARS:
|
|
return False
|
|
if _FINAL_ANSWER_SIGNAL.search(stripped):
|
|
return False
|
|
return _FORCED_REPEAT_PLAN_SIGNAL.search(stripped) is not None
|
|
|
|
|
|
# ── Pre-compiled patterns for GGUF shard detection ───────────
|
|
_SHARD_FULL_RE = re.compile(r"^(.*)-(\d{5})-of-(\d{5})\.gguf$")
|
|
_SHARD_RE = re.compile(r"^(.*)-\d{5}-of-\d{5}\.gguf$")
|
|
|
|
|
|
# ── Sliding-window-pattern resolver ───────────────────────────
|
|
# Resolves the per-layer SWA mask when a GGUF reports a sliding window but
|
|
# no `sliding_window_pattern` field. Tier order in `_resolve_swa_pattern`:
|
|
# GGUF metadata, on-disk cache, bootstrap dict below, transformers
|
|
# introspection, HF Hub config.json, legacy 1/4 fallback. Period N means
|
|
# layer i is SWA iff `(i + 1) % N != 0`, matching transformers. Skipped on
|
|
# purpose: phi3 (no key/val length in GGUF, window >= ctx anyway), qwen2
|
|
# family (converter strips sliding_window when use_sliding_window=False),
|
|
# mistral v0.1/v0.2 (all-SWA can't be a period).
|
|
_BOOTSTRAP_SWA_DEFAULTS: dict[str, int] = {
|
|
"gemma2": 2, # Gemma2Config.sliding_window_pattern
|
|
"gemma3": 6, # Gemma3TextConfig.sliding_window_pattern
|
|
"gemma3n": 5, # text_config.layer_types: SWA*4 + FULL
|
|
"gpt_oss": 2, # text_config.layer_types: alternating
|
|
"cohere2": 4, # Cohere2Config.sliding_window_pattern
|
|
}
|
|
|
|
# Process-wide cache backed by JSON on disk. Values are int period or
|
|
# list[bool] mask. Lazy-loaded.
|
|
_SWA_CACHE: Optional[dict] = None
|
|
_SWA_CACHE_LOCK = threading.Lock()
|
|
|
|
|
|
def _probe_dns_dead(host: str = "huggingface.co", timeout: float = 2.0) -> bool:
|
|
"""Quick DNS check on a daemon thread, so concurrent sockets aren't
|
|
affected by socket.setdefaulttimeout."""
|
|
result: list[Optional[bool]] = [None]
|
|
|
|
def _probe() -> None:
|
|
try:
|
|
socket.gethostbyname(host)
|
|
result[0] = False
|
|
except Exception:
|
|
result[0] = True
|
|
|
|
t = threading.Thread(target = _probe, daemon = True)
|
|
t.start()
|
|
t.join(timeout)
|
|
# Thread still running -> resolver wedged -> dead.
|
|
return True if result[0] is None else result[0]
|
|
|
|
|
|
@contextlib.contextmanager
|
|
def _hf_offline_if_dns_dead():
|
|
"""Set HF_HUB_OFFLINE for this block only when DNS to huggingface.co fails;
|
|
restores env on exit so a transient hiccup can't quarantine the process.
|
|
No-op if the user already set it."""
|
|
if "HF_HUB_OFFLINE" in os.environ:
|
|
yield False
|
|
return
|
|
if not _probe_dns_dead():
|
|
yield False
|
|
return
|
|
|
|
transformers_was_set = "TRANSFORMERS_OFFLINE" in os.environ
|
|
os.environ["HF_HUB_OFFLINE"] = "1"
|
|
if not transformers_was_set:
|
|
os.environ["TRANSFORMERS_OFFLINE"] = "1"
|
|
logger.warning("huggingface.co unreachable; using local HF cache for this load.")
|
|
try:
|
|
yield True
|
|
finally:
|
|
os.environ.pop("HF_HUB_OFFLINE", None)
|
|
if not transformers_was_set:
|
|
os.environ.pop("TRANSFORMERS_OFFLINE", None)
|
|
|
|
|
|
def _swa_cache_path() -> Path:
|
|
home = os.environ.get("UNSLOTH_STUDIO_HOME") or os.environ.get("STUDIO_HOME")
|
|
base = Path(home) if home else Path.home() / ".unsloth" / "studio"
|
|
return base / "swa_cache.json"
|
|
|
|
|
|
def _load_swa_cache() -> dict:
|
|
global _SWA_CACHE
|
|
with _SWA_CACHE_LOCK:
|
|
if _SWA_CACHE is not None:
|
|
return _SWA_CACHE
|
|
try:
|
|
with open(_swa_cache_path()) as f:
|
|
_SWA_CACHE = json.load(f)
|
|
if not isinstance(_SWA_CACHE, dict):
|
|
_SWA_CACHE = {}
|
|
except (FileNotFoundError, json.JSONDecodeError, OSError):
|
|
_SWA_CACHE = {}
|
|
return _SWA_CACHE
|
|
|
|
|
|
def _save_swa_cache(cache: dict) -> None:
|
|
try:
|
|
path = _swa_cache_path()
|
|
path.parent.mkdir(parents = True, exist_ok = True)
|
|
tmp = path.with_suffix(".json.tmp")
|
|
with open(tmp, "w") as f:
|
|
json.dump(cache, f, indent = 2, sort_keys = True)
|
|
tmp.replace(path)
|
|
except OSError:
|
|
pass
|
|
|
|
|
|
def _period_from_layer_types(layer_types: list) -> Optional[int]:
|
|
"""Smallest period N where `(i+1) % N != 0` matches the SWA mask, else None."""
|
|
if not layer_types:
|
|
return None
|
|
is_swa = ["full" not in str(t).lower() for t in layer_types]
|
|
n = len(is_swa)
|
|
for N in range(1, n + 1):
|
|
if all(((i + 1) % N != 0) == is_swa[i] for i in range(n)):
|
|
return N
|
|
return None
|
|
|
|
|
|
def _fetch_swa_entry_from_hf(repo_id: str) -> Optional[object]:
|
|
try:
|
|
from huggingface_hub import hf_hub_download
|
|
cfg_path = hf_hub_download(repo_id, "config.json", repo_type = "model")
|
|
with open(cfg_path) as f:
|
|
cfg = json.load(f)
|
|
except Exception:
|
|
return None
|
|
|
|
src = cfg.get("text_config") if isinstance(cfg.get("text_config"), dict) else cfg
|
|
period = src.get("sliding_window_pattern")
|
|
if isinstance(period, int) and period > 0:
|
|
return period
|
|
lt = src.get("layer_types")
|
|
if isinstance(lt, list) and lt:
|
|
return _period_from_layer_types(lt) or ["full" not in str(t).lower() for t in lt]
|
|
return None
|
|
|
|
|
|
def _arch_aliases(arch: str) -> tuple:
|
|
# GGUF emits `falcon-h1`; HF model_type is `falcon_h1`. Normalise both ways.
|
|
seen = []
|
|
for a in (arch, arch.replace("-", "_"), arch.replace("_", "-")):
|
|
if a and a not in seen:
|
|
seen.append(a)
|
|
return tuple(seen)
|
|
|
|
|
|
def _swa_entry_from_config_obj(cfg) -> Optional[object]:
|
|
src = getattr(cfg, "text_config", None) or cfg
|
|
period = getattr(src, "sliding_window_pattern", None)
|
|
if isinstance(period, int) and period > 0:
|
|
return period
|
|
lt = getattr(src, "layer_types", None)
|
|
if isinstance(lt, list) and lt:
|
|
return _period_from_layer_types(lt) or ["full" not in str(t).lower() for t in lt]
|
|
return None
|
|
|
|
|
|
_SWA_PATTERN_SOURCE_RE = re.compile(r"sliding_window_pattern\s*(?::\s*[\w\[\], ]*)?\s*=\s*(\d+)")
|
|
|
|
|
|
def _resolve_swa_entry_from_transformers(arch: str) -> Optional[object]:
|
|
"""Default-instantiate the matching Config; on failure, regex-parse its
|
|
source for `sliding_window_pattern = N`."""
|
|
try:
|
|
from transformers.models.auto.configuration_auto import (
|
|
CONFIG_MAPPING,
|
|
CONFIG_MAPPING_NAMES,
|
|
)
|
|
except Exception:
|
|
return None
|
|
|
|
cfg_class = None
|
|
for alias in _arch_aliases(arch):
|
|
if alias in CONFIG_MAPPING_NAMES:
|
|
try:
|
|
cfg_class = CONFIG_MAPPING[alias]
|
|
break
|
|
except Exception:
|
|
cfg_class = None
|
|
if cfg_class is None:
|
|
return None
|
|
|
|
try:
|
|
if (entry := _swa_entry_from_config_obj(cfg_class())) is not None:
|
|
return entry
|
|
except Exception:
|
|
pass
|
|
|
|
import inspect
|
|
|
|
candidates = [cfg_class]
|
|
text_cfg_class = getattr(cfg_class, "sub_configs", {}).get("text_config")
|
|
if text_cfg_class is not None:
|
|
candidates.append(text_cfg_class)
|
|
for cls in candidates:
|
|
try:
|
|
src = inspect.getsource(cls)
|
|
except (OSError, TypeError):
|
|
continue
|
|
if m := _SWA_PATTERN_SOURCE_RE.search(src):
|
|
period = int(m.group(1))
|
|
if period > 0:
|
|
return period
|
|
return None
|
|
|
|
|
|
def _resolve_swa_pattern(
|
|
arch: Optional[str],
|
|
n_layers: Optional[int],
|
|
source_repo_candidates: tuple = (),
|
|
*,
|
|
allow_network: Optional[bool] = None,
|
|
) -> Optional[list]:
|
|
if not arch or not n_layers:
|
|
return None
|
|
if allow_network is None:
|
|
allow_network = os.environ.get("UNSLOTH_STUDIO_OFFLINE", "0") not in (
|
|
"1",
|
|
"true",
|
|
"True",
|
|
"yes",
|
|
)
|
|
|
|
cache = _load_swa_cache()
|
|
|
|
def _entry_to_mask(entry):
|
|
if isinstance(entry, int) and entry > 0:
|
|
return [(i + 1) % entry != 0 for i in range(n_layers)]
|
|
if isinstance(entry, list) and entry:
|
|
return [bool(entry[i % len(entry)]) for i in range(n_layers)]
|
|
return None
|
|
|
|
def _persist(entry):
|
|
with _SWA_CACHE_LOCK:
|
|
cache[arch] = entry
|
|
_save_swa_cache(cache)
|
|
|
|
if (entry := cache.get(arch)) is not None:
|
|
if (mask := _entry_to_mask(entry)) is not None:
|
|
return mask
|
|
|
|
if (entry := _BOOTSTRAP_SWA_DEFAULTS.get(arch)) is not None:
|
|
return _entry_to_mask(entry)
|
|
|
|
entry = _resolve_swa_entry_from_transformers(arch)
|
|
if entry is not None:
|
|
_persist(entry)
|
|
return _entry_to_mask(entry)
|
|
|
|
# Tier 3: live HF fetch (result persistently cached)
|
|
if allow_network:
|
|
for repo_id in source_repo_candidates:
|
|
if not repo_id:
|
|
continue
|
|
entry = _fetch_swa_entry_from_hf(repo_id)
|
|
if entry is not None:
|
|
_persist(entry)
|
|
return _entry_to_mask(entry)
|
|
|
|
return None
|
|
|
|
|
|
def _hf_repo_from_url(url: Optional[str]) -> Optional[str]:
|
|
"""Strip `https://huggingface.co/owner/name(/...)` -> `owner/name`."""
|
|
if not url or "huggingface.co/" not in url:
|
|
return None
|
|
tail = url.split("huggingface.co/", 1)[1].rstrip("/")
|
|
parts = tail.split("/")
|
|
if len(parts) < 2:
|
|
return None
|
|
return f"{parts[0]}/{parts[1]}"
|
|
|
|
|
|
# Lazy import to avoid pulling transformers in at module level.
|
|
def _extract_model_size_b(model_id: str):
|
|
from utils.models import extract_model_size_b
|
|
return extract_model_size_b(model_id)
|
|
|
|
|
|
_TOOL_TEMPLATE_MARKERS = (
|
|
"{%- if tools %}",
|
|
"{%- if tools -%}",
|
|
"{% if tools %}",
|
|
"{% if tools -%}",
|
|
'"role" == "tool"',
|
|
"'role' == 'tool'",
|
|
'message.role == "tool"',
|
|
"message.role == 'tool'",
|
|
)
|
|
|
|
|
|
def detect_reasoning_flags(
|
|
chat_template: Optional[str],
|
|
model_identifier: Optional[str] = None,
|
|
*,
|
|
log_source: Optional[str] = None,
|
|
) -> dict:
|
|
"""Classify a chat template's reasoning and tool-calling capabilities.
|
|
|
|
Returns the same five keys as the GGUF sniffer: ``supports_reasoning``,
|
|
``reasoning_style`` (``"enable_thinking"`` | ``"reasoning_effort"``),
|
|
``reasoning_always_on``, ``supports_preserve_thinking``,
|
|
``supports_tools``. Used by both the llama-server backend at load time
|
|
and the safetensors/transformers paths in ``routes/inference`` so they
|
|
agree on what the frontend sees.
|
|
"""
|
|
flags = {
|
|
"supports_reasoning": False,
|
|
"reasoning_style": "enable_thinking",
|
|
"reasoning_always_on": False,
|
|
"supports_preserve_thinking": False,
|
|
"supports_tools": False,
|
|
}
|
|
if not chat_template:
|
|
return flags
|
|
tpl = chat_template
|
|
prefix = f"{log_source}: " if log_source else ""
|
|
|
|
if "enable_thinking" in tpl:
|
|
flags["supports_reasoning"] = True
|
|
flags["reasoning_style"] = "enable_thinking"
|
|
logger.info(f"{prefix}model supports reasoning (enable_thinking)")
|
|
elif "reasoning_effort" in tpl:
|
|
# gpt-oss / Harmony use reasoning_effort
|
|
# ("low" | "medium" | "high"), not a boolean.
|
|
flags["supports_reasoning"] = True
|
|
flags["reasoning_style"] = "reasoning_effort"
|
|
logger.info(f"{prefix}model supports reasoning (reasoning_effort)")
|
|
elif "thinking" in tpl:
|
|
# DeepSeek uses 'thinking', not 'enable_thinking'
|
|
normalized_id = (model_identifier or "").lower()
|
|
if "deepseek" in normalized_id:
|
|
flags["supports_reasoning"] = True
|
|
logger.info(f"{prefix}model supports reasoning (DeepSeek thinking)")
|
|
|
|
# Hardcoded <think> tags or reasoning_content in the template mean
|
|
# thinking is always on (no toggle).
|
|
if not flags["supports_reasoning"]:
|
|
if ("<think>" in tpl and "</think>" in tpl) or "reasoning_content" in tpl:
|
|
flags["supports_reasoning"] = True
|
|
flags["reasoning_always_on"] = True
|
|
logger.info(f"{prefix}model always reasons (<think> tags in template)")
|
|
|
|
# preserve_thinking: independent kwarg on some Qwen templates that
|
|
# keeps historical <think> blocks in prior assistant turns.
|
|
if "preserve_thinking" in tpl:
|
|
flags["supports_preserve_thinking"] = True
|
|
logger.info(f"{prefix}model supports preserve_thinking")
|
|
|
|
if any(marker in tpl for marker in _TOOL_TEMPLATE_MARKERS):
|
|
flags["supports_tools"] = True
|
|
logger.info(f"{prefix}model supports tool calling")
|
|
|
|
return flags
|
|
|
|
|
|
def _is_mtp_model_name(model_identifier: Optional[str], gguf_path: Optional[str] = None) -> bool:
|
|
"""Name-based MTP detector. Fallback for the metadata signal."""
|
|
for cand in (model_identifier, Path(gguf_path).name if gguf_path else None):
|
|
if cand and "-mtp" in cand.lower():
|
|
return True
|
|
return False
|
|
|
|
|
|
def _is_companion_gguf_path(path: str) -> bool:
|
|
"""True for a non-main GGUF: vision mmproj or a separate MTP drafter
|
|
(repo-root ``mtp-*.gguf`` or the ``MTP/`` subdir copies, Gemma 4).
|
|
|
|
Mirrors hub.utils.gguf so variant resolution never picks a companion as
|
|
the main model -- e.g. a Gemma ``Q8_0`` request must not resolve to the
|
|
``MTP/...-Q8_0-MTP.gguf`` drafter, which sorts ahead of the real weight.
|
|
"""
|
|
p = path.lower()
|
|
if not p.endswith(".gguf"):
|
|
return False
|
|
if "mmproj" in p:
|
|
return True
|
|
name = p.rsplit("/", 1)[-1]
|
|
return name.startswith("mtp-") or "/mtp/" in f"/{p}"
|
|
|
|
|
|
# Below this many B params, draft-mtp regresses vs spec-off (bench in
|
|
# _build_speculative_flags); auto mode drops MTP under it.
|
|
_MTP_MIN_SIZE_B = 3.0
|
|
|
|
# Context-fit VRAM budget: tighter than _GPU_PIN_VRAM_FRACTION (0.95) on
|
|
# purpose -- over-promising context OOMs at runtime (#5106).
|
|
_CTX_FIT_VRAM_FRACTION = 0.90
|
|
|
|
# Extra VRAM fraction reserved when MTP will engage: the draft model's
|
|
# weights, KV cache, and compute buffers live outside the main model's
|
|
# estimate. Applied to BOTH the ctx-fit budget and the GPU pin thresholds --
|
|
# tightening only the fit lets a load whose weights land between the two
|
|
# fractions pin without any room for the drafter.
|
|
_MTP_VRAM_RESERVE_FRAC = 0.05
|
|
|
|
|
|
def _auto_mode_drops_mtp(
|
|
req_mode: Optional[str],
|
|
size_b: Optional[float],
|
|
*,
|
|
has_separate_drafter: bool = False,
|
|
) -> bool:
|
|
"""Auto mode drops MTP below _MTP_MIN_SIZE_B for an embedded draft head
|
|
(its per-token cost regresses there); a separate drafter (Gemma) is a tiny
|
|
standalone model that still speeds up below 3B, so it never drops. Forced
|
|
mtp / mtp+ngram engage regardless of size."""
|
|
if has_separate_drafter:
|
|
return False
|
|
return req_mode == "auto" and size_b is not None and size_b < _MTP_MIN_SIZE_B
|
|
|
|
|
|
def _extra_args_set_spec_type(extra_args: Optional[Iterable[str]]) -> bool:
|
|
"""User passed --spec-type / --spec-default? llama-server takes one
|
|
--spec-type (comma-separated to chain), so suppress auto-emit."""
|
|
if not extra_args:
|
|
return False
|
|
for raw in extra_args:
|
|
tok = str(raw)
|
|
if not tok.startswith("--"):
|
|
continue
|
|
flag = tok.split("=", 1)[0]
|
|
if flag in ("--spec-type", "--spec-default"):
|
|
return True
|
|
return False
|
|
|
|
|
|
def _build_ngram_mod_flags(
|
|
caps: Optional[dict],
|
|
n_match: int = 24,
|
|
n_min: int = 48,
|
|
n_max: int = 64,
|
|
) -> list[str]:
|
|
"""Emit the right ngram-mod knob flags for the running llama-server.
|
|
|
|
Post-rename builds expose ``--spec-ngram-mod-n-{match,min,max}``;
|
|
pre-rename builds expose legacy ``--spec-ngram-size-n`` /
|
|
``--draft-min`` / ``--draft-max``. ``caps`` comes from
|
|
``probe_server_capabilities``; ``ngram_mod_flavor`` says which set is
|
|
real (vs a removal-stub). Returns ``[]`` when neither is available so
|
|
the caller can drop ngram-mod entirely.
|
|
"""
|
|
flavor = caps.get("ngram_mod_flavor") if caps else None
|
|
if flavor == "new":
|
|
return [
|
|
"--spec-ngram-mod-n-match",
|
|
str(n_match),
|
|
"--spec-ngram-mod-n-min",
|
|
str(n_min),
|
|
"--spec-ngram-mod-n-max",
|
|
str(n_max),
|
|
]
|
|
if flavor == "legacy":
|
|
# Pre-rename llama.cpp: same knobs lived under --spec-ngram-size-n
|
|
# (lookup length) and generic --draft-min / --draft-max (N range).
|
|
return [
|
|
"--spec-ngram-size-n",
|
|
str(n_match),
|
|
"--draft-min",
|
|
str(n_min),
|
|
"--draft-max",
|
|
str(n_max),
|
|
]
|
|
return []
|
|
|
|
|
|
# Canonical Speculative Decoding modes exposed by the Studio chat UI.
|
|
# Dropdown renders five (auto, mtp, ngram, mtp+ngram, off); the load API
|
|
# also accepts legacy values the original Switch and external callers emit
|
|
# (default, draft-mtp, ngram-mod, ngram-simple).
|
|
_CANONICAL_SPEC_MODES = {"auto", "mtp", "ngram", "mtp+ngram", "off", "ngram-simple"}
|
|
_LEGACY_SPEC_MODE_MAP = {
|
|
"default": "auto",
|
|
"draft-mtp": "mtp",
|
|
"ngram-mod": "ngram",
|
|
}
|
|
|
|
|
|
def _canonicalize_spec_mode(value):
|
|
"""Map any accepted ``speculative_type`` input onto a canonical mode.
|
|
|
|
Returns ``auto``, ``mtp``, ``ngram``, ``mtp+ngram``, ``off``,
|
|
``ngram-simple``, or ``None`` (callers treat ``None`` as ``auto``).
|
|
Unknown strings collapse to ``auto`` so a stale UI value or typo falls
|
|
back to the safe platform-aware path.
|
|
"""
|
|
if value is None:
|
|
return None
|
|
if not isinstance(value, str):
|
|
return None
|
|
stripped = value.strip().lower()
|
|
if not stripped:
|
|
return None
|
|
if stripped in _CANONICAL_SPEC_MODES:
|
|
return stripped
|
|
if stripped in _LEGACY_SPEC_MODE_MAP:
|
|
return _LEGACY_SPEC_MODE_MAP[stripped]
|
|
# Old persisted state emits llama.cpp comma-chains e.g.
|
|
# "ngram-mod,draft-mtp"; collapse the most common one explicitly.
|
|
pieces = [p.strip() for p in stripped.split(",") if p.strip()]
|
|
has_mtp = any(p in ("mtp", "draft-mtp") for p in pieces)
|
|
has_ngram = any(p in ("ngram", "ngram-mod") for p in pieces)
|
|
if has_mtp and has_ngram:
|
|
return "mtp+ngram"
|
|
if has_mtp:
|
|
return "mtp"
|
|
if has_ngram:
|
|
return "ngram"
|
|
return "auto"
|
|
|
|
|
|
def _backfill_usage_from_timings(usage, timings):
|
|
"""Synthesize ``usage`` from llama-server's ``timings`` when the
|
|
OpenAI-style usage block is missing or reports zero tokens.
|
|
|
|
The Studio chat UI computes generation t/s from
|
|
``meta.usage.completion_tokens / totalStreamTime``. llama-server always
|
|
populates ``timings.predicted_n`` (true decoded count) and
|
|
``timings.prompt_n``, but the final SSE chunk's ``usage`` can be absent
|
|
or zero on some server builds / streaming configs, making the UI fall
|
|
back to wall-clock t/s and dilute speculative-decoding speedups.
|
|
"""
|
|
if not timings:
|
|
return usage
|
|
if usage and usage.get("completion_tokens"):
|
|
return usage
|
|
predicted_n = timings.get("predicted_n")
|
|
prompt_n = timings.get("prompt_n")
|
|
if predicted_n is None and prompt_n is None:
|
|
return usage
|
|
out = dict(usage or {})
|
|
if not out.get("completion_tokens") and predicted_n is not None:
|
|
out["completion_tokens"] = predicted_n
|
|
if not out.get("prompt_tokens") and prompt_n is not None:
|
|
out["prompt_tokens"] = prompt_n
|
|
out["total_tokens"] = int(out.get("prompt_tokens") or 0) + int(
|
|
out.get("completion_tokens") or 0
|
|
)
|
|
return out
|
|
|
|
|
|
class LlamaCppBackend:
|
|
"""Manages a llama-server subprocess for GGUF model inference.
|
|
|
|
Lifecycle:
|
|
1. load_model() — start llama-server with the GGUF file
|
|
2. generate_chat_completion() — proxy to /v1/chat/completions, stream back
|
|
3. unload_model() — terminate the subprocess
|
|
"""
|
|
|
|
def __init__(self):
|
|
self._process: Optional[subprocess.Popen] = None
|
|
self._port: Optional[int] = None
|
|
self._model_identifier: Optional[str] = None
|
|
self._gguf_path: Optional[str] = None
|
|
self._hf_repo: Optional[str] = None
|
|
# Separate MTP drafter launched with the current model; reload-dedup
|
|
# key so a drafter that appears next to the weights forces a reload.
|
|
self._mtp_draft_path: Optional[str] = None
|
|
# Why MTP was disabled on the last load that asked for it (auto on an
|
|
# MTP model, or forced mtp / mtp+ngram), else None. Drives the "update
|
|
# llama.cpp" hint in the UI. "binary_no_mtp" / "binary_outdated" ->
|
|
# a newer prebuilt would help; "runtime_error" -> it may not.
|
|
self._spec_fallback_reason: Optional[str] = None
|
|
self._hf_variant: Optional[str] = None
|
|
self._is_vision: bool = False
|
|
# Block-diffusion model (e.g. DiffusionGemma): served by the diffusion
|
|
# runner, not llama-server. Set from the GGUF architecture at load.
|
|
self._architecture: Optional[str] = None
|
|
self._is_diffusion: bool = False
|
|
self._diffusion_visual_bin: Optional[str] = None
|
|
self._healthy = False
|
|
# Set by _classify_gpu_offload after _wait_for_health.
|
|
self._gpu_offload_active: Optional[bool] = None
|
|
self._context_length: Optional[int] = None
|
|
self._effective_context_length: Optional[int] = None
|
|
self._max_context_length: Optional[int] = None
|
|
self._chat_template: Optional[str] = None
|
|
self._chat_template_override: Optional[str] = None
|
|
self._supports_reasoning: bool = False
|
|
self._reasoning_always_on: bool = False
|
|
self._reasoning_style: str = "enable_thinking"
|
|
self._supports_preserve_thinking: bool = False
|
|
self._supports_tools: bool = False
|
|
self._cache_type_kv: Optional[str] = None
|
|
# Whether --split-mode tensor was applied on the active load.
|
|
self._tensor_parallel: bool = False
|
|
self._reasoning_default: bool = True
|
|
self._speculative_type: Optional[str] = None
|
|
# Canonical UI-facing mode the user requested
|
|
# (auto/mtp/ngram/mtp+ngram/off/ngram-simple). Round-tripped through the
|
|
# status API so the dropdown reflects the picked mode, not the resolved
|
|
# flag set (auto on a 27B MTP GGUF resolves to draft-mtp but reads "Auto").
|
|
self._requested_spec_mode: Optional[str] = None
|
|
# User --spec-draft-n-max override (None = platform default).
|
|
self._spec_draft_n_max: Optional[int] = None
|
|
# KV-cache estimation fields (populated by _read_gguf_metadata)
|
|
self._n_layers: Optional[int] = None
|
|
self._n_kv_heads: Optional[int] = None
|
|
self._n_kv_heads_by_layer: Optional[list[int]] = None
|
|
self._n_heads: Optional[int] = None
|
|
self._embedding_length: Optional[int] = None
|
|
# Architecture-aware KV fields for 5-path estimation
|
|
self._kv_key_length: Optional[int] = None
|
|
self._kv_value_length: Optional[int] = None
|
|
self._sliding_window: Optional[int] = None
|
|
self._sliding_window_pattern: Optional[list[bool]] = None
|
|
self._full_attention_interval: Optional[int] = None
|
|
self._kv_lora_rank: Optional[int] = None
|
|
self._key_length_mla: Optional[int] = None
|
|
self._kv_key_length_swa: Optional[int] = None
|
|
self._kv_value_length_swa: Optional[int] = None
|
|
self._ssm_inner_size: Optional[int] = None
|
|
self._ssm_state_size: Optional[int] = None
|
|
# Last N layers reuse earlier layers' KV and don't allocate their own
|
|
# cache (Gemma 3n / Gemma 4: <arch>.attention.shared_kv_layers).
|
|
self._shared_kv_layers: Optional[int] = None
|
|
# MTP head count (llama.cpp #22673); >0 enables --spec-type draft-mtp.
|
|
self._nextn_predict_layers: Optional[int] = None
|
|
self._lock = threading.Lock()
|
|
# Wraps load_model() end-to-end so concurrent loads serialise and never
|
|
# coexist as two llama-server processes (#5401).
|
|
self._serial_load_lock = threading.Lock()
|
|
# Set by the in-app updater while it swaps prebuilt binaries; load_model()
|
|
# rejects fast so no server starts from a half-swapped binary.
|
|
self._llama_update_in_progress = False
|
|
# Last extra_args / requested n_ctx, preserved across unload so the chat
|
|
# UI's /unload+/load Apply path can inherit them (#5401).
|
|
# ``_extra_args_source`` records the (model_identifier, hf_variant) the
|
|
# stored args came from so the route can refuse cross-model inheritance.
|
|
self._extra_args: Optional[List[str]] = None
|
|
self._extra_args_source: Optional[tuple[str, Optional[str]]] = None
|
|
self._requested_n_ctx: int = 0
|
|
self._stdout_lines: list[str] = []
|
|
self._stdout_thread: Optional[threading.Thread] = None
|
|
# llama-server tee log (see _drain_stdout / _kill_process).
|
|
self._llama_log_fh = None
|
|
self._llama_log_path: Optional[Path] = None
|
|
self._cancel_event = threading.Event()
|
|
self._api_key: Optional[str] = None
|
|
# True once a probe has completed; cleared on transient failure.
|
|
self._is_audio: bool = False
|
|
self._audio_type: Optional[str] = None
|
|
self._audio_probed: bool = False
|
|
# Audio INPUT capability (distinct from _is_audio, which is TTS output).
|
|
self._has_audio_input: bool = False
|
|
self._mmproj_has_audio: bool = False # clip.has_audio_encoder, set at load
|
|
# Monotonic timestamp set in _kill_process; read by load_model
|
|
# to decide whether to wait for the VRAM reclaim to finish.
|
|
self._last_kill_monotonic: float = 0.0
|
|
|
|
self._kill_orphaned_servers()
|
|
atexit.register(self._cleanup)
|
|
|
|
# ── Properties ────────────────────────────────────────────────
|
|
|
|
@property
|
|
def is_loaded(self) -> bool:
|
|
return self._process is not None and self._healthy
|
|
|
|
@property
|
|
def is_active(self) -> bool:
|
|
"""True if a llama-server process exists (loading or loaded)."""
|
|
return self._process is not None
|
|
|
|
@property
|
|
def base_url(self) -> str:
|
|
return f"http://127.0.0.1:{self._port}"
|
|
|
|
@property
|
|
def model_identifier(self) -> Optional[str]:
|
|
return self._model_identifier
|
|
|
|
@property
|
|
def is_vision(self) -> bool:
|
|
return self._is_vision
|
|
|
|
@property
|
|
def is_diffusion(self) -> bool:
|
|
"""True when the loaded GGUF is a block-diffusion model (DiffusionGemma)."""
|
|
return self._is_diffusion
|
|
|
|
@property
|
|
def hf_variant(self) -> Optional[str]:
|
|
return self._hf_variant
|
|
|
|
@property
|
|
def gguf_path(self) -> Optional[str]:
|
|
return self._gguf_path
|
|
|
|
@property
|
|
def mtp_draft_path(self) -> Optional[str]:
|
|
return self._mtp_draft_path
|
|
|
|
@property
|
|
def spec_fallback_reason(self) -> Optional[str]:
|
|
"""Why MTP was disabled on the last MTP-requesting load, else None."""
|
|
return self._spec_fallback_reason
|
|
|
|
@property
|
|
def extra_args(self) -> Optional[List[str]]:
|
|
"""Extra llama-server flags from the last load (a copy). None =
|
|
never set, [] = explicitly cleared. Used by the route for
|
|
inheritance."""
|
|
return list(self._extra_args) if self._extra_args is not None else None
|
|
|
|
@property
|
|
def requested_n_ctx(self) -> int:
|
|
"""n_ctx the last load was invoked with (not the effective cap).
|
|
0 means Auto. Used by the route to detect Auto-vs-explicit flips."""
|
|
return self._requested_n_ctx
|
|
|
|
@property
|
|
def extra_args_source(self) -> Optional[tuple[str, Optional[str]]]:
|
|
"""(model_identifier, hf_variant) the stored extra_args came from.
|
|
``None`` if no extras have ever been recorded. Used by the route
|
|
to refuse cross-model inheritance (#5401)."""
|
|
return self._extra_args_source
|
|
|
|
@property
|
|
def context_length(self) -> Optional[int]:
|
|
"""Return the effective context length the server is running at."""
|
|
return self._effective_context_length or self._context_length
|
|
|
|
@property
|
|
def max_context_length(self) -> Optional[int]:
|
|
"""Return the largest context that fits on this hardware at load time.
|
|
|
|
The UI's "safe zone" warning threshold: the ``_fit_context_to_vram``
|
|
binary-search cap for the best GPU subset, or the 4096 fallback if the
|
|
weights exceed 90% of every subset. The slider ceiling is
|
|
``native_context_length``; dragging above this triggers the warning.
|
|
"""
|
|
return self._max_context_length or self._context_length
|
|
|
|
@property
|
|
def native_context_length(self) -> Optional[int]:
|
|
"""Return the model's native context length from GGUF metadata."""
|
|
return self._context_length
|
|
|
|
def load_progress(self) -> Optional[dict]:
|
|
"""Return live model-load progress, or None if not loading.
|
|
|
|
During warm-up llama-server mmaps weight shards into page cache before
|
|
pushing layers to VRAM, a window where status only reports ``loading``
|
|
and the UI spinner looks stuck for minutes on large MoEs. Samples
|
|
``/proc/<pid>/status VmRSS`` against the sum of GGUF shard sizes for a
|
|
real progress bar. Returns ``None`` when no load is in flight.
|
|
|
|
Shape::
|
|
|
|
{
|
|
"phase": "mmap" | "ready",
|
|
"bytes_loaded": int, # VmRSS of the llama-server
|
|
"bytes_total": int, # sum of shard file sizes
|
|
"fraction": float, # bytes_loaded / bytes_total, 0..1
|
|
}
|
|
|
|
Linux-only; returns ``None`` where ``/proc/<pid>/status`` is unavailable.
|
|
"""
|
|
proc = self._process
|
|
if proc is None:
|
|
return None
|
|
pid = proc.pid
|
|
if pid is None:
|
|
return None
|
|
|
|
# Sum shard sizes (primary + any extras alongside).
|
|
bytes_total = 0
|
|
gguf_path = self._gguf_path
|
|
if gguf_path:
|
|
primary = Path(gguf_path)
|
|
try:
|
|
if primary.is_file():
|
|
bytes_total += primary.stat().st_size
|
|
except OSError:
|
|
pass
|
|
# Extra shards share the primary's prefix before the shard index.
|
|
try:
|
|
parent = primary.parent
|
|
stem = primary.name
|
|
m = _SHARD_RE.match(stem)
|
|
prefix = m.group(1) if m else None
|
|
if prefix and parent.is_dir():
|
|
for sibling in parent.iterdir():
|
|
if (
|
|
sibling.is_file()
|
|
and sibling.name.startswith(prefix)
|
|
and sibling.name != stem
|
|
and sibling.suffix == ".gguf"
|
|
):
|
|
try:
|
|
bytes_total += sibling.stat().st_size
|
|
except OSError:
|
|
pass
|
|
except OSError:
|
|
pass
|
|
|
|
# Read VmRSS from /proc/<pid>/status (kilobytes on Linux).
|
|
bytes_loaded = 0
|
|
try:
|
|
with open(f"/proc/{pid}/status", "r", encoding = "utf-8") as f:
|
|
for line in f:
|
|
if line.startswith("VmRSS:"):
|
|
kb = int(line.split()[1])
|
|
bytes_loaded = kb * 1024
|
|
break
|
|
except (FileNotFoundError, PermissionError, ValueError, OSError):
|
|
return None
|
|
|
|
phase = "ready" if self._healthy else "mmap"
|
|
fraction = 0.0
|
|
if bytes_total > 0:
|
|
fraction = min(1.0, bytes_loaded / bytes_total)
|
|
return {
|
|
"phase": phase,
|
|
"bytes_loaded": bytes_loaded,
|
|
"bytes_total": bytes_total,
|
|
"fraction": round(fraction, 4),
|
|
}
|
|
|
|
@property
|
|
def chat_template(self) -> Optional[str]:
|
|
return self._chat_template
|
|
|
|
@property
|
|
def chat_template_override(self) -> Optional[str]:
|
|
return self._chat_template_override
|
|
|
|
@property
|
|
def supports_reasoning(self) -> bool:
|
|
return self._supports_reasoning
|
|
|
|
@property
|
|
def reasoning_always_on(self) -> bool:
|
|
return self._reasoning_always_on
|
|
|
|
@property
|
|
def reasoning_style(self) -> str:
|
|
return self._reasoning_style
|
|
|
|
@property
|
|
def supports_preserve_thinking(self) -> bool:
|
|
return self._supports_preserve_thinking
|
|
|
|
@property
|
|
def reasoning_default(self) -> bool:
|
|
return self._reasoning_default
|
|
|
|
def _reasoning_kwargs(self, enable_thinking: bool) -> dict:
|
|
if self._reasoning_style == "reasoning_effort":
|
|
return {"reasoning_effort": "high" if enable_thinking else "low"}
|
|
return {"enable_thinking": enable_thinking}
|
|
|
|
def _request_reasoning_kwargs(
|
|
self,
|
|
enable_thinking: Optional[bool],
|
|
reasoning_effort: Optional[str] = None,
|
|
preserve_thinking: Optional[bool] = None,
|
|
) -> Optional[dict]:
|
|
"""Build chat_template_kwargs from per-request reasoning fields.
|
|
|
|
Merges the active model's reasoning style (``enable_thinking`` or
|
|
``reasoning_effort``) plus the independent ``preserve_thinking``
|
|
kwarg when the template supports it.
|
|
"""
|
|
kwargs: dict = {}
|
|
# Always-on reasoning models hardcode <think> tags and don't consume
|
|
# enable_thinking / reasoning_effort -- skip.
|
|
if self._supports_reasoning and not self._reasoning_always_on:
|
|
if self._reasoning_style == "reasoning_effort":
|
|
if reasoning_effort in ("none", "low", "medium", "high"):
|
|
kwargs["reasoning_effort"] = reasoning_effort
|
|
elif reasoning_effort == "minimal":
|
|
kwargs["reasoning_effort"] = "low"
|
|
elif enable_thinking is not None:
|
|
kwargs["reasoning_effort"] = "high" if enable_thinking else "low"
|
|
else:
|
|
if enable_thinking is not None:
|
|
kwargs["enable_thinking"] = enable_thinking
|
|
if self._supports_preserve_thinking and preserve_thinking is not None:
|
|
kwargs["preserve_thinking"] = preserve_thinking
|
|
return kwargs or None
|
|
|
|
@property
|
|
def supports_tools(self) -> bool:
|
|
return self._supports_tools
|
|
|
|
@property
|
|
def cache_type_kv(self) -> Optional[str]:
|
|
return self._cache_type_kv
|
|
|
|
@property
|
|
def tensor_parallel(self) -> bool:
|
|
"""Whether --split-mode tensor is active on the loaded server."""
|
|
return self._tensor_parallel
|
|
|
|
@property
|
|
def speculative_type(self) -> Optional[str]:
|
|
return self._speculative_type
|
|
|
|
@property
|
|
def requested_spec_mode(self) -> Optional[str]:
|
|
"""Canonical UI-facing mode the user requested (see field doc)."""
|
|
return self._requested_spec_mode
|
|
|
|
@property
|
|
def spec_draft_n_max(self) -> Optional[int]:
|
|
"""User --spec-draft-n-max override active on the load, or None when
|
|
the platform default (6 GPU / 3 CPU) is in effect."""
|
|
return self._spec_draft_n_max
|
|
|
|
# ── Binary discovery ──────────────────────────────────────────
|
|
|
|
@staticmethod
|
|
def _find_llama_server_binary() -> Optional[str]:
|
|
"""
|
|
Locate the llama-server binary.
|
|
|
|
Search order:
|
|
1. LLAMA_SERVER_PATH environment variable (direct path to binary)
|
|
1b. UNSLOTH_LLAMA_CPP_PATH env var (custom llama.cpp install dir)
|
|
2. ~/.unsloth/llama.cpp/llama-server (make build, root dir)
|
|
3. ~/.unsloth/llama.cpp/build/bin/llama-server (cmake build, Linux)
|
|
4. ~/.unsloth/llama.cpp/build/bin/Release/llama-server.exe (cmake build, Windows)
|
|
5. ./llama.cpp/llama-server (legacy: make build, root dir)
|
|
6. ./llama.cpp/build/bin/llama-server (legacy: cmake in-tree build)
|
|
7. llama-server on PATH (system install)
|
|
8. ./bin/llama-server (legacy: extracted binary)
|
|
"""
|
|
binary_name = "llama-server.exe" if sys.platform == "win32" else "llama-server"
|
|
|
|
# 1. Env var: direct path to binary
|
|
env_path = os.environ.get("LLAMA_SERVER_PATH")
|
|
if env_path and Path(env_path).is_file():
|
|
return env_path
|
|
|
|
# 1b. UNSLOTH_LLAMA_CPP_PATH: custom llama.cpp install dir
|
|
custom_llama_cpp = os.environ.get("UNSLOTH_LLAMA_CPP_PATH")
|
|
if custom_llama_cpp:
|
|
custom_dir = Path(custom_llama_cpp)
|
|
# Root dir (make builds)
|
|
root_bin = custom_dir / binary_name
|
|
if root_bin.is_file():
|
|
return str(root_bin)
|
|
# build/bin/ (cmake on Linux)
|
|
cmake_bin = custom_dir / "build" / "bin" / binary_name
|
|
if cmake_bin.is_file():
|
|
return str(cmake_bin)
|
|
# build/bin/Release/ (cmake on Windows)
|
|
if sys.platform == "win32":
|
|
win_bin = custom_dir / "build" / "bin" / "Release" / binary_name
|
|
if win_bin.is_file():
|
|
return str(win_bin)
|
|
|
|
# 2-4. Match installer layout: env-mode -> $STUDIO_HOME/llama.cpp;
|
|
# default/HOME-redirect -> ~/.unsloth/llama.cpp (sibling of studio).
|
|
legacy_llama = Path.home() / ".unsloth" / "llama.cpp"
|
|
try:
|
|
from utils.paths.storage_roots import studio_root as _sr # noqa: WPS433
|
|
|
|
_resolved_sr = _sr()
|
|
_legacy_studio = Path.home() / ".unsloth" / "studio"
|
|
try:
|
|
_is_legacy = _resolved_sr.resolve() == _legacy_studio.resolve()
|
|
except (OSError, ValueError):
|
|
_is_legacy = _resolved_sr == _legacy_studio
|
|
if _is_legacy:
|
|
search_roots = [legacy_llama]
|
|
else:
|
|
# _kill_orphaned_servers excludes the legacy root in custom
|
|
# mode; discovery must match so we never spawn a server we
|
|
# then refuse to clean up. UNSLOTH_LLAMA_CPP_PATH (handled
|
|
# earlier) is the explicit way to share a build across roots.
|
|
search_roots = [_resolved_sr / "llama.cpp"]
|
|
except (ImportError, OSError, ValueError):
|
|
search_roots = [legacy_llama]
|
|
_seen_roots: set[str] = set()
|
|
_unique_roots: list[Path] = []
|
|
for r in search_roots:
|
|
k = str(r)
|
|
if k not in _seen_roots:
|
|
_seen_roots.add(k)
|
|
_unique_roots.append(r)
|
|
for unsloth_home in _unique_roots:
|
|
home_root = unsloth_home / binary_name
|
|
if home_root.is_file():
|
|
return str(home_root)
|
|
home_linux = unsloth_home / "build" / "bin" / binary_name
|
|
if home_linux.is_file():
|
|
return str(home_linux)
|
|
if sys.platform == "win32":
|
|
home_win = unsloth_home / "build" / "bin" / "Release" / binary_name
|
|
if home_win.is_file():
|
|
return str(home_win)
|
|
|
|
# 5-6. Legacy: in-tree build (older setup.sh / setup.ps1)
|
|
project_root = Path(__file__).resolve().parents[4]
|
|
# Root dir (make builds)
|
|
root_path = project_root / "llama.cpp" / binary_name
|
|
if root_path.is_file():
|
|
return str(root_path)
|
|
# build/bin/ (cmake builds)
|
|
build_path = project_root / "llama.cpp" / "build" / "bin" / binary_name
|
|
if build_path.is_file():
|
|
return str(build_path)
|
|
if sys.platform == "win32":
|
|
win_path = project_root / "llama.cpp" / "build" / "bin" / "Release" / binary_name
|
|
if win_path.is_file():
|
|
return str(win_path)
|
|
|
|
# 7. System PATH
|
|
system_path = shutil.which("llama-server")
|
|
if system_path:
|
|
return system_path
|
|
|
|
# 8. Legacy: extracted to bin/
|
|
bin_path = project_root / "bin" / binary_name
|
|
if bin_path.is_file():
|
|
return str(bin_path)
|
|
|
|
return None
|
|
|
|
# ── llama-server capability probe ─────────────────────────────
|
|
|
|
# Cached on (path, mtime); `unsloth studio update` bumps mtime.
|
|
_capability_cache: dict[tuple[str, int], dict[str, object]] = {}
|
|
|
|
@classmethod
|
|
def probe_server_capabilities(cls, binary: Optional[str] = None) -> dict[str, object]:
|
|
"""Parse `llama-server --help` for feature flags. Returns
|
|
{found, mtp_token, supports_mtp, ngram_mod_flavor,
|
|
supports_ngram_mod, spec_draft_n_max_flag}.
|
|
|
|
``ngram_mod_flavor``: ``"new"`` when the post-rename
|
|
``--spec-ngram-mod-n-match / -n-min / -n-max`` are real args;
|
|
``"legacy"`` when only the pre-rename
|
|
``--spec-ngram-size-n / --draft-min / --draft-max`` are real (the
|
|
rename ships stub removal entries for legacy names, told apart by
|
|
the "argument has been removed" description); ``None`` if neither
|
|
set is usable.
|
|
|
|
``spec_draft_n_max_flag``: the flag the binary accepts --
|
|
``--spec-draft-n-max`` post-rename, ``--draft-max`` on legacy.
|
|
``None`` means n_max cannot be set.
|
|
"""
|
|
bin_path = binary or cls._find_llama_server_binary()
|
|
if not bin_path or not Path(bin_path).is_file():
|
|
return {
|
|
"found": False,
|
|
"mtp_token": None,
|
|
"supports_mtp": False,
|
|
"ngram_mod_flavor": None,
|
|
"supports_ngram_mod": False,
|
|
"spec_draft_n_max_flag": None,
|
|
"supports_kv_unified": False,
|
|
"supports_fit_ctx": False,
|
|
}
|
|
try:
|
|
mtime = int(Path(bin_path).stat().st_mtime)
|
|
except OSError:
|
|
mtime = 0
|
|
cache_key = (bin_path, mtime)
|
|
cached = cls._capability_cache.get(cache_key)
|
|
if cached is not None:
|
|
return cached
|
|
|
|
mtp_token: Optional[str] = None
|
|
ngram_mod_flavor: Optional[str] = None
|
|
spec_draft_n_max_flag: Optional[str] = None
|
|
supports_kv_unified = False
|
|
supports_fit_ctx = False
|
|
try:
|
|
result = subprocess.run(
|
|
[bin_path, "--help"],
|
|
capture_output = True,
|
|
text = True,
|
|
timeout = 10,
|
|
check = False,
|
|
)
|
|
help_text = (result.stdout or "") + "\n" + (result.stderr or "")
|
|
# Split into per-flag blocks (each --flag line + its indented
|
|
# continuation), so the "argument has been removed" description
|
|
# sits with its flag.
|
|
blocks: dict[str, str] = {}
|
|
current_flags: list[str] = []
|
|
current_desc: list[str] = []
|
|
for line in help_text.splitlines():
|
|
stripped = line.strip()
|
|
if stripped.startswith("-") and not line.startswith(" "):
|
|
# New flag line; flush previous.
|
|
if current_flags:
|
|
desc = " ".join(current_desc)
|
|
for f in current_flags:
|
|
blocks[f] = desc
|
|
current_flags = []
|
|
current_desc = [stripped]
|
|
# Extract long-form flag tokens from the DECLARATION
|
|
# prefix only (comma-separated aliases). Stop at the
|
|
# first non-flag token so flag references inside
|
|
# descriptions are ignored.
|
|
for tok in re.split(r"[,\s]+", stripped):
|
|
if tok.startswith("--") and re.match(r"--[A-Za-z][A-Za-z0-9_-]*$", tok):
|
|
current_flags.append(tok)
|
|
elif tok.startswith("-") and len(tok) > 1:
|
|
# short alias like -fa; keep scanning aliases.
|
|
continue
|
|
else:
|
|
# First non-flag token marks end of decl.
|
|
break
|
|
else:
|
|
current_desc.append(stripped)
|
|
if current_flags:
|
|
desc = " ".join(current_desc)
|
|
for f in current_flags:
|
|
blocks[f] = desc
|
|
|
|
def _is_real(flag: str) -> bool:
|
|
"""True if the flag exists AND is not a removal stub."""
|
|
desc = blocks.get(flag)
|
|
if desc is None:
|
|
return False
|
|
return "argument has been removed" not in desc
|
|
|
|
# MTP token from the --spec-type line.
|
|
spec_line = ""
|
|
for line in help_text.splitlines():
|
|
if "--spec-type" in line:
|
|
spec_line = line
|
|
break
|
|
# PR #22673 used draft-mtp; later renamed to mtp.
|
|
if "draft-mtp" in spec_line:
|
|
mtp_token = "draft-mtp"
|
|
elif re.search(r"[|,\[]mtp[|,\]]", spec_line):
|
|
mtp_token = "mtp"
|
|
|
|
# ngram-mod flag flavor. Post-rename builds advertise both new
|
|
# args (real) and legacy ones (stubs); pre-rename builds only
|
|
# have legacy ones as real.
|
|
new_ngram_real = (
|
|
_is_real("--spec-ngram-mod-n-match")
|
|
and _is_real("--spec-ngram-mod-n-min")
|
|
and _is_real("--spec-ngram-mod-n-max")
|
|
)
|
|
legacy_ngram_real = (
|
|
_is_real("--spec-ngram-size-n")
|
|
and _is_real("--draft-max")
|
|
and _is_real("--draft-min")
|
|
)
|
|
if new_ngram_real:
|
|
ngram_mod_flavor = "new"
|
|
elif legacy_ngram_real:
|
|
ngram_mod_flavor = "legacy"
|
|
|
|
# n_max flag: prefer post-rename, fall back to legacy.
|
|
if _is_real("--spec-draft-n-max"):
|
|
spec_draft_n_max_flag = "--spec-draft-n-max"
|
|
elif _is_real("--draft-max"):
|
|
spec_draft_n_max_flag = "--draft-max"
|
|
|
|
supports_kv_unified = _is_real("--kv-unified")
|
|
supports_fit_ctx = _is_real("--fit-ctx")
|
|
except (OSError, subprocess.SubprocessError) as exc:
|
|
logger.debug(f"llama-server --help probe failed: {exc}")
|
|
|
|
info = {
|
|
"found": True,
|
|
"mtp_token": mtp_token,
|
|
"supports_mtp": mtp_token is not None,
|
|
"ngram_mod_flavor": ngram_mod_flavor,
|
|
"supports_ngram_mod": ngram_mod_flavor is not None,
|
|
"spec_draft_n_max_flag": spec_draft_n_max_flag,
|
|
"supports_kv_unified": supports_kv_unified,
|
|
"supports_fit_ctx": supports_fit_ctx,
|
|
}
|
|
cls._capability_cache[cache_key] = info
|
|
return info
|
|
|
|
# ── GPU allocation ────────────────────────────────────────────
|
|
|
|
@staticmethod
|
|
def _get_gguf_size_bytes(model_path: str) -> int:
|
|
"""Total GGUF size in bytes, including split shards."""
|
|
main = Path(model_path)
|
|
total = main.stat().st_size
|
|
|
|
# Check for split shards (e.g. model-00001-of-00003.gguf)
|
|
m = _SHARD_FULL_RE.match(main.name)
|
|
if m:
|
|
prefix, _, num_total = m.group(1), m.group(2), m.group(3)
|
|
sibling_pat = re.compile(
|
|
r"^" + re.escape(prefix) + r"-\d{5}-of-" + re.escape(num_total) + r"\.gguf$"
|
|
)
|
|
for sibling in main.parent.iterdir():
|
|
if sibling != main and sibling_pat.match(sibling.name):
|
|
total += sibling.stat().st_size
|
|
|
|
return total
|
|
|
|
@staticmethod
|
|
def _amd_apu_wants_unified_memory() -> bool:
|
|
"""True only for AMD unified-memory APUs (gfx1150/gfx1151), where
|
|
GGML_CUDA_ENABLE_UNIFIED_MEMORY lets llama.cpp use shared system RAM.
|
|
False elsewhere (the env hurts discrete GPUs). ROCm reuses torch.cuda.*;
|
|
gcnArchName suffix is stripped."""
|
|
try:
|
|
import torch
|
|
|
|
if getattr(torch.version, "hip", None) is None:
|
|
return False
|
|
if not (hasattr(torch, "cuda") and torch.cuda.is_available()):
|
|
return False
|
|
for _i in range(torch.cuda.device_count()):
|
|
try:
|
|
_arch = getattr(torch.cuda.get_device_properties(_i), "gcnArchName", "") or ""
|
|
except Exception:
|
|
continue
|
|
if _arch.split(":")[0].strip().lower() in {"gfx1150", "gfx1151"}:
|
|
return True
|
|
except Exception:
|
|
return False
|
|
return False
|
|
|
|
# Datacenter / professional NVIDIA parts that benefit from the llama.cpp
|
|
# FP32-accum / P2P tunings. Whole-word (\b) so short markers don't match
|
|
# workstation parts as substrings: "a100" must not fire on "RTX A1000".
|
|
_DATACENTER_GPU_RE = re.compile(
|
|
r"\b(?:a100|a30|h100|h200|h800|gh200|b200|b100|b300|gb200|gb300|"
|
|
r"l40s?|l4|rtx pro 6000|rtx 6000 ada)\b"
|
|
)
|
|
|
|
@staticmethod
|
|
def _is_datacenter_gpu(gpu_indices = None) -> bool:
|
|
"""True iff every selected NVIDIA GPU is a datacenter/professional part.
|
|
NVIDIA-only, fails open to False (consumer GeForce, ROCm, CPU and errors
|
|
are left untouched); a mixed DC+consumer selection counts as non-DC.
|
|
|
|
gpu_indices are PHYSICAL ids (see _get_gpu_free_memory), but
|
|
get_device_properties wants mask-relative ordinals, so we rebuild the
|
|
ordinal->physical map from CUDA_VISIBLE_DEVICES and key names by physical
|
|
id. Otherwise a masked host (CUDA_VISIBLE_DEVICES=4,5,6,7, selection [4,5])
|
|
would drop the tuning or probe the wrong GPU."""
|
|
try:
|
|
import torch
|
|
|
|
if getattr(torch.version, "hip", None) is not None:
|
|
return False # ROCm reuses torch.cuda.*; not a CUDA part
|
|
if not (hasattr(torch, "cuda") and torch.cuda.is_available()):
|
|
return False
|
|
count = torch.cuda.device_count()
|
|
|
|
# Mirror _get_gpu_free_memory: map visible ordinal -> physical id via
|
|
# CUDA_VISIBLE_DEVICES; unset/unparsable leaves physical id == ordinal.
|
|
physical_ids: Optional[list[int]] = None
|
|
cvd = os.environ.get("CUDA_VISIBLE_DEVICES")
|
|
if cvd is not None:
|
|
try:
|
|
physical_ids = [int(x.strip()) for x in cvd.split(",") if x.strip()]
|
|
except ValueError:
|
|
physical_ids = None
|
|
|
|
pattern = LlamaCppBackend._DATACENTER_GPU_RE
|
|
names_by_id: dict[int, str] = {}
|
|
for ordinal in range(count):
|
|
try:
|
|
name = (torch.cuda.get_device_properties(ordinal).name or "").lower()
|
|
except Exception:
|
|
continue
|
|
pid = (
|
|
physical_ids[ordinal]
|
|
if physical_ids is not None and ordinal < len(physical_ids)
|
|
else ordinal
|
|
)
|
|
names_by_id[pid] = name
|
|
|
|
indices = list(gpu_indices) if gpu_indices else list(names_by_id)
|
|
saw = False
|
|
for _i in indices:
|
|
name = names_by_id.get(_i)
|
|
if name is None:
|
|
continue # not visible -> skip (fail conservative)
|
|
saw = True
|
|
if not pattern.search(name):
|
|
return False
|
|
return saw
|
|
except Exception:
|
|
return False
|
|
|
|
@staticmethod
|
|
def _effective_gpu_count(gpu_indices = None) -> int:
|
|
"""GPUs llama-server will use: len(selection), else the visible CUDA
|
|
device count (None = every visible GPU). 0 on error so multi-GPU tuning
|
|
stays off when the count is unknown."""
|
|
if gpu_indices is not None:
|
|
return len(gpu_indices)
|
|
try:
|
|
import torch
|
|
if hasattr(torch, "cuda") and torch.cuda.is_available():
|
|
return torch.cuda.device_count()
|
|
except Exception:
|
|
return 0
|
|
return 0
|
|
|
|
@staticmethod
|
|
def _apply_datacenter_env(env: dict, gpu_indices = None) -> bool:
|
|
"""Inject DC llama.cpp tuning into env in place via setdefault (user
|
|
values win); return whether the box qualified. Opt out with
|
|
UNSLOTH_DISABLE_DC_TUNING=1; only datacenter NVIDIA parts qualify
|
|
(consumer/ROCm/CPU/error are a no-op). Sets GGML_CUDA_FORCE_CUBLAS_COMPUTE_32F
|
|
for any qualifying GPU (FP32 accum: ~0% cost on B200, real cost on GeForce),
|
|
plus GGML_CUDA_P2P + CUDA_SCALE_LAUNCH_QUEUES=4x for multi-GPU (+33-51% pp
|
|
tensor-split, +8-16% pipeline split on B200)."""
|
|
if os.environ.get("UNSLOTH_DISABLE_DC_TUNING") == "1":
|
|
return False
|
|
if not LlamaCppBackend._is_datacenter_gpu(gpu_indices):
|
|
return False
|
|
env.setdefault("GGML_CUDA_FORCE_CUBLAS_COMPUTE_32F", "1")
|
|
if LlamaCppBackend._effective_gpu_count(gpu_indices) > 1:
|
|
env.setdefault("GGML_CUDA_P2P", "1")
|
|
env.setdefault("CUDA_SCALE_LAUNCH_QUEUES", "4x")
|
|
return True
|
|
|
|
@staticmethod
|
|
def _get_gpu_free_memory() -> list[tuple[int, int]]:
|
|
"""Query free memory per GPU.
|
|
|
|
Order:
|
|
1. ``nvidia-smi`` (NVIDIA CUDA hosts) -- respects
|
|
``CUDA_VISIBLE_DEVICES``.
|
|
2. ``torch.cuda.mem_get_info`` -- universal fallback that works
|
|
on AMD ROCm too (HIP runtime reuses the ``torch.cuda.*``
|
|
namespace). Covers the AMD case for issue #5106 (nvidia-smi
|
|
probe returned [] on AMD) and NVIDIA hosts missing
|
|
``nvidia-smi`` from PATH.
|
|
|
|
Returns list of (gpu_index, free_mib) sorted by index; empty if no
|
|
supported GPU is reachable.
|
|
"""
|
|
# ── NVIDIA via nvidia-smi ────────────────────────────────────
|
|
try:
|
|
result = subprocess.run(
|
|
[
|
|
"nvidia-smi",
|
|
"--query-gpu=index,memory.free",
|
|
"--format=csv,noheader,nounits",
|
|
],
|
|
capture_output = True,
|
|
text = True,
|
|
timeout = 10,
|
|
env = child_env_without_native_path_secret(),
|
|
**_windows_hidden_subprocess_kwargs(),
|
|
)
|
|
if result.returncode == 0:
|
|
allowed: Optional[set[int]] = None
|
|
cvd = os.environ.get("CUDA_VISIBLE_DEVICES")
|
|
if cvd is not None:
|
|
try:
|
|
# `if x.strip()` filters trailing-comma masks ("0,1,").
|
|
# Empty mask (CVD="") yields an empty set -> all GPUs
|
|
# filtered out, per codebase convention.
|
|
allowed = set(int(x.strip()) for x in cvd.split(",") if x.strip())
|
|
except ValueError:
|
|
pass
|
|
gpus: list[tuple[int, int]] = []
|
|
for line in result.stdout.strip().splitlines():
|
|
parts = line.split(",")
|
|
if len(parts) == 2:
|
|
idx = int(parts[0].strip())
|
|
free_mib = int(parts[1].strip())
|
|
if allowed is not None and idx not in allowed:
|
|
continue
|
|
gpus.append((idx, free_mib))
|
|
# Match the docstring's sort-by-id guarantee (driver order isn't).
|
|
gpus.sort(key = lambda g: g[0])
|
|
if gpus:
|
|
return gpus
|
|
except Exception as e:
|
|
logger.debug(f"nvidia-smi probe failed: {e}")
|
|
|
|
# ── Torch fallback (covers AMD ROCm and missing nvidia-smi) ──
|
|
try:
|
|
import torch
|
|
|
|
if not hasattr(torch, "cuda") or not torch.cuda.is_available():
|
|
return []
|
|
if not hasattr(torch.cuda, "mem_get_info"):
|
|
return []
|
|
# torch.cuda enumerates GPUs RELATIVE to the visibility mask. We
|
|
# feed these IDs back into the subprocess as CVD, so visible ordinals
|
|
# must be translated to physical indices first; otherwise CVD=2,3
|
|
# gets rewritten to 0,1 and targets the wrong GPUs.
|
|
physical_ids: Optional[list[int]] = None
|
|
# Match utils/hardware/hardware.py::_get_parent_visible_gpu_spec:
|
|
# treat an empty mask (HIP_VISIBLE_DEVICES="") as "no GPUs" rather
|
|
# than falling through. ``or`` would coerce "" to the wrong source.
|
|
if getattr(torch.version, "hip", None) is not None:
|
|
hip_v = os.environ.get("HIP_VISIBLE_DEVICES")
|
|
rocr_v = os.environ.get("ROCR_VISIBLE_DEVICES")
|
|
cvd = (
|
|
hip_v
|
|
if hip_v is not None
|
|
else rocr_v
|
|
if rocr_v is not None
|
|
else os.environ.get("CUDA_VISIBLE_DEVICES")
|
|
)
|
|
else:
|
|
cvd = os.environ.get("CUDA_VISIBLE_DEVICES")
|
|
if cvd is not None:
|
|
try:
|
|
# Empty mask (CVD="") yields an empty list -> no GPUs,
|
|
# consistent with the nvidia-smi path.
|
|
physical_ids = [int(x.strip()) for x in cvd.split(",") if x.strip()]
|
|
except ValueError:
|
|
physical_ids = None
|
|
gpus = []
|
|
for ordinal in range(torch.cuda.device_count()):
|
|
free_bytes, _total_bytes = torch.cuda.mem_get_info(ordinal)
|
|
idx = (
|
|
physical_ids[ordinal]
|
|
if physical_ids is not None and ordinal < len(physical_ids)
|
|
else ordinal
|
|
)
|
|
gpus.append((idx, free_bytes // (1024 * 1024)))
|
|
# Match the nvidia-smi path's docstring guarantee of sorted-by-id.
|
|
return sorted(gpus, key = lambda g: g[0])
|
|
except Exception as e:
|
|
logger.debug(f"torch GPU probe failed: {e}")
|
|
return []
|
|
|
|
# Skip the wait when the last kill is older than this; the driver has
|
|
# already reclaimed the prior process's allocations.
|
|
_VRAM_SETTLE_WINDOW_S: float = 15.0
|
|
|
|
@staticmethod
|
|
def _wait_for_vram_settle(
|
|
max_wait: float = 2.0,
|
|
interval: float = 0.25,
|
|
tolerance_mib: int = 256,
|
|
since_kill: float = 0.0,
|
|
) -> None:
|
|
"""Poll ``_get_gpu_free_memory`` until free VRAM stabilises.
|
|
|
|
The driver reclaims a dead process's allocations asynchronously, so
|
|
sampling free memory in the kill-to-spawn window reads artificially low
|
|
and pushes GPU selection toward needless CPU offload (the Apply-reload
|
|
OOM bare-shell launches never see).
|
|
|
|
Short-circuits on cold start, stale kill (older than
|
|
``_VRAM_SETTLE_WINDOW_S``), CPU-only hosts, probe exceptions, and GPU-set
|
|
changes. ``max_wait`` bounds wall-clock time so a wedged ``nvidia-smi``
|
|
can't extend the reload.
|
|
"""
|
|
now = time.monotonic()
|
|
if since_kill <= 0.0:
|
|
return
|
|
if now - since_kill > LlamaCppBackend._VRAM_SETTLE_WINDOW_S:
|
|
return
|
|
deadline = now + max_wait
|
|
|
|
def _probe_or_none():
|
|
if time.monotonic() >= deadline:
|
|
return None
|
|
try:
|
|
return LlamaCppBackend._get_gpu_free_memory()
|
|
except Exception:
|
|
return None
|
|
|
|
prev = _probe_or_none()
|
|
if prev is None or not prev:
|
|
return
|
|
while time.monotonic() < deadline:
|
|
remaining = deadline - time.monotonic()
|
|
if remaining <= 0:
|
|
return
|
|
# Clip the nap so a near-zero ``max_wait`` is respected.
|
|
time.sleep(min(interval, remaining))
|
|
curr = _probe_or_none()
|
|
if curr is None or not curr or len(curr) != len(prev):
|
|
return
|
|
prev_map = dict(prev)
|
|
stable = True
|
|
for idx, free in curr:
|
|
if idx not in prev_map:
|
|
stable = False
|
|
break
|
|
prev_free = prev_map[idx]
|
|
# Adaptive: 2% of the larger sample dominates the 256 MiB floor.
|
|
per_gpu_tol = max(tolerance_mib, int(max(free, prev_free) * 0.02))
|
|
if abs(free - prev_free) >= per_gpu_tol:
|
|
stable = False
|
|
break
|
|
if stable:
|
|
return
|
|
prev = curr
|
|
|
|
# Free-VRAM fraction at which Studio pins the GPU directly instead of
|
|
# deferring to ``--fit on``. 5% headroom covers CUDA context + compute
|
|
# buffers; 0.90 dropped 91-94% fits to CPU offload (#5106).
|
|
_GPU_PIN_VRAM_FRACTION = 0.95
|
|
|
|
# Per-GPU compute-graph buffer to reserve in tensor mode (MiB). This is the
|
|
# logits buffer (n_batch x vocab) + activation scratch that llama.cpp sizes
|
|
# via graph_reserve -- it is roughly EQUAL on every device (not proportional
|
|
# to the tensor split) and independent of context. Measured ~2.3 GB
|
|
# (gemma-3-27B) to ~3.8 GB (gemma-4-31B) on a 256k-vocab model; we reserve a
|
|
# conservative headroom above that. It is (a) subtracted from each GPU's free
|
|
# VRAM before computing --tensor-split, so the roomier GPU absorbs more
|
|
# weight and the smallest GPU keeps room for KV, and (b) reserved per device
|
|
# when capping context. The auto-fallback to layer split covers any
|
|
# underestimate. NOTE: scales with the model's vocab / batch size; tune if a
|
|
# large-vocab model OOMs at load.
|
|
_TENSOR_PARALLEL_BUFFER_RESERVE_MIB = 5120
|
|
|
|
# KV cache types llama.cpp accepts in tensor mode. A quantized KV cache
|
|
# aborts a --split-mode tensor load, so it's dropped for the tensor attempt.
|
|
_TENSOR_PARALLEL_KV_TYPES = frozenset({"f16", "bf16", "f32"})
|
|
|
|
@staticmethod
|
|
def _windows_pip_nvidia_dll_dirs(prefix: str) -> list[str]:
|
|
"""Return DLL dirs from pip-installed CUDA wheels under
|
|
``<prefix>/Lib/site-packages/`` so llama-server.exe can load
|
|
``cudart64_X.dll`` / ``cublas64_X.dll`` without a system CUDA toolkit.
|
|
Mirrors the Linux ``nvidia/cu*/lib`` LD_LIBRARY_PATH block, covering the
|
|
Windows wheel layouts seen in the wild:
|
|
* ``nvidia/<pkg>/bin`` -- legacy modular wheels.
|
|
* ``nvidia/<pkg>/bin/x86_64`` and ``.../bin/x64`` -- CUDA 13 layout
|
|
for unsuffixed packages (#5106).
|
|
* ``nvidia/<pkg>/Library/bin`` (and arch subdirs) -- conda repacks.
|
|
* ``torch/lib`` -- PyTorch's CUDA-bundled wheel can ship
|
|
``cudart64_*.dll`` here; mirrors install_llama_prebuilt.py.
|
|
|
|
Walks with ``Path.iterdir`` not ``glob.glob`` so it's safe against
|
|
Windows paths containing ``[`` or ``]`` (valid in usernames)."""
|
|
site_packages = Path(prefix) / "Lib" / "site-packages"
|
|
out: list[str] = []
|
|
seen: set[str] = set()
|
|
|
|
def _add(path: Path) -> None:
|
|
if not path.is_dir():
|
|
return
|
|
key = os.path.normcase(os.path.abspath(str(path)))
|
|
if key in seen:
|
|
return
|
|
seen.add(key)
|
|
out.append(str(path))
|
|
|
|
nvidia_root = site_packages / "nvidia"
|
|
if nvidia_root.is_dir():
|
|
for pkg_dir in nvidia_root.iterdir():
|
|
if not pkg_dir.is_dir():
|
|
continue
|
|
# Arch-specific subdirs first so the explicit cudart64_X.dll
|
|
# location wins over an empty sibling ``bin``.
|
|
for sub in (
|
|
pkg_dir / "bin" / "x86_64",
|
|
pkg_dir / "bin" / "x64",
|
|
pkg_dir / "bin",
|
|
pkg_dir / "Library" / "bin" / "x86_64",
|
|
pkg_dir / "Library" / "bin" / "x64",
|
|
pkg_dir / "Library" / "bin",
|
|
):
|
|
_add(sub)
|
|
_add(site_packages / "torch" / "lib")
|
|
return out
|
|
|
|
@staticmethod
|
|
def _build_windows_path_dirs(binary_dir: str, prefix: str, cuda_path: str) -> list[str]:
|
|
"""Ordered PATH entries prepended so llama-server.exe resolves cudart /
|
|
cublas DLLs: binary_dir, pip nvidia wheels, CUDA_PATH/bin, .../bin/x64.
|
|
Extracted so test_windows_gpu_detection_mock tests the real logic. #5106."""
|
|
path_dirs = [binary_dir]
|
|
path_dirs.extend(LlamaCppBackend._windows_pip_nvidia_dll_dirs(prefix))
|
|
if cuda_path:
|
|
cuda_bin = os.path.join(cuda_path, "bin")
|
|
if os.path.isdir(cuda_bin):
|
|
path_dirs.append(cuda_bin)
|
|
cuda_bin_x64 = os.path.join(cuda_path, "bin", "x64")
|
|
if os.path.isdir(cuda_bin_x64):
|
|
path_dirs.append(cuda_bin_x64)
|
|
return path_dirs
|
|
|
|
@staticmethod
|
|
def _select_gpus(
|
|
model_size_bytes: int,
|
|
gpus: list[tuple[int, int]],
|
|
usable_fraction: Optional[float] = None,
|
|
) -> tuple[Optional[list[int]], bool]:
|
|
"""Pick GPU(s) for a model from estimated VRAM and free memory.
|
|
|
|
``model_size_bytes`` should include weights and estimated KV cache.
|
|
``usable_fraction`` (default ``_GPU_PIN_VRAM_FRACTION``) provides
|
|
headroom for compute buffers, CUDA context, and other runtime
|
|
overhead; callers lower it when MTP reserves VRAM for a draft model.
|
|
|
|
Returns (gpu_indices, use_fit):
|
|
- ([1], False) fits on 1 GPU at the headroom threshold
|
|
- ([1, 2], False) needs 2 GPUs
|
|
- (None, True) too large, let --fit handle it
|
|
"""
|
|
if not gpus:
|
|
return None, True
|
|
|
|
model_size_mib = model_size_bytes / (1024 * 1024)
|
|
if usable_fraction is None:
|
|
usable_fraction = LlamaCppBackend._GPU_PIN_VRAM_FRACTION
|
|
|
|
# Sort GPUs by free memory descending
|
|
ranked = sorted(gpus, key = lambda g: g[1], reverse = True)
|
|
|
|
# Try 1 GPU at the usable-VRAM threshold.
|
|
if ranked[0][1] * usable_fraction >= model_size_mib:
|
|
return [ranked[0][0]], False
|
|
|
|
# Try N GPUs (accumulate free memory from most-free)
|
|
cumulative = 0
|
|
selected = []
|
|
for idx, free_mib in ranked:
|
|
selected.append(idx)
|
|
cumulative += free_mib * usable_fraction
|
|
if cumulative >= model_size_mib:
|
|
return sorted(selected), False
|
|
|
|
# Too large even for all GPUs; let --fit handle it
|
|
logger.debug(
|
|
"Model does not fit in available GPU memory, falling back to --fit",
|
|
model_size_mib = round(model_size_mib, 2),
|
|
ranked_gpus = ranked,
|
|
)
|
|
return None, True
|
|
|
|
# ── KV cache VRAM estimation ─────────────────────────────────────
|
|
|
|
def _can_estimate_kv(self) -> bool:
|
|
"""True if we have enough GGUF metadata to estimate KV cache size."""
|
|
if self._n_layers is None:
|
|
return False
|
|
# MLA: kv_lora_rank suffices (K-only cache).
|
|
if self._kv_lora_rank is not None:
|
|
return True
|
|
# New-style: need explicit key AND value dimensions.
|
|
if self._kv_key_length is not None and self._kv_value_length is not None:
|
|
return True
|
|
# Legacy: need embedding_length + a head count (scalar or per-layer).
|
|
return self._embedding_length is not None and (
|
|
self._n_kv_heads is not None
|
|
or self._n_heads is not None
|
|
or self._n_kv_heads_by_layer is not None
|
|
)
|
|
|
|
def _kv_heads_for_layer(self, layer_idx: int, fallback: int) -> int:
|
|
if self._n_kv_heads_by_layer is not None and layer_idx < len(self._n_kv_heads_by_layer):
|
|
return self._n_kv_heads_by_layer[layer_idx]
|
|
return fallback
|
|
|
|
def _estimate_kv_cache_bytes(
|
|
self,
|
|
n_ctx: int,
|
|
cache_type_kv: Optional[str] = None,
|
|
*,
|
|
swa_full: bool = False,
|
|
n_parallel: int = 1,
|
|
kv_unified: bool = True,
|
|
ctx_checkpoints: int = 0,
|
|
) -> int:
|
|
"""Estimate KV cache VRAM for a given context length.
|
|
|
|
5-path architecture-aware estimation:
|
|
1. MLA -- compressed KV latent + RoPE, K-only (no separate V)
|
|
2. Hybrid -- only attention layers need KV (Mamba layers don't)
|
|
3. SWA -- sliding-window layers cache min(ctx, window) tokens
|
|
4. GQA -- standard full KV with explicit key/value dimensions
|
|
5. Legacy -- fallback using embed // n_heads
|
|
|
|
Server-flag knobs (mirror llama-server's CLI):
|
|
swa_full -- ``--swa-full``: force SWA layers to cache full
|
|
``n_ctx`` (collapses path 3 to path 4 for them).
|
|
n_parallel -- ``--parallel`` slots: non-SWA layers stay constant
|
|
(cells split across slots), SWA layers scale linearly.
|
|
kv_unified -- ``--kv-unified`` (default on): no-op for memory math;
|
|
kept for API forward-compat.
|
|
ctx_checkpoints -- ``--ctx-checkpoints`` (PR #15293): N SWA snapshots
|
|
per slot, one sliding-window of state per SWA layer.
|
|
|
|
Returns 0 if metadata is insufficient.
|
|
"""
|
|
if not self._can_estimate_kv() or n_ctx <= 0:
|
|
return 0
|
|
|
|
n_layers = self._n_layers # type: ignore[assignment]
|
|
# Gemma 3n / Gemma 4 reuse earlier KV in the last ``shared_kv_layers``
|
|
# blocks (no cache). Floor at 1 so a bad GGUF can't zero out KV.
|
|
shared = self._shared_kv_layers or 0
|
|
n_layers_kv = max(1, n_layers - shared)
|
|
n_kv = self._n_kv_heads or self._n_heads or 1 # type: ignore[assignment]
|
|
|
|
# Bytes per element depends on KV cache quantization
|
|
bpe = {
|
|
"f32": 4.0,
|
|
"f16": 2.0,
|
|
"bf16": 2.0,
|
|
"q8_0": 34 / 32,
|
|
"q5_1": 0.75,
|
|
"q5_0": 0.6875,
|
|
"q4_1": 0.625,
|
|
"q4_0": 0.5625,
|
|
"iq4_nl": 0.5625,
|
|
}.get(cache_type_kv or "f16", 2.0)
|
|
|
|
slots = max(1, n_parallel)
|
|
|
|
# Path 1: MLA (DeepSeek-V2/V3, GLM-4.7, GLM-5, Kimi-K2.5)
|
|
# One compressed KV latent per token/layer (shared across heads); V is
|
|
# reconstructed from it, no separate V cache. key_length = kv_lora_rank
|
|
# + rope_dim. MLA GGUFs set head_count_kv=1; default to 1 if absent to
|
|
# avoid falling back to n_heads (e.g. 128 for DeepSeek-V3) which 128x's.
|
|
if self._kv_lora_rank is not None:
|
|
n_kv_mla = self._n_kv_heads or 1
|
|
rope_dim = self._key_length_mla or 64
|
|
key_len = self._kv_key_length or (self._kv_lora_rank + rope_dim)
|
|
return int(n_layers_kv * n_ctx * n_kv_mla * key_len * bpe)
|
|
|
|
key_len = self._kv_key_length
|
|
val_len = self._kv_value_length
|
|
|
|
# Path 2: Hybrid Mamba/Attention (Qwen3.5-27B, Qwen3.5-35B-A3B)
|
|
# Only 1 in N layers is attention; the rest are Mamba (no KV cache).
|
|
if self._ssm_inner_size is not None and self._full_attention_interval is not None:
|
|
fai = self._full_attention_interval
|
|
n_attn = -(-n_layers // fai) if fai > 0 else n_layers # ceiling division
|
|
if key_len is not None and val_len is not None:
|
|
return int(n_attn * n_ctx * n_kv * (key_len + val_len) * bpe)
|
|
head_dim = self._embedding_length // self._n_heads if self._n_heads else 128 # type: ignore[operator]
|
|
return int(n_attn * n_ctx * n_kv * 2 * head_dim * bpe)
|
|
|
|
# Path 3: Sliding window (Gemma 2/3/3n/4, gpt-oss, Cohere2 ...).
|
|
# Pattern filled by the resolver at parse time; if absent, falls through
|
|
# to the legacy 1/4-global heuristic below. Per-layer-type --parallel N
|
|
# accounting (verified against llama-server):
|
|
# * non-SWA layers: total cells = n_ctx split across slots -> CONSTANT.
|
|
# * SWA layers: per-slot cells = 2*sliding_window (capped at n_ctx
|
|
# and per_slot_ctx) -> grows LINEARLY in slots.
|
|
# --swa-full forces full n_ctx for SWA layers; --ctx-checkpoints N adds
|
|
# N snapshots per SWA layer per slot.
|
|
if (
|
|
self._sliding_window is not None
|
|
and self._sliding_window > 0
|
|
and key_len is not None
|
|
and val_len is not None
|
|
):
|
|
swa = self._sliding_window
|
|
per_slot_ctx = max(1, n_ctx // slots)
|
|
# --swa-full caches full context like non-SWA (per-slot cells =
|
|
# per_slot_ctx, collapsing to constant n_ctx total); otherwise SWA
|
|
# caches 2*sliding_window per slot, clamped at per-slot ctx.
|
|
swa_cells_per_slot = per_slot_ctx if swa_full else min(n_ctx, 2 * swa, per_slot_ctx)
|
|
key_len_swa = self._kv_key_length_swa or key_len
|
|
val_len_swa = self._kv_value_length_swa or val_len
|
|
if self._sliding_window_pattern is not None:
|
|
global_bytes = 0.0 # constant across slots
|
|
swa_bytes_per_slot = 0.0 # multiplied by slots
|
|
checkpoint_extra_per_slot = 0.0
|
|
# Only layers that allocate their own KV; trailing shared layers
|
|
# reuse earlier caches.
|
|
for layer_idx in range(n_layers_kv):
|
|
layer_n_kv = self._kv_heads_for_layer(layer_idx, n_kv)
|
|
is_swa = (
|
|
layer_idx < len(self._sliding_window_pattern)
|
|
and self._sliding_window_pattern[layer_idx]
|
|
)
|
|
if is_swa:
|
|
swa_bytes_per_slot += (
|
|
swa_cells_per_slot * layer_n_kv * (key_len_swa + val_len_swa) * bpe
|
|
)
|
|
if ctx_checkpoints > 0 and not swa_full:
|
|
checkpoint_extra_per_slot += (
|
|
ctx_checkpoints
|
|
* swa
|
|
* layer_n_kv
|
|
* (key_len_swa + val_len_swa)
|
|
* bpe
|
|
)
|
|
else:
|
|
global_bytes += n_ctx * layer_n_kv * (key_len + val_len) * bpe
|
|
return int(global_bytes + slots * (swa_bytes_per_slot + checkpoint_extra_per_slot))
|
|
n_global = max(1, n_layers_kv // 4)
|
|
n_swa = n_layers_kv - n_global
|
|
kv_per_token = n_kv * (key_len + val_len) * bpe
|
|
kv_per_token_swa = n_kv * (key_len_swa + val_len_swa) * bpe
|
|
global_bytes = n_global * n_ctx * kv_per_token
|
|
swa_bytes_per_slot = n_swa * swa_cells_per_slot * kv_per_token_swa
|
|
checkpoint_extra_per_slot = (
|
|
ctx_checkpoints * n_swa * swa * kv_per_token_swa
|
|
if ctx_checkpoints > 0 and not swa_full
|
|
else 0.0
|
|
)
|
|
return int(global_bytes + slots * (swa_bytes_per_slot + checkpoint_extra_per_slot))
|
|
|
|
# Path 4: Standard GQA with explicit key/value dimensions
|
|
if key_len is not None and val_len is not None:
|
|
return int(n_layers_kv * n_ctx * n_kv * (key_len + val_len) * bpe)
|
|
|
|
# Path 5: Legacy fallback (old GGUFs without explicit dimensions)
|
|
head_dim = self._embedding_length // self._n_heads if self._n_heads else 128 # type: ignore[operator]
|
|
return int(2 * n_kv * head_dim * n_layers_kv * n_ctx * bpe)
|
|
|
|
def _fit_context_to_vram(
|
|
self,
|
|
requested_ctx: int,
|
|
available_mib: int,
|
|
model_size_bytes: int,
|
|
cache_type_kv: Optional[str] = None,
|
|
min_ctx: int = 4096,
|
|
*,
|
|
swa_full: bool = False,
|
|
n_parallel: int = 1,
|
|
kv_unified: bool = True,
|
|
ctx_checkpoints: int = 0,
|
|
kv_on_gpu: bool = True,
|
|
mtp_engaged: bool = False,
|
|
budget_frac: Optional[float] = None,
|
|
) -> int:
|
|
"""Return the largest context length that fits in GPU VRAM.
|
|
|
|
Uses 90% of available VRAM as the ctx-fit budget -- tighter than
|
|
``_GPU_PIN_VRAM_FRACTION`` on purpose (over-promising context OOMs at
|
|
runtime). If the weights alone don't fit, returns ``requested_ctx``.
|
|
|
|
``kv_on_gpu`` mirrors ``--kv-offload`` (default on); when False the KV
|
|
cache lives in CPU RAM and the requested context is honored verbatim.
|
|
Other keyword args mirror ``_estimate_kv_cache_bytes``.
|
|
|
|
``mtp_engaged`` reserves extra VRAM for the MTP draft model's KV cache +
|
|
compute buffers, else tight tiers (e.g. 32 GB) spill to a slower path.
|
|
"""
|
|
if not self._can_estimate_kv():
|
|
logger.debug(
|
|
"Skipping context fit because KV cache metadata is unavailable",
|
|
requested_ctx = requested_ctx,
|
|
available_mib = available_mib,
|
|
)
|
|
return requested_ctx
|
|
|
|
# KV lives off-GPU: no VRAM accounting needed for the cache itself.
|
|
if not kv_on_gpu:
|
|
return requested_ctx
|
|
|
|
kv_kwargs = dict(
|
|
swa_full = swa_full,
|
|
n_parallel = n_parallel,
|
|
kv_unified = kv_unified,
|
|
ctx_checkpoints = ctx_checkpoints,
|
|
)
|
|
|
|
# MTP engaged: carve the drafter's reserve out of the fit budget. Callers
|
|
# can override outright (tensor-parallel mode passes a fatter margin), so
|
|
# only compute a default when none was supplied.
|
|
if budget_frac is None:
|
|
budget_frac = _CTX_FIT_VRAM_FRACTION - (_MTP_VRAM_RESERVE_FRAC if mtp_engaged else 0.0)
|
|
budget_bytes = available_mib * 1024 * 1024 * budget_frac
|
|
model_footprint = model_size_bytes
|
|
|
|
# Already fits?
|
|
kv = self._estimate_kv_cache_bytes(requested_ctx, cache_type_kv, **kv_kwargs)
|
|
if model_footprint + kv <= budget_bytes:
|
|
return requested_ctx
|
|
|
|
# Weights alone exceed budget -- reducing ctx can't help; --fit handles it.
|
|
if model_footprint >= budget_bytes:
|
|
logger.debug(
|
|
"Model footprint exceeds GPU budget before KV cache",
|
|
requested_ctx = requested_ctx,
|
|
available_mib = available_mib,
|
|
model_size_gb = round(model_footprint / (1024**3), 2),
|
|
)
|
|
return requested_ctx
|
|
|
|
# Binary search for max context that fits
|
|
remaining = budget_bytes - model_footprint
|
|
effective_min = min(min_ctx, requested_ctx)
|
|
lo, hi = effective_min, requested_ctx
|
|
best = effective_min
|
|
while lo <= hi:
|
|
mid = (lo + hi) // 2
|
|
kv = self._estimate_kv_cache_bytes(mid, cache_type_kv, **kv_kwargs)
|
|
if kv <= remaining:
|
|
best = mid
|
|
lo = mid + 1
|
|
else:
|
|
hi = mid - 1
|
|
|
|
# Round down to nearest 256 for alignment, never above requested_ctx
|
|
best = (best // 256) * 256
|
|
best = max(effective_min, best)
|
|
best = min(best, requested_ctx)
|
|
return best
|
|
|
|
# ── Variant fallback ────────────────────────────────────────────
|
|
|
|
@staticmethod
|
|
def _find_smallest_fitting_variant(
|
|
hf_repo: str,
|
|
free_bytes: int,
|
|
hf_token: Optional[str] = None,
|
|
) -> Optional[tuple[str, int]]:
|
|
"""Find the smallest GGUF variant (including all shards) that fits.
|
|
|
|
Groups split shards by variant prefix and sums their sizes (e.g.
|
|
UD-Q4_K_XL with 9 shards of 50 GB each = 450 GB total).
|
|
|
|
Returns (first_shard_filename, total_size_bytes) or None.
|
|
"""
|
|
try:
|
|
from huggingface_hub import get_paths_info, list_repo_files
|
|
|
|
files = list_repo_files(hf_repo, token = hf_token)
|
|
gguf_files = [
|
|
f for f in files if f.endswith(".gguf") and not _is_companion_gguf_path(f)
|
|
]
|
|
if not gguf_files:
|
|
return None
|
|
|
|
# Sizes for all GGUF files
|
|
path_infos = list(get_paths_info(hf_repo, gguf_files, token = hf_token))
|
|
size_map = {p.path: (p.size or 0) for p in path_infos}
|
|
|
|
# Group by variant: shards share a prefix before -NNNNN-of-NNNNN
|
|
variants: dict[str, list[str]] = {}
|
|
for f in gguf_files:
|
|
m = _SHARD_RE.match(f)
|
|
key = m.group(1) if m else f
|
|
variants.setdefault(key, []).append(f)
|
|
|
|
# Sum shard sizes per variant, track the first shard (for download)
|
|
variant_sizes: list[tuple[str, int, list[str]]] = []
|
|
for key, shard_files in variants.items():
|
|
total = sum(size_map.get(f, 0) for f in shard_files)
|
|
first = sorted(shard_files)[0]
|
|
variant_sizes.append((first, total, shard_files))
|
|
|
|
# Smallest that fits
|
|
variant_sizes.sort(key = lambda x: x[1])
|
|
for first_file, total_size, _ in variant_sizes:
|
|
if total_size > 0 and total_size <= free_bytes:
|
|
return first_file, total_size
|
|
|
|
return None
|
|
except Exception:
|
|
return None
|
|
|
|
# ── Port allocation ───────────────────────────────────────────
|
|
|
|
@staticmethod
|
|
def _find_free_port() -> int:
|
|
"""Find an available TCP port."""
|
|
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
|
s.bind(("127.0.0.1", 0))
|
|
return s.getsockname()[1]
|
|
|
|
# ── Stdout drain (prevents pipe deadlock on Windows) ─────────
|
|
|
|
def _drain_stdout(self):
|
|
"""Read subprocess stdout lines in a background thread.
|
|
|
|
Prevents a pipe-buffer deadlock on Windows (~4 KB buffer): without
|
|
draining, llama-server blocks on writes and never becomes healthy.
|
|
Each line is also teed to ``self._llama_log_fh`` when set, so a
|
|
post-mortem has the full output even if the crash predates the
|
|
drain-thread join in ``_wait_for_health``.
|
|
"""
|
|
try:
|
|
for line in self._process.stdout:
|
|
line = line.rstrip()
|
|
if line:
|
|
self._stdout_lines.append(line)
|
|
logger.debug(f"[llama-server] {line}")
|
|
fh = getattr(self, "_llama_log_fh", None)
|
|
if fh is not None:
|
|
try:
|
|
fh.write(line + "\n")
|
|
fh.flush()
|
|
except (ValueError, OSError):
|
|
# Log file closed under us; tee silently.
|
|
pass
|
|
except (ValueError, OSError):
|
|
# Pipe closed -- process terminating.
|
|
pass
|
|
|
|
# GGUF KV type sizes for fast skipping
|
|
_GGUF_TYPE_SIZE = {
|
|
0: 1,
|
|
1: 1,
|
|
2: 2,
|
|
3: 2,
|
|
4: 4,
|
|
5: 4,
|
|
6: 4,
|
|
7: 1,
|
|
10: 8,
|
|
11: 8,
|
|
12: 8,
|
|
}
|
|
|
|
@staticmethod
|
|
def _gguf_skip_value(f, vtype: int) -> None:
|
|
"""Skip a GGUF KV value without reading it."""
|
|
sz = LlamaCppBackend._GGUF_TYPE_SIZE.get(vtype)
|
|
if sz is not None:
|
|
f.seek(sz, 1)
|
|
elif vtype == 8: # STRING
|
|
slen = struct.unpack("<Q", f.read(8))[0]
|
|
f.seek(slen, 1)
|
|
elif vtype == 9: # ARRAY
|
|
atype = struct.unpack("<I", f.read(4))[0]
|
|
alen = struct.unpack("<Q", f.read(8))[0]
|
|
elem_sz = LlamaCppBackend._GGUF_TYPE_SIZE.get(atype)
|
|
if elem_sz is not None:
|
|
f.seek(elem_sz * alen, 1)
|
|
elif atype == 8:
|
|
for _ in range(alen):
|
|
slen = struct.unpack("<Q", f.read(8))[0]
|
|
f.seek(slen, 1)
|
|
else:
|
|
for _ in range(alen):
|
|
LlamaCppBackend._gguf_skip_value(f, atype)
|
|
|
|
@staticmethod
|
|
def _gguf_read_array_value(f, atype: int, alen: int) -> Optional[list]:
|
|
if atype == 4: # UINT32
|
|
return [struct.unpack("<I", f.read(4))[0] for _ in range(alen)]
|
|
if atype == 5: # INT32
|
|
return [struct.unpack("<i", f.read(4))[0] for _ in range(alen)]
|
|
if atype == 7: # BOOL
|
|
return [struct.unpack("<?", f.read(1))[0] for _ in range(alen)]
|
|
|
|
for _ in range(alen):
|
|
LlamaCppBackend._gguf_skip_value(f, atype)
|
|
return None
|
|
|
|
def _read_gguf_metadata(self, gguf_path: str) -> None:
|
|
"""Read context_length, architecture params, and chat_template from a GGUF header.
|
|
|
|
Parses only the KV pairs we need (~30ms even for multi-GB files).
|
|
For split GGUFs, metadata is always in shard 1.
|
|
"""
|
|
# Reset metadata so stale flags (e.g. _supports_reasoning) don't
|
|
# carry over when switching models.
|
|
self._context_length = None
|
|
self._chat_template = None
|
|
self._supports_reasoning = False
|
|
self._reasoning_always_on = False
|
|
self._reasoning_style = "enable_thinking"
|
|
self._reasoning_default = True
|
|
self._supports_preserve_thinking = False
|
|
self._supports_tools = False
|
|
self._n_layers = None
|
|
self._n_kv_heads = None
|
|
self._n_kv_heads_by_layer = None
|
|
self._n_heads = None
|
|
self._embedding_length = None
|
|
self._kv_key_length = None
|
|
self._kv_value_length = None
|
|
self._sliding_window = None
|
|
self._sliding_window_pattern = None
|
|
self._full_attention_interval = None
|
|
self._kv_lora_rank = None
|
|
self._key_length_mla = None
|
|
self._kv_key_length_swa = None
|
|
self._kv_value_length_swa = None
|
|
self._ssm_inner_size = None
|
|
self._ssm_state_size = None
|
|
self._shared_kv_layers = None
|
|
self._nextn_predict_layers = None
|
|
self._architecture = None
|
|
self._is_diffusion = False
|
|
|
|
try:
|
|
canvas_seen = False
|
|
WANTED = {
|
|
"general.architecture",
|
|
"tokenizer.chat_template",
|
|
# Block-diffusion marker (DiffusionGemma); routes to the diffusion runner.
|
|
"diffusion.canvas_length",
|
|
# Source-repo hints for the SWA resolver's HF fallback.
|
|
"general.source.huggingface.repository",
|
|
"general.source.url",
|
|
"general.source.repo_url",
|
|
"general.base_model.0.repo_url",
|
|
"general.base_model.0.organization",
|
|
"general.base_model.0.name",
|
|
"general.basename",
|
|
"general.organization",
|
|
"general.size_label",
|
|
"general.finetune",
|
|
}
|
|
# Arch-specific keys added dynamically once we know the arch.
|
|
arch_keys: dict[str, str] = {} # gguf_key -> attribute name
|
|
arch = None
|
|
sliding_window_pattern_period: Optional[int] = None
|
|
general: dict[str, str] = {}
|
|
|
|
with open(gguf_path, "rb") as f:
|
|
magic = struct.unpack("<I", f.read(4))[0]
|
|
if magic != 0x46554747: # b"GGUF" as little-endian u32
|
|
return
|
|
_version = struct.unpack("<I", f.read(4))[0]
|
|
_tensor_count, kv_count = struct.unpack("<QQ", f.read(16))
|
|
|
|
for _ in range(kv_count):
|
|
# Tolerate truncated input (e.g. a partial header from an
|
|
# HTTP byte-range fetch): bail out so the resolver
|
|
# fallback runs on whatever we parsed.
|
|
try:
|
|
key_len_bytes = f.read(8)
|
|
if len(key_len_bytes) < 8:
|
|
break
|
|
key_len = struct.unpack("<Q", key_len_bytes)[0]
|
|
key_bytes = f.read(key_len)
|
|
if len(key_bytes) < key_len:
|
|
break
|
|
key = key_bytes.decode("utf-8")
|
|
vtype_bytes = f.read(4)
|
|
if len(vtype_bytes) < 4:
|
|
break
|
|
vtype = struct.unpack("<I", vtype_bytes)[0]
|
|
except (struct.error, UnicodeDecodeError):
|
|
break
|
|
|
|
try:
|
|
if key in WANTED or key in arch_keys:
|
|
if vtype == 8: # STRING
|
|
slen = struct.unpack("<Q", f.read(8))[0]
|
|
val_s = f.read(slen).decode("utf-8")
|
|
if key.startswith("general.") and key != "general.architecture":
|
|
general[key] = val_s
|
|
if key == "general.architecture":
|
|
arch = val_s
|
|
self._architecture = val_s
|
|
arch_keys = {
|
|
f"{arch}.context_length": "context_length",
|
|
f"{arch}.block_count": "n_layers",
|
|
f"{arch}.attention.head_count_kv": "n_kv_heads",
|
|
f"{arch}.attention.head_count": "n_heads",
|
|
f"{arch}.embedding_length": "embedding_length",
|
|
f"{arch}.attention.key_length": "kv_key_length",
|
|
f"{arch}.attention.value_length": "kv_value_length",
|
|
f"{arch}.attention.sliding_window": "sliding_window",
|
|
f"{arch}.attention.sliding_window_pattern": "sliding_window_pattern",
|
|
f"{arch}.full_attention_interval": "full_attention_interval",
|
|
f"{arch}.attention.kv_lora_rank": "kv_lora_rank",
|
|
f"{arch}.attention.key_length_mla": "key_length_mla",
|
|
f"{arch}.attention.key_length_swa": "kv_key_length_swa",
|
|
f"{arch}.attention.value_length_swa": "kv_value_length_swa",
|
|
f"{arch}.attention.shared_kv_layers": "shared_kv_layers",
|
|
f"{arch}.ssm.inner_size": "ssm_inner_size",
|
|
f"{arch}.ssm.state_size": "ssm_state_size",
|
|
f"{arch}.nextn_predict_layers": "nextn_predict_layers",
|
|
}
|
|
elif key == "tokenizer.chat_template":
|
|
self._chat_template = val_s
|
|
elif vtype in (4, 10): # UINT32 or UINT64
|
|
val_i = (
|
|
struct.unpack("<I", f.read(4))[0]
|
|
if vtype == 4
|
|
else struct.unpack("<Q", f.read(8))[0]
|
|
)
|
|
if key == "diffusion.canvas_length":
|
|
canvas_seen = True
|
|
attr = arch_keys.get(key)
|
|
if attr:
|
|
if attr == "sliding_window_pattern":
|
|
sliding_window_pattern_period = val_i
|
|
else:
|
|
setattr(self, f"_{attr}", val_i)
|
|
elif vtype == 9: # ARRAY
|
|
atype = struct.unpack("<I", f.read(4))[0]
|
|
alen = struct.unpack("<Q", f.read(8))[0]
|
|
val_a = self._gguf_read_array_value(f, atype, alen)
|
|
attr = arch_keys.get(key)
|
|
if attr == "n_kv_heads" and val_a is not None:
|
|
self._n_kv_heads_by_layer = [int(x) for x in val_a]
|
|
if self._n_kv_heads is None and val_a:
|
|
self._n_kv_heads = max(int(x) for x in val_a)
|
|
elif attr == "sliding_window_pattern" and val_a is not None:
|
|
self._sliding_window_pattern = [bool(x) for x in val_a]
|
|
sliding_window_pattern_period = None
|
|
else:
|
|
self._gguf_skip_value(f, vtype)
|
|
else:
|
|
self._gguf_skip_value(f, vtype)
|
|
except (struct.error, UnicodeDecodeError):
|
|
# Truncated input (e.g. HTTP byte-range header
|
|
# fetch); break so the resolver fallback runs on
|
|
# what we have.
|
|
break
|
|
|
|
# Expand a scalar period straight from the GGUF first.
|
|
if (
|
|
self._sliding_window_pattern is None
|
|
and sliding_window_pattern_period
|
|
and self._n_layers
|
|
):
|
|
self._sliding_window_pattern = [
|
|
(i + 1) % sliding_window_pattern_period != 0 for i in range(self._n_layers)
|
|
]
|
|
|
|
# Otherwise hand off to the resolver (cache / bootstrap /
|
|
# transformers / HF); see `_resolve_swa_pattern`.
|
|
if self._sliding_window_pattern is None and self._sliding_window and self._n_layers:
|
|
hf_repo_candidates = (
|
|
general.get("general.source.huggingface.repository"),
|
|
_hf_repo_from_url(general.get("general.source.url")),
|
|
_hf_repo_from_url(general.get("general.source.repo_url")),
|
|
_hf_repo_from_url(general.get("general.base_model.0.repo_url")),
|
|
(
|
|
f"{general['general.base_model.0.organization']}/"
|
|
f"{general['general.base_model.0.name']}".replace(" ", "-")
|
|
if general.get("general.base_model.0.organization")
|
|
and general.get("general.base_model.0.name")
|
|
else None
|
|
),
|
|
(
|
|
f"{general['general.organization']}/{general['general.basename']}".replace(
|
|
" ", "-"
|
|
)
|
|
if general.get("general.organization") and general.get("general.basename")
|
|
else None
|
|
),
|
|
)
|
|
self._sliding_window_pattern = _resolve_swa_pattern(
|
|
arch,
|
|
self._n_layers,
|
|
hf_repo_candidates,
|
|
)
|
|
|
|
# Block-diffusion models (DiffusionGemma) report a diffusion arch
|
|
# and/or a diffusion.canvas_length KV; they need the diffusion runner.
|
|
self._is_diffusion = bool(
|
|
(arch and arch.lower().startswith("diffusion")) or canvas_seen
|
|
)
|
|
if self._is_diffusion:
|
|
logger.info(
|
|
f"GGUF metadata: diffusion model detected (architecture={arch}); "
|
|
"will serve via the diffusion runner"
|
|
)
|
|
|
|
if self._context_length:
|
|
logger.info(f"GGUF metadata: context_length={self._context_length}")
|
|
if self._chat_template:
|
|
logger.info(f"GGUF metadata: chat_template={len(self._chat_template)} chars")
|
|
# Detect thinking/reasoning support from chat template.
|
|
flags = detect_reasoning_flags(
|
|
self._chat_template,
|
|
self._model_identifier,
|
|
log_source = "GGUF metadata",
|
|
)
|
|
self._supports_reasoning = flags["supports_reasoning"]
|
|
self._reasoning_style = flags["reasoning_style"]
|
|
self._reasoning_always_on = flags["reasoning_always_on"]
|
|
self._supports_preserve_thinking = flags["supports_preserve_thinking"]
|
|
self._supports_tools = flags["supports_tools"]
|
|
except Exception as e:
|
|
logger.warning(f"Failed to read GGUF metadata: {e}")
|
|
|
|
# ── Diffusion runner (DiffusionGemma) ──
|
|
|
|
def _find_diffusion_assets(self) -> Optional[tuple[list, str, Optional[str]]]:
|
|
"""Resolve how to launch the DiffusionGemma runner: (shim argv prefix,
|
|
visual-server binary, optional extra PYTHONPATH dir for the file override).
|
|
|
|
Shim: UNSLOTH_DG_SHIM (a .py file) first, else the installed
|
|
unsloth_zoo.diffusion_studio.shim. Binary: DG_VISUAL_BIN first, else
|
|
alongside llama-server. Returns None if neither can be found.
|
|
"""
|
|
import importlib.util
|
|
import os
|
|
import sys
|
|
|
|
# Visual-server binary: env override, else next to llama-server or in the
|
|
# install's build/bin (where the prebuilt/installer puts it). .exe on Windows.
|
|
visual_bin = os.environ.get("DG_VISUAL_BIN")
|
|
if not visual_bin:
|
|
name = "llama-diffusion-gemma-visual-server" + (".exe" if os.name == "nt" else "")
|
|
base = self._find_llama_server_binary()
|
|
if base:
|
|
base_dir = Path(base).parent
|
|
for cand in (
|
|
base_dir / name,
|
|
base_dir / "build" / "bin" / name,
|
|
base_dir / "build" / "bin" / "Release" / name,
|
|
):
|
|
if cand.is_file():
|
|
visual_bin = str(cand)
|
|
break
|
|
if not (visual_bin and Path(visual_bin).is_file()):
|
|
return None
|
|
|
|
# Shim: a file override (its dir goes on PYTHONPATH), else the zoo package via -m.
|
|
shim_file = os.environ.get("UNSLOTH_DG_SHIM")
|
|
if shim_file and Path(shim_file).is_file():
|
|
return ([sys.executable, shim_file], visual_bin, str(Path(shim_file).parent))
|
|
|
|
# Find the installed shim without importing the heavy unsloth_zoo package
|
|
# (find_spec on the top-level package does not run its __init__).
|
|
try:
|
|
spec = importlib.util.find_spec("unsloth_zoo")
|
|
except Exception:
|
|
spec = None
|
|
if spec is not None and spec.submodule_search_locations:
|
|
pkg_dir = Path(list(spec.submodule_search_locations)[0])
|
|
if (pkg_dir / "diffusion_studio" / "shim.py").is_file():
|
|
return (
|
|
[sys.executable, "-m", "unsloth_zoo.diffusion_studio.shim"],
|
|
visual_bin,
|
|
None,
|
|
)
|
|
|
|
return None
|
|
|
|
def _start_diffusion_server(
|
|
self,
|
|
*,
|
|
model_path: str,
|
|
gguf_path: Optional[str],
|
|
hf_repo: Optional[str],
|
|
hf_variant: Optional[str],
|
|
model_identifier: str,
|
|
n_ctx: int,
|
|
extra_args: Optional[List[str]],
|
|
) -> bool:
|
|
"""Launch the OpenAI-compat diffusion shim (which drives the on-device
|
|
visual decoder) and wait for health. Presents the same /v1 + /health
|
|
interface as llama-server, so the rest of Studio is unchanged.
|
|
"""
|
|
import os
|
|
|
|
assets = self._find_diffusion_assets()
|
|
if assets is None:
|
|
raise RuntimeError(
|
|
"DiffusionGemma runner not found. Install unsloth_zoo (which ships "
|
|
"unsloth_zoo.diffusion_studio.shim) or set UNSLOTH_DG_SHIM to a shim "
|
|
"file, and provide the visual-server binary via DG_VISUAL_BIN or next "
|
|
"to llama-server in the install tree."
|
|
)
|
|
shim_cmd, visual_bin, extra_pythonpath = assets
|
|
self._diffusion_visual_bin = visual_bin
|
|
|
|
self._kill_process()
|
|
self._port = self._find_free_port()
|
|
# Auto-size (0): the visual server probes the largest context that fits this GPU's VRAM
|
|
# (capped at the training context). An explicit in-range n_ctx overrides it.
|
|
maxtok = n_ctx if (n_ctx and 0 < n_ctx <= 65536) else 0
|
|
gpu = os.environ.get("DG_GPU", "0")
|
|
|
|
cmd = list(shim_cmd) + [
|
|
"--gguf",
|
|
model_path,
|
|
"--host",
|
|
"127.0.0.1",
|
|
"--port",
|
|
str(self._port),
|
|
"--gpu",
|
|
gpu,
|
|
"--maxtok",
|
|
str(maxtok),
|
|
]
|
|
|
|
env = child_env_without_native_path_secret()
|
|
env["DG_VISUAL_BIN"] = visual_bin
|
|
env["DG_GPU"] = gpu
|
|
# The file-override shim imports its sibling visual_engine; put its dir on PYTHONPATH.
|
|
# (The zoo-package shim is an installed module and needs no PYTHONPATH change.)
|
|
if extra_pythonpath:
|
|
existing = env.get("PYTHONPATH")
|
|
env["PYTHONPATH"] = (
|
|
(extra_pythonpath + os.pathsep + existing) if existing else extra_pythonpath
|
|
)
|
|
|
|
logger.info(f"Starting DiffusionGemma runner: {' '.join(cmd)}")
|
|
self._stdout_lines = []
|
|
self._llama_log_fh = None
|
|
self._llama_log_path = None
|
|
try:
|
|
log_dir = _swa_cache_path().parent / "logs" / "diffusion-server"
|
|
log_dir.mkdir(parents = True, exist_ok = True)
|
|
self._llama_log_path = log_dir / f"diffusion-{int(time.time())}-port-{self._port}.log"
|
|
self._llama_log_fh = open(self._llama_log_path, "w", encoding = "utf-8", buffering = 1)
|
|
logger.info(f"diffusion runner stdout/stderr -> {self._llama_log_path}")
|
|
except OSError as e:
|
|
logger.debug(f"Could not open diffusion runner log file: {e}")
|
|
|
|
# PR_SET_PDEATHSIG: the shim (and its visual server) die with this backend
|
|
# process, so a Studio crash/restart never orphans a GPU process.
|
|
popen_kwargs = dict(_windows_hidden_subprocess_kwargs())
|
|
if sys.platform.startswith("linux"): # prctl/libc.so.6 are Linux-only
|
|
|
|
def _pdeathsig():
|
|
try:
|
|
import ctypes
|
|
import signal as _signal
|
|
ctypes.CDLL("libc.so.6", use_errno = True).prctl(1, _signal.SIGTERM)
|
|
except Exception:
|
|
pass
|
|
|
|
popen_kwargs["preexec_fn"] = _pdeathsig
|
|
|
|
self._process = subprocess.Popen(
|
|
cmd,
|
|
stdout = subprocess.PIPE,
|
|
stderr = subprocess.STDOUT,
|
|
text = True,
|
|
env = env,
|
|
**popen_kwargs,
|
|
)
|
|
self._stdout_thread = threading.Thread(
|
|
target = self._drain_stdout, daemon = True, name = "diffusion-stdout"
|
|
)
|
|
self._stdout_thread.start()
|
|
|
|
# Publish state before the health wait (mirrors the llama-server path).
|
|
self._gguf_path = model_path
|
|
self._hf_repo = hf_repo
|
|
self._is_vision = False
|
|
self._is_audio = False # clear any prior TTS/audio model's routing flag
|
|
self._model_identifier = model_identifier
|
|
self._cache_type_kv = None
|
|
self._gpu_offload_active = True
|
|
if hf_variant:
|
|
self._hf_variant = hf_variant
|
|
elif gguf_path:
|
|
try:
|
|
from utils.models.model_config import _extract_quant_label
|
|
self._hf_variant = _extract_quant_label(gguf_path)
|
|
except Exception:
|
|
self._hf_variant = None
|
|
else:
|
|
self._hf_variant = None
|
|
# Provisional until the server reports the budget it resolved (auto-size picks it from VRAM).
|
|
self._effective_context_length = maxtok or self._context_length
|
|
self._max_context_length = self._context_length or maxtok or None
|
|
|
|
healthy = self._wait_for_health(timeout = 600.0)
|
|
if healthy:
|
|
self._healthy = True
|
|
self._gpu_offload_active = True
|
|
if extra_args is not None:
|
|
self._extra_args = list(extra_args)
|
|
self._extra_args_source = (model_identifier, hf_variant)
|
|
# The visual server logs "MAXTOK=<N>" with the context budget it actually resolved
|
|
# (auto-sized to VRAM). Read it back so the UI context bar shows the real budget.
|
|
chosen = maxtok
|
|
try:
|
|
import re as _re
|
|
for _ln in reversed(self._stdout_lines):
|
|
_m = _re.search(r"MAXTOK=(\d+)", _ln)
|
|
if _m:
|
|
chosen = int(_m.group(1))
|
|
break
|
|
except Exception:
|
|
pass
|
|
if chosen and chosen > 0:
|
|
self._effective_context_length = chosen
|
|
self._max_context_length = chosen
|
|
self._requested_n_ctx = int(n_ctx)
|
|
else:
|
|
self._healthy = False
|
|
logger.error("DiffusionGemma runner failed to become healthy")
|
|
return healthy
|
|
|
|
# ── HF download (no lock held) ───────────────────────────────
|
|
|
|
def _download_gguf(
|
|
self,
|
|
*,
|
|
hf_repo: str,
|
|
hf_variant: Optional[str] = None,
|
|
hf_token: Optional[str] = None,
|
|
) -> str:
|
|
"""Download GGUF file(s) from HuggingFace. Returns local path.
|
|
|
|
Runs WITHOUT self._lock so unload_model() can set _cancel_event at
|
|
any time; checks it between each shard download.
|
|
"""
|
|
try:
|
|
from huggingface_hub import hf_hub_download
|
|
except ImportError:
|
|
raise RuntimeError(
|
|
"huggingface_hub is required for HF model loading. "
|
|
"Install it with: pip install huggingface_hub"
|
|
)
|
|
|
|
# Resolve the filename from the variant
|
|
gguf_filename = None
|
|
gguf_extra_shards: list[str] = []
|
|
if hf_variant:
|
|
try:
|
|
from huggingface_hub import list_repo_files
|
|
|
|
files = list_repo_files(hf_repo, token = hf_token)
|
|
variant_lower = hf_variant.lower()
|
|
boundary = re.compile(
|
|
r"(?<![a-zA-Z0-9])" + re.escape(variant_lower) + r"(?![a-zA-Z0-9])"
|
|
)
|
|
gguf_files = sorted(
|
|
f
|
|
for f in files
|
|
if f.endswith(".gguf")
|
|
and boundary.search(f.lower())
|
|
and not _is_companion_gguf_path(f)
|
|
)
|
|
if gguf_files:
|
|
gguf_filename = gguf_files[0]
|
|
m = _SHARD_FULL_RE.match(gguf_filename)
|
|
if m:
|
|
prefix = m.group(1)
|
|
total = m.group(3)
|
|
sibling_pat = re.compile(
|
|
r"^" + re.escape(prefix) + r"-\d{5}-of-" + re.escape(total) + r"\.gguf$"
|
|
)
|
|
gguf_extra_shards = [f for f in gguf_files[1:] if sibling_pat.match(f)]
|
|
except Exception as e:
|
|
logger.warning(f"Could not list repo files: {e}")
|
|
|
|
# Offline: resolve variant -> filename from the local HF cache.
|
|
# The heuristic below assumes filenames echo the repo name, which
|
|
# breaks for e.g. Qwen3.6-27B-MTP-GGUF (no "MTP" in file). Match
|
|
# against the rel path (not just basename) so subdir layouts like
|
|
# ``BF16/foo.gguf`` are findable.
|
|
if not gguf_filename:
|
|
try:
|
|
from utils.models.model_config import _iter_hf_cache_snapshots
|
|
boundary = re.compile(
|
|
r"(?<![a-zA-Z0-9])" + re.escape(hf_variant.lower()) + r"(?![a-zA-Z0-9])"
|
|
)
|
|
for snap in _iter_hf_cache_snapshots(hf_repo):
|
|
matches = sorted(
|
|
p.relative_to(snap).as_posix()
|
|
for p in snap.rglob("*.gguf")
|
|
if not _is_companion_gguf_path(p.relative_to(snap).as_posix())
|
|
and boundary.search(p.relative_to(snap).as_posix().lower())
|
|
)
|
|
if not matches:
|
|
continue
|
|
gguf_filename = matches[0]
|
|
m = _SHARD_FULL_RE.match(Path(gguf_filename).name)
|
|
if m:
|
|
prefix = m.group(1)
|
|
total = m.group(3)
|
|
sibling_pat = re.compile(
|
|
r"^"
|
|
+ re.escape(prefix)
|
|
+ r"-\d{5}-of-"
|
|
+ re.escape(total)
|
|
+ r"\.gguf$"
|
|
)
|
|
gguf_extra_shards = [
|
|
f for f in matches[1:] if sibling_pat.match(Path(f).name)
|
|
]
|
|
logger.info(
|
|
"Resolved variant %s -> %s from local HF cache",
|
|
hf_variant,
|
|
gguf_filename,
|
|
)
|
|
break
|
|
except Exception as e:
|
|
logger.debug(f"Offline cache lookup for variant failed: {e}")
|
|
|
|
if not gguf_filename:
|
|
repo_name = hf_repo.split("/")[-1].replace("-GGUF", "")
|
|
gguf_filename = f"{repo_name}-{hf_variant}.gguf"
|
|
|
|
# Check disk space; fall back to a smaller variant if needed
|
|
all_gguf_files = [gguf_filename] + gguf_extra_shards
|
|
try:
|
|
from huggingface_hub import get_paths_info, try_to_load_from_cache
|
|
|
|
path_infos = list(get_paths_info(hf_repo, all_gguf_files, token = hf_token))
|
|
total_bytes = sum((p.size or 0) for p in path_infos)
|
|
|
|
# Subtract bytes already in the HF cache so we only preflight
|
|
# against what we must download. Without this, re-loading a
|
|
# cached large model (e.g. MiniMax-M2.7-GGUF at 131 GB) fails
|
|
# cold whenever free disk is below the full weight footprint,
|
|
# even though nothing needs downloading.
|
|
already_cached_bytes = 0
|
|
for p in path_infos:
|
|
if not p.size:
|
|
continue
|
|
try:
|
|
cached_path = try_to_load_from_cache(hf_repo, p.path)
|
|
except Exception:
|
|
cached_path = None
|
|
if isinstance(cached_path, str) and os.path.exists(cached_path):
|
|
try:
|
|
on_disk = os.path.getsize(cached_path)
|
|
except OSError:
|
|
on_disk = 0
|
|
# Satisfied only when the full blob is present.
|
|
if on_disk >= p.size:
|
|
already_cached_bytes += p.size
|
|
|
|
total_download_bytes = max(0, total_bytes - already_cached_bytes)
|
|
|
|
if total_download_bytes > 0:
|
|
cache_dir = os.environ.get(
|
|
"HF_HUB_CACHE",
|
|
str(Path.home() / ".cache" / "huggingface" / "hub"),
|
|
)
|
|
Path(cache_dir).mkdir(parents = True, exist_ok = True)
|
|
free_bytes = shutil.disk_usage(cache_dir).free
|
|
|
|
total_gb = total_download_bytes / (1024**3)
|
|
free_gb = free_bytes / (1024**3)
|
|
cached_gb = already_cached_bytes / (1024**3)
|
|
|
|
logger.info(
|
|
f"GGUF download: {total_gb:.1f} GB needed "
|
|
f"({cached_gb:.1f} GB already cached), "
|
|
f"{free_gb:.1f} GB free on disk"
|
|
)
|
|
|
|
if total_download_bytes > free_bytes:
|
|
smaller = self._find_smallest_fitting_variant(
|
|
hf_repo,
|
|
free_bytes,
|
|
hf_token,
|
|
)
|
|
if smaller:
|
|
fallback_file, fallback_size = smaller
|
|
logger.info(
|
|
f"Selected variant too large ({total_gb:.1f} GB), "
|
|
f"falling back to {fallback_file} ({fallback_size / (1024**3):.1f} GB)"
|
|
)
|
|
gguf_filename = fallback_file
|
|
_m = _SHARD_RE.match(gguf_filename)
|
|
_prefix = _m.group(1) if _m else None
|
|
if _prefix:
|
|
gguf_extra_shards = sorted(
|
|
f
|
|
for f in all_gguf_files
|
|
if f.startswith(_prefix)
|
|
and f != gguf_filename
|
|
and not _is_companion_gguf_path(f)
|
|
)
|
|
else:
|
|
gguf_extra_shards = []
|
|
else:
|
|
raise RuntimeError(
|
|
f"Not enough disk space to download any variant. "
|
|
f"Only {free_gb:.1f} GB free in {cache_dir}"
|
|
)
|
|
except RuntimeError:
|
|
raise
|
|
except Exception as e:
|
|
logger.warning(f"Could not check disk space: {e}")
|
|
|
|
gguf_label = f"{hf_repo}/{gguf_filename}" + (
|
|
f" (+{len(gguf_extra_shards)} shards)" if gguf_extra_shards else ""
|
|
)
|
|
logger.info(f"Resolving GGUF: {gguf_label}")
|
|
try:
|
|
if self._cancel_event.is_set():
|
|
raise RuntimeError("Cancelled")
|
|
dl_start = time.monotonic()
|
|
local_path = hf_hub_download(
|
|
repo_id = hf_repo,
|
|
filename = gguf_filename,
|
|
token = hf_token,
|
|
)
|
|
for shard in gguf_extra_shards:
|
|
if self._cancel_event.is_set():
|
|
raise RuntimeError("Cancelled")
|
|
logger.info(f"Resolving GGUF shard: {shard}")
|
|
hf_hub_download(
|
|
repo_id = hf_repo,
|
|
filename = shard,
|
|
token = hf_token,
|
|
)
|
|
except RuntimeError as e:
|
|
if "Cancelled" in str(e):
|
|
raise
|
|
raise RuntimeError(
|
|
f"Failed to download GGUF file '{gguf_filename}' from {hf_repo}: {e}"
|
|
)
|
|
except Exception as e:
|
|
raise RuntimeError(
|
|
f"Failed to download GGUF file '{gguf_filename}' from {hf_repo}: {e}"
|
|
)
|
|
|
|
dl_elapsed = time.monotonic() - dl_start
|
|
if dl_elapsed < 2.0:
|
|
logger.info(f"GGUF resolved from cache: {local_path}")
|
|
else:
|
|
logger.info(f"GGUF downloaded in {dl_elapsed:.1f}s: {local_path}")
|
|
return local_path
|
|
|
|
def _download_companion_gguf(
|
|
self,
|
|
*,
|
|
hf_repo: str,
|
|
hf_token: Optional[str],
|
|
pick: Callable[[list[str]], Optional[str]],
|
|
label: str,
|
|
) -> Optional[str]:
|
|
"""Resolve and fetch a companion GGUF (mmproj / MTP drafter) by name.
|
|
|
|
Tries the live repo file list, then the local HF cache snapshots
|
|
(offline, same fallback as _download_gguf), then hf_hub_download.
|
|
Runs WITHOUT self._lock (like _download_gguf); honors _cancel_event so
|
|
an /unload between the main download and here skips the fetch.
|
|
"""
|
|
if self._cancel_event.is_set():
|
|
return None
|
|
|
|
target: Optional[str] = None
|
|
try:
|
|
from huggingface_hub import list_repo_files
|
|
target = pick(list_repo_files(hf_repo, token = hf_token))
|
|
except Exception as e:
|
|
logger.debug(f"Could not list repo files for {label}: {e}")
|
|
|
|
if target is None:
|
|
try:
|
|
from utils.models.model_config import _iter_hf_cache_snapshots
|
|
for snap in _iter_hf_cache_snapshots(hf_repo):
|
|
rel_files = [p.relative_to(snap).as_posix() for p in snap.rglob("*.gguf")]
|
|
target = pick(rel_files)
|
|
if target is not None:
|
|
logger.info("Resolved %s %s from local HF cache", label, target)
|
|
break
|
|
except Exception as e:
|
|
logger.debug(f"Offline cache lookup for {label} failed: {e}")
|
|
|
|
if target is None or self._cancel_event.is_set():
|
|
return None
|
|
|
|
try:
|
|
from huggingface_hub import hf_hub_download
|
|
logger.info(f"Downloading {label}: {hf_repo}/{target}")
|
|
return hf_hub_download(
|
|
repo_id = hf_repo,
|
|
filename = target,
|
|
token = hf_token,
|
|
)
|
|
except Exception as e:
|
|
logger.warning(f"Could not download {label}: {e}")
|
|
return None
|
|
|
|
def _download_mmproj(
|
|
self,
|
|
*,
|
|
hf_repo: str,
|
|
hf_token: Optional[str] = None,
|
|
) -> Optional[str]:
|
|
"""Download the mmproj (vision projection) file from a GGUF repo.
|
|
|
|
Prefers mmproj-F16.gguf, else any mmproj*.gguf. Returns the local
|
|
path, or None if none exists.
|
|
"""
|
|
|
|
def _pick_mmproj(candidates: list[str]) -> Optional[str]:
|
|
mmproj_files = sorted(
|
|
f
|
|
for f in candidates
|
|
if f.lower().endswith(".gguf") and "mmproj" in Path(f).name.lower()
|
|
)
|
|
if not mmproj_files:
|
|
return None
|
|
for f in mmproj_files:
|
|
if f.lower().endswith("-f16.gguf"):
|
|
return f
|
|
return mmproj_files[0]
|
|
|
|
return self._download_companion_gguf(
|
|
hf_repo = hf_repo,
|
|
hf_token = hf_token,
|
|
pick = _pick_mmproj,
|
|
label = "mmproj",
|
|
)
|
|
|
|
def _download_mtp(
|
|
self,
|
|
*,
|
|
hf_repo: str,
|
|
hf_token: Optional[str] = None,
|
|
) -> Optional[str]:
|
|
"""Download the separate MTP drafter (speculative head) from a GGUF repo.
|
|
|
|
Targets the repo-root ``mtp-*.gguf`` companion -- the Q8_0 drafter
|
|
unsloth mirrors there for llama.cpp ``-hf`` auto-discovery (smallest,
|
|
recommended for speculation). Repos that bake the MTP head into the
|
|
main GGUF (e.g. Qwen) ship no such sibling and this returns None. The
|
|
higher-precision copies under ``MTP/`` are for explicit selection and
|
|
are intentionally skipped. Returns the local path, or None.
|
|
"""
|
|
|
|
def _pick_mtp(candidates: list[str]) -> Optional[str]:
|
|
mtp_files = sorted(
|
|
f
|
|
for f in candidates
|
|
if f.lower().endswith(".gguf") and Path(f).name.lower().startswith("mtp-")
|
|
)
|
|
return mtp_files[0] if mtp_files else None
|
|
|
|
return self._download_companion_gguf(
|
|
hf_repo = hf_repo,
|
|
hf_token = hf_token,
|
|
pick = _pick_mtp,
|
|
label = "MTP drafter",
|
|
)
|
|
|
|
def _resolve_launch_mmproj_path(
|
|
self, *, model_path: str, mmproj_path: Optional[str]
|
|
) -> Optional[str]:
|
|
"""Return mmproj_path iff it exists on disk AND matches the model family.
|
|
|
|
None if mmproj_path is None, missing, or family-mismatched.
|
|
"""
|
|
if not mmproj_path:
|
|
return None
|
|
|
|
mmproj = Path(mmproj_path)
|
|
if not mmproj.is_file():
|
|
logger.warning(f"mmproj file not found: {mmproj_path}")
|
|
return None
|
|
|
|
from utils.models.model_config import mmproj_matches_model_family
|
|
|
|
if not mmproj_matches_model_family(model_path, str(mmproj)):
|
|
logger.warning(
|
|
f"mmproj does not match model family: model={Path(model_path).name} "
|
|
f"mmproj={mmproj.name}"
|
|
)
|
|
return None
|
|
|
|
return str(mmproj)
|
|
|
|
def _resolve_launch_mtp_path(self, *, mtp_draft_path: Optional[str]) -> Optional[str]:
|
|
"""Return mtp_draft_path iff it exists on disk, else None.
|
|
|
|
No family check needed: the drafter is only ever auto-resolved from
|
|
the same repo as the main GGUF (see _download_mtp).
|
|
"""
|
|
if not mtp_draft_path:
|
|
return None
|
|
if not Path(mtp_draft_path).is_file():
|
|
logger.warning(f"MTP drafter file not found: {mtp_draft_path}")
|
|
return None
|
|
return str(mtp_draft_path)
|
|
|
|
# ── Lifecycle ─────────────────────────────────────────────────
|
|
|
|
# GGUF ``general.architecture`` values for diffusion / image models.
|
|
# llama.cpp has no such architectures, so loading one as a chat model dies
|
|
# with "unknown model architecture: '<arch>'". These match the patched
|
|
# stable-diffusion.cpp / ComfyUI-GGUF enums. Unsloth publishes FLUX and
|
|
# Qwen-Image GGUFs under
|
|
# https://huggingface.co/collections/unsloth/unsloth-diffusion-ggufs.
|
|
# Matched exactly (not a substring) so a chat arch containing "wan"/"sd1"
|
|
# (e.g. "taiwan") isn't misrouted to Images.
|
|
_DIFFUSION_ARCHES = frozenset(
|
|
(
|
|
"qwen_image",
|
|
"flux",
|
|
"sd1",
|
|
"sdxl",
|
|
"sd3",
|
|
"aura",
|
|
"hidream",
|
|
"cosmos",
|
|
"ltxv",
|
|
"hyvid",
|
|
"wan",
|
|
"lumina2",
|
|
)
|
|
)
|
|
|
|
@staticmethod
|
|
def _classify_llama_start_failure(
|
|
output: str, gguf_path: Optional[str], model_identifier: Optional[str]
|
|
) -> str:
|
|
"""Explain *why* llama-server failed to start, from its output.
|
|
|
|
Several distinct failures otherwise collapse into the same opaque
|
|
"invalid GGUF or out of memory" message. Worst case: a diffusion GGUF
|
|
loaded as a chat model -- valid file, plenty of memory, but llama.cpp
|
|
has no such architecture, so the user is told to free memory that was
|
|
never the problem (#5842). Pick the most specific message we can.
|
|
"""
|
|
lowered = (output or "").lower()
|
|
|
|
# Tensor parallelism (--split-mode tensor) is arch-gated in llama.cpp;
|
|
# unsupported architectures abort the load with this marker. Point the
|
|
# user at the toggle instead of a generic invalid-GGUF/OOM message.
|
|
if "split_mode_tensor not implemented" in lowered:
|
|
return (
|
|
"Tensor parallelism is not supported for this model's "
|
|
"architecture. Turn off Tensor Parallelism in the model "
|
|
"settings and reload."
|
|
)
|
|
|
|
# Detect Ollama source up front so the arch branch can keep the
|
|
# Ollama hint instead of the generic "unsupported arch" message.
|
|
gguf = gguf_path or ""
|
|
is_ollama = (
|
|
".studio_links" in gguf
|
|
or os.sep + "ollama_links" + os.sep in gguf
|
|
or os.sep + ".cache" + os.sep + "ollama" + os.sep in gguf
|
|
or (model_identifier or "").startswith("ollama/")
|
|
)
|
|
|
|
# "unknown model architecture: '<arch>'": diffusion -> Images page,
|
|
# Ollama -> Ollama hint, else a precise "unsupported" message. Exact
|
|
# match so chat archs aren't misrouted.
|
|
arch_match = re.search(r"unknown model architecture:\s*'([^']+)'", lowered)
|
|
if arch_match:
|
|
arch = arch_match.group(1)
|
|
if arch in LlamaCppBackend._DIFFUSION_ARCHES:
|
|
return (
|
|
f"'{arch}' is a diffusion (image-generation) GGUF, which "
|
|
"llama-server cannot run as a chat/completion model. Use "
|
|
"Studio's Images page to generate with local diffusion "
|
|
"GGUFs such as FLUX and Qwen-Image."
|
|
)
|
|
if is_ollama:
|
|
return (
|
|
"Some Ollama models do not work with llama.cpp. Try a "
|
|
"different model, or use this model directly through "
|
|
"Ollama instead."
|
|
)
|
|
return (
|
|
f"llama.cpp does not support this GGUF's model architecture "
|
|
f"('{arch}'). The file is valid, but this model type cannot "
|
|
"be run with llama-server."
|
|
)
|
|
|
|
# Other Ollama compat failures that don't name an arch. Only when
|
|
# the output shows a GGUF compat issue, not OOM / missing binaries.
|
|
if is_ollama:
|
|
gguf_compat_hints = (
|
|
"key not found",
|
|
"unknown model architecture",
|
|
"failed to load model",
|
|
)
|
|
if any(h in lowered for h in gguf_compat_hints):
|
|
return (
|
|
"Some Ollama models do not work with llama.cpp. Try a "
|
|
"different model, or use this model directly through "
|
|
"Ollama instead."
|
|
)
|
|
|
|
# Fallback: genuinely unknown failure (OOM, missing binary ...).
|
|
return (
|
|
"llama-server failed to start. "
|
|
"Check that the GGUF file is valid and you have enough memory."
|
|
)
|
|
|
|
def _plan_tensor_parallel(
|
|
self,
|
|
gpus: list[tuple[int, int]],
|
|
model_size: int,
|
|
target_ctx: int,
|
|
cache_type_kv: Optional[str] = None,
|
|
n_parallel: int = 1,
|
|
mtp_engaged: bool = False,
|
|
max_target_ctx: Optional[int] = None,
|
|
) -> tuple[int, int, list[int], Optional[list[int]]]:
|
|
"""Plan a ``--split-mode tensor`` load. Pure: no model or GPU needed.
|
|
|
|
``gpus`` is a list of ``(gpu_index, free_mib)``; ``model_size`` is the
|
|
weight size in bytes; ``target_ctx`` is the context to fit (the explicit
|
|
request, or the model's native length for auto). ``max_target_ctx`` is
|
|
the native/hardware ceiling used only for the UI bound (defaults to
|
|
``target_ctx``). Returns
|
|
``(effective_ctx, max_available_ctx, gpu_indices, tensor_split)``.
|
|
|
|
Policy (assumes >= 2 GPUs; the caller drops the toggle below that):
|
|
- Cap context to the KV that fits the pooled VRAM after the weights and
|
|
one per-device compute-graph buffer (``_TENSOR_PARALLEL_BUFFER_RESERVE_MIB``).
|
|
llama.cpp's ``--fit`` is a no-op in tensor mode, so this is the only
|
|
cap, honored even for an explicit ``-c``. It is more accurate than the
|
|
0.80 whole-pool heuristic, which over-reserves and leaves VRAM unused.
|
|
- ``tensor_split`` is None (llama.cpp's even default, safe for every arch
|
|
incl. Gemma 3n which GGML_ASSERTs on a weighted split) when an even
|
|
share fits the smallest GPU; otherwise it is weighted by
|
|
``(free - buffer)`` so the roomier GPU absorbs more weight and the
|
|
smallest GPU keeps room for KV.
|
|
"""
|
|
# Drop GPUs that can't hold the per-device compute-graph buffer; they'd
|
|
# OOM in tensor mode. load_model already filters before calling, so this
|
|
# is defense-in-depth that also keeps the pure function self-contained
|
|
# (and unit-testable without a GPU).
|
|
reserve_mib = self._TENSOR_PARALLEL_BUFFER_RESERVE_MIB
|
|
usable_gpus = [g for g in gpus if g[1] >= reserve_mib]
|
|
gpu_indices = sorted(idx for idx, _ in usable_gpus)
|
|
if len(gpu_indices) < 2:
|
|
# Tensor parallelism is meaningless on <2 GPUs (the caller drops the
|
|
# toggle before this); be defensive and never emit a split here.
|
|
return (
|
|
target_ctx if target_ctx > 0 else 4096,
|
|
target_ctx if target_ctx > 0 else 4096,
|
|
gpu_indices,
|
|
None,
|
|
)
|
|
free_by_idx = {idx: free for idx, free in usable_gpus}
|
|
pool_mib = sum(free_by_idx.values())
|
|
kv_budget_b = (pool_mib - len(gpu_indices) * reserve_mib) * 1024 * 1024 - model_size
|
|
if mtp_engaged:
|
|
# MTP keeps a draft model + its own KV cache on GPU.
|
|
kv_budget_b -= 2 * 1024**3
|
|
|
|
def _fit_ctx(ctx: int) -> int:
|
|
# Largest context whose KV fits the pooled budget. Floors small, but
|
|
# never raises an explicit ctx above what was asked.
|
|
if self._can_estimate_kv() and ctx > 0:
|
|
ctx_floor = min(2048, ctx)
|
|
if kv_budget_b <= 0:
|
|
# Weights + buffers exceed the pool -> floor; the load then
|
|
# falls back to layer split.
|
|
return ctx_floor
|
|
kv_at = self._estimate_kv_cache_bytes(ctx, cache_type_kv, n_parallel = n_parallel)
|
|
if kv_at <= kv_budget_b:
|
|
return ctx
|
|
return max(ctx_floor, int(ctx * kv_budget_b / kv_at))
|
|
# KV size unknown -> can't prove a safe cap; floor.
|
|
return min(4096, ctx) if ctx > 0 else 4096
|
|
|
|
# max_available_ctx is the hardware ceiling for the UI bound, sized from
|
|
# the native context independent of an explicit small -c (which only
|
|
# caps effective_ctx).
|
|
max_ctx_target = max_target_ctx if (max_target_ctx and max_target_ctx > 0) else target_ctx
|
|
max_available_ctx = _fit_ctx(max_ctx_target)
|
|
effective_ctx = min(_fit_ctx(target_ctx), max_available_ctx)
|
|
|
|
min_free_mib = min(free_by_idx.values())
|
|
kv_bytes = (
|
|
self._estimate_kv_cache_bytes(effective_ctx, cache_type_kv, n_parallel = n_parallel)
|
|
if (self._can_estimate_kv() and effective_ctx > 0)
|
|
else 0
|
|
)
|
|
even_share_mib = (model_size + kv_bytes) / len(gpu_indices) / (1024 * 1024)
|
|
tensor_split: Optional[list[int]] = None
|
|
if even_share_mib > (min_free_mib - reserve_mib):
|
|
adj = [max(0, int(free_by_idx[i] - reserve_mib)) for i in gpu_indices]
|
|
if sum(adj) > 0:
|
|
tensor_split = adj
|
|
return effective_ctx, max_available_ctx, gpu_indices, tensor_split
|
|
|
|
@staticmethod
|
|
def _is_projector_incompatibility(output: str) -> bool:
|
|
"""True when llama-server aborted because it cannot load the model's
|
|
vision/audio projector (mmproj), typically an installed llama.cpp
|
|
that predates the projector format. Conservative: only matches
|
|
projector-format errors so unrelated failures (OOM, bad GGUF, port
|
|
bind, ...) keep their own handling, and a bare 'clip'/'mmproj'
|
|
mention in a normal startup log does not match.
|
|
"""
|
|
text = (output or "").lower()
|
|
if any(
|
|
m in text
|
|
for m in (
|
|
"unknown projector type",
|
|
"unsupported projector",
|
|
"unsupported mmproj",
|
|
)
|
|
):
|
|
return True
|
|
# Builds that phrase it via clip.cpp without the exact words above.
|
|
return (
|
|
"clip" in text
|
|
and "projector" in text
|
|
and ("unknown" in text or "unsupported" in text or "not supported" in text)
|
|
)
|
|
|
|
@staticmethod
|
|
def _strip_mmproj_args(cmd: list[str]) -> list[str]:
|
|
"""Return cmd without the '--mmproj <path>' pair (text-only retry).
|
|
Every other flag is preserved; a no-op when --mmproj is absent.
|
|
"""
|
|
out: list[str] = []
|
|
skip_value = False
|
|
for tok in cmd:
|
|
if skip_value:
|
|
skip_value = False
|
|
continue
|
|
if tok == "--mmproj":
|
|
skip_value = True
|
|
continue
|
|
out.append(tok)
|
|
return out
|
|
|
|
def _start_llama_process(self, cmd: list[str], env: dict) -> None:
|
|
"""Spawn llama-server from cmd and start draining its output.
|
|
|
|
Caller holds self._lock. Resets the stdout buffer, opens a fresh
|
|
per-attempt tee log, launches the process, and starts the drain
|
|
thread. Used for the initial start and the text-only mmproj retry.
|
|
"""
|
|
# Defensive kill: if a concurrent load slipped past Phase 1
|
|
# (because its `self._process` was None at the time) and already
|
|
# stored a Popen handle here, drop that orphan before we overwrite
|
|
# the reference. See issue #5161.
|
|
self._kill_process()
|
|
|
|
self._stdout_lines = []
|
|
# Tee llama-server output to a dedicated log file so a post-mortem
|
|
# in CI (or after a remote-debug session) has the full subprocess
|
|
# trail even when the parent only stored the last 50 lines.
|
|
self._llama_log_fh = None
|
|
try:
|
|
log_dir = _swa_cache_path().parent / "logs" / "llama-server"
|
|
log_dir.mkdir(parents = True, exist_ok = True)
|
|
self._llama_log_path = log_dir / f"llama-{int(time.time())}-port-{self._port}.log"
|
|
self._llama_log_fh = open(
|
|
self._llama_log_path,
|
|
"w",
|
|
encoding = "utf-8",
|
|
buffering = 1,
|
|
)
|
|
logger.info(f"llama-server stdout/stderr -> {self._llama_log_path}")
|
|
except OSError as e:
|
|
# Best-effort; never block the load on logging.
|
|
logger.debug(f"Could not open llama-server log file: {e}")
|
|
self._llama_log_path = None
|
|
|
|
# Log the argv per attempt (the text-only mmproj retry re-enters here
|
|
# with --mmproj stripped), redacting the API key.
|
|
_log_cmd = list(cmd)
|
|
if "--api-key" in _log_cmd:
|
|
_ki = _log_cmd.index("--api-key") + 1
|
|
if _ki < len(_log_cmd):
|
|
_log_cmd[_ki] = "<redacted>"
|
|
logger.info(f"Starting llama-server: {' '.join(_log_cmd)}")
|
|
|
|
self._process = subprocess.Popen(
|
|
cmd,
|
|
stdout = subprocess.PIPE,
|
|
stderr = subprocess.STDOUT,
|
|
text = True,
|
|
env = env,
|
|
**_windows_hidden_subprocess_kwargs(),
|
|
)
|
|
|
|
# Start background thread to drain stdout and prevent pipe deadlock
|
|
self._stdout_thread = threading.Thread(
|
|
target = self._drain_stdout, daemon = True, name = "llama-stdout"
|
|
)
|
|
self._stdout_thread.start()
|
|
|
|
def load_model(
|
|
self,
|
|
*,
|
|
# Local mode: pass a path to a .gguf file
|
|
gguf_path: Optional[str] = None,
|
|
# Vision projection (mmproj) for local vision models
|
|
mmproj_path: Optional[str] = None,
|
|
# Separate MTP drafter for local Gemma loads (HF loads auto-resolve it)
|
|
mtp_draft_path: Optional[str] = None,
|
|
# HF mode: let llama-server download via -hf "repo:quant"
|
|
hf_repo: Optional[str] = None,
|
|
hf_variant: Optional[str] = None,
|
|
hf_token: Optional[str] = None,
|
|
# Common
|
|
model_identifier: str,
|
|
is_vision: bool = False,
|
|
n_ctx: int = 4096,
|
|
chat_template_override: Optional[str] = None,
|
|
cache_type_kv: Optional[str] = None,
|
|
speculative_type: Optional[str] = None,
|
|
spec_draft_n_max: Optional[int] = None,
|
|
tensor_parallel: bool = False,
|
|
n_threads: Optional[int] = None,
|
|
n_gpu_layers: Optional[int] = None, # caller compat, unused
|
|
n_parallel: int = 1,
|
|
extra_args: Optional[List[str]] = None,
|
|
) -> bool:
|
|
"""Start llama-server with a GGUF model.
|
|
|
|
Two modes:
|
|
- Local: ``gguf_path="/path/to/model.gguf"`` → uses ``-m``
|
|
- HF: ``hf_repo="...-GGUF", hf_variant="Q4_K_M"`` → uses ``-hf``
|
|
|
|
Returns True if the server started and the health check passed.
|
|
"""
|
|
# Serialise the whole load so concurrent /load calls never leave two
|
|
# llama-server processes alive (#5401 / #5161). Doesn't block /unload.
|
|
with self._serial_load_lock:
|
|
# In-app update swapping binaries: refuse fast (set under this lock,
|
|
# so any in-flight load has drained) instead of using a half-swapped one.
|
|
if getattr(self, "_llama_update_in_progress", False):
|
|
raise RuntimeError("llama.cpp is updating; try again in a moment.")
|
|
# Duplicate /load that raced past the route check: do nothing if the
|
|
# live server already satisfies this request.
|
|
if self._already_in_target_state(
|
|
gguf_path = gguf_path,
|
|
mtp_draft_path = mtp_draft_path,
|
|
model_identifier = model_identifier,
|
|
hf_variant = hf_variant,
|
|
n_ctx = n_ctx,
|
|
cache_type_kv = cache_type_kv,
|
|
speculative_type = speculative_type,
|
|
spec_draft_n_max = spec_draft_n_max,
|
|
tensor_parallel = tensor_parallel,
|
|
chat_template_override = chat_template_override,
|
|
extra_args = extra_args,
|
|
is_vision = is_vision,
|
|
):
|
|
logger.info(
|
|
f"load_model: backend already in target state for "
|
|
f"'{model_identifier}', skipping reload"
|
|
)
|
|
# Retry probe only if a prior attempt didn't finish.
|
|
if not self._audio_probed:
|
|
try:
|
|
detected = self._detect_audio_type_strict()
|
|
self._audio_probed = True
|
|
except Exception as exc:
|
|
logger.debug("Fast-path audio probe failed: %s", exc)
|
|
detected = None
|
|
if detected in ("snac", "bicodec", "dac"):
|
|
with self._lock:
|
|
if not self._healthy:
|
|
return False
|
|
try:
|
|
self.init_audio_codec(detected)
|
|
self._is_audio = True
|
|
self._audio_type = detected
|
|
except Exception as exc:
|
|
logger.warning(
|
|
"Failed to init audio codec '%s': %s",
|
|
detected,
|
|
exc,
|
|
)
|
|
self._audio_probed = False
|
|
return False
|
|
elif detected:
|
|
# csm / whisper / audio_vlm: track type but keep
|
|
# _is_audio False -- GGUF TTS routing only fires for
|
|
# snac/bicodec/dac.
|
|
with self._lock:
|
|
if not self._healthy:
|
|
return False
|
|
self._audio_type = detected
|
|
# Re-derive after a retried probe (_mmproj_has_audio persists).
|
|
from utils.models.model_config import is_audio_input_type
|
|
|
|
self._has_audio_input = bool(is_audio_input_type(self._audio_type)) or bool(
|
|
self._mmproj_has_audio
|
|
)
|
|
if not self._healthy:
|
|
return False
|
|
return True
|
|
|
|
self._cancel_event.clear()
|
|
|
|
# ── Phase 1: kill old process (under lock, fast) ──────────
|
|
with self._lock:
|
|
self._kill_process()
|
|
|
|
# Resolve llama-server now but defer a not-found error: a block-diffusion
|
|
# GGUF uses the diffusion runner, and its arch is only known after the header.
|
|
binary = self._find_llama_server_binary()
|
|
|
|
# ── Phase 2: download (NO lock held, so cancel can proceed) ──
|
|
# mtp_draft_path arrives set for local Gemma loads (detected
|
|
# sibling); for -hf loads it's None here and resolved just below.
|
|
# Scope HF_HUB_OFFLINE to the download block only when DNS is
|
|
# dead; cleanup runs even on exception so a transient hiccup
|
|
# can't quarantine future loads.
|
|
if hf_repo:
|
|
with _hf_offline_if_dns_dead():
|
|
model_path = self._download_gguf(
|
|
hf_repo = hf_repo,
|
|
hf_variant = hf_variant,
|
|
hf_token = hf_token,
|
|
)
|
|
# Auto-download mmproj for vision models unless opted out.
|
|
if is_vision and not mmproj_path and not extra_args_disable_mmproj(extra_args):
|
|
mmproj_path = self._download_mmproj(
|
|
hf_repo = hf_repo,
|
|
hf_token = hf_token,
|
|
)
|
|
# Auto-download the separate MTP drafter (e.g. Gemma) when
|
|
# the requested spec mode can use it. Repos with the head
|
|
# baked into the main GGUF (Qwen) have no mtp- sibling and
|
|
# this no-ops, so the size gate stays out of it: a separate
|
|
# drafter speeds up even sub-3B (Gemma E2B), and the resolver
|
|
# below decides the final emission. Skipped only when the
|
|
# user disabled MTP or drives --spec-type manually.
|
|
_spec_canon = _canonicalize_spec_mode(speculative_type) or "auto"
|
|
if (
|
|
not mtp_draft_path
|
|
and _spec_canon in ("auto", "mtp", "mtp+ngram")
|
|
and not _extra_args_set_spec_type(extra_args)
|
|
):
|
|
mtp_draft_path = self._download_mtp(
|
|
hf_repo = hf_repo,
|
|
hf_token = hf_token,
|
|
)
|
|
elif gguf_path:
|
|
if not Path(gguf_path).is_file():
|
|
raise FileNotFoundError(f"GGUF file not found: {gguf_path}")
|
|
model_path = gguf_path
|
|
else:
|
|
raise ValueError("Either gguf_path or hf_repo must be provided")
|
|
|
|
# Set identifier early so _read_gguf_metadata can use it (DeepSeek).
|
|
self._model_identifier = model_identifier
|
|
|
|
# Read GGUF metadata (context_length, chat_template); header-only.
|
|
self._read_gguf_metadata(model_path)
|
|
|
|
if self._cancel_event.is_set():
|
|
logger.info("Load cancelled after download phase")
|
|
return False
|
|
|
|
# Block-diffusion GGUFs (DiffusionGemma) cannot run on llama-server;
|
|
# serve them with the diffusion runner (same OpenAI-compat interface).
|
|
if self._is_diffusion:
|
|
with self._lock:
|
|
if self._cancel_event.is_set():
|
|
logger.info("Load cancelled before diffusion server start")
|
|
return False
|
|
return self._start_diffusion_server(
|
|
model_path = model_path,
|
|
gguf_path = gguf_path,
|
|
hf_repo = hf_repo,
|
|
hf_variant = hf_variant,
|
|
model_identifier = model_identifier,
|
|
n_ctx = n_ctx,
|
|
extra_args = extra_args,
|
|
)
|
|
|
|
if not binary:
|
|
raise RuntimeError(
|
|
"llama-server binary not found. "
|
|
"Run setup.sh to build it, install llama.cpp, "
|
|
"or set LLAMA_SERVER_PATH environment variable."
|
|
)
|
|
|
|
# Outside ``self._lock`` so /unload, /cancel, /status aren't
|
|
# blocked. ``unload_model`` also records the kill, so the
|
|
# frontend /unload+/load Apply path engages the wait here even
|
|
# without an in-process kill.
|
|
self._wait_for_vram_settle(since_kill = self._last_kill_monotonic)
|
|
|
|
# ── Phase 3: start llama-server (under lock) ──────────────
|
|
with self._lock:
|
|
# Re-check cancel inside lock
|
|
if self._cancel_event.is_set():
|
|
logger.info("Load cancelled before server start")
|
|
return False
|
|
|
|
self._port = self._find_free_port()
|
|
|
|
# Select GPU(s) from model size + estimated KV cache. Seed
|
|
# safe defaults before probing so the except path has valid
|
|
# state to publish.
|
|
ctx_override = parse_ctx_override(extra_args)
|
|
requested_ctx = resolve_requested_ctx(extra_args, n_ctx)
|
|
cache_override = parse_cache_override(extra_args)
|
|
cache_type_kv = resolve_cache_type_kv(extra_args, cache_type_kv)
|
|
# A user --split-mode in extras last-wins-overrides the
|
|
# toggle, so reconcile it back into tensor_parallel state.
|
|
split_mode_override = parse_split_mode_override(extra_args)
|
|
tensor_parallel = resolve_tensor_parallel(extra_args, tensor_parallel)
|
|
# Tensor mode aborts on a quantized KV cache, so drop it for the
|
|
# tensor attempt (and strip any inherited/explicit --cache-type
|
|
# that would re-impose it when appended last). The layer-split
|
|
# fallback re-runs with tensor_parallel False and keeps the type.
|
|
if (
|
|
tensor_parallel
|
|
and cache_type_kv
|
|
and cache_type_kv.strip().lower() not in self._TENSOR_PARALLEL_KV_TYPES
|
|
):
|
|
logger.info(
|
|
"Tensor parallelism requires a non-quantized KV cache; "
|
|
"ignoring cache type %s for the tensor attempt.",
|
|
cache_type_kv,
|
|
)
|
|
cache_type_kv = None
|
|
if extra_args:
|
|
extra_args = strip_shadowing_flags(
|
|
extra_args,
|
|
strip_context = False,
|
|
strip_cache = True,
|
|
strip_spec = False,
|
|
strip_template = False,
|
|
strip_split_mode = False,
|
|
)
|
|
if ctx_override is not None and ctx_override > 0:
|
|
logger.info(f"User --ctx-size {ctx_override} honored; skipping auto-reduce")
|
|
if cache_override is not None:
|
|
logger.info(f"User --cache-type-k/-v {cache_override} honored for KV estimate")
|
|
if split_mode_override is not None:
|
|
logger.info(
|
|
f"User --split-mode {split_mode_override} honored; "
|
|
"reconciled into tensor_parallel state"
|
|
)
|
|
effective_ctx = requested_ctx if requested_ctx > 0 else (self._context_length or 0)
|
|
max_available_ctx = self._context_length or effective_ctx
|
|
gpus: list[tuple[int, int]] = []
|
|
try:
|
|
model_size = self._get_gguf_size_bytes(model_path)
|
|
gpus = self._get_gpu_free_memory()
|
|
|
|
# Resolve effective context: 0 means let llama-server use
|
|
# the model's native length. Only expand to a known native
|
|
# length if metadata exists; else keep 0 as a sentinel.
|
|
if requested_ctx > 0:
|
|
effective_ctx = requested_ctx
|
|
elif self._context_length is not None:
|
|
effective_ctx = self._context_length
|
|
else:
|
|
effective_ctx = 0
|
|
original_ctx = effective_ctx
|
|
# Default UI ceiling to the native context length;
|
|
# GPU/VRAM-fit logic below may shrink it on limited HW.
|
|
max_available_ctx = self._context_length or effective_ctx
|
|
|
|
# Will MTP engage on this load? If so, auto-fit reserves
|
|
# extra VRAM for the draft model. Mirrors
|
|
# _build_speculative_flags' resolver: forced mtp / mtp+ngram
|
|
# always engage; auto only on an MTP model >= 3B; ngram /
|
|
# ngram-simple / off never engage MTP. A separate drafter
|
|
# (Gemma) counts as an MTP model just like a baked-in head.
|
|
_mtp_canonical = _canonicalize_spec_mode(speculative_type)
|
|
_mtp_effective = _mtp_canonical or "auto"
|
|
_mtp_size_for_fit = _extract_model_size_b(model_identifier)
|
|
# Sub-3B drops MTP only for an embedded head; a separate
|
|
# drafter (Gemma) engages and needs its VRAM reserved.
|
|
_mtp_sub_3b_for_fit = (
|
|
_mtp_size_for_fit is not None
|
|
and _mtp_size_for_fit < _MTP_MIN_SIZE_B
|
|
and not bool(mtp_draft_path)
|
|
)
|
|
_mtp_will_engage = bool(
|
|
not _extra_args_set_spec_type(extra_args)
|
|
and (
|
|
_mtp_effective in ("mtp", "mtp+ngram")
|
|
or (
|
|
_mtp_effective == "auto"
|
|
and (
|
|
bool(self._nextn_predict_layers)
|
|
or _is_mtp_model_name(model_identifier, model_path)
|
|
or bool(mtp_draft_path)
|
|
)
|
|
and not _mtp_sub_3b_for_fit
|
|
)
|
|
)
|
|
)
|
|
|
|
# Auto-cap context to fit GPU VRAM and select GPUs. Two
|
|
# policies by whether the user set n_ctx:
|
|
# Explicit n_ctx: honor it. Try the full context with
|
|
# _select_gpus (as many GPUs as needed); cap only if it
|
|
# doesn't fit on any combination.
|
|
# Auto n_ctx=0 (native): prefer fewer GPUs with reduced
|
|
# context, since multi-GPU is slower.
|
|
gpu_indices, use_fit = None, True
|
|
# Per-GPU weight proportions for tensor mode (None = even).
|
|
tp_tensor_split: Optional[list[int]] = None
|
|
explicit_ctx = requested_ctx > 0
|
|
# MTP draft model lives outside the main estimates; carve
|
|
# its reserve out of every fit budget and pin threshold so
|
|
# a load can't pin into the drafter's headroom.
|
|
_mtp_reserve = _MTP_VRAM_RESERVE_FRAC if _mtp_will_engage else 0.0
|
|
_pin_fraction = self._GPU_PIN_VRAM_FRACTION - _mtp_reserve
|
|
|
|
# Tensor mode allocates a compute-graph buffer on every
|
|
# participating GPU, so a GPU with less free VRAM than that
|
|
# reserve can't host it and would OOM at load. Drop those
|
|
# from the tensor-parallel set up front (gpu_indices below
|
|
# becomes the CUDA_VISIBLE_DEVICES mask, so they're excluded
|
|
# from llama-server entirely, not just given zero weight).
|
|
tp_gpus = gpus
|
|
if tensor_parallel:
|
|
reserve_mib = self._TENSOR_PARALLEL_BUFFER_RESERVE_MIB
|
|
tp_gpus = [g for g in gpus if g[1] >= reserve_mib]
|
|
|
|
if tensor_parallel and len(tp_gpus) < 2:
|
|
# Tensor parallelism needs >= 2 usable GPUs. On a single
|
|
# GPU --split-mode tensor is a no-op; with 0 GPUs (CPU-only
|
|
# or probe failed) it must not reach llama-server; and a
|
|
# GPU below the buffer reserve can't participate. Drop the
|
|
# flag and fall through to normal layer/CPU allocation.
|
|
logger.info(
|
|
"Tensor parallelism requested but only %d of %d GPU(s) "
|
|
"have enough free VRAM for the compute buffer; "
|
|
"ignoring (needs >= 2).",
|
|
len(tp_gpus),
|
|
len(gpus),
|
|
)
|
|
tensor_parallel = False
|
|
# A user --split-mode tensor in extras is appended after
|
|
# Studio's flags, so it would still reach llama-server and
|
|
# fail here; strip it so the downgrade actually applies.
|
|
extra_args = strip_split_mode_only(extra_args)
|
|
|
|
if tensor_parallel and tp_gpus:
|
|
# Tensor-parallel allocation: use all usable GPUs, weight
|
|
# the split by (free - buffer), and cap context to the
|
|
# pooled VRAM after weights + per-device compute-graph
|
|
# buffers. See _plan_tensor_parallel for the policy.
|
|
target_ctx = (
|
|
effective_ctx
|
|
if explicit_ctx
|
|
else (self._context_length or effective_ctx)
|
|
)
|
|
(
|
|
effective_ctx,
|
|
max_available_ctx,
|
|
gpu_indices,
|
|
tp_tensor_split,
|
|
) = self._plan_tensor_parallel(
|
|
tp_gpus,
|
|
model_size,
|
|
target_ctx,
|
|
cache_type_kv = cache_type_kv,
|
|
n_parallel = n_parallel,
|
|
mtp_engaged = _mtp_will_engage,
|
|
# Report the UI ceiling from native ctx, not the
|
|
# explicit small request.
|
|
max_target_ctx = self._context_length or target_ctx,
|
|
)
|
|
use_fit = False
|
|
elif gpus and self._can_estimate_kv() and effective_ctx > 0:
|
|
# Compute the largest hardware-aware cap from the model's
|
|
# native context across all usable GPU subsets (for UI
|
|
# bounds), independent of the currently requested context.
|
|
native_ctx_for_cap = self._context_length or effective_ctx
|
|
if native_ctx_for_cap > 0:
|
|
ranked_for_cap = sorted(gpus, key = lambda g: g[1], reverse = True)
|
|
best_cap = 0
|
|
for n_gpus in range(1, len(ranked_for_cap) + 1):
|
|
subset = ranked_for_cap[:n_gpus]
|
|
pool_mib = sum(free for _, free in subset)
|
|
capped = self._fit_context_to_vram(
|
|
native_ctx_for_cap,
|
|
pool_mib,
|
|
model_size,
|
|
cache_type_kv,
|
|
n_parallel = n_parallel,
|
|
mtp_engaged = _mtp_will_engage,
|
|
)
|
|
kv = self._estimate_kv_cache_bytes(
|
|
capped, cache_type_kv, n_parallel = n_parallel
|
|
)
|
|
total_mib = (model_size + kv) / (1024 * 1024)
|
|
if total_mib <= pool_mib * (_CTX_FIT_VRAM_FRACTION - _mtp_reserve):
|
|
best_cap = max(best_cap, capped)
|
|
if best_cap > 0:
|
|
max_available_ctx = best_cap
|
|
else:
|
|
# Weights exceed 90% of every GPU subset, so no
|
|
# context fits. Anchor the UI "safe zone" at 4096
|
|
# so the slider warns above the fallback.
|
|
max_available_ctx = min(4096, native_ctx_for_cap)
|
|
|
|
if explicit_ctx:
|
|
# Honor the requested context verbatim. If it fits,
|
|
# pin GPUs and skip --fit; else ship -c <ctx> --fit
|
|
# on and let llama-server flex -ngl (CPU offload).
|
|
requested_total = model_size + self._estimate_kv_cache_bytes(
|
|
effective_ctx, cache_type_kv, n_parallel = n_parallel
|
|
)
|
|
gpu_indices, use_fit = self._select_gpus(
|
|
requested_total, gpus, usable_fraction = _pin_fraction
|
|
)
|
|
# No silent shrink: effective_ctx stays == requested_ctx.
|
|
else:
|
|
# Auto context: prefer fewer GPUs, cap to fit. Same
|
|
# headroom threshold as _select_gpus (#5106).
|
|
ranked = sorted(gpus, key = lambda g: g[1], reverse = True)
|
|
pin_fraction = _pin_fraction
|
|
for n_gpus in range(1, len(ranked) + 1):
|
|
subset = ranked[:n_gpus]
|
|
pool_mib = sum(free for _, free in subset)
|
|
capped = self._fit_context_to_vram(
|
|
effective_ctx,
|
|
pool_mib,
|
|
model_size,
|
|
cache_type_kv,
|
|
n_parallel = n_parallel,
|
|
mtp_engaged = _mtp_will_engage,
|
|
)
|
|
kv = self._estimate_kv_cache_bytes(
|
|
capped, cache_type_kv, n_parallel = n_parallel
|
|
)
|
|
total_mib = (model_size + kv) / (1024 * 1024)
|
|
if total_mib <= pool_mib * pin_fraction:
|
|
effective_ctx = capped
|
|
gpu_indices = sorted(idx for idx, _ in subset)
|
|
use_fit = False
|
|
break
|
|
else:
|
|
# Native ctx doesn't fit. Drop to 4096 and
|
|
# re-check before --fit on: a model overflowing
|
|
# at 131k may pin fine with a 4096 KV (#5106).
|
|
effective_ctx = min(4096, effective_ctx)
|
|
if effective_ctx > 0:
|
|
for n_gpus in range(1, len(ranked) + 1):
|
|
subset = ranked[:n_gpus]
|
|
pool_mib = sum(free for _, free in subset)
|
|
kv = self._estimate_kv_cache_bytes(
|
|
effective_ctx,
|
|
cache_type_kv,
|
|
n_parallel = n_parallel,
|
|
)
|
|
total_mib = (model_size + kv) / (1024 * 1024)
|
|
if total_mib <= pool_mib * pin_fraction:
|
|
gpu_indices = sorted(idx for idx, _ in subset)
|
|
use_fit = False
|
|
break
|
|
|
|
elif gpus:
|
|
# Can't estimate KV -- file-size-only check; keep the
|
|
# ceiling at native context (already the default).
|
|
logger.debug(
|
|
"Falling back to file-size-only GPU selection",
|
|
model_size_gb = round(model_size / (1024**3), 2),
|
|
)
|
|
gpu_indices, use_fit = self._select_gpus(
|
|
model_size, gpus, usable_fraction = _pin_fraction
|
|
)
|
|
if use_fit and not explicit_ctx:
|
|
# Weights don't fit on any subset; default UI to 4096
|
|
# so the slider isn't on an unusable native ctx.
|
|
effective_ctx = min(4096, effective_ctx) if effective_ctx > 0 else 4096
|
|
|
|
if effective_ctx < original_ctx:
|
|
kv_est = self._estimate_kv_cache_bytes(
|
|
effective_ctx, cache_type_kv, n_parallel = n_parallel
|
|
)
|
|
logger.info(
|
|
f"Context auto-reduced: {original_ctx} -> {effective_ctx} "
|
|
f"(model: {model_size / (1024**3):.1f} GB, "
|
|
f"est. KV cache: {kv_est / (1024**3):.1f} GB)"
|
|
)
|
|
|
|
kv_cache_bytes = self._estimate_kv_cache_bytes(
|
|
effective_ctx, cache_type_kv, n_parallel = n_parallel
|
|
)
|
|
logger.info(
|
|
f"GGUF size: {model_size / (1024**3):.1f} GB, "
|
|
f"est. KV cache: {kv_cache_bytes / (1024**3):.1f} GB, "
|
|
f"context: {effective_ctx}, "
|
|
f"GPUs free: {gpus}, selected: {gpu_indices}, fit: {use_fit}"
|
|
)
|
|
except Exception as e:
|
|
logger.warning(f"GPU selection failed ({e}), using --fit on")
|
|
gpu_indices, use_fit = None, True
|
|
tp_tensor_split = None
|
|
effective_ctx = requested_ctx # fall back to original
|
|
|
|
launch_mmproj_path = None
|
|
if not extra_args_disable_mmproj(extra_args):
|
|
launch_mmproj_path = self._resolve_launch_mmproj_path(
|
|
model_path = model_path,
|
|
mmproj_path = mmproj_path,
|
|
)
|
|
# Need both a resolved mmproj AND the config vision flag; a stray
|
|
# mmproj passing the family-name heuristic must not flip a non-VLM
|
|
# GGUF into vision mode.
|
|
effective_is_vision = bool(launch_mmproj_path) and bool(is_vision)
|
|
if is_vision and not effective_is_vision:
|
|
logger.warning(
|
|
"Vision-capable GGUF loaded without a usable mmproj; "
|
|
"image input will be disabled for this session"
|
|
)
|
|
|
|
# Audio input straight from the mmproj (clip.has_audio_encoder),
|
|
# independent of token names.
|
|
self._mmproj_has_audio = False
|
|
if launch_mmproj_path:
|
|
try:
|
|
from utils.models.gguf_metadata import (
|
|
read_mmproj_audio_capability,
|
|
)
|
|
self._mmproj_has_audio = bool(
|
|
read_mmproj_audio_capability(launch_mmproj_path)
|
|
)
|
|
except Exception as e:
|
|
logger.debug(f"mmproj audio-capability read failed: {e}")
|
|
|
|
cmd = [
|
|
binary,
|
|
"-m",
|
|
model_path,
|
|
"--port",
|
|
str(self._port),
|
|
"-c",
|
|
str(effective_ctx) if effective_ctx > 0 else "0",
|
|
"--parallel",
|
|
str(n_parallel),
|
|
"--flash-attn",
|
|
"on", # Force flash attention for speed
|
|
# Error out at n_ctx instead of silently rotating the KV cache; frontend catches it and points the user at "Context Length".
|
|
"--no-context-shift",
|
|
]
|
|
|
|
if use_fit:
|
|
cmd.extend(["--fit", "on"])
|
|
elif gpu_indices is not None:
|
|
# Fits on selected GPU(s) -- offload all layers
|
|
cmd.extend(["-ngl", "-1"])
|
|
|
|
cmd.extend(
|
|
self._ctx_integrity_flags(
|
|
n_parallel,
|
|
use_fit,
|
|
requested_ctx,
|
|
effective_ctx,
|
|
self.probe_server_capabilities(binary),
|
|
)
|
|
)
|
|
|
|
# -1 = llama.cpp auto-detect (physical cores). Pass explicitly
|
|
# so we don't inherit llama-server's internal default, which
|
|
# has varied (hardware concurrency incl. hyperthreads on some
|
|
# builds).
|
|
cmd.extend(["--threads", str(n_threads if n_threads is not None else -1)])
|
|
|
|
# Enable Jinja chat template rendering
|
|
cmd.extend(["--jinja"])
|
|
|
|
# KV cache data type
|
|
_valid_cache_types = {
|
|
"f16",
|
|
"bf16",
|
|
"q8_0",
|
|
"q4_0",
|
|
"q4_1",
|
|
"q5_0",
|
|
"q5_1",
|
|
"iq4_nl",
|
|
"f32",
|
|
}
|
|
if cache_type_kv and cache_type_kv in _valid_cache_types:
|
|
cmd.extend(
|
|
[
|
|
"--cache-type-k",
|
|
cache_type_kv,
|
|
"--cache-type-v",
|
|
cache_type_kv,
|
|
]
|
|
)
|
|
self._cache_type_kv = cache_type_kv
|
|
logger.info(f"KV cache type: {cache_type_kv}")
|
|
else:
|
|
self._cache_type_kv = None
|
|
|
|
# Tensor parallelism: split the model across GPUs by tensor
|
|
# rather than by layer. Multi-GPU only -- a no-op on a single
|
|
# GPU. Default (layer split) is left implicit by omitting the
|
|
# flag. See llama.cpp --split-mode.
|
|
if tensor_parallel:
|
|
cmd.extend(["--split-mode", "tensor"])
|
|
if tp_tensor_split and len(tp_tensor_split) > 1:
|
|
cmd.extend(
|
|
[
|
|
"--tensor-split",
|
|
",".join(str(int(x)) for x in tp_tensor_split),
|
|
]
|
|
)
|
|
self._tensor_parallel = True
|
|
logger.info(
|
|
"Tensor parallelism: --split-mode tensor, --tensor-split %s",
|
|
tp_tensor_split,
|
|
)
|
|
else:
|
|
self._tensor_parallel = False
|
|
|
|
# Speculative decoding. See _build_speculative_flags for the
|
|
# mode resolution, benchmarks, and llama.cpp references.
|
|
launch_mtp_draft_path = self._resolve_launch_mtp_path(
|
|
mtp_draft_path = mtp_draft_path,
|
|
)
|
|
spec_flags = self._build_speculative_flags(
|
|
speculative_type = speculative_type,
|
|
spec_draft_n_max = spec_draft_n_max,
|
|
extra_args = extra_args,
|
|
model_identifier = model_identifier,
|
|
model_path = model_path,
|
|
gpus = bool(gpus),
|
|
binary = binary,
|
|
mtp_draft_path = launch_mtp_draft_path,
|
|
)
|
|
# Remember where the spec block sits so a drafter-load failure
|
|
# can be retried with these flags swapped out (see below).
|
|
_spec_start = len(cmd)
|
|
cmd.extend(spec_flags)
|
|
|
|
# Apply custom chat template override if provided.
|
|
self._chat_template_override = chat_template_override
|
|
if chat_template_override:
|
|
import tempfile
|
|
|
|
flags = detect_reasoning_flags(
|
|
chat_template_override,
|
|
self._model_identifier,
|
|
log_source = "GGUF chat template override",
|
|
)
|
|
self._supports_reasoning = flags["supports_reasoning"]
|
|
self._reasoning_style = flags["reasoning_style"]
|
|
self._reasoning_always_on = flags["reasoning_always_on"]
|
|
self._supports_preserve_thinking = flags["supports_preserve_thinking"]
|
|
self._supports_tools = flags["supports_tools"]
|
|
|
|
self._chat_template_file = tempfile.NamedTemporaryFile(
|
|
mode = "w",
|
|
encoding = "utf-8",
|
|
suffix = ".jinja",
|
|
delete = False,
|
|
prefix = "unsloth_chat_template_",
|
|
)
|
|
self._chat_template_file.write(chat_template_override)
|
|
self._chat_template_file.close()
|
|
cmd.extend(["--chat-template-file", self._chat_template_file.name])
|
|
logger.info(f"Using custom chat template file: {self._chat_template_file.name}")
|
|
|
|
# Default thinking mode for reasoning models. Qwen3.5/3.6 below
|
|
# 9B disable thinking by default; 9B+ enable it. Always-on
|
|
# templates ignore the kwarg, so skip.
|
|
if self._supports_reasoning and not self._reasoning_always_on:
|
|
thinking_default = True
|
|
mid = (model_identifier or "").lower()
|
|
if "qwen3.5" in mid or "qwen3.6" in mid:
|
|
size_val = _extract_model_size_b(mid)
|
|
if size_val is not None and size_val < 9:
|
|
thinking_default = False
|
|
self._reasoning_default = thinking_default
|
|
reasoning_kw = self._reasoning_kwargs(thinking_default)
|
|
# preserve_thinking is an independent kwarg. Default it OFF
|
|
# at launch so direct OpenAI-compatible callers that omit the
|
|
# field match the UI's default-off behavior (the bundled
|
|
# gemma-4 template also defaults it false; the frontend sends
|
|
# preserve_thinking per request once toggled on).
|
|
if self._supports_preserve_thinking:
|
|
reasoning_kw["preserve_thinking"] = False
|
|
cmd.extend(
|
|
[
|
|
"--chat-template-kwargs",
|
|
json.dumps(reasoning_kw),
|
|
]
|
|
)
|
|
logger.info(f"Reasoning model: {reasoning_kw} by default")
|
|
|
|
if launch_mmproj_path and effective_is_vision:
|
|
cmd.extend(["--mmproj", launch_mmproj_path])
|
|
logger.info(f"Using mmproj for vision: {launch_mmproj_path}")
|
|
|
|
# Option C: --api-key for direct client access when enabled
|
|
import os as _os
|
|
import secrets as _secrets
|
|
|
|
if _os.getenv("UNSLOTH_DIRECT_STREAM", "0") == "1":
|
|
self._api_key = _secrets.token_urlsafe(32)
|
|
cmd.extend(["--api-key", self._api_key])
|
|
logger.info("llama-server started with --api-key for direct streaming")
|
|
else:
|
|
self._api_key = None
|
|
|
|
# User pass-through args go last so llama.cpp's last-wins parsing
|
|
# lets the user override Studio's auto-set flags. Already
|
|
# validated by the route via validate_extra_args().
|
|
if extra_args:
|
|
cmd.extend(str(a) for a in extra_args)
|
|
logger.info(f"Appending user extra args to llama-server: {list(extra_args)}")
|
|
|
|
_log_cmd = list(cmd)
|
|
if "--api-key" in _log_cmd:
|
|
_ki = _log_cmd.index("--api-key") + 1
|
|
if _ki < len(_log_cmd):
|
|
_log_cmd[_ki] = "<redacted>"
|
|
logger.info(f"Starting llama-server: {' '.join(_log_cmd)}")
|
|
|
|
# Library paths so llama-server finds its shared libs and CUDA DLLs.
|
|
import os
|
|
import sys
|
|
|
|
env = child_env_without_native_path_secret()
|
|
binary_dir = str(Path(binary).parent)
|
|
|
|
# AMD unified-memory APUs (gfx1150/gfx1151): let llama.cpp use
|
|
# shared system RAM. setdefault so a user value wins.
|
|
if self._amd_apu_wants_unified_memory():
|
|
env.setdefault("GGML_CUDA_ENABLE_UNIFIED_MEMORY", "1")
|
|
logger.info("AMD unified-memory APU: set GGML_CUDA_ENABLE_UNIFIED_MEMORY=1")
|
|
|
|
# DC NVIDIA GPUs: FP32 accum (+ P2P / launch queues for multi-GPU).
|
|
# See _apply_datacenter_env; opt out with UNSLOTH_DISABLE_DC_TUNING=1.
|
|
if self._apply_datacenter_env(env, gpu_indices):
|
|
multi_gpu = self._effective_gpu_count(gpu_indices) > 1
|
|
logger.info(
|
|
f"Data-center GPU detected: applied DC llama.cpp env tuning (multi_gpu={multi_gpu})"
|
|
)
|
|
|
|
if sys.platform == "win32":
|
|
# Ordering: see _build_windows_path_dirs. #5106.
|
|
path_dirs = self._build_windows_path_dirs(
|
|
binary_dir,
|
|
sys.prefix,
|
|
os.environ.get("CUDA_PATH", ""),
|
|
)
|
|
existing_path = env.get("PATH", "")
|
|
env["PATH"] = ";".join(path_dirs) + ";" + existing_path
|
|
|
|
# ROCm: the prebuilt bundles rocblas.dll but NOT the Tensile
|
|
# kernel files (rocblas/library/*.dat + *.hsaco); the DLL
|
|
# searches <binary_dir>/rocblas/library/ which doesn't exist
|
|
# -> silent crash on the first GEMM. ROCBLAS_TENSILE_LIBPATH
|
|
# repoints that search at the ROCm install.
|
|
_hip_path = os.environ.get("HIP_PATH", os.environ.get("ROCM_PATH", ""))
|
|
if _hip_path:
|
|
_rocblas_lib = os.path.join(_hip_path, "bin", "rocblas", "library")
|
|
if os.path.isdir(_rocblas_lib):
|
|
env.setdefault("ROCBLAS_TENSILE_LIBPATH", _rocblas_lib)
|
|
else:
|
|
# Linux: LD_LIBRARY_PATH for shared libs next to the binary
|
|
# plus CUDA runtime libs (libcudart, libcublas, etc.)
|
|
import platform
|
|
|
|
lib_dirs = []
|
|
# WSL: system HIP before the bundle's (which segfaults on
|
|
# /dev/dxg). Mirror install_llama_prebuilt.binary_env, which
|
|
# validates the prebuilt with this same ordering.
|
|
for _wsl_rocm in _wsl_system_rocm_lib_dirs():
|
|
lib_dirs.append(_wsl_rocm)
|
|
if lib_dirs:
|
|
env.setdefault("HSA_ENABLE_DXG_DETECTION", "1")
|
|
lib_dirs.append(binary_dir)
|
|
_arch = platform.machine() # x86_64, aarch64, etc.
|
|
|
|
# Pip-installed nvidia CUDA runtime libs. The prebuilt
|
|
# binary links libcudart.so.13 / libcublas.so.13 which live
|
|
# here, not in /usr/local/cuda.
|
|
import glob as _glob
|
|
|
|
for _nv_pattern in [
|
|
os.path.join(
|
|
sys.prefix,
|
|
"lib",
|
|
"python*",
|
|
"site-packages",
|
|
"nvidia",
|
|
"cu*",
|
|
"lib",
|
|
),
|
|
os.path.join(
|
|
sys.prefix,
|
|
"lib",
|
|
"python*",
|
|
"site-packages",
|
|
"nvidia",
|
|
"cudnn",
|
|
"lib",
|
|
),
|
|
os.path.join(
|
|
sys.prefix,
|
|
"lib",
|
|
"python*",
|
|
"site-packages",
|
|
"nvidia",
|
|
"nvjitlink",
|
|
"lib",
|
|
),
|
|
]:
|
|
for _nv_dir in _glob.glob(_nv_pattern):
|
|
if os.path.isdir(_nv_dir):
|
|
lib_dirs.append(_nv_dir)
|
|
|
|
for cuda_lib in [
|
|
"/usr/local/cuda/lib64",
|
|
f"/usr/local/cuda/targets/{_arch}-linux/lib",
|
|
# Fallback CUDA compat paths (e.g. binary built with
|
|
# CUDA 12 where default /usr/local/cuda is CUDA 13+).
|
|
"/usr/local/cuda-12/lib64",
|
|
"/usr/local/cuda-12.8/lib64",
|
|
f"/usr/local/cuda-12/targets/{_arch}-linux/lib",
|
|
f"/usr/local/cuda-12.8/targets/{_arch}-linux/lib",
|
|
]:
|
|
if os.path.isdir(cuda_lib):
|
|
lib_dirs.append(cuda_lib)
|
|
existing_ld = env.get("LD_LIBRARY_PATH", "")
|
|
new_ld = ":".join(lib_dirs)
|
|
env["LD_LIBRARY_PATH"] = f"{new_ld}:{existing_ld}" if existing_ld else new_ld
|
|
|
|
# Pin to selected GPU(s). On ROCm, narrowing only
|
|
# CUDA_VISIBLE_DEVICES leaves an AMD child seeing the full
|
|
# set, so set HIP_VISIBLE_DEVICES too.
|
|
if gpu_indices is not None:
|
|
pinned = ",".join(str(i) for i in gpu_indices)
|
|
env["CUDA_VISIBLE_DEVICES"] = pinned
|
|
try:
|
|
import torch as _torch
|
|
if getattr(_torch.version, "hip", None) is not None:
|
|
env["HIP_VISIBLE_DEVICES"] = pinned
|
|
# Do NOT also set ROCR_VISIBLE_DEVICES to the same
|
|
# value. ROCR_VISIBLE_DEVICES filters at the HSA/ROCr
|
|
# layer and HIP_VISIBLE_DEVICES at the HIP layer, so
|
|
# setting both with the same physical indices applies
|
|
# the mask twice: ROCR reduces the visible set and
|
|
# re-indexes it from 0, then HIP indexes into the
|
|
# already-reduced set. A single non-zero pin (e.g.
|
|
# "1") then points out of range at the HIP layer, HIP
|
|
# enumerates 0 devices, and llama.cpp falls back to
|
|
# CPU ("ggml_cuda_init: no ROCm-capable device is
|
|
# detected"). The HIP mask alone narrows correctly;
|
|
# clear any inherited ROCR mask so it can't double up.
|
|
env.pop("ROCR_VISIBLE_DEVICES", None)
|
|
except Exception as e:
|
|
logger.debug("Failed to set ROCm visibility env vars for child: %s", e)
|
|
|
|
# Captured before any text-only fallback strips it from cmd.
|
|
launched_with_mmproj = "--mmproj" in cmd
|
|
|
|
# One-shot --fit off retry: recent llama.cpp runs a "fitting
|
|
# params to device memory" step by default (--fit defaults to
|
|
# 'on') even when -ngl is explicit. That step has aborted on
|
|
# some ROCm hosts (ggml-cuda.cu ROCm error during worst-case
|
|
# estimation, e.g. MTP + mmproj models on gfx1151). When
|
|
# Studio's own VRAM math already placed the model
|
|
# (use_fit=False), the step is redundant second-guessing --
|
|
# retry once with --fit off before declaring the load failed.
|
|
# Never retry when fit was requested (use_fit) or the caller
|
|
# passed an explicit fit flag via extra args.
|
|
def _spawn_and_wait(run_cmd, *, label = ""):
|
|
"""Start llama-server with run_cmd and wait for health.
|
|
|
|
Retries once with --fit off when the first attempt
|
|
crashes during startup and run_cmd is eligible (see
|
|
_fit_off_retry_eligible).
|
|
"""
|
|
_fit_retry_allowed = self._fit_off_retry_eligible(run_cmd, use_fit)
|
|
for _spawn_attempt in (0, 1):
|
|
# Defensive kill: drop an orphan Popen a concurrent load may
|
|
# have stored before we overwrite the reference (#5161).
|
|
# Also reaps the crashed first attempt on the retry pass.
|
|
self._kill_process()
|
|
|
|
self._stdout_lines = []
|
|
# Tee llama-server output to a dedicated log file so a
|
|
# post-mortem has the full trail even when the parent only
|
|
# kept the last 50 lines. Path is under the studio home.
|
|
# ``label`` (MTP fallback) and the attempt index (--fit
|
|
# off retry) keep a respawn within the same epoch second
|
|
# from truncating the crash log a retry warning just
|
|
# pointed the user at.
|
|
self._llama_log_fh = None
|
|
try:
|
|
log_dir = _swa_cache_path().parent / "logs" / "llama-server"
|
|
log_dir.mkdir(parents = True, exist_ok = True)
|
|
self._llama_log_path = log_dir / (
|
|
f"llama-{int(time.time())}{label}-port-{self._port}"
|
|
f"-try{_spawn_attempt}.log"
|
|
)
|
|
self._llama_log_fh = open(
|
|
self._llama_log_path,
|
|
"w",
|
|
encoding = "utf-8",
|
|
buffering = 1,
|
|
)
|
|
logger.info(f"llama-server stdout/stderr -> {self._llama_log_path}")
|
|
except OSError as e:
|
|
# Best-effort; never block the load on logging.
|
|
logger.debug(f"Could not open llama-server log file: {e}")
|
|
self._llama_log_path = None
|
|
self._process = subprocess.Popen(
|
|
run_cmd,
|
|
stdout = subprocess.PIPE,
|
|
stderr = subprocess.STDOUT,
|
|
text = True,
|
|
env = env,
|
|
**_windows_hidden_subprocess_kwargs(),
|
|
)
|
|
|
|
# Background thread to drain stdout (prevents pipe deadlock)
|
|
self._stdout_thread = threading.Thread(
|
|
target = self._drain_stdout, daemon = True, name = "llama-stdout"
|
|
)
|
|
self._stdout_thread.start()
|
|
if self._wait_for_health(timeout = 600.0):
|
|
return True
|
|
_startup_crashed = (
|
|
self._process.poll() is not None and self._process.returncode != 0
|
|
)
|
|
if _spawn_attempt == 0 and _fit_retry_allowed and _startup_crashed:
|
|
logger.warning(
|
|
"llama-server crashed during startup (exit code %s) "
|
|
"with the default memory-fit step enabled; Studio "
|
|
"already verified the model fits, retrying once "
|
|
"with --fit off. Crash log: %s",
|
|
self._process.returncode,
|
|
self._llama_log_path,
|
|
)
|
|
run_cmd = [*run_cmd, "--fit", "off"]
|
|
continue
|
|
return False
|
|
|
|
# Store the resolved on-disk path, not the caller's kwarg: in
|
|
# HF mode gguf_path is None and ``model_path`` is what
|
|
# llama-server mmap's, which downstream consumers need. Must be
|
|
# set BEFORE the spawn: load_progress() reads _gguf_path for
|
|
# the mmap progress total while the health wait runs.
|
|
self._gguf_path = model_path
|
|
self._hf_repo = hf_repo
|
|
self._mtp_draft_path = launch_mtp_draft_path
|
|
# For local GGUF files, extract variant from filename if absent
|
|
if hf_variant:
|
|
self._hf_variant = hf_variant
|
|
elif gguf_path:
|
|
try:
|
|
from utils.models.model_config import _extract_quant_label
|
|
self._hf_variant = _extract_quant_label(gguf_path)
|
|
except Exception:
|
|
self._hf_variant = None
|
|
else:
|
|
self._hf_variant = None
|
|
self._is_vision = effective_is_vision
|
|
self._model_identifier = model_identifier
|
|
|
|
# Store the effective (possibly capped) context separately; do
|
|
# NOT overwrite _context_length (the native length for display).
|
|
self._effective_context_length = (
|
|
effective_ctx if effective_ctx > 0 else self._context_length
|
|
)
|
|
self._reconcile_effective_ctx_with_server()
|
|
self._max_context_length = (
|
|
max_available_ctx if max_available_ctx > 0 else self._effective_context_length
|
|
)
|
|
|
|
healthy = _spawn_and_wait(cmd)
|
|
# Any MTP request can abort the server: a separate drafter
|
|
# (Gemma) on a binary that predates its arch, or an embedded
|
|
# head (Qwen) the binary cannot build. Retry once with the
|
|
# spec slice replaced by --spec-default so the main model still
|
|
# loads. Gate on the spec block (not the drafter path, which
|
|
# off/ngram local loads also carry) and keep
|
|
# _requested_spec_mode so a duplicate /load doesn't thrash. The
|
|
# cancel check stops an /unload-killed attempt respawning.
|
|
_spec_requested_mtp = any("mtp" in str(t).lower() for t in spec_flags)
|
|
if not healthy and _spec_requested_mtp and not self._cancel_event.is_set():
|
|
# Blame the binary only when the output shows MTP itself
|
|
# failing (unknown arch / draft or context build); an
|
|
# unrelated crash (e.g. OOM) gets a neutral message.
|
|
_lo = "\n".join(self._stdout_lines).lower()
|
|
# Only an unknown architecture proves the prebuilt predates
|
|
# this MTP model (an update fixes it). The memory/context
|
|
# build failures are generic (VRAM / ctx pressure), where an
|
|
# update may not help, so classify those as runtime_error.
|
|
_arch_unsupported = "unknown model architecture" in _lo
|
|
if (
|
|
_arch_unsupported
|
|
or "failed to measure draft model memory" in _lo
|
|
or "failed to measure mtp context memory" in _lo
|
|
or "failed to create llama_context" in _lo
|
|
):
|
|
_retry_reason = (
|
|
"the prebuilt may predate it; retrying without "
|
|
"speculative decoding -- run `unsloth studio "
|
|
"update` for MTP"
|
|
)
|
|
self._spec_fallback_reason = (
|
|
"binary_outdated" if _arch_unsupported else "runtime_error"
|
|
)
|
|
else:
|
|
_retry_reason = (
|
|
"retrying without speculative decoding in case MTP is the cause"
|
|
)
|
|
self._spec_fallback_reason = "runtime_error"
|
|
_drafter = (
|
|
Path(launch_mtp_draft_path).name
|
|
if launch_mtp_draft_path
|
|
else "embedded head"
|
|
)
|
|
logger.warning(
|
|
"llama-server failed to start with MTP (%s); %s.",
|
|
_drafter,
|
|
_retry_reason,
|
|
)
|
|
self._kill_process()
|
|
fallback_cmd = (
|
|
cmd[:_spec_start]
|
|
+ ["--spec-default"]
|
|
+ cmd[_spec_start + len(spec_flags) :]
|
|
)
|
|
healthy = _spawn_and_wait(fallback_cmd, label = "-retry")
|
|
if healthy:
|
|
self._speculative_type = "default"
|
|
|
|
# A vision GGUF launched with --mmproj can abort when the
|
|
# installed llama.cpp is too old for the model's projector
|
|
# ("Unknown projector type"); in that one case retry once
|
|
# text-only rather than failing the whole load.
|
|
if not healthy:
|
|
out = "\n".join(self._stdout_lines[-50:])
|
|
self._kill_process()
|
|
if launched_with_mmproj and self._is_projector_incompatibility(out):
|
|
logger.warning(
|
|
"llama-server could not load this model's vision "
|
|
"projector (--mmproj). The installed llama.cpp build is "
|
|
"likely too old for it. Loading text-only for this "
|
|
"session; run 'unsloth studio update' to enable vision."
|
|
)
|
|
cmd = self._strip_mmproj_args(cmd)
|
|
self._is_vision = False
|
|
self._mmproj_has_audio = False
|
|
self._start_llama_process(cmd, env)
|
|
if not self._wait_for_health(timeout = 600.0):
|
|
self._kill_process()
|
|
raise RuntimeError(
|
|
"Vision projector incompatible with this llama.cpp "
|
|
"build, and the text-only retry also failed: "
|
|
+ self._classify_llama_start_failure(
|
|
"\n".join(self._stdout_lines[-50:]),
|
|
gguf_path,
|
|
self._model_identifier,
|
|
)
|
|
)
|
|
else:
|
|
raise RuntimeError(
|
|
self._classify_llama_start_failure(
|
|
out,
|
|
gguf_path,
|
|
self._model_identifier,
|
|
)
|
|
)
|
|
|
|
self._healthy = True
|
|
|
|
# Commit caller intent only after _healthy=True so a failed start
|
|
# can't poison the next inheritance check. None keeps prior, []
|
|
# clears, list sets. Source records hf_variant for the route's
|
|
# same_source check.
|
|
if extra_args is not None:
|
|
self._extra_args = list(extra_args)
|
|
self._extra_args_source = (model_identifier, hf_variant)
|
|
self._requested_n_ctx = int(n_ctx)
|
|
|
|
# Catch silent CPU fallback when GPU was intended (#5106).
|
|
self._gpu_offload_active = self._classify_gpu_offload(
|
|
gpu_indices is not None or use_fit, gpus or []
|
|
)
|
|
if self._gpu_offload_active is False:
|
|
logger.warning(
|
|
"llama-server appears to have loaded the model entirely "
|
|
"on CPU even though Studio detected at least one GPU. "
|
|
"This usually means the prebuilt binary's GPU backend "
|
|
"failed to load -- on Windows, cudart64_X.dll / "
|
|
"cublas64_X.dll could not be resolved. Reinstall the "
|
|
"Studio llama.cpp prebuilt or install a matching CUDA "
|
|
"toolkit (issue unslothai/unsloth#5106).",
|
|
)
|
|
|
|
logger.info(
|
|
f"llama-server ready on port {self._port} for model '{model_identifier}'"
|
|
)
|
|
|
|
# Probe outside _lock (interruptible by /unload); init inside.
|
|
self._is_audio = False
|
|
self._audio_type = None
|
|
self._audio_probed = False
|
|
self._has_audio_input = False
|
|
try:
|
|
detected = self._detect_audio_type_strict()
|
|
self._audio_probed = True
|
|
except Exception as exc:
|
|
logger.debug("Audio probe failed: %s", exc)
|
|
detected = None
|
|
if detected in ("snac", "bicodec", "dac"):
|
|
with self._lock:
|
|
if not self._healthy:
|
|
return False
|
|
try:
|
|
self.init_audio_codec(detected)
|
|
self._is_audio = True
|
|
self._audio_type = detected
|
|
except Exception as exc:
|
|
# Surface as HTTP 500 (matches pre-PR contract).
|
|
logger.warning(
|
|
"Failed to init audio codec '%s': %s",
|
|
detected,
|
|
exc,
|
|
)
|
|
self._audio_probed = False
|
|
return False
|
|
elif detected:
|
|
# csm / whisper / audio_vlm: track type but keep _is_audio
|
|
# False -- GGUF TTS routing only fires for snac/bicodec/dac.
|
|
with self._lock:
|
|
if not self._healthy:
|
|
return False
|
|
self._audio_type = detected
|
|
|
|
# Audio input = token probe (audio_vlm/whisper) OR mmproj encoder.
|
|
from utils.models.model_config import is_audio_input_type
|
|
|
|
self._has_audio_input = bool(is_audio_input_type(self._audio_type)) or bool(
|
|
self._mmproj_has_audio
|
|
)
|
|
|
|
if not self._healthy:
|
|
return False
|
|
return True
|
|
|
|
def _build_speculative_flags(
|
|
self,
|
|
*,
|
|
speculative_type: Optional[str],
|
|
spec_draft_n_max: Optional[int],
|
|
extra_args: Optional[List[str]],
|
|
model_identifier: str,
|
|
model_path: Optional[str],
|
|
gpus: bool,
|
|
binary: Optional[str],
|
|
mtp_draft_path: Optional[str] = None,
|
|
) -> List[str]:
|
|
"""Return the llama-server flag list for the requested spec mode.
|
|
|
|
Side effects: sets ``self._speculative_type`` (resolved internal
|
|
emit), ``self._requested_spec_mode`` (canonical UI mode for the
|
|
status round-trip), and ``self._spec_draft_n_max`` (user override
|
|
only; None when the platform default applies).
|
|
|
|
Speculative decoding (n-gram self-speculation, zero VRAM):
|
|
ngram-mod uses a ~16 MB shared hash pool, constant memory /
|
|
complexity, variable draft lengths. Helps most when the model
|
|
repeats existing text (code refactor, summarisation, reasoning);
|
|
for low-repetition chat, overhead is ~5 ms.
|
|
|
|
Benchmarks from upstream llama.cpp speculative-decoding PRs:
|
|
Scenario | Without | With | Speedup
|
|
gpt-oss-120b code refactor | 181 t/s | 446 t/s | 2.5x
|
|
Qwen3-235B offloaded | 12 t/s | 21 t/s | 1.8x
|
|
gpt-oss-120b repeat (92% accept)| 181 t/s | 814 t/s | 4.5x
|
|
Refs: https://github.com/ggml-org/llama.cpp/blob/master/docs/speculative.md
|
|
https://github.com/ggml-org/llama.cpp/pull/19164
|
|
https://github.com/ggml-org/llama.cpp/pull/18471
|
|
MTP guide: unsloth.ai/docs/models/qwen3.6#mtp-guide
|
|
|
|
Sub-3B dense MTP regresses vs spec-off when the head is baked into the
|
|
main GGUF (Qwen): the draft head's per-token cost exceeds the
|
|
acceptance savings at this scale. Q4_K_XL clean bench (each prompt once
|
|
after an unrelated warmup) on B200 + x86 CPU:
|
|
0.8B GPU: draft-mtp n=2 = 0.58x vs OFF; ngram-only = 1.10x
|
|
2B GPU: draft-mtp n=2 = 0.82x vs OFF; OFF or ngram = 1.00x
|
|
0.8B CPU: chained n=2 = 0.86x vs OFF; ngram-only = 1.19x
|
|
2B CPU: chained n=2 = 0.83x vs OFF; ngram-only = 1.01x
|
|
4B+ GPU/CPU: spec on is a net win (1.08x-1.46x).
|
|
A separate drafter (Gemma's root mtp-*.gguf) is a different, cheaper
|
|
mechanism that wins even below 3B, so it is exempt from the sub-3B drop
|
|
(``mtp_draft_path`` set -> not too small). B200 Q4_K_XL bench, draft-mtp
|
|
n=2 vs OFF: gemma-4-E2B (2B) = 1.21x, accept ~0.65 (vs ngram = 1.00x);
|
|
gemma-4-E4B (4B) and 12B engage as usual.
|
|
Auto falls back to ngram-mod (zero-VRAM, near-zero idle cost on
|
|
diverse content) for an embedded sub-3B head; forced MTP on a model
|
|
with no head/drafter defaults back (mtp -> spec-default, mtp+ngram ->
|
|
ngram-mod) since llama-server aborts otherwise; a drafter the binary
|
|
cannot build (older prebuilt, or a CUDA kernel limit) aborts the spawn
|
|
and the load retries once without speculative decoding.
|
|
"""
|
|
flags: List[str] = []
|
|
# Reset; emit branches re-set on the resolved emission.
|
|
self._spec_draft_n_max = None
|
|
self._speculative_type = None
|
|
self._spec_fallback_reason = None
|
|
|
|
# Canonical UI-facing requested mode (legacy values mapped via
|
|
# _canonicalize_spec_mode).
|
|
canonical_mode = _canonicalize_spec_mode(speculative_type)
|
|
# MTP signals: head baked into the main GGUF (Qwen, via metadata or
|
|
# name), or a separate drafter resolved from the repo (Gemma).
|
|
is_mtp_model = (
|
|
bool(self._nextn_predict_layers)
|
|
or _is_mtp_model_name(model_identifier, model_path)
|
|
or bool(mtp_draft_path)
|
|
)
|
|
user_owns_spec_type = _extra_args_set_spec_type(extra_args)
|
|
_mtp_size_b = _extract_model_size_b(model_identifier)
|
|
# The sub-3B regression is an embedded-head cost; a separate drafter
|
|
# (Gemma) is a cheap standalone model that wins below 3B, so exempt it.
|
|
_mtp_too_small = (
|
|
_mtp_size_b is not None and _mtp_size_b < _MTP_MIN_SIZE_B and not bool(mtp_draft_path)
|
|
)
|
|
|
|
if user_owns_spec_type:
|
|
# User --spec-type wins outright; suppress auto-emit to avoid a
|
|
# duplicate spec block.
|
|
self._requested_spec_mode = None
|
|
return flags
|
|
|
|
effective_mode = canonical_mode or "auto"
|
|
self._requested_spec_mode = effective_mode
|
|
|
|
def _resolved_draft_n_max() -> int:
|
|
# User override wins; else platform default (the B200 / x86
|
|
# clean-sweep sweet spot from PR #5582 is n=2 GPU, n=3 CPU;
|
|
# past 3 regresses on essay-style low-acceptance prompts).
|
|
if spec_draft_n_max is not None:
|
|
n = int(spec_draft_n_max)
|
|
self._spec_draft_n_max = n
|
|
return n
|
|
return 2 if gpus else 3
|
|
|
|
def _emit_mtp(*, chain_ngram: bool) -> bool:
|
|
"""Append --spec-type mtp[/draft-mtp][,ngram-mod] + n-max."""
|
|
caps = self.probe_server_capabilities(binary)
|
|
mtp_token = caps.get("mtp_token") if caps else None
|
|
if not mtp_token:
|
|
logger.warning(
|
|
"Requested MTP speculative decoding but "
|
|
"llama-server lacks --spec-type mtp/draft-mtp; "
|
|
"run `unsloth studio update`. Loading without "
|
|
"speculative decoding."
|
|
)
|
|
self._spec_fallback_reason = "binary_no_mtp"
|
|
return False
|
|
draft_n_max = _resolved_draft_n_max()
|
|
n_max_flag = caps.get("spec_draft_n_max_flag") or "--spec-draft-n-max"
|
|
# Separate-file drafter (Gemma): point llama-server at it. Baked-in
|
|
# heads (Qwen) pass no path -- llama-server reads them from the
|
|
# main GGUF.
|
|
if mtp_draft_path:
|
|
flags.extend(["--model-draft", mtp_draft_path])
|
|
logger.info(f"Using separate MTP drafter: {mtp_draft_path}")
|
|
if chain_ngram:
|
|
ngram_knobs = _build_ngram_mod_flags(caps)
|
|
if ngram_knobs:
|
|
spec_value = f"ngram-mod,{mtp_token}"
|
|
else:
|
|
logger.warning(
|
|
"llama-server lacks ngram-mod tuning "
|
|
"flags; loading MTP only (no ngram chain)"
|
|
)
|
|
spec_value = mtp_token
|
|
flags.extend(
|
|
[
|
|
"--spec-type",
|
|
spec_value,
|
|
n_max_flag,
|
|
str(draft_n_max),
|
|
]
|
|
)
|
|
flags.extend(ngram_knobs)
|
|
else:
|
|
flags.extend(
|
|
[
|
|
"--spec-type",
|
|
mtp_token,
|
|
n_max_flag,
|
|
str(draft_n_max),
|
|
]
|
|
)
|
|
self._speculative_type = "draft-mtp"
|
|
chain_label = "chained ngram-mod" if chain_ngram else "MTP-only"
|
|
logger.info(f"Spec decoding: {mtp_token} ({chain_label})")
|
|
return True
|
|
|
|
def _emit_ngram_mod() -> bool:
|
|
"""Append --spec-type ngram-mod + flag-set knobs."""
|
|
ngram_caps = self.probe_server_capabilities(binary)
|
|
ngram_knobs = _build_ngram_mod_flags(ngram_caps)
|
|
flags.extend(["--spec-type", "ngram-mod"])
|
|
if not ngram_knobs:
|
|
logger.warning(
|
|
"llama-server lacks ngram-mod tuning "
|
|
"flags; loading without --spec-ngram-mod-* knobs"
|
|
)
|
|
flags.extend(ngram_knobs)
|
|
self._speculative_type = "ngram-mod"
|
|
logger.info("Spec decoding: ngram-mod")
|
|
return True
|
|
|
|
if effective_mode == "off":
|
|
return flags # nothing to emit
|
|
if effective_mode == "ngram-simple":
|
|
flags.extend(["--spec-type", "ngram-simple"])
|
|
self._speculative_type = "ngram-simple"
|
|
return flags
|
|
if effective_mode == "ngram":
|
|
_emit_ngram_mod()
|
|
return flags
|
|
if effective_mode == "mtp":
|
|
if not is_mtp_model:
|
|
# No head and no drafter: llama-server aborts on draft-mtp
|
|
# instead of no-op'ing, so default back.
|
|
logger.warning(
|
|
"MTP requested but this GGUF has no MTP head or drafter; "
|
|
"loading without speculative decoding."
|
|
)
|
|
flags.append("--spec-default")
|
|
self._speculative_type = "default"
|
|
return flags
|
|
if _mtp_too_small:
|
|
logger.warning(
|
|
f"Forcing MTP on a {_mtp_size_b:.1f}B model; "
|
|
"the bench shows draft-mtp regresses below 3B. "
|
|
"Engaging anyway (user override)."
|
|
)
|
|
_emit_mtp(chain_ngram = False)
|
|
return flags
|
|
if effective_mode == "mtp+ngram":
|
|
if not is_mtp_model:
|
|
# No head/drafter: keep the ngram half (needs no head),
|
|
# drop the draft-mtp chain that would abort the server.
|
|
logger.warning(
|
|
"MTP+Ngram requested but this GGUF has no MTP head or "
|
|
"drafter; loading ngram-mod only."
|
|
)
|
|
_emit_ngram_mod()
|
|
return flags
|
|
if _mtp_too_small:
|
|
logger.warning(
|
|
f"Forcing MTP+Ngram on a {_mtp_size_b:.1f}B model; "
|
|
"the bench shows the chain regresses below 3B. "
|
|
"Engaging anyway (user override)."
|
|
)
|
|
_emit_mtp(chain_ngram = True)
|
|
return flags
|
|
|
|
# effective_mode == "auto": the promotion path. llama.cpp #22673:
|
|
# MTP is compatible with mmproj, so there's no vision gate.
|
|
if is_mtp_model and not _mtp_too_small:
|
|
# GPU: MTP-only. CPU/Mac: chain ngram-mod + MTP.
|
|
_emit_mtp(chain_ngram = not gpus)
|
|
elif is_mtp_model and _mtp_too_small:
|
|
# Sub-3B fallback: drop the MTP draft head, keep ngram-mod when
|
|
# the binary supports it.
|
|
_small_caps = self.probe_server_capabilities(binary)
|
|
if _small_caps.get("supports_ngram_mod"):
|
|
logger.info(
|
|
f"MTP GGUF detected but model size {_mtp_size_b:.1f}B "
|
|
"is below the 3B speedup threshold; using ngram-mod "
|
|
"only (zero-VRAM, no draft head). Override via "
|
|
"--spec-type or the Studio Speculative Decoding "
|
|
"dropdown."
|
|
)
|
|
_emit_ngram_mod()
|
|
else:
|
|
logger.info(
|
|
f"MTP GGUF detected but model size {_mtp_size_b:.1f}B "
|
|
"is below the 3B speedup threshold and the bundled "
|
|
"llama-server does not advertise ngram-mod; "
|
|
"auto-disabling speculative decoding."
|
|
)
|
|
else:
|
|
# Non-MTP model: let llama-server choose its default strategy.
|
|
flags.append("--spec-default")
|
|
self._speculative_type = "default"
|
|
return flags
|
|
|
|
def _already_in_target_state(
|
|
self,
|
|
*,
|
|
model_identifier: str,
|
|
hf_variant: Optional[str],
|
|
n_ctx: int,
|
|
cache_type_kv: Optional[str],
|
|
speculative_type: Optional[str],
|
|
chat_template_override: Optional[str],
|
|
extra_args: Optional[List[str]],
|
|
is_vision: bool,
|
|
gguf_path: Optional[str] = None,
|
|
spec_draft_n_max: Optional[int] = None,
|
|
tensor_parallel: bool = False,
|
|
mtp_draft_path: Optional[str] = None,
|
|
) -> bool:
|
|
"""True iff the live server already satisfies these load kwargs.
|
|
|
|
Mirrors ``routes/inference.py:_request_matches_loaded_settings`` but
|
|
compares raw kwargs so ``load_model`` can short-circuit a duplicate
|
|
/load that raced past the route-level check (#5401).
|
|
"""
|
|
if not self.is_loaded:
|
|
return False
|
|
if (self._model_identifier or "").lower() != (model_identifier or "").lower():
|
|
return False
|
|
# Direct-file loads pass hf_variant=None while the backend stores an
|
|
# extracted filename label; compare paths to keep the guard symmetric.
|
|
if gguf_path is not None and self._gguf_path:
|
|
try:
|
|
if Path(self._gguf_path).resolve() != Path(gguf_path).resolve():
|
|
return False
|
|
except OSError:
|
|
return False
|
|
elif (self._hf_variant or "").lower() != (hf_variant or "").lower():
|
|
return False
|
|
if self._requested_n_ctx != int(n_ctx):
|
|
return False
|
|
|
|
def _norm(value):
|
|
if value is None:
|
|
return None
|
|
if isinstance(value, str):
|
|
stripped = value.strip().lower()
|
|
return stripped or None
|
|
return value
|
|
|
|
if _norm(self._cache_type_kv) != _norm(cache_type_kv):
|
|
return False
|
|
|
|
# Reconcile a user --split-mode in extras (load_model does the same), so
|
|
# an extras-driven tensor load isn't seen as a mismatch that needlessly
|
|
# kills/reloads a healthy server.
|
|
if self._tensor_parallel != resolve_tensor_parallel(extra_args, tensor_parallel):
|
|
return False
|
|
|
|
# Compare on the canonical requested mode. With --spec-type in
|
|
# extra_args the backend stores None; mirror that here.
|
|
if _extra_args_set_spec_type(extra_args):
|
|
req_mode = None
|
|
else:
|
|
req_mode = _canonicalize_spec_mode(speculative_type) or "auto"
|
|
backend_mode = self._requested_spec_mode
|
|
if req_mode != backend_mode:
|
|
return False
|
|
|
|
# spec_draft_n_max only matters when an MTP variant is engaged. Compare
|
|
# on the resolved spec so an Auto request promoted to draft-mtp still
|
|
# bounces a reload when n_max changes.
|
|
if (
|
|
self._speculative_type == "draft-mtp"
|
|
and spec_draft_n_max is not None
|
|
and int(spec_draft_n_max) != (self._spec_draft_n_max or 0)
|
|
):
|
|
return False
|
|
|
|
if (self._chat_template_override or None) != (chat_template_override or None):
|
|
return False
|
|
|
|
# A drafter appearing/disappearing next to a local GGUF changes the
|
|
# launch command (--model-draft) when the mode can use it; without
|
|
# this, adding mtp-*.gguf after a load is deduped away and MTP can't
|
|
# engage short of an unload. HF loads resolve the drafter inside
|
|
# load_model (gguf_path is None here), so only local paths compare;
|
|
# the route-level probe covers HF cache repos. No sub-3B gate: both
|
|
# sides come from the same config detection, so a sub-3B mismatch
|
|
# only happens when a drafter genuinely appeared (one benign reload,
|
|
# then the stored path converges).
|
|
if (
|
|
gguf_path is not None
|
|
and req_mode in ("auto", "mtp", "mtp+ngram")
|
|
and (mtp_draft_path or None) != (self._mtp_draft_path or None)
|
|
):
|
|
return False
|
|
|
|
# extra_args=None means "no opinion" (inherit handled at the route
|
|
# layer); only an explicit list forces equality.
|
|
if extra_args is not None:
|
|
current = list(self._extra_args) if self._extra_args is not None else []
|
|
if list(extra_args) != current:
|
|
return False
|
|
return True
|
|
|
|
def _classify_gpu_offload(
|
|
self, expected_gpu: bool, detected_gpus: list[tuple[int, int]]
|
|
) -> Optional[bool]:
|
|
"""True if a GPU model buffer was allocated, False if only CPU
|
|
buffers landed despite GPU intent, None when there's no signal (no
|
|
GPU detected, no buffer-size lines, etc.)."""
|
|
if not detected_gpus or not expected_gpu:
|
|
return None
|
|
# llama-server logs one "model buffer size = N MiB" line per backend
|
|
# buffer; CUDA/ROCm/Metal/Vulkan/OpenCL/SYCL are GPU, CPU* are not.
|
|
gpu_markers = ("CUDA", "ROCm", "Metal", "Vulkan", "OpenCL", "SYCL")
|
|
saw_buffer_line = False
|
|
saw_gpu_buffer = False
|
|
for line in self._stdout_lines:
|
|
if "model buffer size" not in line:
|
|
continue
|
|
saw_buffer_line = True
|
|
if any(marker in line for marker in gpu_markers):
|
|
saw_gpu_buffer = True
|
|
break
|
|
if not saw_buffer_line:
|
|
return None
|
|
return saw_gpu_buffer
|
|
|
|
def load_cancelled(self) -> bool:
|
|
"""True if a load was cancelled (e.g. via unload/_cancel_event) and not
|
|
yet consumed by the next load_model. Lets the tensor->layer fallback
|
|
avoid restarting a load the user just cancelled."""
|
|
return self._cancel_event.is_set()
|
|
|
|
def unload_model(self) -> bool:
|
|
"""Terminate the subprocess and cancel any in-flight download."""
|
|
self._cancel_event.set()
|
|
with self._lock:
|
|
self._kill_process()
|
|
logger.info(f"Unloaded GGUF model: {self._model_identifier}")
|
|
self._model_identifier = None
|
|
self._gguf_path = None
|
|
self._hf_repo = None
|
|
self._mtp_draft_path = None
|
|
self._spec_fallback_reason = None
|
|
self._hf_variant = None
|
|
self._is_vision = False
|
|
self._is_audio = False
|
|
self._audio_type = None
|
|
self._audio_probed = False
|
|
self._has_audio_input = False
|
|
self._mmproj_has_audio = False
|
|
self._port = None
|
|
self._healthy = False
|
|
self._context_length = None
|
|
self._effective_context_length = None
|
|
self._max_context_length = None
|
|
self._chat_template = None
|
|
self._chat_template_override = None
|
|
self._supports_reasoning = False
|
|
self._reasoning_always_on = False
|
|
self._reasoning_style = "enable_thinking"
|
|
self._reasoning_default = True
|
|
self._supports_preserve_thinking = False
|
|
self._supports_tools = False
|
|
self._cache_type_kv = None
|
|
self._tensor_parallel = False
|
|
self._speculative_type = None
|
|
self._requested_spec_mode = None
|
|
self._spec_draft_n_max = None
|
|
self._n_layers = None
|
|
self._n_kv_heads = None
|
|
self._n_kv_heads_by_layer = None
|
|
self._n_heads = None
|
|
self._embedding_length = None
|
|
self._kv_key_length = None
|
|
self._kv_value_length = None
|
|
self._sliding_window = None
|
|
self._sliding_window_pattern = None
|
|
self._full_attention_interval = None
|
|
self._kv_lora_rank = None
|
|
self._key_length_mla = None
|
|
self._kv_key_length_swa = None
|
|
self._kv_value_length_swa = None
|
|
self._ssm_inner_size = None
|
|
self._ssm_state_size = None
|
|
self._shared_kv_layers = None
|
|
self._nextn_predict_layers = None
|
|
# Clean up temp chat template file.
|
|
if hasattr(self, "_chat_template_file") and self._chat_template_file:
|
|
try:
|
|
import os
|
|
os.unlink(self._chat_template_file.name)
|
|
except Exception:
|
|
pass
|
|
self._chat_template_file = None
|
|
# Free audio codec GPU memory.
|
|
if LlamaCppBackend._codec_mgr is not None:
|
|
LlamaCppBackend._codec_mgr.unload()
|
|
LlamaCppBackend._codec_mgr = None
|
|
import torch
|
|
|
|
if torch.cuda.is_available():
|
|
torch.cuda.empty_cache()
|
|
return True
|
|
|
|
def _kill_process(self):
|
|
"""Terminate the subprocess if running."""
|
|
if self._process is None:
|
|
return
|
|
try:
|
|
self._process.terminate()
|
|
self._process.wait(timeout = 5)
|
|
except subprocess.TimeoutExpired:
|
|
logger.warning("llama-server did not exit on SIGTERM, sending SIGKILL")
|
|
self._process.kill()
|
|
self._process.wait(timeout = 5)
|
|
except Exception as e:
|
|
logger.warning(f"Error killing llama-server process: {e}")
|
|
finally:
|
|
self._process = None
|
|
# Clear healthy so a /load during the replacement's warm-up can't
|
|
# short-circuit against the previous server's health (#5401).
|
|
self._healthy = False
|
|
# Drives _wait_for_vram_settle in the next load_model; set in finally
|
|
# so both in-process and frontend Apply paths record the kill.
|
|
self._last_kill_monotonic = time.monotonic()
|
|
if self._stdout_thread is not None:
|
|
self._stdout_thread.join(timeout = 2)
|
|
self._stdout_thread = None
|
|
fh = getattr(self, "_llama_log_fh", None)
|
|
if fh is not None:
|
|
try:
|
|
fh.close()
|
|
except Exception:
|
|
pass
|
|
self._llama_log_fh = None
|
|
|
|
@staticmethod
|
|
def _kill_orphaned_servers():
|
|
"""Kill orphaned llama-server processes started by studio.
|
|
|
|
Only kills processes whose resolved binary lives under a known
|
|
Studio install dir (or matches an exact env-var override), to avoid
|
|
terminating unrelated llama-server instances. Mirrors every location
|
|
_find_llama_server_binary() can return, so orphans from any
|
|
supported install path are cleaned up.
|
|
|
|
Uses psutil for cross-platform support (Linux, macOS, Windows);
|
|
falls back to pgrep + /proc/<pid>/exe on Linux when psutil is
|
|
absent.
|
|
"""
|
|
try:
|
|
# -- Build the ownership allowlist --------------------------------
|
|
# exact_binaries -- env var overrides (exact path match).
|
|
# install_roots -- Studio-owned dir trees (binary must be under one).
|
|
install_roots: list[Path] = []
|
|
|
|
# Env-mode custom root (mirrors _find_llama_server_binary).
|
|
_is_custom_root = False
|
|
try:
|
|
from utils.paths.storage_roots import studio_root as _sr # noqa: WPS433
|
|
|
|
_resolved_sr = _sr()
|
|
_legacy_studio = Path.home() / ".unsloth" / "studio"
|
|
try:
|
|
_is_custom_root = _resolved_sr.resolve() != _legacy_studio.resolve()
|
|
except (OSError, ValueError):
|
|
_is_custom_root = _resolved_sr != _legacy_studio
|
|
if _is_custom_root:
|
|
install_roots.append(_resolved_sr / "llama.cpp")
|
|
except (ImportError, OSError, ValueError):
|
|
pass
|
|
|
|
# Primary install dir (default mode only). Env-mode skips this so a
|
|
# custom-root Studio can't kill a default-install Studio's server.
|
|
if not _is_custom_root:
|
|
install_roots.append(Path.home() / ".unsloth" / "llama.cpp")
|
|
|
|
# Legacy in-tree build dirs (older setup.sh)
|
|
project_root = Path(__file__).resolve().parents[4]
|
|
install_roots.append(project_root / "llama.cpp")
|
|
|
|
# Legacy: extracted binary
|
|
install_roots.append(project_root / "bin")
|
|
|
|
# UNSLOTH_LLAMA_CPP_PATH env var (custom install dir)
|
|
custom_dir = os.environ.get("UNSLOTH_LLAMA_CPP_PATH")
|
|
if custom_dir:
|
|
install_roots.append(Path(custom_dir))
|
|
|
|
# LLAMA_SERVER_PATH env var (exact binary path)
|
|
exact_binaries: list[Path] = []
|
|
env_binary = os.environ.get("LLAMA_SERVER_PATH")
|
|
if env_binary:
|
|
try:
|
|
exact_binaries.append(Path(env_binary).resolve())
|
|
except OSError:
|
|
pass
|
|
|
|
# Resolve all roots so is_relative_to works reliably.
|
|
resolved_roots: list[Path] = []
|
|
for root in install_roots:
|
|
try:
|
|
resolved_roots.append(root.resolve())
|
|
except OSError:
|
|
pass
|
|
|
|
my_pid = os.getpid()
|
|
|
|
# -- Enumerate processes -------------------------------------------
|
|
# Prefer psutil (cross-platform); fall back to pgrep + /proc on
|
|
# Linux when psutil is absent.
|
|
try:
|
|
import psutil
|
|
has_psutil = True
|
|
except ImportError:
|
|
has_psutil = False
|
|
|
|
if has_psutil:
|
|
for proc in psutil.process_iter(["pid", "name", "exe"]):
|
|
try:
|
|
if proc.info["pid"] == my_pid:
|
|
continue
|
|
|
|
name = proc.info.get("name") or ""
|
|
if not name.lower().startswith("llama-server"):
|
|
continue
|
|
|
|
exe = proc.info.get("exe")
|
|
if not exe:
|
|
continue
|
|
|
|
exe_path = Path(exe).resolve()
|
|
|
|
# Ownership: exact match OR binary under a known root.
|
|
is_ours = exe_path in exact_binaries or any(
|
|
exe_path.is_relative_to(root) for root in resolved_roots
|
|
)
|
|
if not is_ours:
|
|
continue
|
|
|
|
proc.kill()
|
|
logger.info(
|
|
f"Killed orphaned llama-server process (pid={proc.info['pid']})"
|
|
)
|
|
except (
|
|
psutil.NoSuchProcess,
|
|
psutil.AccessDenied,
|
|
psutil.ZombieProcess,
|
|
):
|
|
pass
|
|
else:
|
|
# -- Fallback: pgrep + /proc/<pid>/exe (Linux only) -----------
|
|
if sys.platform != "linux":
|
|
return
|
|
result = subprocess.run(
|
|
["pgrep", "-a", "-f", "llama-server"],
|
|
capture_output = True,
|
|
text = True,
|
|
timeout = 5,
|
|
env = child_env_without_native_path_secret(),
|
|
)
|
|
if result.returncode != 0:
|
|
return
|
|
|
|
for line in result.stdout.strip().splitlines():
|
|
parts = line.strip().split(None, 1)
|
|
if len(parts) < 2:
|
|
continue
|
|
pid = int(parts[0])
|
|
if pid == my_pid:
|
|
continue
|
|
|
|
# /proc/<pid>/exe symlinks the real binary, avoiding
|
|
# cmdline-parsing ambiguities; fall back to the first
|
|
# cmdline token when /proc is unavailable.
|
|
proc_exe = Path(f"/proc/{pid}/exe")
|
|
try:
|
|
binary = proc_exe.resolve(strict = True)
|
|
except (OSError, ValueError):
|
|
cmdline = parts[1]
|
|
token = cmdline.split()[0] if cmdline.strip() else ""
|
|
if not token:
|
|
continue
|
|
binary = Path(token).resolve(strict = False)
|
|
|
|
owned = binary in exact_binaries or any(
|
|
binary.is_relative_to(root) for root in resolved_roots
|
|
)
|
|
if not owned:
|
|
continue
|
|
|
|
try:
|
|
os.kill(pid, signal.SIGKILL)
|
|
logger.info(f"Killed orphaned llama-server process (pid={pid})")
|
|
except ProcessLookupError:
|
|
pass
|
|
except PermissionError:
|
|
pass
|
|
except Exception:
|
|
logger.warning("Error during orphan server cleanup", exc_info = True)
|
|
|
|
def _cleanup(self):
|
|
"""atexit handler to ensure llama-server is terminated."""
|
|
self._kill_process()
|
|
|
|
@staticmethod
|
|
def _fit_off_retry_eligible(cmd: "list[str]", use_fit: bool) -> bool:
|
|
"""Whether a llama-server startup crash may be retried with --fit off.
|
|
|
|
Only when Studio's own VRAM math placed the model (use_fit=False)
|
|
and nothing on the command line set the fit mode explicitly
|
|
(-fit / --fit, space- or equals-form). --fit-ctx / --fit-target /
|
|
-fitc / -fitt tune the fit step but do not select the mode, so
|
|
they do not block the retry.
|
|
"""
|
|
if use_fit:
|
|
return False
|
|
for a in cmd:
|
|
if a in ("-fit", "--fit") or a.startswith(("-fit=", "--fit=")):
|
|
return False
|
|
return True
|
|
|
|
def _wait_for_health(
|
|
self,
|
|
timeout: float = 120.0,
|
|
interval: float = 0.5,
|
|
) -> bool:
|
|
"""Poll llama-server's /health until 200; also detect early exit/crash."""
|
|
deadline = time.monotonic() + timeout
|
|
url = f"http://127.0.0.1:{self._port}/health"
|
|
|
|
while time.monotonic() < deadline:
|
|
# Process crashed?
|
|
if self._process.poll() is not None:
|
|
# Let the drain thread collect final output.
|
|
if self._stdout_thread is not None:
|
|
self._stdout_thread.join(timeout = 2)
|
|
output = "\n".join(self._stdout_lines[-50:])
|
|
# Keep the TAIL: crash details (abort reason, ROCm/CUDA error
|
|
# text) print last, after the long startup banner. Head
|
|
# truncation has cut off exactly the diagnostic line before.
|
|
_log_hint = (
|
|
f" Full log: {self._llama_log_path}"
|
|
if getattr(self, "_llama_log_path", None)
|
|
else ""
|
|
)
|
|
logger.error(
|
|
f"llama-server exited with code {self._process.returncode}. "
|
|
f"Output (tail): {output[-2000:]}{_log_hint}"
|
|
)
|
|
return False
|
|
|
|
try:
|
|
resp = httpx.get(url, timeout = 2.0)
|
|
if resp.status_code == 200:
|
|
return True
|
|
except (
|
|
httpx.ConnectError,
|
|
httpx.TimeoutException,
|
|
# ReadError covers TCP RST mid-read while still binding the port
|
|
# (Windows: WinError 10054); the crash branch catches real exits.
|
|
httpx.ReadError,
|
|
httpx.RemoteProtocolError,
|
|
httpx.WriteError,
|
|
):
|
|
pass
|
|
|
|
time.sleep(interval)
|
|
|
|
logger.error(f"llama-server health check timed out after {timeout}s")
|
|
return False
|
|
|
|
@staticmethod
|
|
def _ctx_integrity_flags(
|
|
n_parallel: int, use_fit: bool, requested_ctx: int, effective_ctx: int, caps: dict
|
|
) -> list[str]:
|
|
"""Flags that keep the per-request window equal to the advertised ctx.
|
|
|
|
Explicit ``--parallel`` disables llama-server's auto-slots
|
|
``--kv-unified`` default, silently splitting ``-c`` into per-slot
|
|
windows of ``-c / N``; restore the shared pool so one request can use
|
|
the full context. With ``--fit on``, ``--fit-ctx`` floors the fit step
|
|
at an explicitly requested ctx (default floor is 4096) so it offloads
|
|
or fails instead of silently shrinking the window.
|
|
"""
|
|
flags: list[str] = []
|
|
if n_parallel > 1 and caps.get("supports_kv_unified"):
|
|
flags.append("--kv-unified")
|
|
if use_fit and requested_ctx > 0 and effective_ctx > 0 and caps.get("supports_fit_ctx"):
|
|
flags.extend(["--fit-ctx", str(effective_ctx)])
|
|
return flags
|
|
|
|
def _query_server_n_ctx(self) -> Optional[int]:
|
|
"""Per-slot context llama-server actually allocated, from ``/props``.
|
|
|
|
The memory-fit step or ``--parallel`` slot split can leave this below
|
|
the requested ``-c``; requests are validated against this value.
|
|
"""
|
|
url = f"http://127.0.0.1:{self._port}/props"
|
|
try:
|
|
resp = httpx.get(url, timeout = 5.0)
|
|
if resp.status_code != 200:
|
|
return None
|
|
settings = resp.json().get("default_generation_settings") or {}
|
|
n_ctx = settings.get("n_ctx")
|
|
return int(n_ctx) if n_ctx else None
|
|
except Exception:
|
|
return None
|
|
|
|
def _reconcile_effective_ctx_with_server(self) -> None:
|
|
"""Adopt the server's real ``n_ctx`` when it is below Studio's value.
|
|
|
|
Keeps ``context_length`` (load response, status route, passthrough
|
|
``max_tokens`` ceiling) honest; clients sized to the requested value
|
|
would otherwise hit ``exceed_context_size_error`` 400s early.
|
|
"""
|
|
actual_n_ctx = self._query_server_n_ctx()
|
|
if not actual_n_ctx or actual_n_ctx <= 0:
|
|
return
|
|
if self._effective_context_length and actual_n_ctx < self._effective_context_length:
|
|
logger.warning(
|
|
"llama-server allocated a smaller per-request context than "
|
|
f"requested ({self._effective_context_length} -> {actual_n_ctx}; "
|
|
"memory fit or --parallel slot split); clients must treat "
|
|
f"{actual_n_ctx} as the real context window."
|
|
)
|
|
self._effective_context_length = actual_n_ctx
|
|
elif not self._effective_context_length:
|
|
self._effective_context_length = actual_n_ctx
|
|
|
|
# ── Message building (OpenAI format) ──────────────────────────
|
|
|
|
@staticmethod
|
|
def _parse_tool_calls_from_text(content: str, *, allow_incomplete: bool = True) -> list[dict]:
|
|
"""Thin wrapper around the shared parser in tool_call_parser
|
|
so safetensors and llama_cpp pick up the same fixes."""
|
|
return _shared_parse_tool_calls_from_text(
|
|
content,
|
|
allow_incomplete = allow_incomplete,
|
|
)
|
|
|
|
@staticmethod
|
|
def _build_openai_messages(messages: list[dict], image_b64: Optional[str] = None) -> list[dict]:
|
|
"""Build OpenAI-format messages, optionally injecting an image_url part
|
|
into the last user message for vision models. As-is if no image."""
|
|
if not image_b64:
|
|
return messages
|
|
|
|
# Convert the last user message to multimodal content parts
|
|
result = [msg.copy() for msg in messages]
|
|
last_user_idx = None
|
|
for i, msg in enumerate(result):
|
|
if msg["role"] == "user":
|
|
last_user_idx = i
|
|
|
|
if last_user_idx is not None:
|
|
text_content = result[last_user_idx].get("content", "")
|
|
result[last_user_idx]["content"] = [
|
|
{"type": "text", "text": text_content},
|
|
{
|
|
"type": "image_url",
|
|
"image_url": {
|
|
"url": f"data:image/png;base64,{image_b64}",
|
|
},
|
|
},
|
|
]
|
|
|
|
return result
|
|
|
|
# ── Generation (proxy to llama-server) ────────────────────────
|
|
|
|
@staticmethod
|
|
def _iter_text_cancellable(
|
|
response: "httpx.Response", cancel_event: Optional[threading.Event] = None
|
|
) -> Generator[str, None, None]:
|
|
"""Iterate an httpx streaming response with cancel support.
|
|
|
|
Checks cancel_event between chunks and on ReadTimeout; the
|
|
_stream_with_retry watcher also closes the response on cancel.
|
|
"""
|
|
text_iter = response.iter_text()
|
|
while True:
|
|
if cancel_event is not None and cancel_event.is_set():
|
|
response.close()
|
|
return
|
|
try:
|
|
chunk = next(text_iter)
|
|
yield chunk
|
|
except StopIteration:
|
|
return
|
|
except httpx.ReadTimeout:
|
|
# No data within the timeout window -- loop back and re-check
|
|
# cancel_event.
|
|
continue
|
|
|
|
@staticmethod
|
|
@contextlib.contextmanager
|
|
def _stream_with_retry(
|
|
client: "httpx.Client",
|
|
url: str,
|
|
payload: dict,
|
|
cancel_event: Optional[threading.Event] = None,
|
|
headers: Optional[dict] = None,
|
|
):
|
|
"""Open an httpx streaming POST with cancel support.
|
|
|
|
Sends once with a long read timeout (120 s) so prefill finishes without
|
|
a retry storm (the old 0.5 s timeout caused duplicate POSTs every half
|
|
second). A watcher thread cancels by closing the response. httpx can't
|
|
interrupt a blocked read before the response exists, so cancel during
|
|
the header wait (1-5 s prefill) is deferred until headers arrive.
|
|
"""
|
|
if cancel_event is not None and cancel_event.is_set():
|
|
raise GeneratorExit
|
|
|
|
# Background watcher: close the response if cancel is requested.
|
|
# Only effective after response headers arrive (httpx limitation).
|
|
_cancel_closed = threading.Event()
|
|
_response_ref: list = [None]
|
|
|
|
def _cancel_watcher():
|
|
while not _cancel_closed.is_set():
|
|
if cancel_event.wait(timeout = 0.3):
|
|
# Cancel requested. Poll until the response object exists
|
|
# so we can close it, or until the main thread finishes
|
|
# (_cancel_closed set in finally).
|
|
while not _cancel_closed.is_set():
|
|
r = _response_ref[0]
|
|
if r is not None:
|
|
try:
|
|
r.close()
|
|
return
|
|
except Exception as e:
|
|
logger.debug(f"Error closing response in cancel watcher: {e}")
|
|
# Response not created yet -- wait briefly and retry
|
|
_cancel_closed.wait(timeout = 0.1)
|
|
return
|
|
|
|
watcher = None
|
|
if cancel_event is not None:
|
|
watcher = threading.Thread(target = _cancel_watcher, daemon = True, name = "prefill-cancel")
|
|
watcher.start()
|
|
|
|
try:
|
|
# Long read timeout so prefill can finish without a retry storm.
|
|
# Cancel during prefill and streaming is handled by the watcher
|
|
# thread closing the response, unblocking any httpx read.
|
|
prefill_timeout = httpx.Timeout(
|
|
connect = 30,
|
|
read = 120.0,
|
|
write = 10,
|
|
pool = 10,
|
|
)
|
|
with client.stream(
|
|
"POST",
|
|
url,
|
|
json = payload,
|
|
timeout = prefill_timeout,
|
|
headers = headers,
|
|
) as response:
|
|
_response_ref[0] = response
|
|
if cancel_event is not None and cancel_event.is_set():
|
|
raise GeneratorExit
|
|
yield response
|
|
return
|
|
except (httpx.ReadError, httpx.RemoteProtocolError, httpx.CloseError):
|
|
# Response was closed by the cancel watcher
|
|
if cancel_event is not None and cancel_event.is_set():
|
|
raise GeneratorExit
|
|
raise
|
|
finally:
|
|
_cancel_closed.set()
|
|
|
|
def generate_chat_completion(
|
|
self,
|
|
messages: list[dict],
|
|
image_b64: Optional[str] = None,
|
|
temperature: float = 0.6,
|
|
top_p: float = 0.95,
|
|
top_k: int = 20,
|
|
min_p: float = 0.01,
|
|
max_tokens: Optional[int] = None,
|
|
repetition_penalty: float = 1.0,
|
|
presence_penalty: float = 0.0,
|
|
stop: Optional[list[str]] = None,
|
|
cancel_event: Optional[threading.Event] = None,
|
|
enable_thinking: Optional[bool] = None,
|
|
reasoning_effort: Optional[str] = None,
|
|
preserve_thinking: Optional[bool] = None,
|
|
seed: Optional[int] = None,
|
|
) -> Generator[str | dict, None, None]:
|
|
"""
|
|
Send a chat completion to llama-server and stream tokens back.
|
|
|
|
Uses /v1/chat/completions -- llama-server applies the chat template
|
|
and handles vision (multimodal image_url parts) natively.
|
|
|
|
Yields cumulative text (matching InferenceBackend's convention).
|
|
"""
|
|
if not self.is_loaded:
|
|
raise RuntimeError("llama-server is not loaded")
|
|
|
|
openai_messages = self._build_openai_messages(messages, image_b64)
|
|
|
|
payload = {
|
|
"messages": openai_messages,
|
|
"stream": True,
|
|
"temperature": temperature,
|
|
"top_p": top_p,
|
|
"top_k": top_k if top_k >= 0 else 0,
|
|
"min_p": min_p,
|
|
"repeat_penalty": repetition_penalty,
|
|
"presence_penalty": presence_penalty,
|
|
}
|
|
# Per-request enable_thinking / reasoning_effort / preserve_thinking
|
|
_reasoning_kw = self._request_reasoning_kwargs(
|
|
enable_thinking, reasoning_effort, preserve_thinking
|
|
)
|
|
if _reasoning_kw is not None:
|
|
payload["chat_template_kwargs"] = _reasoning_kw
|
|
# Cap to the effective context length when known, else the floor.
|
|
# The wall-clock backstop below stops a stuck model regardless.
|
|
payload["max_tokens"] = (
|
|
max_tokens
|
|
if max_tokens is not None
|
|
else (self._effective_context_length or _DEFAULT_MAX_TOKENS_FLOOR)
|
|
)
|
|
payload["t_max_predict_ms"] = _DEFAULT_T_MAX_PREDICT_MS
|
|
if stop:
|
|
payload["stop"] = stop
|
|
if seed is not None:
|
|
payload["seed"] = seed
|
|
payload["stream_options"] = {"include_usage": True}
|
|
|
|
url = f"{self.base_url}/v1/chat/completions"
|
|
cumulative = ""
|
|
in_thinking = False
|
|
_stream_done = False
|
|
_metadata_usage = None
|
|
_metadata_timings = None
|
|
_metadata_finish_reason = None
|
|
|
|
try:
|
|
# _stream_with_retry uses a 120 s read timeout so prefill can
|
|
# finish. Cancel during streaming is handled by the watcher
|
|
# thread (closes the response on cancel_event).
|
|
stream_timeout = httpx.Timeout(connect = 10, read = 0.5, write = 10, pool = 10)
|
|
_auth_headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else None
|
|
with httpx.Client(
|
|
timeout = stream_timeout, limits = httpx.Limits(max_keepalive_connections = 0)
|
|
) as client:
|
|
with self._stream_with_retry(
|
|
client,
|
|
url,
|
|
payload,
|
|
cancel_event,
|
|
headers = _auth_headers,
|
|
) as response:
|
|
if response.status_code != 200:
|
|
error_body = response.read().decode()
|
|
raise RuntimeError(
|
|
f"llama-server returned {response.status_code}: {error_body}"
|
|
)
|
|
|
|
buffer = ""
|
|
has_content_tokens = False
|
|
reasoning_text = ""
|
|
for raw_chunk in self._iter_text_cancellable(response, cancel_event):
|
|
buffer += raw_chunk
|
|
while "\n" in buffer:
|
|
line, buffer = buffer.split("\n", 1)
|
|
line = line.strip()
|
|
|
|
if not line:
|
|
continue
|
|
if line == "data: [DONE]":
|
|
if in_thinking:
|
|
if has_content_tokens:
|
|
# Real thinking + content: close the tag
|
|
cumulative += "</think>"
|
|
yield cumulative
|
|
else:
|
|
# Only reasoning_content, no content:
|
|
# model put its whole reply in reasoning
|
|
# (e.g. Qwen3 always-think). Show it as
|
|
# the main response, not a thinking block.
|
|
cumulative = reasoning_text
|
|
yield cumulative
|
|
_stream_done = True
|
|
break # exit inner while
|
|
if not line.startswith("data: "):
|
|
continue
|
|
|
|
try:
|
|
data = json.loads(line[6:])
|
|
# Diffusion frame (per-step canvas) from the shim: forward untouched so
|
|
# the frontend renders it in place. No assistant text, so it never enters
|
|
# the cumulative content.
|
|
if data.get("type") == "diffusion_frame":
|
|
yield data
|
|
continue
|
|
# Capture server timings/usage from final chunks.
|
|
_chunk_timings = data.get("timings")
|
|
if _chunk_timings:
|
|
_metadata_timings = _chunk_timings
|
|
_chunk_usage = data.get("usage")
|
|
if _chunk_usage:
|
|
_metadata_usage = _chunk_usage
|
|
choices = data.get("choices", [])
|
|
if choices:
|
|
delta = choices[0].get("delta", {})
|
|
_fr = choices[0].get("finish_reason")
|
|
if _fr:
|
|
_metadata_finish_reason = _fr
|
|
|
|
# Reasoning/thinking tokens: llama-server
|
|
# sends these as "reasoning_content"; wrap
|
|
# in <think> tags for the frontend parser.
|
|
reasoning = delta.get("reasoning_content", "")
|
|
if reasoning:
|
|
reasoning_text += reasoning
|
|
if not in_thinking:
|
|
cumulative += "<think>"
|
|
in_thinking = True
|
|
cumulative += reasoning
|
|
yield cumulative
|
|
|
|
token = delta.get("content", "")
|
|
if token:
|
|
has_content_tokens = True
|
|
if in_thinking:
|
|
cumulative += "</think>"
|
|
in_thinking = False
|
|
cumulative += token
|
|
yield cumulative
|
|
except json.JSONDecodeError:
|
|
logger.debug(f"Skipping malformed SSE line: {line[:100]}")
|
|
if _stream_done:
|
|
break # exit outer for
|
|
if _metadata_usage or _metadata_timings or _metadata_finish_reason:
|
|
_metadata_usage = _backfill_usage_from_timings(
|
|
_metadata_usage, _metadata_timings
|
|
)
|
|
yield {
|
|
"type": "metadata",
|
|
# Never None: a finish-only metadata event (no usage,
|
|
# no timings) would otherwise crash consumers that do
|
|
# usage.get(...) on the non-streaming paths.
|
|
"usage": _metadata_usage or {},
|
|
"timings": _metadata_timings,
|
|
"finish_reason": _metadata_finish_reason,
|
|
}
|
|
|
|
except httpx.ConnectError:
|
|
raise RuntimeError("Lost connection to llama-server")
|
|
except Exception as e:
|
|
if cancel_event is not None and cancel_event.is_set():
|
|
return
|
|
raise
|
|
|
|
# ── Tool-calling agentic loop ──────────────────────────────
|
|
|
|
def generate_chat_completion_with_tools(
|
|
self,
|
|
messages: list[dict],
|
|
tools: list[dict],
|
|
temperature: float = 0.6,
|
|
top_p: float = 0.95,
|
|
top_k: int = 20,
|
|
min_p: float = 0.01,
|
|
max_tokens: Optional[int] = None,
|
|
repetition_penalty: float = 1.0,
|
|
presence_penalty: float = 0.0,
|
|
stop: Optional[list[str]] = None,
|
|
cancel_event: Optional[threading.Event] = None,
|
|
enable_thinking: Optional[bool] = None,
|
|
reasoning_effort: Optional[str] = None,
|
|
preserve_thinking: Optional[bool] = None,
|
|
max_tool_iterations: int = 25,
|
|
auto_heal_tool_calls: bool = True,
|
|
tool_call_timeout: int = 300,
|
|
session_id: Optional[str] = None,
|
|
rag_scope: Optional[dict] = None,
|
|
seed: Optional[int] = None,
|
|
disable_parallel_tool_use: bool = False,
|
|
confirm_tool_calls: bool = False,
|
|
) -> Generator[dict, None, None]:
|
|
"""
|
|
Agentic loop: let the model call tools, execute them, and continue.
|
|
|
|
Yields dicts:
|
|
{"type": "status", "text": "Searching: ..."/"Reading: ..."} -- tool status updates
|
|
{"type": "content", "text": "token"} -- streamed content tokens (cumulative)
|
|
{"type": "reasoning", "text": "token"} -- streamed reasoning tokens (cumulative)
|
|
"""
|
|
from core.inference.tools import build_rag_autoinject, execute_tool
|
|
|
|
if not self.is_loaded:
|
|
raise RuntimeError("llama-server is not loaded")
|
|
|
|
conversation = list(messages)
|
|
|
|
# Forced first-pass RAG so a doc question doesn't lose to web_search. Emits
|
|
# the same tool card + citations a real call would.
|
|
_auto = None if confirm_tool_calls else build_rag_autoinject(conversation, rag_scope)
|
|
if _auto:
|
|
for _ev in _auto["events"]:
|
|
yield _ev
|
|
conversation.extend(_auto["messages"])
|
|
|
|
url = f"{self.base_url}/v1/chat/completions"
|
|
_accumulated_completion_tokens = 0
|
|
_accumulated_predicted_ms = 0.0
|
|
_accumulated_predicted_n = 0
|
|
|
|
def _strip_tool_markup(
|
|
text: str,
|
|
*,
|
|
final: bool = False,
|
|
force: bool = False,
|
|
) -> str:
|
|
if not (auto_heal_tool_calls or force):
|
|
return text
|
|
return strip_tool_call_markup(text, final = final)
|
|
|
|
def _strip_tool_markup_streaming(text: str, *, force: bool = False) -> str:
|
|
if not (auto_heal_tool_calls or force):
|
|
return text
|
|
for pat in _TOOL_ALL_PATS:
|
|
text = pat.sub("", text)
|
|
return text
|
|
|
|
tool_controller = ToolLoopController(
|
|
tools = tools,
|
|
auto_heal_tool_calls = auto_heal_tool_calls,
|
|
)
|
|
|
|
def _tool_succeeded(tool_name: str) -> bool:
|
|
key_prefix = f"{tool_name}:"
|
|
return any(
|
|
record.executed and not record.is_error and record.key.startswith(key_prefix)
|
|
for record in tool_controller.history
|
|
)
|
|
|
|
_MAX_BUFFER_CHARS = 32
|
|
_append_budget_exhausted_nudge = True
|
|
# RAG: cap knowledge-base searches per assistant turn. The controller is
|
|
# tool-agnostic, so this gate stays in the loop.
|
|
_kb_search_count = 0
|
|
|
|
# ── Re-prompt on plan-without-action ─────────────────
|
|
# When the model describes what it intends to do (forward-looking
|
|
# language) without calling a tool, re-prompt once. Only triggers on
|
|
# responses signaling intent/planning -- a direct answer like "4" or
|
|
# "Hello!" won't match. Pattern compiled at module level
|
|
# (_INTENT_SIGNAL).
|
|
_reprompt_count = 0
|
|
_forced_tool_call_pending = False
|
|
|
|
# Reserve extra iterations for re-prompts so they don't consume the
|
|
# caller's tool-call budget; only when tool iterations are allowed.
|
|
_extra = _MAX_REPROMPTS if max_tool_iterations > 0 else 0
|
|
for iteration in range(max_tool_iterations + _extra):
|
|
if cancel_event is not None and cancel_event.is_set():
|
|
return
|
|
|
|
active_tools = tool_controller.active_tools()
|
|
if not active_tools:
|
|
_append_budget_exhausted_nudge = False
|
|
break
|
|
_tool_xml_signals = TOOL_XML_SIGNALS if active_tools else ()
|
|
|
|
# Build payload -- stream: True so we detect tool signals
|
|
# in the first 1-2 chunks without a non-streaming penalty.
|
|
payload = {
|
|
"messages": conversation,
|
|
"stream": True,
|
|
"stream_options": {"include_usage": True},
|
|
"temperature": temperature,
|
|
"top_p": top_p,
|
|
"top_k": top_k if top_k >= 0 else 0,
|
|
"min_p": min_p,
|
|
"repeat_penalty": repetition_penalty,
|
|
"presence_penalty": presence_penalty,
|
|
"tools": active_tools,
|
|
"tool_choice": "auto",
|
|
}
|
|
_reasoning_kw = self._request_reasoning_kwargs(
|
|
enable_thinking, reasoning_effort, preserve_thinking
|
|
)
|
|
if _reasoning_kw is not None:
|
|
payload["chat_template_kwargs"] = _reasoning_kw
|
|
payload["max_tokens"] = (
|
|
max_tokens
|
|
if max_tokens is not None
|
|
else (self._effective_context_length or _DEFAULT_MAX_TOKENS_FLOOR)
|
|
)
|
|
payload["t_max_predict_ms"] = _DEFAULT_T_MAX_PREDICT_MS
|
|
if stop:
|
|
payload["stop"] = stop
|
|
if seed is not None:
|
|
payload["seed"] = seed
|
|
|
|
try:
|
|
_auth_headers = (
|
|
{"Authorization": f"Bearer {self._api_key}"} if self._api_key else None
|
|
)
|
|
|
|
# ── Speculative buffer state machine ──────────────────
|
|
# BUFFERING: accumulate content, check for tool signals
|
|
# STREAMING: no tool detected, yield tokens to caller
|
|
# DRAINING: tool signal found, silently consume rest
|
|
_S_BUFFERING = 0
|
|
_S_STREAMING = 1
|
|
_S_DRAINING = 2
|
|
|
|
detect_state = _S_BUFFERING
|
|
content_buffer = "" # Raw content held during BUFFERING
|
|
content_accum = "" # All content tokens (for tool parsing)
|
|
reasoning_accum = ""
|
|
cumulative_display = "" # Cumulative yielded text (with <think>)
|
|
in_thinking = False
|
|
has_content_tokens = False
|
|
tool_calls_acc = {} # Structured delta.tool_calls fragments
|
|
has_structured_tc = False
|
|
_iter_usage = None
|
|
_iter_timings = None
|
|
_iter_finish_reason = None
|
|
_stream_done = False
|
|
_last_emitted = ""
|
|
provisional_render_html_tool_call_ids = set()
|
|
_suppress_visible_output = _forced_tool_call_pending
|
|
|
|
stream_timeout = httpx.Timeout(
|
|
connect = 10,
|
|
read = 0.5,
|
|
write = 10,
|
|
pool = 10,
|
|
)
|
|
with httpx.Client(
|
|
timeout = stream_timeout,
|
|
limits = httpx.Limits(max_keepalive_connections = 0),
|
|
) as client:
|
|
with self._stream_with_retry(
|
|
client,
|
|
url,
|
|
payload,
|
|
cancel_event,
|
|
headers = _auth_headers,
|
|
) as response:
|
|
if response.status_code != 200:
|
|
error_body = response.read().decode()
|
|
raise RuntimeError(
|
|
f"llama-server returned {response.status_code}: {error_body}"
|
|
)
|
|
|
|
raw_buf = ""
|
|
for raw_chunk in self._iter_text_cancellable(
|
|
response,
|
|
cancel_event,
|
|
):
|
|
raw_buf += raw_chunk
|
|
while "\n" in raw_buf:
|
|
line, raw_buf = raw_buf.split("\n", 1)
|
|
line = line.strip()
|
|
|
|
if not line:
|
|
continue
|
|
if line == "data: [DONE]":
|
|
# Flush thinking state for STREAMING
|
|
if detect_state == _S_STREAMING and in_thinking:
|
|
if has_content_tokens:
|
|
cumulative_display += "</think>"
|
|
if not _suppress_visible_output:
|
|
yield {
|
|
"type": "content",
|
|
"text": _strip_tool_markup(
|
|
cumulative_display,
|
|
final = True,
|
|
),
|
|
}
|
|
else:
|
|
cumulative_display = reasoning_accum
|
|
if not _suppress_visible_output:
|
|
yield {
|
|
"type": "content",
|
|
"text": cumulative_display,
|
|
}
|
|
_stream_done = True
|
|
break # exit inner while
|
|
if not line.startswith("data: "):
|
|
continue
|
|
|
|
try:
|
|
chunk_data = json.loads(line[6:])
|
|
_ct = chunk_data.get("timings")
|
|
if _ct:
|
|
_iter_timings = _ct
|
|
_cu = chunk_data.get("usage")
|
|
if _cu:
|
|
_iter_usage = _cu
|
|
|
|
choices = chunk_data.get("choices", [])
|
|
if not choices:
|
|
continue
|
|
|
|
delta = choices[0].get("delta", {})
|
|
_fr = choices[0].get("finish_reason")
|
|
if _fr:
|
|
_iter_finish_reason = _fr
|
|
|
|
# ── Structured tool_calls ──
|
|
tc_deltas = delta.get("tool_calls")
|
|
if tc_deltas:
|
|
# llama-server can emit visible assistant
|
|
# preface content before native structured
|
|
# tool_calls. Preserve content_accum as
|
|
# the assistant pre-tool text and still
|
|
# drain/execute the structured call.
|
|
has_structured_tc = True
|
|
detect_state = _S_DRAINING
|
|
for tc_d in tc_deltas:
|
|
idx = tc_d.get("index", 0)
|
|
if idx not in tool_calls_acc:
|
|
tool_calls_acc[idx] = {
|
|
"id": tc_d.get("id", f"call_{idx}"),
|
|
"type": "function",
|
|
"function": {
|
|
"name": "",
|
|
"arguments": "",
|
|
},
|
|
}
|
|
elif tc_d.get("id"):
|
|
# Update ID if a real one
|
|
# arrives on a later delta.
|
|
tool_calls_acc[idx]["id"] = tc_d["id"]
|
|
func = tc_d.get("function", {})
|
|
if func.get("name"):
|
|
tool_calls_acc[idx]["function"]["name"] += func[
|
|
"name"
|
|
]
|
|
if func.get("arguments"):
|
|
tool_calls_acc[idx]["function"]["arguments"] += (
|
|
func["arguments"]
|
|
)
|
|
current_name = tool_calls_acc[idx]["function"].get(
|
|
"name", ""
|
|
)
|
|
fallback_id = f"call_{idx}"
|
|
current_id = tool_calls_acc[idx].get("id", fallback_id)
|
|
already_started = (
|
|
current_id in provisional_render_html_tool_call_ids
|
|
)
|
|
has_real_id = current_id != fallback_id
|
|
if (
|
|
current_name == "render_html"
|
|
and not _tool_succeeded("render_html")
|
|
and any(
|
|
(
|
|
(tool.get("function") or {}).get("name")
|
|
== "render_html"
|
|
)
|
|
for tool in active_tools
|
|
)
|
|
and not already_started
|
|
and not provisional_render_html_tool_call_ids
|
|
and has_real_id
|
|
):
|
|
provisional_render_html_tool_call_ids.add(
|
|
current_id
|
|
)
|
|
yield {
|
|
"type": "tool_start",
|
|
"tool_name": "render_html",
|
|
"tool_call_id": current_id,
|
|
"arguments": {},
|
|
"provenance": tool_event_provenance(
|
|
provisional = True,
|
|
),
|
|
}
|
|
continue
|
|
|
|
# ── Reasoning tokens ──
|
|
# Yield only in STREAMING. In BUFFERING and
|
|
# DRAINING, accumulate silently so we don't
|
|
# corrupt the consumer's prev_text tracker
|
|
# (routes/inference.py never resets it
|
|
# between tool iterations).
|
|
reasoning = delta.get("reasoning_content", "")
|
|
if reasoning:
|
|
reasoning_accum += reasoning
|
|
if detect_state == _S_STREAMING:
|
|
if not in_thinking:
|
|
cumulative_display += "<think>"
|
|
in_thinking = True
|
|
cumulative_display += reasoning
|
|
if not _suppress_visible_output:
|
|
yield {
|
|
"type": "content",
|
|
"text": cumulative_display,
|
|
}
|
|
|
|
# ── Content tokens ──
|
|
token = delta.get("content", "")
|
|
if token:
|
|
has_content_tokens = True
|
|
content_accum += token
|
|
|
|
if detect_state == _S_DRAINING:
|
|
pass # accumulate silently
|
|
|
|
elif detect_state == _S_STREAMING:
|
|
if in_thinking:
|
|
cumulative_display += "</think>"
|
|
in_thinking = False
|
|
cumulative_display += token
|
|
cleaned = _strip_tool_markup_streaming(
|
|
cumulative_display
|
|
)
|
|
if len(cleaned) > len(_last_emitted):
|
|
_last_emitted = cleaned
|
|
if not _suppress_visible_output:
|
|
yield {
|
|
"type": "content",
|
|
"text": cleaned,
|
|
}
|
|
|
|
elif detect_state == _S_BUFFERING:
|
|
content_buffer += token
|
|
stripped_buf = content_buffer.lstrip()
|
|
if not stripped_buf:
|
|
continue
|
|
|
|
# Check tool signal prefixes.
|
|
is_prefix = False
|
|
is_match = False
|
|
for sig in _tool_xml_signals:
|
|
if stripped_buf.startswith(sig):
|
|
is_match = True
|
|
break
|
|
if sig.startswith(stripped_buf):
|
|
is_prefix = True
|
|
break
|
|
|
|
if is_match:
|
|
# Tool signal -- flush any visible
|
|
# prefix before DRAINING so the
|
|
# route sends it before tool_start.
|
|
if reasoning_accum:
|
|
cumulative_display += "<think>"
|
|
cumulative_display += reasoning_accum
|
|
cumulative_display += "</think>"
|
|
cumulative_display += content_buffer
|
|
cleaned = _strip_tool_markup_streaming(
|
|
cumulative_display,
|
|
force = True,
|
|
)
|
|
if len(cleaned) > len(_last_emitted):
|
|
_last_emitted = cleaned
|
|
if not _suppress_visible_output:
|
|
yield {
|
|
"type": "content",
|
|
"text": cleaned,
|
|
}
|
|
detect_state = _S_DRAINING
|
|
elif (
|
|
is_prefix and len(stripped_buf) < _MAX_BUFFER_CHARS
|
|
):
|
|
pass # keep buffering
|
|
else:
|
|
# Not a tool -- flush buffer
|
|
detect_state = _S_STREAMING
|
|
# Flush reasoning accumulated
|
|
# during BUFFERING.
|
|
if reasoning_accum:
|
|
cumulative_display += "<think>"
|
|
cumulative_display += reasoning_accum
|
|
cumulative_display += "</think>"
|
|
cumulative_display += content_buffer
|
|
cleaned = _strip_tool_markup(
|
|
cumulative_display,
|
|
)
|
|
if len(cleaned) > len(_last_emitted):
|
|
_last_emitted = cleaned
|
|
if not _suppress_visible_output:
|
|
yield {
|
|
"type": "content",
|
|
"text": cleaned,
|
|
}
|
|
|
|
except json.JSONDecodeError:
|
|
logger.debug(f"Skipping malformed SSE line: {line[:100]}")
|
|
if _stream_done:
|
|
break # exit outer for
|
|
|
|
# ── Resolve BUFFERING at stream end ──
|
|
if detect_state == _S_BUFFERING:
|
|
stripped_buf = content_buffer.lstrip()
|
|
if stripped_buf and any(s in stripped_buf for s in _tool_xml_signals):
|
|
detect_state = _S_DRAINING
|
|
elif content_accum or reasoning_accum:
|
|
detect_state = _S_STREAMING
|
|
if content_buffer:
|
|
# Flush reasoning first.
|
|
if reasoning_accum:
|
|
cumulative_display += "<think>"
|
|
cumulative_display += reasoning_accum
|
|
cumulative_display += "</think>"
|
|
cumulative_display += content_buffer
|
|
if not _suppress_visible_output:
|
|
yield {
|
|
"type": "content",
|
|
"text": _strip_tool_markup(
|
|
cumulative_display,
|
|
final = True,
|
|
),
|
|
}
|
|
elif reasoning_accum and not has_content_tokens:
|
|
# Reasoning-only response: show reasoning as plain
|
|
# text, matching the final streaming pass for
|
|
# models that put everything in reasoning.
|
|
cumulative_display = reasoning_accum
|
|
if not _suppress_visible_output:
|
|
yield {
|
|
"type": "content",
|
|
"text": cumulative_display,
|
|
}
|
|
else:
|
|
return
|
|
|
|
# ── STREAMING path: no tool call ──
|
|
if detect_state == _S_STREAMING:
|
|
# Safety net: check for XML tool signals in content. The
|
|
# route layer resets prev_text on tool_start, so post-tool
|
|
# synthesis streams correctly even if content was emitted
|
|
# before the tool XML.
|
|
_safety_tc = None
|
|
if any(s in content_accum for s in _tool_xml_signals):
|
|
_safety_tc = self._parse_tool_calls_from_text(
|
|
content_accum,
|
|
allow_incomplete = auto_heal_tool_calls,
|
|
)
|
|
if not _safety_tc:
|
|
# ── Re-prompt on plan-without-action ──
|
|
# If the model described its intent (forward-looking
|
|
# language) without calling a tool, nudge it to act.
|
|
# Fires at most once per request, only on short
|
|
# responses with intent signals -- "4" or "Hello!"
|
|
# won't trigger it. Use content if available, else
|
|
# fall back to reasoning text (reasoning-only stalls).
|
|
_stripped = content_accum.strip()
|
|
if not _stripped:
|
|
_stripped = reasoning_accum.strip()
|
|
_render_html_already_done_intent = _tool_succeeded(
|
|
"render_html"
|
|
) and re.search(
|
|
r"(?i)\brender[_\s-]?html\b",
|
|
_stripped,
|
|
)
|
|
if (
|
|
auto_heal_tool_calls
|
|
and active_tools
|
|
and not _render_html_already_done_intent
|
|
and _reprompt_count < _MAX_REPROMPTS
|
|
and _is_short_intent_without_action(_stripped)
|
|
):
|
|
_reprompt_count += 1
|
|
logger.info(
|
|
f"Re-prompt {_reprompt_count}/{_MAX_REPROMPTS}: "
|
|
f"model responded without calling tools "
|
|
f"({len(_stripped)} chars)"
|
|
)
|
|
conversation.append(
|
|
{
|
|
"role": "assistant",
|
|
"content": _stripped,
|
|
}
|
|
)
|
|
available_tool_names = [
|
|
(tool.get("function") or {}).get("name")
|
|
for tool in active_tools
|
|
if isinstance(tool, dict) and isinstance(tool.get("function"), dict)
|
|
]
|
|
available_tool_names = [name for name in available_tool_names if name]
|
|
tool_hint = " or ".join(available_tool_names) or "an available tool"
|
|
_forced_tool_call_pending = True
|
|
conversation.append(
|
|
{
|
|
"role": "user",
|
|
"content": (
|
|
"You have access to enabled tools. If a tool is needed to satisfy "
|
|
"the user's request or complete the action you described, call "
|
|
f"{tool_hint} now. If no tool is needed, provide the final answer "
|
|
"and follow the user's requested format."
|
|
),
|
|
}
|
|
)
|
|
# Accumulate tokens and timing from this iteration.
|
|
_fu_r = _backfill_usage_from_timings(_iter_usage, _iter_timings) or {}
|
|
_accumulated_completion_tokens += _fu_r.get("completion_tokens", 0)
|
|
_it_r = _iter_timings or {}
|
|
_accumulated_predicted_ms += _it_r.get("predicted_ms", 0)
|
|
_accumulated_predicted_n += _it_r.get("predicted_n", 0)
|
|
yield {"type": "status", "text": ""}
|
|
continue
|
|
|
|
if _forced_tool_call_pending:
|
|
_forced_tool_call_pending = False
|
|
if not _should_suppress_forced_no_tool_output(_stripped):
|
|
if cumulative_display:
|
|
forced_visible_text = _strip_tool_markup(
|
|
cumulative_display,
|
|
final = True,
|
|
)
|
|
elif content_accum:
|
|
forced_visible_text = _strip_tool_markup(
|
|
content_accum,
|
|
final = True,
|
|
)
|
|
else:
|
|
forced_visible_text = reasoning_accum
|
|
if forced_visible_text:
|
|
yield {
|
|
"type": "content",
|
|
"text": forced_visible_text,
|
|
}
|
|
|
|
# Content was already streamed. Yield metadata.
|
|
yield {"type": "status", "text": ""}
|
|
_fu = _backfill_usage_from_timings(_iter_usage, _iter_timings) or {}
|
|
_fc = _fu.get("completion_tokens", 0)
|
|
_fp = _fu.get("prompt_tokens", 0)
|
|
_tc = _fc + _accumulated_completion_tokens
|
|
if _iter_usage or _iter_timings or _accumulated_completion_tokens:
|
|
_mt = dict(_iter_timings) if _iter_timings else {}
|
|
if _accumulated_predicted_ms or _accumulated_predicted_n:
|
|
_mt["predicted_ms"] = (
|
|
_mt.get("predicted_ms", 0) + _accumulated_predicted_ms
|
|
)
|
|
_tn = _mt.get("predicted_n", 0) + _accumulated_predicted_n
|
|
_mt["predicted_n"] = _tn
|
|
_tms = _mt["predicted_ms"]
|
|
if _tms > 0:
|
|
_mt["predicted_per_second"] = _tn / (_tms / 1000.0)
|
|
yield {
|
|
"type": "metadata",
|
|
"usage": {
|
|
"prompt_tokens": _fp,
|
|
"completion_tokens": _tc,
|
|
"total_tokens": _fp + _tc,
|
|
},
|
|
"timings": _mt,
|
|
"finish_reason": _iter_finish_reason,
|
|
}
|
|
return
|
|
|
|
# Safety net caught tool XML -- treat as tool call.
|
|
tool_calls = _safety_tc
|
|
content_text = _strip_tool_markup(
|
|
content_accum,
|
|
final = True,
|
|
force = True,
|
|
)
|
|
logger.info(
|
|
f"Safety net: parsed {len(tool_calls)} tool call(s) from streamed content"
|
|
)
|
|
else:
|
|
# ── DRAINING path: assemble tool_calls ──
|
|
tool_calls = None
|
|
content_text = content_accum
|
|
if has_structured_tc:
|
|
# Drop incomplete fragments (e.g. from max_tokens
|
|
# truncation or disconnect).
|
|
tool_calls = [
|
|
tool_calls_acc[i]
|
|
for i in sorted(tool_calls_acc)
|
|
if (tool_calls_acc[i].get("function", {}).get("name", "").strip())
|
|
] or None
|
|
if not tool_calls and any(s in content_accum for s in _tool_xml_signals):
|
|
tool_calls = self._parse_tool_calls_from_text(
|
|
content_accum,
|
|
allow_incomplete = auto_heal_tool_calls,
|
|
)
|
|
if tool_calls and not has_structured_tc:
|
|
content_text = _strip_tool_markup(
|
|
content_text,
|
|
final = True,
|
|
force = True,
|
|
)
|
|
if tool_calls:
|
|
logger.info(
|
|
f"Parsed {len(tool_calls)} tool call(s) from "
|
|
f"{'structured delta' if has_structured_tc else 'content text'}"
|
|
)
|
|
if not tool_calls:
|
|
# DRAINING but no tool calls (false positive). Merge
|
|
# accumulated metrics from prior tool iterations so
|
|
# they aren't silently dropped.
|
|
yield {"type": "status", "text": ""}
|
|
if content_accum:
|
|
# Strip leaked tool-call XML before yielding.
|
|
content_accum = _strip_tool_markup(content_accum, final = True)
|
|
if content_accum:
|
|
yield {"type": "content", "text": content_accum}
|
|
_fu = _backfill_usage_from_timings(_iter_usage, _iter_timings) or {}
|
|
_fc = _fu.get("completion_tokens", 0)
|
|
_fp = _fu.get("prompt_tokens", 0)
|
|
_tc = _fc + _accumulated_completion_tokens
|
|
if _iter_usage or _iter_timings or _accumulated_completion_tokens:
|
|
_mt = dict(_iter_timings) if _iter_timings else {}
|
|
if _accumulated_predicted_ms or _accumulated_predicted_n:
|
|
_mt["predicted_ms"] = (
|
|
_mt.get("predicted_ms", 0) + _accumulated_predicted_ms
|
|
)
|
|
_tn = _mt.get("predicted_n", 0) + _accumulated_predicted_n
|
|
_mt["predicted_n"] = _tn
|
|
_tms = _mt["predicted_ms"]
|
|
if _tms > 0:
|
|
_mt["predicted_per_second"] = _tn / (_tms / 1000.0)
|
|
yield {
|
|
"type": "metadata",
|
|
"usage": {
|
|
"prompt_tokens": _fp,
|
|
"completion_tokens": _tc,
|
|
"total_tokens": _fp + _tc,
|
|
},
|
|
"timings": _mt,
|
|
"finish_reason": _iter_finish_reason,
|
|
}
|
|
return
|
|
|
|
# ── Execute tool calls ──
|
|
_accumulated_completion_tokens += (
|
|
_backfill_usage_from_timings(_iter_usage, _iter_timings) or {}
|
|
).get("completion_tokens", 0)
|
|
_it = _iter_timings or {}
|
|
_accumulated_predicted_ms += _it.get("predicted_ms", 0)
|
|
_accumulated_predicted_n += _it.get("predicted_n", 0)
|
|
|
|
# disable_parallel_tool_use: execute only the first tool call
|
|
# this turn. Truncate before building assistant_msg so the
|
|
# conversation stays consistent and extra calls are never executed.
|
|
if disable_parallel_tool_use and tool_calls and len(tool_calls) > 1:
|
|
tool_calls = tool_calls[:1]
|
|
|
|
assistant_msg: dict = {"role": "assistant", "content": content_text}
|
|
assistant_appended = False
|
|
|
|
for tc in tool_calls or []:
|
|
func = tc.get("function", {})
|
|
tool_name = func.get("name", "")
|
|
provisional_render_html_match = (
|
|
tool_name == "render_html"
|
|
and tc.get("id") in provisional_render_html_tool_call_ids
|
|
)
|
|
decision = tool_controller.prepare_call(
|
|
tc,
|
|
forced = _forced_tool_call_pending,
|
|
provisional = provisional_render_html_match,
|
|
)
|
|
|
|
if not decision.should_execute:
|
|
if content_text and not assistant_appended:
|
|
conversation.append(assistant_msg)
|
|
assistant_appended = True
|
|
completion = tool_controller.record_noop(decision)
|
|
conversation.append(completion.model_message())
|
|
if _forced_tool_call_pending:
|
|
_forced_tool_call_pending = False
|
|
logger.info(
|
|
"Suppressed local GGUF tool call as internal no-op: "
|
|
f"action={decision.action} tool={decision.tool_name}"
|
|
)
|
|
break
|
|
|
|
if not assistant_appended:
|
|
assistant_msg["tool_calls"] = [decision.as_assistant_tool_call()]
|
|
conversation.append(assistant_msg)
|
|
assistant_appended = True
|
|
else:
|
|
assistant_msg.setdefault("tool_calls", []).append(
|
|
decision.as_assistant_tool_call()
|
|
)
|
|
|
|
needs_confirm = bool(confirm_tool_calls)
|
|
approval_id = new_approval_id() if needs_confirm else ""
|
|
decision_slot = (
|
|
begin_tool_decision(session_id, approval_id) if needs_confirm else None
|
|
)
|
|
start_event = decision.tool_start_event()
|
|
start_event["approval_id"] = approval_id
|
|
start_event["awaiting_confirmation"] = needs_confirm
|
|
|
|
try:
|
|
yield {"type": "status", "text": decision.status_text}
|
|
yield start_event
|
|
|
|
if (
|
|
decision_slot is not None
|
|
and wait_tool_decision(
|
|
decision_slot,
|
|
approval_id,
|
|
cancel_event = cancel_event,
|
|
)
|
|
== "deny"
|
|
):
|
|
decision_slot = None
|
|
yield {
|
|
"type": "tool_end",
|
|
"tool_name": decision.tool_name,
|
|
"tool_call_id": decision.tool_call_id,
|
|
"result": TOOL_REJECTED_MESSAGE,
|
|
"provenance": decision.provenance,
|
|
}
|
|
denied_message = {
|
|
"role": "tool",
|
|
"name": decision.tool_name,
|
|
"content": TOOL_REJECTED_MESSAGE,
|
|
}
|
|
if decision.tool_call_id:
|
|
denied_message["tool_call_id"] = decision.tool_call_id
|
|
conversation.append(denied_message)
|
|
if _forced_tool_call_pending:
|
|
_forced_tool_call_pending = False
|
|
continue
|
|
decision_slot = None
|
|
finally:
|
|
if decision_slot is not None:
|
|
abort_tool_decision(decision_slot, approval_id)
|
|
|
|
_effective_timeout = None if tool_call_timeout >= 9999 else tool_call_timeout
|
|
# RAG: cap paraphrased KB re-searches that slip past the dup guard.
|
|
if (
|
|
decision.tool_name == "search_knowledge_base"
|
|
and _kb_search_count >= RAG_MAX_SEARCHES_PER_TURN
|
|
):
|
|
result = RAG_SEARCH_CAP_NUDGE
|
|
else:
|
|
result = execute_tool(
|
|
decision.tool_name,
|
|
decision.arguments,
|
|
cancel_event = cancel_event,
|
|
timeout = _effective_timeout,
|
|
session_id = session_id,
|
|
rag_scope = rag_scope,
|
|
)
|
|
if decision.tool_name == "search_knowledge_base":
|
|
_kb_search_count += 1
|
|
completion = tool_controller.record_result(decision, result)
|
|
yield completion.tool_end_event()
|
|
conversation.append(completion.tool_message())
|
|
|
|
if _forced_tool_call_pending:
|
|
_forced_tool_call_pending = False
|
|
|
|
# Clear tool status badge before next generation/final pass.
|
|
yield {"type": "status", "text": ""}
|
|
if tool_controller.force_final_answer or not tool_controller.active_tools():
|
|
_append_budget_exhausted_nudge = False
|
|
break
|
|
continue
|
|
|
|
except httpx.ConnectError:
|
|
raise RuntimeError("Lost connection to llama-server")
|
|
except Exception as e:
|
|
if cancel_event is not None and cancel_event.is_set():
|
|
return
|
|
raise
|
|
|
|
# ── Tool iteration cap reached -- synthesize final answer ──
|
|
# The model used all iterations without a final text response. Nudge
|
|
# the final streaming pass to produce a useful answer instead of
|
|
# continuing to request tools.
|
|
if max_tool_iterations > 0 and _append_budget_exhausted_nudge:
|
|
conversation.append(
|
|
{
|
|
"role": "user",
|
|
"content": (
|
|
"You have used all available tool calls. Based on "
|
|
"everything you have found so far, provide your final "
|
|
"answer now. Do not call any more tools."
|
|
),
|
|
}
|
|
)
|
|
|
|
# Clear status.
|
|
yield {"type": "status", "text": ""}
|
|
|
|
# Final streaming pass with the full conversation context.
|
|
stream_payload = {
|
|
"messages": conversation,
|
|
"stream": True,
|
|
"temperature": temperature,
|
|
"top_p": top_p,
|
|
"top_k": top_k if top_k >= 0 else 0,
|
|
"min_p": min_p,
|
|
"repeat_penalty": repetition_penalty,
|
|
"presence_penalty": presence_penalty,
|
|
}
|
|
_reasoning_kw = self._request_reasoning_kwargs(
|
|
enable_thinking, reasoning_effort, preserve_thinking
|
|
)
|
|
if _reasoning_kw is not None:
|
|
stream_payload["chat_template_kwargs"] = _reasoning_kw
|
|
stream_payload["max_tokens"] = (
|
|
max_tokens
|
|
if max_tokens is not None
|
|
else (self._effective_context_length or _DEFAULT_MAX_TOKENS_FLOOR)
|
|
)
|
|
stream_payload["t_max_predict_ms"] = _DEFAULT_T_MAX_PREDICT_MS
|
|
if stop:
|
|
stream_payload["stop"] = stop
|
|
if seed is not None:
|
|
stream_payload["seed"] = seed
|
|
stream_payload["stream_options"] = {"include_usage": True}
|
|
|
|
cumulative = ""
|
|
_last_emitted = ""
|
|
in_thinking = False
|
|
has_content_tokens = False
|
|
reasoning_text = ""
|
|
_metadata_usage = None
|
|
_metadata_timings = None
|
|
_metadata_finish_reason = None
|
|
_stream_done = False
|
|
|
|
try:
|
|
stream_timeout = httpx.Timeout(connect = 10, read = 0.5, write = 10, pool = 10)
|
|
_auth_headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else None
|
|
with httpx.Client(
|
|
timeout = stream_timeout, limits = httpx.Limits(max_keepalive_connections = 0)
|
|
) as client:
|
|
with self._stream_with_retry(
|
|
client,
|
|
url,
|
|
stream_payload,
|
|
cancel_event,
|
|
headers = _auth_headers,
|
|
) as response:
|
|
if response.status_code != 200:
|
|
error_body = response.read().decode()
|
|
raise RuntimeError(
|
|
f"llama-server returned {response.status_code}: {error_body}"
|
|
)
|
|
|
|
buffer = ""
|
|
for raw_chunk in self._iter_text_cancellable(response, cancel_event):
|
|
buffer += raw_chunk
|
|
while "\n" in buffer:
|
|
line, buffer = buffer.split("\n", 1)
|
|
line = line.strip()
|
|
|
|
if not line:
|
|
continue
|
|
if line == "data: [DONE]":
|
|
if in_thinking:
|
|
if has_content_tokens:
|
|
cumulative += "</think>"
|
|
yield {
|
|
"type": "content",
|
|
"text": _strip_tool_markup(cumulative, final = True),
|
|
}
|
|
else:
|
|
cumulative = reasoning_text
|
|
yield {"type": "content", "text": cumulative}
|
|
_stream_done = True
|
|
break # exit inner while
|
|
if not line.startswith("data: "):
|
|
continue
|
|
|
|
try:
|
|
chunk_data = json.loads(line[6:])
|
|
# Capture server timings/usage from final chunks.
|
|
_chunk_timings = chunk_data.get("timings")
|
|
if _chunk_timings:
|
|
_metadata_timings = _chunk_timings
|
|
_chunk_usage = chunk_data.get("usage")
|
|
if _chunk_usage:
|
|
_metadata_usage = _chunk_usage
|
|
choices = chunk_data.get("choices", [])
|
|
if choices:
|
|
delta = choices[0].get("delta", {})
|
|
_fr = choices[0].get("finish_reason")
|
|
if _fr:
|
|
_metadata_finish_reason = _fr
|
|
|
|
reasoning = delta.get("reasoning_content", "")
|
|
if reasoning:
|
|
reasoning_text += reasoning
|
|
if not in_thinking:
|
|
cumulative += "<think>"
|
|
in_thinking = True
|
|
cumulative += reasoning
|
|
yield {"type": "content", "text": cumulative}
|
|
|
|
token = delta.get("content", "")
|
|
if token:
|
|
has_content_tokens = True
|
|
if in_thinking:
|
|
cumulative += "</think>"
|
|
in_thinking = False
|
|
cumulative += token
|
|
cleaned = _strip_tool_markup(cumulative)
|
|
# Emit only when cleaned text grows (monotonic).
|
|
if len(cleaned) > len(_last_emitted):
|
|
_last_emitted = cleaned
|
|
yield {"type": "content", "text": cleaned}
|
|
except json.JSONDecodeError:
|
|
logger.debug(f"Skipping malformed SSE line: {line[:100]}")
|
|
if _stream_done:
|
|
break # exit outer for
|
|
_final_usage = _metadata_usage or {}
|
|
_final_completion = _final_usage.get("completion_tokens", 0)
|
|
_final_prompt = _final_usage.get("prompt_tokens", 0)
|
|
_total_completion = _final_completion + _accumulated_completion_tokens
|
|
if _metadata_usage or _metadata_timings or _metadata_finish_reason:
|
|
_merged_timings = dict(_metadata_timings) if _metadata_timings else {}
|
|
if _accumulated_predicted_ms or _accumulated_predicted_n:
|
|
_merged_timings["predicted_ms"] = (
|
|
_merged_timings.get("predicted_ms", 0) + _accumulated_predicted_ms
|
|
)
|
|
_total_predicted_n = (
|
|
_merged_timings.get("predicted_n", 0) + _accumulated_predicted_n
|
|
)
|
|
_merged_timings["predicted_n"] = _total_predicted_n
|
|
_total_predicted_ms = _merged_timings["predicted_ms"]
|
|
if _total_predicted_ms > 0:
|
|
_merged_timings["predicted_per_second"] = _total_predicted_n / (
|
|
_total_predicted_ms / 1000.0
|
|
)
|
|
yield {
|
|
"type": "metadata",
|
|
"usage": {
|
|
"prompt_tokens": _final_prompt,
|
|
"completion_tokens": _total_completion,
|
|
"total_tokens": _final_prompt + _total_completion,
|
|
},
|
|
"timings": _merged_timings,
|
|
"finish_reason": _metadata_finish_reason,
|
|
}
|
|
|
|
except httpx.ConnectError:
|
|
raise RuntimeError("Lost connection to llama-server")
|
|
except Exception as e:
|
|
if cancel_event is not None and cancel_event.is_set():
|
|
return
|
|
raise
|
|
|
|
# ── Prompt token counting ──────────────────────────────────
|
|
|
|
def count_chat_tokens(
|
|
self,
|
|
messages,
|
|
system = None,
|
|
tools = None,
|
|
strict: bool = False,
|
|
) -> int:
|
|
"""Count prompt tokens for a chat request via llama-server.
|
|
|
|
Non-strict callers keep the historical best-effort behavior and receive
|
|
0 when a count cannot be determined. Strict callers (public count_tokens
|
|
endpoints) get an exception instead of a successful-looking zero when
|
|
tokenizer/template calls fail or a multimodal prompt would fall back to a
|
|
text-only approximation.
|
|
"""
|
|
if not self.is_loaded:
|
|
if strict:
|
|
raise RuntimeError("llama-server is not loaded")
|
|
return 0
|
|
|
|
def _has_non_text_content(content) -> bool:
|
|
if isinstance(content, list):
|
|
for block in content:
|
|
if isinstance(block, str):
|
|
continue
|
|
if not isinstance(block, dict):
|
|
return True
|
|
if block.get("type") == "text" and isinstance(block.get("text"), str):
|
|
continue
|
|
if isinstance(block.get("text"), str):
|
|
continue
|
|
return True
|
|
return False
|
|
|
|
def _has_non_text_prompt_parts() -> bool:
|
|
if _has_non_text_content(system):
|
|
return True
|
|
for msg in messages or []:
|
|
if isinstance(msg, dict) and _has_non_text_content(msg.get("content", "")):
|
|
return True
|
|
return False
|
|
|
|
def _block_text(content) -> str:
|
|
if isinstance(content, str):
|
|
return content
|
|
if isinstance(content, list):
|
|
parts = []
|
|
for block in content:
|
|
if isinstance(block, dict):
|
|
if block.get("type") == "text" and isinstance(block.get("text"), str):
|
|
parts.append(block["text"])
|
|
elif isinstance(block.get("text"), str):
|
|
parts.append(block["text"])
|
|
elif isinstance(block, str):
|
|
parts.append(block)
|
|
return "".join(parts)
|
|
return ""
|
|
|
|
# Normalize system into a leading message / plain text.
|
|
system_text = ""
|
|
if isinstance(system, str):
|
|
system_text = system
|
|
elif isinstance(system, list):
|
|
system_text = _block_text(system)
|
|
|
|
try:
|
|
_auth_headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else None
|
|
with httpx.Client(timeout = 10, headers = _auth_headers) as client:
|
|
|
|
def _tokenize(text: str) -> int:
|
|
r = client.post(
|
|
f"{self.base_url}/tokenize",
|
|
json = {"content": text, "add_special": True},
|
|
)
|
|
if r.status_code != 200:
|
|
if strict:
|
|
raise RuntimeError("llama-server tokenizer failed")
|
|
return 0
|
|
tokens = r.json().get("tokens", [])
|
|
if not isinstance(tokens, list):
|
|
if strict:
|
|
raise RuntimeError("llama-server tokenizer returned invalid tokens")
|
|
return 0
|
|
return len(tokens)
|
|
|
|
# 1. Try /apply-template to render the real chat prompt.
|
|
template_messages = list(messages) if messages else []
|
|
if system_text:
|
|
template_messages = [
|
|
{"role": "system", "content": system_text}
|
|
] + template_messages
|
|
apply_template_failed = False
|
|
try:
|
|
# llama-server's /apply-template renders tool declarations
|
|
# into the prompt when ``tools`` is supplied, so pass them
|
|
# through — otherwise tool-schema tokens go uncounted.
|
|
template_body = {"messages": template_messages}
|
|
if tools:
|
|
template_body["tools"] = tools
|
|
resp = client.post(
|
|
f"{self.base_url}/apply-template",
|
|
json = template_body,
|
|
)
|
|
if resp.status_code == 200:
|
|
prompt = resp.json().get("prompt", "")
|
|
if isinstance(prompt, str):
|
|
return _tokenize(prompt)
|
|
apply_template_failed = True
|
|
except Exception:
|
|
apply_template_failed = True
|
|
|
|
if strict and apply_template_failed and _has_non_text_prompt_parts():
|
|
raise RuntimeError(
|
|
"cannot fall back to text-only token counting for multimodal messages"
|
|
)
|
|
|
|
# 2. Fallback: concatenate plain text and tokenize. Append a
|
|
# serialized form of the tools so they still contribute to the
|
|
# count when /apply-template is unavailable.
|
|
parts = []
|
|
if system_text:
|
|
parts.append(system_text)
|
|
for msg in messages or []:
|
|
if isinstance(msg, dict):
|
|
parts.append(_block_text(msg.get("content", "")))
|
|
if tools:
|
|
try:
|
|
parts.append(json.dumps(tools, ensure_ascii = False))
|
|
except Exception:
|
|
pass
|
|
return _tokenize("\n".join(p for p in parts if p))
|
|
except Exception:
|
|
if strict:
|
|
raise
|
|
return 0
|
|
|
|
# ── TTS support ────────────────────────────────────────────
|
|
|
|
def detect_audio_type(self) -> Optional[str]:
|
|
"""Detect audio/TTS codec; swallows errors (use _strict to distinguish)."""
|
|
try:
|
|
return self._detect_audio_type_strict()
|
|
except Exception as e:
|
|
logger.debug(f"Audio type detection failed: {e}")
|
|
return None
|
|
|
|
def _detect_audio_type_strict(self) -> Optional[str]:
|
|
"""Codec name on match, None on non-audio, raises on transport/JSON errors."""
|
|
if not self.is_loaded:
|
|
return None
|
|
_auth_headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else None
|
|
with httpx.Client(timeout = 10, headers = _auth_headers) as client:
|
|
|
|
def _detok(tid: int) -> str:
|
|
# Non-200 means "marker not in vocab" -- keep probing.
|
|
# Transport / JSON errors still raise.
|
|
r = client.post(f"{self.base_url}/detokenize", json = {"tokens": [tid]})
|
|
if r.status_code != 200:
|
|
return ""
|
|
return r.json().get("content", "")
|
|
|
|
def _tok(text: str) -> list[int]:
|
|
r = client.post(
|
|
f"{self.base_url}/tokenize",
|
|
json = {"content": text, "add_special": False},
|
|
)
|
|
if r.status_code != 200:
|
|
return []
|
|
return r.json().get("tokens", [])
|
|
|
|
# Codec-specific tokens (not generic ones that non-audio models may have)
|
|
if "<custom_token_" in _detok(128258) and "<custom_token_" in _detok(128259):
|
|
return "snac"
|
|
if len(_tok("<|AUDIO|>")) == 1 and len(_tok("<|audio_eos|>")) == 1:
|
|
return "csm"
|
|
if len(_tok("<|startoftranscript|>")) == 1:
|
|
return "whisper"
|
|
# Gemma 3n: <audio_soft_token>; Gemma 4: <|audio|> (not csm's <|AUDIO|>).
|
|
if len(_tok("<audio_soft_token>")) == 1 or len(_tok("<|audio|>")) == 1:
|
|
return "audio_vlm"
|
|
if len(_tok("<|bicodec_semantic_0|>")) == 1 and len(_tok("<|bicodec_global_0|>")) == 1:
|
|
return "bicodec"
|
|
if len(_tok("<|c1_0|>")) == 1 and len(_tok("<|c2_0|>")) == 1:
|
|
return "dac"
|
|
return None
|
|
|
|
# Prompt format per codec: (template, stop_tokens, needs_token_ids).
|
|
# Matches InferenceBackend._generate_snac/bicodec/dac.
|
|
_TTS_PROMPTS = {
|
|
"snac": (
|
|
"<custom_token_3>{text}<|eot_id|><custom_token_4>",
|
|
["<custom_token_2>"],
|
|
True,
|
|
),
|
|
"bicodec": (
|
|
"<|task_tts|><|start_content|>{text}<|end_content|><|start_global_token|>",
|
|
["<|im_end|>", "</s>"],
|
|
False,
|
|
),
|
|
"dac": (
|
|
"<|im_start|>\n<|text_start|>{text}<|text_end|>\n<|audio_start|><|global_features_start|>\n",
|
|
["<|im_end|>", "<|audio_end|>"],
|
|
False,
|
|
),
|
|
}
|
|
|
|
_codec_mgr = None # Shared AudioCodecManager instance
|
|
|
|
def init_audio_codec(self, audio_type: str) -> None:
|
|
"""Load the audio codec at model load time (mirrors the non-GGUF path)."""
|
|
import torch
|
|
from core.inference.audio_codecs import AudioCodecManager
|
|
|
|
if LlamaCppBackend._codec_mgr is None:
|
|
LlamaCppBackend._codec_mgr = AudioCodecManager()
|
|
|
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
model_repo_path = None
|
|
|
|
# BiCodec needs a repo with BiCodec/ weights -- download canonical SparkTTS
|
|
if audio_type == "bicodec":
|
|
from huggingface_hub import snapshot_download
|
|
import os
|
|
|
|
repo_path = snapshot_download("unsloth/Spark-TTS-0.5B", local_dir = "Spark-TTS-0.5B")
|
|
model_repo_path = os.path.abspath(repo_path)
|
|
|
|
LlamaCppBackend._codec_mgr.load_codec(audio_type, device, model_repo_path = model_repo_path)
|
|
logger.info(f"Loaded audio codec for GGUF TTS: {audio_type}")
|
|
|
|
def generate_audio_response(
|
|
self,
|
|
text: str,
|
|
audio_type: str,
|
|
temperature: float = 0.6,
|
|
top_p: float = 0.95,
|
|
top_k: int = 50,
|
|
min_p: float = 0.0,
|
|
max_new_tokens: int = 2048,
|
|
repetition_penalty: float = 1.1,
|
|
) -> tuple:
|
|
"""
|
|
Generate TTS audio via llama-server /completion + codec decode.
|
|
Returns (wav_bytes, sample_rate).
|
|
"""
|
|
if audio_type not in self._TTS_PROMPTS:
|
|
raise RuntimeError(f"GGUF TTS does not support '{audio_type}' codec.")
|
|
|
|
tpl, stop, need_ids = self._TTS_PROMPTS[audio_type]
|
|
|
|
payload: dict = {
|
|
"prompt": tpl.format(text = text),
|
|
"stream": False,
|
|
"n_predict": max_new_tokens,
|
|
"temperature": temperature,
|
|
"top_p": top_p,
|
|
"top_k": top_k if top_k >= 0 else 0,
|
|
"min_p": min_p,
|
|
"repeat_penalty": repetition_penalty,
|
|
}
|
|
if stop:
|
|
payload["stop"] = stop
|
|
if need_ids:
|
|
payload["n_probs"] = 1
|
|
|
|
_auth_headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else None
|
|
with httpx.Client(timeout = httpx.Timeout(300, connect = 10), headers = _auth_headers) as client:
|
|
resp = client.post(f"{self.base_url}/completion", json = payload)
|
|
if resp.status_code != 200:
|
|
raise RuntimeError(f"llama-server returned {resp.status_code}: {resp.text}")
|
|
|
|
data = resp.json()
|
|
token_ids = (
|
|
[p["id"] for p in data.get("completion_probabilities", []) if "id" in p]
|
|
if need_ids
|
|
else None
|
|
)
|
|
|
|
import torch
|
|
|
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
return LlamaCppBackend._codec_mgr.decode(
|
|
audio_type, device, token_ids = token_ids, text = data.get("content", "")
|
|
)
|