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 <arch>.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>
This commit is contained in:
Daniel Han 2026-05-18 00:15:42 -07:00 committed by GitHub
commit 5a6e94d422
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 556 additions and 7 deletions

View file

@ -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: <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).
@ -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:

View file

@ -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(

View file

@ -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("<Q", len(b)) + b
def _enc_kv_string(key: str, value: str) -> bytes:
return _enc_string(key) + struct.pack("<I", _VTYPE_STRING) + _enc_string(value)
def _enc_kv_uint32(key: str, value: int) -> bytes:
return (
_enc_string(key) + struct.pack("<I", _VTYPE_UINT32) + struct.pack("<I", value)
)
def _write_minimal_gguf(
path: Path,
*,
arch: str,
nextn: int | None,
extra_uint32: dict[str, int] | None = None,
) -> 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("<IIQQ", _GGUF_MAGIC, 3, 0, kv_count)
path.write_bytes(header + body)
return path
# _is_mtp_model_name helper.
@pytest.mark.parametrize(
"identifier",
[
"unsloth/Qwen3.6-27B-MTP-GGUF",
"unsloth/Qwen3.6-35B-A3B-MTP-GGUF",
"unsloth/qwen3.6-27b-mtp-gguf",
"unsloth/Qwen3.6-27B-Mtp-GGUF",
"unsloth/Qwen3.6-27B-MTP-GGUF:UD-Q4_K_XL",
],
)
def test_is_mtp_model_name_detects_marker_in_identifier(identifier):
assert _is_mtp_model_name(identifier) is True
@pytest.mark.parametrize(
"identifier",
[
"unsloth/Qwen3-27B-GGUF",
"unsloth/Llama-3.1-8B-Instruct-GGUF",
"google/gemma-3-4b-it",
# mtp inside an org name should not match.
"mtp-research/foo",
"MTPower/bar",
],
)
def test_is_mtp_model_name_does_not_overmatch(identifier):
assert _is_mtp_model_name(identifier) is False
def test_is_mtp_model_name_handles_none():
assert _is_mtp_model_name(None) is False
assert _is_mtp_model_name(None, None) is False
assert _is_mtp_model_name("", "") is False
def test_is_mtp_model_name_detects_marker_in_filename(tmp_path):
gguf = tmp_path / "Qwen3.6-27B-MTP-Q4_K_M.gguf"
gguf.write_bytes(b"")
assert _is_mtp_model_name("local-model", str(gguf)) is True
def test_is_mtp_model_name_filename_case_insensitive(tmp_path):
gguf = tmp_path / "qwen3.6-35b-a3b-mtp-q4_k_m.gguf"
gguf.write_bytes(b"")
assert _is_mtp_model_name(None, str(gguf)) is True
def test_is_mtp_model_name_ignores_non_mtp_filename(tmp_path):
gguf = tmp_path / "Qwen3.6-27B-Q4_K_M.gguf"
gguf.write_bytes(b"")
assert _is_mtp_model_name("local-model", str(gguf)) is False
# _already_in_target_state MTP promotion.
class _FakeProcess:
"""Minimal stand-in so is_loaded returns True."""
def terminate(self):
pass
def wait(self, timeout = None):
return 0
def kill(self):
pass
def poll(self):
return 0
def _mtp_backend(**overrides):
"""MTP-named GGUF backend that's already running with draft-mtp."""
backend = LlamaCppBackend()
backend._process = _FakeProcess()
backend._healthy = True
backend._model_identifier = "unsloth/Qwen3.6-27B-MTP-GGUF"
backend._hf_variant = "Q4_K_M"
backend._requested_n_ctx = 8192
backend._cache_type_kv = None
backend._speculative_type = "draft-mtp"
backend._chat_template_override = None
backend._is_vision = False
backend._extra_args = None
backend._extra_args_source = None
backend._gguf_path = None
for key, value in overrides.items():
setattr(backend, key, value)
return backend
def test_already_in_target_state_matches_when_request_omits_spec_for_mtp_model():
# Duplicate /load with no spec must match a running draft-mtp backend.
backend = _mtp_backend()
assert (
backend._already_in_target_state(
gguf_path = None,
model_identifier = "unsloth/Qwen3.6-27B-MTP-GGUF",
hf_variant = "Q4_K_M",
n_ctx = 8192,
cache_type_kv = None,
speculative_type = None,
chat_template_override = None,
extra_args = None,
is_vision = False,
)
is True
)
def test_already_in_target_state_matches_when_request_uses_default_for_mtp_model():
backend = _mtp_backend()
assert (
backend._already_in_target_state(
gguf_path = None,
model_identifier = "unsloth/Qwen3.6-27B-MTP-GGUF",
hf_variant = "Q4_K_M",
n_ctx = 8192,
cache_type_kv = None,
speculative_type = "default",
chat_template_override = None,
extra_args = None,
is_vision = False,
)
is True
)
def test_already_in_target_state_non_mtp_model_unaffected():
# Promotion is gated on the name; non-MTP must still mismatch req=None.
backend = _mtp_backend(_model_identifier = "unsloth/Qwen3.6-27B-GGUF")
assert (
backend._already_in_target_state(
gguf_path = None,
model_identifier = "unsloth/Qwen3.6-27B-GGUF",
hf_variant = "Q4_K_M",
n_ctx = 8192,
cache_type_kv = None,
speculative_type = None,
chat_template_override = None,
extra_args = None,
is_vision = False,
)
is False
)
def test_already_in_target_state_explicit_off_still_mismatches_mtp_backend():
backend = _mtp_backend()
assert (
backend._already_in_target_state(
gguf_path = None,
model_identifier = "unsloth/Qwen3.6-27B-MTP-GGUF",
hf_variant = "Q4_K_M",
n_ctx = 8192,
cache_type_kv = None,
speculative_type = "off",
chat_template_override = None,
extra_args = None,
is_vision = False,
)
is False
)
# User override via extra_args (unsloth run / unsloth studio run).
@pytest.mark.parametrize(
"extra_args",
[
["--spec-type", "none"],
["--spec-type", "ngram-mod"],
["--spec-type", "draft-mtp"],
["--spec-type=none"],
["--top-k", "20", "--spec-type", "ngram-simple", "--seed", "42"],
["--spec-default"],
],
)
def test_extra_args_set_spec_type_detects_user_override(extra_args):
assert _extra_args_set_spec_type(extra_args) is True
@pytest.mark.parametrize(
"extra_args",
[
None,
[],
# Scalar tuning knobs compose safely with auto-emitted --spec-type.
["--spec-draft-n-max", "4"],
["--spec-ngram-mod-n-match", "32"],
["--draft-max", "32"],
["--top-k", "20", "--seed", "42"],
],
)
def test_extra_args_set_spec_type_passes_on_non_spec_type_args(extra_args):
assert _extra_args_set_spec_type(extra_args) is False
def test_already_in_target_state_user_spec_type_override_matches_clean_backend():
# User --spec-type none suppressed auto-MTP; repeat /load must not re-promote.
backend = _mtp_backend(
_speculative_type = None,
_extra_args = ["--spec-type", "none"],
)
assert (
backend._already_in_target_state(
gguf_path = None,
model_identifier = "unsloth/Qwen3.6-27B-MTP-GGUF",
hf_variant = "Q4_K_M",
n_ctx = 8192,
cache_type_kv = None,
speculative_type = None,
chat_template_override = None,
extra_args = ["--spec-type", "none"],
is_vision = False,
)
is True
)
def test_already_in_target_state_local_file_mtp_match(tmp_path):
# Local-file load: -MTP marker comes from the filename.
gguf = tmp_path / "Qwen3.6-35B-A3B-MTP-Q4_K_M.gguf"
gguf.write_bytes(b"")
backend = _mtp_backend(
_model_identifier = "local-qwen-mtp",
_gguf_path = str(gguf),
_hf_variant = None,
)
assert (
backend._already_in_target_state(
gguf_path = str(gguf),
model_identifier = "local-qwen-mtp",
hf_variant = None,
n_ctx = 8192,
cache_type_kv = None,
speculative_type = None,
chat_template_override = None,
extra_args = None,
is_vision = False,
)
is True
)
# GGUF-metadata-based detection (nextn_predict_layers).
@pytest.mark.parametrize(
"arch, nextn",
[
# Verified against real Unsloth MTP GGUFs (qwen35 / qwen35moe).
("qwen35", 1),
("qwen35moe", 1),
# Future-proofing: any arch + n>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

View file

@ -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.