unsloth/studio/backend/tests/test_mlx_inference_backend.py
Daniel Han bb4eb88fdc
Studio: tools, thinking blocks, code execution and web search for safetensors (#5520)
Adds tools, thinking blocks, code execution, and web search support to the safetensors / transformers and MLX inference backends in Studio, bringing them to parity with the GGUF path.

What ships
- safetensors / transformers agentic tool loop with cumulative-text state machine, tool-call XML parser, and template kwarg forwarding (tools / enable_thinking / reasoning_effort / preserve_thinking).
- MLX backend: same kwargs accepted on Apple Silicon; chat_template_info shipped through worker IPC; pills enable for Qwen / Qwen3 / Qwen3.5 / Gemma reasoning.
- Capability classifier (_detect_safetensors_features) gates supports_tools on actual parser-compatible emission markers (<tool_call> / <function=) so Llama-3 / Mistral / Gemma 4 do not advertise toggles the parser cannot honour.
- gpt-oss override stays: reasoning on, tools off (Harmony channel, not <tool_call> XML).
- CWE-209 hygiene: safetensors SSE error path emits a constant message and logs the trace server-side.

Validation
- 256 unit tests green (43 tool-loop, 11 capability advertise, 7 MLX backend, 5 main-added, 190 adjacent inference / anthropic / openai regression).
- Cross-OS staging CI green on ubuntu-latest / macos-14 / windows-latest plus a dedicated MLX cartesian probe against real unsloth/Qwen3.5-0.8B on macos-14 (CI 26098107440).
- Capability parity verified across Qwen3 / Qwen3.5 / Llama-3 / Mistral / Gemma / DeepSeek-R1 / gpt-oss (incl. BF16).
- Manual confirmation from Imagineer99 on Qwen3.5-2B: think + search + code exec working.

Closes the safetensors / MLX gap with the GGUF backend.
2026-05-19 06:30:17 -07:00

254 lines
7.7 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
import sys
import types
from types import SimpleNamespace
class _DummyMetal:
@staticmethod
def is_available():
return False
class _DummyMX:
metal = _DummyMetal()
@staticmethod
def set_wired_limit(_limit):
return None
@staticmethod
def device_info():
return {"max_recommended_working_set_size": 1024}
class _DummyTokenizer:
pass
class _DummyProcessor:
tokenizer = _DummyTokenizer()
class _DummyModel:
pass
def _install_fake_mlx(monkeypatch):
mlx_pkg = types.ModuleType("mlx")
mlx_core = types.ModuleType("mlx.core")
mlx_core.metal = _DummyMetal()
mlx_core.set_wired_limit = _DummyMX.set_wired_limit
mlx_core.device_info = _DummyMX.device_info
mlx_pkg.core = mlx_core
monkeypatch.setitem(sys.modules, "mlx", mlx_pkg)
monkeypatch.setitem(sys.modules, "mlx.core", mlx_core)
def _install_fake_fast_mlx(monkeypatch, calls):
class _FastMLXModel:
@staticmethod
def from_pretrained(*args, **kwargs):
calls.append((args, kwargs))
if kwargs["text_only"] is False:
return _DummyModel(), _DummyProcessor()
return _DummyModel(), _DummyTokenizer()
unsloth_zoo_pkg = types.ModuleType("unsloth_zoo")
mlx_pkg = types.ModuleType("unsloth_zoo.mlx")
mlx_loader = types.ModuleType("unsloth_zoo.mlx.loader")
mlx_loader.FastMLXModel = _FastMLXModel
unsloth_zoo_pkg.mlx = mlx_pkg
mlx_pkg.loader = mlx_loader
monkeypatch.setitem(sys.modules, "unsloth_zoo", unsloth_zoo_pkg)
monkeypatch.setitem(sys.modules, "unsloth_zoo.mlx", mlx_pkg)
monkeypatch.setitem(sys.modules, "unsloth_zoo.mlx.loader", mlx_loader)
def test_mlx_inference_text_load_forwards_studio_settings(monkeypatch):
_install_fake_mlx(monkeypatch)
calls = []
_install_fake_fast_mlx(monkeypatch, calls)
from core.inference.mlx_inference import MLXInferenceBackend
backend = MLXInferenceBackend()
config = SimpleNamespace(identifier = "fake/text", is_vision = False, is_lora = False)
assert backend.load_model(
config,
max_seq_length = 4096,
load_in_4bit = False,
hf_token = "hf-token",
trust_remote_code = True,
dtype = "float16",
)
assert calls == [
(
("fake/text",),
{
"max_seq_length": 4096,
"dtype": "float16",
"load_in_4bit": False,
"token": "hf-token",
"trust_remote_code": True,
"text_only": True,
},
)
]
assert backend._is_vlm is False
assert isinstance(backend._tokenizer, _DummyTokenizer)
def test_mlx_inference_vlm_lora_uses_unsloth_loader_without_native_adapter_rewrite(
monkeypatch,
tmp_path,
):
_install_fake_mlx(monkeypatch)
calls = []
_install_fake_fast_mlx(monkeypatch, calls)
def _native_vlm_load(*_args, **_kwargs):
raise AssertionError("Studio MLX VLM inference must use FastMLXModel")
mlx_vlm = types.ModuleType("mlx_vlm")
mlx_vlm.load = _native_vlm_load
monkeypatch.setitem(sys.modules, "mlx_vlm", mlx_vlm)
adapter_dir = tmp_path / "adapter"
adapter_dir.mkdir()
cfg_path = adapter_dir / "adapter_config.json"
original_cfg = '{"base_model_name_or_path": "fake/base", "rank": 8}\n'
cfg_path.write_text(original_cfg)
from core.inference.mlx_inference import MLXInferenceBackend
backend = MLXInferenceBackend()
config = SimpleNamespace(
identifier = str(adapter_dir),
is_vision = True,
is_lora = True,
base_model = "fake/base",
)
assert backend.load_model(
config,
max_seq_length = 8192,
load_in_4bit = True,
hf_token = "hf-token",
trust_remote_code = True,
)
assert calls == [
(
(str(adapter_dir),),
{
"max_seq_length": 8192,
"dtype": None,
"load_in_4bit": True,
"token": "hf-token",
"trust_remote_code": True,
"text_only": False,
},
)
]
assert cfg_path.read_text() == original_cfg
assert backend._is_vlm is True
assert isinstance(backend._processor, _DummyProcessor)
assert isinstance(backend._tokenizer, _DummyTokenizer)
# Regression: MLXInferenceBackend.generate_chat_response must accept the
# four template kwargs (tools / enable_thinking / reasoning_effort /
# preserve_thinking) so the route layer can forward what the user
# toggled in the UI. The previous signature raised
# "got an unexpected keyword argument 'tools'" on Mac.
def test_mlx_generate_chat_response_accepts_template_kwargs():
import inspect
from core.inference.mlx_inference import MLXInferenceBackend
sig = inspect.signature(MLXInferenceBackend.generate_chat_response)
params = sig.parameters
for name in ("tools", "enable_thinking", "reasoning_effort", "preserve_thinking"):
assert name in params, (
f"MLX.generate_chat_response is missing the {name!r} kwarg; "
"the route layer forwards this and a missing kwarg raises "
"TypeError on Mac"
)
assert (
params[name].default is None
), f"{name!r} must default to None so existing callers stay valid"
def test_mlx_generate_text_forwards_kwargs_into_template_helper(monkeypatch):
"""The Mac text path must route through apply_chat_template_for_
generation so reasoning / tool kwargs reach the tokenizer."""
_install_fake_mlx(monkeypatch)
from core.inference.mlx_inference import MLXInferenceBackend
captured = {}
def _fake_apply(tokenizer, messages, **kwargs):
captured["tokenizer"] = tokenizer
captured["messages"] = messages
captured["kwargs"] = kwargs
return "<rendered prompt>"
monkeypatch.setattr(
"core.inference.chat_template_helpers." "apply_chat_template_for_generation",
_fake_apply,
raising = True,
)
# mlx_lm.stream_generate yields response objects with .token; make a
# one-token generator so _generate_text returns without touching the
# real stack.
import types as _types
mlx_lm_pkg = _types.ModuleType("mlx_lm")
mlx_lm_sample = _types.ModuleType("mlx_lm.sample_utils")
mlx_lm_sample.make_sampler = lambda **_kw: object()
mlx_lm_sample.make_logits_processors = lambda **_kw: None
class _Resp:
def __init__(self, tok):
self.token = tok
def _stream_generate(_model, _tokenizer, **_kw):
yield _Resp(1)
mlx_lm_pkg.stream_generate = _stream_generate
monkeypatch.setitem(sys.modules, "mlx_lm", mlx_lm_pkg)
monkeypatch.setitem(sys.modules, "mlx_lm.sample_utils", mlx_lm_sample)
class _Tok:
chat_template = "x"
def decode(self, ids, skip_special_tokens = False):
return "hi"
backend = MLXInferenceBackend()
backend._model = object()
backend._tokenizer = _Tok()
backend._is_vlm = False
out = list(
backend.generate_chat_response(
messages = [{"role": "user", "content": "ping"}],
tools = [{"function": {"name": "web_search"}}],
enable_thinking = True,
reasoning_effort = "medium",
preserve_thinking = True,
max_new_tokens = 1,
)
)
assert out == ["hi"]
# The kwargs the user toggled must reach the chat-template helper.
assert captured["kwargs"]["tools"] == [{"function": {"name": "web_search"}}]
assert captured["kwargs"]["enable_thinking"] is True
assert captured["kwargs"]["reasoning_effort"] == "medium"
assert captured["kwargs"]["preserve_thinking"] is True