From 5a6e94d422f8fdcd3813790b7ad2c3f8e69d1b63 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Mon, 18 May 2026 00:15:42 -0700 Subject: [PATCH] Studio: auto-enable MTP speculative decoding for MTP GGUFs (#5527) * Studio: auto-enable MTP speculative decoding for MTP GGUFs Detect Unsloth's MTP (multi-token-prediction) GGUFs and auto-emit the right --spec-type draft-mtp flags for llama-server (llama.cpp PR #22673), so users get the speedup without configuration. Detection prefers the GGUF metadata field .nextn_predict_layers (verified on Qwen3.6-27B-MTP-GGUF / qwen35 and Qwen3.6-35B-A3B-MTP-GGUF / qwen35moe). Falls back to a -MTP marker in the identifier / filename so HF-mode loads can detect MTP from the repo name before the GGUF is downloaded. Flag presets follow the Unsloth MTP guide: GPU: --spec-type draft-mtp --spec-draft-n-max 6 CPU/Mac: --spec-type draft-mtp --spec-draft-n-max 3 \ --spec-type ngram-mod --spec-ngram-mod-n-match 24 \ --spec-ngram-mod-n-min 48 --spec-ngram-mod-n-max 6 User overrides win: if the caller passes --spec-type / --spec-default via unsloth run / unsloth studio run pass-through (or HTTP llama_extra_args), the auto-emit steps aside so llama-server only sees the user's flag. Scalar tuning knobs like --spec-draft-n-max compose with the auto preset via llama-server's last-wins parsing. _already_in_target_state mirrors the same promotion so a repeat /load with unchanged settings against an MTP backend running draft-mtp short-circuits cleanly instead of forcing a reload. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> --- studio/backend/core/inference/llama_cpp.py | 100 ++++- .../core/inference/llama_server_args.py | 6 + .../tests/test_llama_cpp_mtp_detection.py | 411 ++++++++++++++++++ .../backend/tests/test_llama_server_args.py | 46 ++ 4 files changed, 556 insertions(+), 7 deletions(-) create mode 100644 studio/backend/tests/test_llama_cpp_mtp_detection.py diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py index 0b3c9f958e..e41edbba35 100644 --- a/studio/backend/core/inference/llama_cpp.py +++ b/studio/backend/core/inference/llama_cpp.py @@ -23,7 +23,7 @@ import sys import threading import time from pathlib import Path -from typing import Generator, List, Optional +from typing import Generator, Iterable, List, Optional from urllib.parse import urlparse import httpx @@ -459,6 +459,32 @@ def detect_reasoning_flags( 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 _extra_args_set_spec_type(extra_args: Optional[Iterable[str]]) -> bool: + """User passed --spec-type / --spec-default? llama-server accumulates + repeated --spec-type, so we suppress auto-emit when this is true.""" + 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 + + class LlamaCppBackend: """ Manages a llama-server subprocess for GGUF model inference. @@ -514,6 +540,8 @@ class LlamaCppBackend: # Last N layers reuse KV from earlier layers and don't allocate # their own cache (Gemma 3n / Gemma 4: .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). @@ -1638,6 +1666,7 @@ class LlamaCppBackend: self._ssm_inner_size = None self._ssm_state_size = None self._shared_kv_layers = None + self._nextn_predict_layers = None try: WANTED = { @@ -1720,6 +1749,7 @@ class LlamaCppBackend: 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 @@ -2510,18 +2540,65 @@ class LlamaCppBackend: # ref: https://github.com/ggml-org/llama.cpp/blob/master/docs/speculative.md # ref: https://github.com/ggml-org/llama.cpp/pull/19164 # ref: https://github.com/ggml-org/llama.cpp/pull/18471 - # ``"default"`` -> let llama-server pick a sensible spec - # config via ``--spec-default``. Explicit type names are - # passed through with the manual draft tuning we've shipped - # historically so power users keep their overrides. - _valid_spec_types = {"ngram-simple", "ngram-mod"} + # draft-mtp: MTP heads on Unsloth's *-MTP GGUFs + # (llama.cpp #22673). Auto-enabled via nextn_predict_layers, + # fallback to -MTP in name. GPU: MTP-only. CPU/Mac: chain + # with ngram-mod. See unsloth.ai/docs/models/qwen3.6#mtp-guide. + _valid_spec_types = {"ngram-simple", "ngram-mod", "draft-mtp"} normalized_spec = ( speculative_type.lower().strip() if speculative_type else None ) + is_mtp_model = bool(self._nextn_predict_layers) or ( + _is_mtp_model_name(model_identifier, model_path) + ) + user_owns_spec_type = _extra_args_set_spec_type(extra_args) + # Auto-promote unset/"default" to draft-mtp on MTP GGUFs. + if ( + is_mtp_model + and not is_vision + and not user_owns_spec_type + and normalized_spec in (None, "", "default") + ): + normalized_spec = "draft-mtp" + if user_owns_spec_type: + # User --spec-type wins (it accumulates if repeated). + normalized_spec = None + self._speculative_type = None if normalized_spec and normalized_spec != "off" and not is_vision: if normalized_spec == "default": cmd.append("--spec-default") self._speculative_type = "default" + elif normalized_spec == "draft-mtp": + if gpus: + cmd.extend( + [ + "--spec-type", + "draft-mtp", + "--spec-draft-n-max", + "6", + ] + ) + else: + cmd.extend( + [ + "--spec-type", + "draft-mtp", + "--spec-draft-n-max", + "3", + "--spec-type", + "ngram-mod", + "--spec-ngram-mod-n-match", + "24", + "--spec-ngram-mod-n-min", + "48", + "--spec-ngram-mod-n-max", + "6", + ] + ) + self._speculative_type = "draft-mtp" + logger.info( + f"Spec decoding: draft-mtp ({'GPU' if gpus else 'CPU/Mac'})" + ) elif normalized_spec in _valid_spec_types: cmd.extend(["--spec-type", normalized_spec]) if normalized_spec == "ngram-mod": @@ -2941,7 +3018,15 @@ class LlamaCppBackend: if self._is_vision or is_vision: req_spec = "off" else: - req_spec = _norm(speculative_type) or "off" + raw_spec = _norm(speculative_type) + req_spec = raw_spec or "off" + # Mirror load_model's auto-promotion so repeat /load matches. + if ( + raw_spec in (None, "default") + and _is_mtp_model_name(model_identifier, gguf_path) + and not _extra_args_set_spec_type(extra_args) + ): + req_spec = "draft-mtp" backend_spec = _norm(self._speculative_type) or "off" if req_spec != backend_spec: return False @@ -3029,6 +3114,7 @@ class LlamaCppBackend: 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: diff --git a/studio/backend/core/inference/llama_server_args.py b/studio/backend/core/inference/llama_server_args.py index 0f6927fc5a..572ac2ceda 100644 --- a/studio/backend/core/inference/llama_server_args.py +++ b/studio/backend/core/inference/llama_server_args.py @@ -145,6 +145,12 @@ _SPEC_FLAGS: frozenset[str] = frozenset( "--spec-ngram-size", "--draft-min", "--draft-max", + # MTP path (llama.cpp #22673). + "--spec-draft-n-max", + "--spec-draft-n-min", + "--spec-ngram-mod-n-match", + "--spec-ngram-mod-n-min", + "--spec-ngram-mod-n-max", } ) _TEMPLATE_FLAGS: frozenset[str] = frozenset( diff --git a/studio/backend/tests/test_llama_cpp_mtp_detection.py b/studio/backend/tests/test_llama_cpp_mtp_detection.py new file mode 100644 index 0000000000..7ae245a1da --- /dev/null +++ b/studio/backend/tests/test_llama_cpp_mtp_detection.py @@ -0,0 +1,411 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Tests for the MTP auto-detection path (llama.cpp #22673). + +Pins three contracts: name-based detector, user-override detector, and +the _already_in_target_state mirror that prevents needless reloads. +""" + +from __future__ import annotations + +import struct +import sys +import types as _types +from pathlib import Path + +_BACKEND_DIR = str(Path(__file__).resolve().parent.parent) +if _BACKEND_DIR not in sys.path: + sys.path.insert(0, _BACKEND_DIR) + +_loggers_stub = _types.ModuleType("loggers") +_loggers_stub.get_logger = lambda name: __import__("logging").getLogger(name) +sys.modules.setdefault("loggers", _loggers_stub) + +_structlog_stub = _types.ModuleType("structlog") +_structlog_stub.get_logger = lambda *a, **k: __import__("logging").getLogger("stub") +sys.modules.setdefault("structlog", _structlog_stub) + +_httpx_stub = _types.ModuleType("httpx") +for _exc in ( + "ConnectError", + "TimeoutException", + "ReadTimeout", + "ReadError", + "RemoteProtocolError", + "CloseError", +): + setattr(_httpx_stub, _exc, type(_exc, (Exception,), {})) +_httpx_stub.Timeout = type("T", (), {"__init__": lambda s, *a, **k: None}) +_httpx_stub.Client = type( + "C", + (), + { + "__init__": lambda s, **kw: None, + "__enter__": lambda s: s, + "__exit__": lambda s, *a: None, + }, +) +sys.modules.setdefault("httpx", _httpx_stub) + +import pytest + +from core.inference.llama_cpp import ( + LlamaCppBackend, + _extra_args_set_spec_type, + _is_mtp_model_name, +) + + +# Synthetic GGUF helper (mirrors test_gguf_metadata.py). + +_GGUF_MAGIC = 0x46554747 +_VTYPE_STRING = 8 +_VTYPE_UINT32 = 4 + + +def _enc_string(s: str) -> bytes: + b = s.encode("utf-8") + return struct.pack(" bytes: + return _enc_string(key) + struct.pack(" bytes: + return ( + _enc_string(key) + struct.pack(" Path: + """Header-only GGUF with arch + optional nextn_predict_layers.""" + extra_uint32 = dict(extra_uint32 or {}) + body = _enc_kv_string("general.architecture", arch) + kv_count = 1 + if nextn is not None: + body += _enc_kv_uint32(f"{arch}.nextn_predict_layers", nextn) + kv_count += 1 + for k, v in extra_uint32.items(): + body += _enc_kv_uint32(k, v) + kv_count += 1 + header = struct.pack("0 should match. + ("qwen3moe", 2), + ("hypothetical_future_arch", 4), + ], +) +def test_read_gguf_metadata_captures_nextn_predict_layers(tmp_path, arch, nextn): + gguf = _write_minimal_gguf( + tmp_path / "model.gguf", + arch = arch, + nextn = nextn, + extra_uint32 = {f"{arch}.block_count": 4}, + ) + backend = LlamaCppBackend() + backend._read_gguf_metadata(str(gguf)) + assert backend._nextn_predict_layers == nextn + + +def test_read_gguf_metadata_leaves_nextn_unset_for_non_mtp_arch(tmp_path): + gguf = _write_minimal_gguf( + tmp_path / "model.gguf", + arch = "qwen3", + nextn = None, + extra_uint32 = {"qwen3.block_count": 4}, + ) + backend = LlamaCppBackend() + backend._read_gguf_metadata(str(gguf)) + assert backend._nextn_predict_layers is None + + +def test_read_gguf_metadata_zero_nextn_is_falsy(tmp_path): + # bool(0) is False, so the spec block short-circuits. + gguf = _write_minimal_gguf( + tmp_path / "model.gguf", + arch = "qwen35", + nextn = 0, + extra_uint32 = {"qwen35.block_count": 4}, + ) + backend = LlamaCppBackend() + backend._read_gguf_metadata(str(gguf)) + assert backend._nextn_predict_layers == 0 + assert bool(backend._nextn_predict_layers) is False + + +def test_unload_resets_nextn_predict_layers(): + # MTP state from a previous load must not bleed into the next load. + backend = LlamaCppBackend() + backend._nextn_predict_layers = 1 + backend.unload_model() + assert backend._nextn_predict_layers is None diff --git a/studio/backend/tests/test_llama_server_args.py b/studio/backend/tests/test_llama_server_args.py index 3013acfdb8..f4dabfcf08 100644 --- a/studio/backend/tests/test_llama_server_args.py +++ b/studio/backend/tests/test_llama_server_args.py @@ -42,6 +42,23 @@ from core.inference.llama_server_args import ( ["--chat-template-kwargs", '{"reasoning_effort":"high"}'], ["--spec-type", "ngram-mod"], ["--spec-default"], + # MTP path (llama.cpp #22673). + ["--spec-type", "draft-mtp"], + ["--spec-type", "draft-mtp", "--spec-draft-n-max", "6"], + [ + "--spec-type", + "draft-mtp", + "--spec-draft-n-max", + "3", + "--spec-type", + "ngram-mod", + "--spec-ngram-mod-n-match", + "24", + "--spec-ngram-mod-n-min", + "48", + "--spec-ngram-mod-n-max", + "6", + ], # Reasoning controls ["--reasoning-format", "deepseek"], ["-rea", "auto"], @@ -266,6 +283,35 @@ def test_strip_shadowing_flags_keeps_spec_when_spec_disabled(): ] +def test_strip_shadowing_flags_drops_mtp_flags_when_requested(): + # MTP / draft-mtp flags must be stripped when speculative_type is re-applied. + out = strip_shadowing_flags( + [ + "--spec-type", + "draft-mtp", + "--spec-draft-n-max", + "6", + "--spec-ngram-mod-n-match", + "24", + "--spec-ngram-mod-n-min", + "48", + "--spec-ngram-mod-n-max", + "6", + "--top-k", + "20", + ], + strip_spec = True, + ) + assert out == ["--top-k", "20"] + + +def test_is_managed_flag_false_for_mtp_pass_through(): + assert is_managed_flag("--spec-draft-n-max") is False + assert is_managed_flag("--spec-ngram-mod-n-match") is False + assert is_managed_flag("--spec-ngram-mod-n-min") is False + assert is_managed_flag("--spec-ngram-mod-n-max") is False + + def test_strip_shadowing_flags_boolean_does_not_consume_next_token(): # --spec-default is a boolean shadowing flag; the value-skipping # heuristic must skip just the flag, not the following positional.