unsloth/studio/backend/core/inference/mlx_inference.py
DoubleMathew a932294627
MLX training support for Studio on Apple Silicon (#5340)
* mlx fixes

* Fix studio integration, local dataset files, chat templates without the torch gpu imports

* pass grad norm in mlx worker

* fix(studio): pass MLX grad clipping settings

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

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

* mlx: update grad value

* fix(mlx): address ci and clipping review

* fix backward compatibility and CI tests

* unsloth local is mlx function

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

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

* dont reference runtime

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

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

* studio mlx: hardcode value clipping, drop max_grad_value from frontend

Simplifies the MLX grad-clipping plumbing now that we are standardising on
elementwise value clipping at [-5, 5] for the compiled MLX path and norm
clipping disabled. The MLX worker no longer reads max_grad_norm /
max_grad_value from the request; both are pinned in one place. Frontend
stops sending the field at all, and the TypeScript request type drops it
to match. Non-MLX (CUDA/AMD/Intel) is untouched and continues to pick up
HF TrainingArguments' default max_grad_norm = 1.0.

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Daniel Han <danielhanchen@gmail.com>
2026-05-14 05:24:20 -07:00

417 lines
14 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
"""MLX inference backend for Apple Silicon.
Drop-in replacement for InferenceBackend — same interface, uses mlx-lm/mlx-vlm
instead of torch/transformers for model loading and generation.
"""
import threading
from typing import Optional, Generator
from loggers import get_logger
logger = get_logger(__name__)
class MLXInferenceBackend:
def __init__(self):
self.models = {}
self.active_model_name = None
self.loading_models = set()
self.loaded_local_models = []
self.device = "mlx"
self._generation_lock = threading.Lock()
# MLX state
self._model = None
self._tokenizer = None
self._processor = None
self._is_vlm = False
self._config = {}
# Recorded for unload to release pinned memory back to the OS.
self._memory_limits_applied = {}
def _configure_memory_limits(self):
"""Apply Metal memory caps before loading a model.
Mirrors MLXTrainer._configure_memory_limits's defaults:
memory_limit = 85% of recommended working-set,
wired_limit = min(recommended, memory_limit). Recorded so unload
can lower wired_limit back to release pinned RAM.
"""
import mlx.core as mx
if not mx.metal.is_available():
return
info = mx.device_info()
rec_bytes = info.get("max_recommended_working_set_size")
if not rec_bytes or rec_bytes <= 0:
return
rec_gb = rec_bytes / 1e9
memory_limit_gb = rec_gb * 0.85
wired_limit_gb = min(rec_gb, memory_limit_gb)
mx.set_memory_limit(int(memory_limit_gb * 1e9))
mx.set_wired_limit(int(wired_limit_gb * 1e9))
self._memory_limits_applied = {
"memory_limit_gb": memory_limit_gb,
"wired_limit_gb": wired_limit_gb,
"recommended_gb": rec_gb,
}
logger.info(
"MLX memory caps: memory_limit=%.2f GB, wired_limit=%.2f GB",
memory_limit_gb,
wired_limit_gb,
)
def load_model(
self,
config,
max_seq_length = 2048,
load_in_4bit = True,
hf_token = None,
trust_remote_code = False,
gpu_ids = None,
dtype = None,
) -> bool:
import mlx.core as mx
model_name = config.identifier if hasattr(config, "identifier") else str(config)
is_vision = getattr(config, "is_vision", False)
# GGUF guard. GGUF models are served via llama-server in the
# parent process, NOT via mlx-lm in this MLX subprocess. The
# route at studio/backend/routes/inference.py:592 (`if config.
# is_gguf:`) is responsible for sending GGUF traffic to the
# llama-server backend before reaching the MLX orchestrator.
# If we end up here with is_gguf=True, the route's
# `detect_gguf_model_remote` returned None on its first call
# (transient HF Hub flake) but the subprocess re-detection
# succeeded. The subprocess cannot reach into the parent's
# llama-server, so all we can do is raise loudly so the caller
# gets a clear error instead of a cryptic
# "config.json does not exist" from mlx_lm.utils.load_model.
if getattr(config, "is_gguf", False):
raise RuntimeError(
f"MLXInferenceBackend cannot load GGUF model '{model_name}': "
f"GGUF models must be served by llama-server in the parent "
f"process. The /api/inference/load route should have "
f"detected this repo as GGUF before dispatching to the MLX "
f"orchestrator -- this fallback indicates a transient HF "
f"Hub failure during initial detection. Retry the request."
)
if hf_token:
import os
os.environ["HF_TOKEN"] = hf_token
self._configure_memory_limits()
is_lora = getattr(config, "is_lora", False)
logger.info(
"Loading %s via %s (is_lora=%s)",
model_name,
"mlx-vlm" if is_vision else "mlx-lm",
is_lora,
)
try:
from unsloth_zoo.mlx.loader import FastMLXModel
except ImportError as e:
raise ImportError(
"Unsloth: MLX inference requires unsloth-zoo with the MLX modules "
"(unsloth_zoo.mlx.loader). Reinstall via install.sh on Apple Silicon."
) from e
model, tokenizer_or_processor = FastMLXModel.from_pretrained(
model_name,
max_seq_length = max_seq_length,
dtype = dtype,
load_in_4bit = load_in_4bit,
token = hf_token,
trust_remote_code = trust_remote_code,
text_only = False if is_vision else True,
)
if is_vision:
processor = tokenizer_or_processor
self._model = model
self._processor = processor
self._tokenizer = getattr(processor, "tokenizer", processor)
self._is_vlm = True
else:
tokenizer = tokenizer_or_processor
self._model = model
self._tokenizer = tokenizer
self._processor = None
self._is_vlm = False
self.active_model_name = model_name
self.models[model_name] = {
"model": self._model,
"tokenizer": self._tokenizer,
"processor": self._processor,
"is_vision": is_vision,
"is_lora": getattr(config, "is_lora", False),
"is_audio": False,
"audio_type": None,
"has_audio_input": False,
}
logger.info("Model %s loaded successfully", model_name)
return True
def unload_model(self, model_name: str) -> bool:
import mlx.core as mx
import gc
if model_name in self.models:
del self.models[model_name]
self._model = None
self._tokenizer = None
self._processor = None
if self.active_model_name == model_name:
self.active_model_name = None
gc.collect()
mx.clear_cache()
if mx.metal.is_available() and self._memory_limits_applied and not self.models:
try:
mx.set_wired_limit(0)
logger.info("MLX wired_limit released back to OS on unload")
except Exception as e:
logger.warning("Failed to release wired_limit: %s", e)
self._memory_limits_applied = {}
logger.info("Model %s unloaded", model_name)
return True
def generate_chat_response(
self,
messages,
system_prompt = "",
image = None,
temperature = 0.7,
top_p = 0.9,
top_k = 40,
min_p = 0.0,
max_new_tokens = 256,
repetition_penalty = 1.0,
cancel_event = None,
) -> Generator[str, None, None]:
if self._model is None:
raise RuntimeError("No model loaded")
# Build messages with system prompt
full_messages = []
if system_prompt:
full_messages.append({"role": "system", "content": system_prompt})
full_messages.extend(messages)
# Inject image into the last user message for VLM
if self._is_vlm and image is not None:
for msg in reversed(full_messages):
if msg.get("role") == "user":
content = msg.get("content", "")
if isinstance(content, str):
msg["content"] = [
{"type": "image"},
{"type": "text", "text": content},
]
elif isinstance(content, list):
# Prepend image if not already there
has_image = any(
p.get("type") == "image"
for p in content
if isinstance(p, dict)
)
if not has_image:
content.insert(0, {"type": "image"})
break
if self._is_vlm:
yield from self._generate_vlm(
full_messages,
image,
temperature,
top_p,
top_k,
min_p,
max_new_tokens,
repetition_penalty,
cancel_event,
)
else:
yield from self._generate_text(
full_messages,
temperature,
top_p,
top_k,
min_p,
max_new_tokens,
repetition_penalty,
cancel_event,
)
def _generate_text(
self,
messages,
temperature,
top_p,
top_k,
min_p,
max_new_tokens,
repetition_penalty,
cancel_event,
):
from mlx_lm import stream_generate
from mlx_lm.sample_utils import make_sampler, make_logits_processors
prompt = self._tokenizer.apply_chat_template(
messages,
tokenize = False,
add_generation_prompt = True,
)
if prompt is None:
raise RuntimeError(
"apply_chat_template returned None — tokenizer may be incompatible"
)
sampler = make_sampler(
temp = temperature,
top_p = top_p,
top_k = int(top_k or 0),
min_p = float(min_p or 0.0),
min_tokens_to_keep = 1,
)
# Only build a logits processor when we actually have a non-trivial
# repetition penalty (1.0 is the no-op value).
logits_processors = None
if repetition_penalty is not None and float(repetition_penalty) not in (
0.0,
1.0,
):
logits_processors = make_logits_processors(
repetition_penalty = float(repetition_penalty),
)
token_ids = []
logger.info(
"Generating: prompt_len=%d, max_tokens=%d, model=%s, tokenizer=%s",
len(prompt),
max_new_tokens,
type(self._model).__name__,
type(self._tokenizer).__name__,
)
with self._generation_lock:
try:
gen_kwargs = dict(
prompt = prompt,
max_tokens = max_new_tokens,
sampler = sampler,
)
if logits_processors is not None:
gen_kwargs["logits_processors"] = logits_processors
for response in stream_generate(
self._model,
self._tokenizer,
**gen_kwargs,
):
token_ids.append(response.token)
# Decode full sequence with skip_special_tokens — same as GPU
cumulative = self._tokenizer.decode(
token_ids,
skip_special_tokens = True,
)
yield cumulative
if cancel_event and cancel_event.is_set():
break
except Exception as e:
import traceback
logger.error("stream_generate failed:\n%s", traceback.format_exc())
raise
def _generate_vlm(
self,
messages,
image,
temperature,
top_p,
top_k,
min_p,
max_new_tokens,
repetition_penalty,
cancel_event,
):
from mlx_vlm import stream_generate as vlm_stream
# Apply chat template
chat_fn = getattr(self._processor, "apply_chat_template", None)
if (
chat_fn is None
or not hasattr(self._processor, "chat_template")
or self._processor.chat_template is None
):
tok = getattr(self._processor, "tokenizer", self._processor)
chat_fn = tok.apply_chat_template
prompt = chat_fn(messages, tokenize = False, add_generation_prompt = True)
# For VLM: always use mlx_vlm's stream_generate which handles
# pixel_values properly (passes None for text-only, image for VLM)
images = [image] if image is not None else None
cumulative = ""
logger.info(
"VLM generating: prompt_len=%d, has_image=%s",
len(prompt),
image is not None,
)
# mlx_vlm.stream_generate forwards **kwargs into generate_step, which
# accepts temp/top_p/top_k/repetition_penalty (and builds the sampler
# + logits_processors internally). Pass them through.
# NOTE: mlx_vlm.generate_step expects ``temperature=`` (long form) —
# passing ``temp=`` silently falls into **kwargs and is ignored,
# leaving generation stuck at the default 0.0 (greedy).
vlm_kwargs = dict(
max_tokens = max_new_tokens,
temperature = temperature,
top_p = top_p,
top_k = int(top_k or 0),
min_p = float(min_p or 0.0),
)
if repetition_penalty is not None and float(repetition_penalty) not in (
0.0,
1.0,
):
vlm_kwargs["repetition_penalty"] = float(repetition_penalty)
with self._generation_lock:
for response in vlm_stream(
self._model,
self._processor,
prompt,
images,
**vlm_kwargs,
):
token_text = (
response.text if hasattr(response, "text") else str(response)
)
cumulative += token_text
yield cumulative
if cancel_event and cancel_event.is_set():
break
def generate_with_adapter_control(
self, use_adapter = None, cancel_event = None, **gen_kwargs
) -> Generator[str, None, None]:
# MLX LoRA adapter toggling not yet supported — generate normally
yield from self.generate_chat_response(cancel_event = cancel_event, **gen_kwargs)
def reset_generation_state(self):
import mlx.core as mx
import gc
gc.collect()
mx.clear_cache()