"""Standalone MLX inference worker: one process = one model. Started by the Model Manager: python -m llm_api.worker.mlx_worker --model-path P --port N --model-id ID [--vision] [--task text|embedding|reranking] Exposes on 127.0.0.1: GET /health loading|ready|error + memory POST /v1/chat/completions OpenAI compatible (stream or not) POST /v1/completions POST /v1/embeddings POST /v1/rerank POST /tokenize POST /warmup POST /clear-cache POST /shutdown Killing the process is the memory-release mechanism: everything (weights, KV cache, Metal heaps) goes away with it. """ from __future__ import annotations import argparse import asyncio import base64 import json import logging import math import os import queue import sys import tempfile import threading import time from pathlib import Path from typing import Any import psutil from fastapi import FastAPI, Request from fastapi.responses import JSONResponse, StreamingResponse from .openai_types import (HARMONY_MARKERS, HARMONY_THINK_END, HARMONY_THINK_START, MarkerStripper, StopMatcher, ThinkSplitter, chat_chunk, new_id, normalize_messages, parse_tool_calls, sse) log = logging.getLogger("mlx_worker") GB = 1024**3 # --------------------------------------------------------------------------- # State # --------------------------------------------------------------------------- class WorkerState: def __init__(self, args: argparse.Namespace): self.args = args self.status = "loading" self.error: str | None = None self.started = time.time() self.loaded_at: float | None = None self.load_ms: float | None = None self.model = None self.tokenizer = None self.processor = None # mlx_vlm self.config: dict = {} self.gen_config: dict = {} self.vision = bool(args.vision) self.task = args.task self.max_context = args.max_context self.gen_lock = threading.Lock() self.requests = 0 self.tokens_generated = 0 self.last_used = time.time() # prompt cache (single conversation) self.cache_tokens: list[int] = [] self.cache_obj = None self.warm_ttft_ms: float | None = None self.template_text = "" self.harmony = False self.thinks = False class MLXThread: """All MLX work (import, load, generation, embeddings) runs on this single thread. MLX streams/command buffers are thread-affine and MLX is not thread-safe.""" def __init__(self): self._q: "queue.Queue[tuple]" = queue.Queue() self._th = threading.Thread(target=self._run, name="mlx", daemon=True) self._th.start() def _run(self): while True: fn, args, fut = self._q.get() if fn is None: return try: res = fn(*args) if fut is not None: fut.set_result(res) except BaseException as e: # noqa: BLE001 if fut is not None: fut.set_exception(e) else: log.exception("mlx thread task failed") def submit(self, fn, *args): import concurrent.futures fut: concurrent.futures.Future = concurrent.futures.Future() self._q.put((fn, args, fut)) return fut def alive(self) -> bool: return self._th.is_alive() MLX_THREAD = MLXThread() STATE: WorkerState | None = None app = FastAPI(title="llm-api mlx worker") def _mx(): import mlx.core as mx return mx def memory_info() -> dict: out: dict[str, Any] = {"rss_gb": round(psutil.Process().memory_info().rss / GB, 3)} try: mx = _mx() out.update({ "active_gb": round(mx.get_active_memory() / GB, 3), "peak_gb": round(mx.get_peak_memory() / GB, 3), "cache_gb": round(mx.get_cache_memory() / GB, 3), }) except Exception: pass return out # --------------------------------------------------------------------------- # Loading # --------------------------------------------------------------------------- def _load_model(st: WorkerState) -> None: t0 = time.time() try: mx = _mx() path = st.args.model_path cfg_p = Path(path) / "config.json" if cfg_p.exists(): st.config = json.loads(cfg_p.read_text()) gc_p = Path(path) / "generation_config.json" if gc_p.exists(): try: st.gen_config = json.loads(gc_p.read_text()) except Exception: st.gen_config = {} tpl_p = Path(path) / "chat_template.jinja" if tpl_p.exists(): st.template_text = tpl_p.read_text(errors="replace") else: tc_p = Path(path) / "tokenizer_config.json" if tc_p.exists(): try: t = json.loads(tc_p.read_text()).get("chat_template") or "" st.template_text = t if isinstance(t, str) else json.dumps(t) except Exception: pass st.harmony = "<|channel|>" in st.template_text st.thinks = st.harmony or "" in st.template_text or "enable_thinking" in st.template_text if st.vision: try: import mlx_vlm from mlx_vlm.utils import load_config st.model, st.processor = mlx_vlm.load(path) st.config = load_config(path) tok = getattr(st.processor, "tokenizer", st.processor) st.tokenizer = tok log.info("loaded with mlx_vlm") except Exception as e: # fall back to text-only log.warning("mlx_vlm load failed (%s); falling back to mlx_lm text-only", e) st.vision = False if not st.vision: from mlx_lm import load st.model, st.tokenizer = load(path) # Touch the weights so lazily-loaded parameters are materialized try: from mlx.utils import tree_flatten params = tree_flatten(st.model.parameters()) mx.eval([p for _, p in params]) except Exception: pass st.load_ms = (time.time() - t0) * 1000 st.loaded_at = time.time() st.status = "ready" log.info("model ready in %.0f ms (%s)", st.load_ms, memory_info()) except Exception as e: st.status = "error" st.error = f"{type(e).__name__}: {e}" log.exception("model load failed") # --------------------------------------------------------------------------- # Prompt building # --------------------------------------------------------------------------- def _apply_template(st: WorkerState, messages: list[dict], tools: list | None, kwargs: dict) -> list[int]: tok = st.tokenizer if st.vision and st.processor is not None: # handled by vision path raise RuntimeError("use vision path") has_tpl = getattr(tok, "has_chat_template", True) if has_tpl: try: out = tok.apply_chat_template(messages, tools=tools, add_generation_prompt=True, tokenize=True, **kwargs) except TypeError: out = tok.apply_chat_template(messages, add_generation_prompt=True, tokenize=True) if isinstance(out, dict): out = out.get("input_ids") if hasattr(out, "tolist"): out = out.tolist() if out and isinstance(out[0], list): out = out[0] return list(out) # No chat template: simple fallback text = "" for m in messages: text += f"{m.get('role','user').capitalize()}: {m.get('content','')}\n" text += "Assistant:" return list(tok.encode(text)) def _sampler_and_processors(body: dict, st: WorkerState): from mlx_lm.sample_utils import make_logits_processors, make_sampler gc = st.gen_config or {} temp = body.get("temperature") if temp is None: temp = gc.get("temperature", 0.7) top_p = body.get("top_p") if top_p is None: top_p = gc.get("top_p", 0.95) top_k = body.get("top_k") if top_k is None: top_k = gc.get("top_k", 0) or 0 min_p = body.get("min_p", 0.0) or 0.0 sampler = make_sampler(temp=float(temp), top_p=float(top_p) if top_p and top_p < 1.0 else 0.0, min_p=float(min_p), top_k=int(top_k)) logit_bias = body.get("logit_bias") lb = None if isinstance(logit_bias, dict) and logit_bias: lb = {} for k, v in logit_bias.items(): try: lb[int(k)] = float(v) except (TypeError, ValueError): pass rep = body.get("repetition_penalty") pres = body.get("presence_penalty") or None freq = body.get("frequency_penalty") or None procs = make_logits_processors(logit_bias=lb, repetition_penalty=float(rep) if rep else None, presence_penalty=float(pres) if pres else None, frequency_penalty=float(freq) if freq else None) return sampler, procs def _max_tokens(body: dict, prompt_len: int, st: WorkerState) -> int: mt = body.get("max_completion_tokens") or body.get("max_tokens") if mt is None: mt = int(st.args.default_max_tokens) ctx = st.max_context room = ctx - prompt_len if room < 16: raise ValueError(f"prompt has {prompt_len} tokens; the loaded context window is {ctx} tokens.") return max(1, min(int(mt), room)) def _prepare_cache(st: WorkerState, prompt: list[int]): """Reuse KV cache for the common prefix of the previous conversation.""" from mlx_lm.models.cache import can_trim_prompt_cache, make_prompt_cache, trim_prompt_cache if st.cache_obj is not None and st.cache_tokens and can_trim_prompt_cache(st.cache_obj): n = 0 for a, b in zip(st.cache_tokens, prompt): if a != b: break n += 1 # always leave at least one token to process n = min(n, len(prompt) - 1) if n > 0: to_trim = len(st.cache_tokens) - n if to_trim > 0: trim_prompt_cache(st.cache_obj, to_trim) st.cache_tokens = prompt[:n] return st.cache_obj, prompt[n:], n st.cache_obj = make_prompt_cache(st.model) st.cache_tokens = [] return st.cache_obj, prompt, 0 # --------------------------------------------------------------------------- # Generation (runs in a thread; yields events into a queue) # --------------------------------------------------------------------------- def _generate_thread(st: WorkerState, prompt: list[int], body: dict, chat: bool, out: "queue.Queue[dict]", cancel: threading.Event) -> None: mx = _mx() try: from mlx_lm import stream_generate if body.get("seed") is not None: mx.random.seed(int(body["seed"])) sampler, procs = _sampler_and_processors(body, st) max_tokens = _max_tokens(body, len(prompt), st) cache, rest, cached = _prepare_cache(st, prompt) t0 = time.time() first = None n_gen = 0 gen_tokens: list[int] = [] prompt_tps = 0.0 gen_tps = 0.0 peak = 0.0 finish = "length" eos_ids = getattr(st.tokenizer, "eos_token_ids", None) or set() kwargs: dict[str, Any] = {} if st.args.kv_bits: kwargs["kv_bits"] = int(st.args.kv_bits) kwargs["quantized_kv_start"] = 4096 for g in stream_generate(st.model, st.tokenizer, rest, max_tokens=max_tokens, sampler=sampler, logits_processors=procs, prompt_cache=cache, prefill_step_size=2048, **kwargs): if first is None: first = time.time() n_gen += 1 gen_tokens.append(g.token) prompt_tps = g.prompt_tps gen_tps = g.generation_tps peak = g.peak_memory if g.finish_reason: finish = g.finish_reason out.put({"text": g.text, "token": g.token}) if cancel.is_set(): finish = "cancelled" break if g.finish_reason: break t1 = time.time() # Remember exactly the tokens the KV cache holds (the last sampled token is never fed back) all_tokens = list(prompt) + gen_tokens try: off = int(cache[0].offset) except Exception: off = len(all_tokens) - 1 st.cache_tokens = all_tokens[:max(0, min(off, len(all_tokens)))] st.tokens_generated += n_gen out.put({ "done": True, "finish_reason": finish, "prompt_tokens": len(prompt), "completion_tokens": n_gen, "cached_tokens": cached, "timings": { "ttft_ms": round(((first or t1) - t0) * 1000, 1), "prompt_ms": round(((first or t1) - t0) * 1000, 1), "generation_ms": round((t1 - (first or t1)) * 1000, 1), "total_ms": round((t1 - t0) * 1000, 1), "prompt_tps": round(prompt_tps, 1), "generation_tps": round(gen_tps, 2), "peak_memory_gb": round(peak, 3), }, }) except Exception as e: log.exception("generation failed") # A failed generation may leave the cache inconsistent st.cache_obj = None st.cache_tokens = [] out.put({"error": f"{type(e).__name__}: {e}"}) finally: try: mx.clear_cache() except Exception: pass def _vision_generate_thread(st: WorkerState, messages: list[dict], images: list[dict], body: dict, tk: dict, tools: list | None, out: "queue.Queue[dict]", cancel: threading.Event) -> None: mx = _mx() tmpfiles: list[str] = [] try: from mlx_vlm import stream_generate as vlm_stream from mlx_vlm.prompt_utils import apply_chat_template paths: list[str] = [] for part in images: url = part.get("image_url", {}).get("url") if isinstance(part.get("image_url"), dict) else part.get("image_url") or part.get("url") if not url: continue if url.startswith("data:"): header, b64 = url.split(",", 1) ext = ".png" if "png" in header else ".jpg" fd, p = tempfile.mkstemp(suffix=ext) with os.fdopen(fd, "wb") as f: f.write(base64.b64decode(b64)) tmpfiles.append(p) paths.append(p) else: paths.append(url) tpl_kwargs = dict(tk or {}) if tools: tpl_kwargs["tools"] = tools try: prompt = apply_chat_template(st.processor, st.config, messages, num_images=len(paths), **tpl_kwargs) except TypeError: prompt = apply_chat_template(st.processor, st.config, messages, num_images=len(paths)) if body.get("seed") is not None: mx.random.seed(int(body["seed"])) mt = body.get("max_completion_tokens") or body.get("max_tokens") or int(st.args.default_max_tokens) gc = st.gen_config or {} temp = body.get("temperature") if temp is None: temp = gc.get("temperature", 0.7) top_p = body.get("top_p") if top_p is None: top_p = gc.get("top_p", 0.95) gen_kwargs: dict[str, Any] = {"max_tokens": int(mt), "temperature": float(temp), "top_p": float(top_p)} top_k = body.get("top_k", gc.get("top_k")) if top_k: gen_kwargs["top_k"] = int(top_k) if body.get("min_p"): gen_kwargs["min_p"] = float(body["min_p"]) for k in ("repetition_penalty", "presence_penalty", "frequency_penalty"): if body.get(k): gen_kwargs[k] = float(body[k]) if isinstance(body.get("logit_bias"), dict) and body["logit_bias"]: try: gen_kwargs["logit_bias"] = {int(k): float(v) for k, v in body["logit_bias"].items()} except (TypeError, ValueError): pass if st.args.kv_bits: gen_kwargs["kv_bits"] = int(st.args.kv_bits) t0 = time.time() first = None n = 0 ptoks = 0 ptps = gtps = peak = 0.0 finish = "length" for g in vlm_stream(st.model, st.processor, prompt, image=paths or None, **gen_kwargs): if first is None: first = time.time() n += 1 ptoks = getattr(g, "prompt_tokens", ptoks) ptps = getattr(g, "prompt_tps", ptps) gtps = getattr(g, "generation_tps", gtps) peak = getattr(g, "peak_memory", peak) fr = getattr(g, "finish_reason", None) if fr: finish = fr out.put({"text": g.text, "token": getattr(g, "token", 0)}) if cancel.is_set(): finish = "cancelled" break if fr: break t1 = time.time() st.tokens_generated += n out.put({"done": True, "finish_reason": finish, "prompt_tokens": int(ptoks), "completion_tokens": n, "cached_tokens": 0, "timings": {"ttft_ms": round(((first or t1) - t0) * 1000, 1), "total_ms": round((t1 - t0) * 1000, 1), "prompt_tps": round(ptps, 1), "generation_tps": round(gtps, 2), "peak_memory_gb": round(peak, 3)}}) except Exception as e: log.exception("vision generation failed") out.put({"error": f"{type(e).__name__}: {e}"}) finally: for p in tmpfiles: try: os.unlink(p) except OSError: pass try: mx.clear_cache() except Exception: pass class _Task: """Thread-like wrapper around a Future running on the MLX thread.""" def __init__(self, fut): self.fut = fut def is_alive(self) -> bool: return not self.fut.done() def join(self, timeout: float | None = None) -> None: try: self.fut.result(timeout=timeout) except Exception: pass async def _run_generation(st: WorkerState, target, *args) -> tuple["queue.Queue[dict]", threading.Event, _Task]: out: queue.Queue[dict] = queue.Queue() cancel = threading.Event() fut = MLX_THREAD.submit(target, st, *args, out, cancel) return out, cancel, _Task(fut) async def _next_event(q: "queue.Queue[dict]", th: "_Task | None" = None) -> dict: """Next event from the generation thread; detects a dead thread instead of waiting forever.""" loop = asyncio.get_running_loop() def _get(): while True: try: return q.get(timeout=1.0) except queue.Empty: if th is not None and not th.is_alive(): try: return q.get_nowait() except queue.Empty: return {"error": "generation thread exited unexpectedly"} return await loop.run_in_executor(None, _get) # --------------------------------------------------------------------------- # Endpoints # --------------------------------------------------------------------------- def _err(status: int, message: str, code: str = "WORKER_ERROR", etype: str = "runtime_error") -> JSONResponse: return JSONResponse(status_code=status, content={"error": {"message": message, "type": etype, "code": code}}) @app.get("/health") async def health(): st = STATE assert st return { "status": st.status, "error": st.error, "model_id": st.args.model_id, "runtime": "mlx", "vision": st.vision, "task": st.task, "elapsed_seconds": round(time.time() - st.started, 1), "load_ms": st.load_ms, "memory": memory_info(), "requests": st.requests, "tokens_generated": st.tokens_generated, "max_context": st.max_context, "pid": os.getpid(), "busy": st.gen_lock.locked(), "warm_ttft_ms": st.warm_ttft_ms, "mlx_thread_alive": MLX_THREAD.alive(), } @app.post("/clear-cache") async def clear_cache(): st = STATE assert st st.cache_obj = None st.cache_tokens = [] try: _mx().clear_cache() except Exception: pass return memory_info() @app.post("/shutdown") async def shutdown(): async def _exit(): await asyncio.sleep(0.2) os._exit(0) asyncio.create_task(_exit()) return {"ok": True} @app.post("/tokenize") async def tokenize(req: Request): st = STATE assert st body = await req.json() if st.status != "ready": return _err(503, "model not ready", "MODEL_NOT_READY") if "messages" in body: msgs, _ = normalize_messages(body["messages"]) try: toks = _apply_template(st, msgs, body.get("tools"), body.get("chat_template_kwargs") or {}) except RuntimeError: toks = list(st.tokenizer.encode(" ".join(m.get("content", "") for m in msgs))) else: toks = list(st.tokenizer.encode(body.get("text", "") or body.get("prompt", ""))) return {"tokens": len(toks), "max_context": st.max_context} @app.post("/warmup") async def warmup(): st = STATE assert st if st.status != "ready": return _err(503, "model not ready", "MODEL_NOT_READY") t0 = time.time() if st.task in ("embedding",): r = await embeddings_impl({"input": "warm up"}) st.warm_ttft_ms = round((time.time() - t0) * 1000, 1) return {"ok": True, "ttft_ms": st.warm_ttft_ms, "kind": "embedding", "dims": len(r["data"][0]["embedding"])} if st.task == "reranking": r = await rerank_impl({"query": "warm", "documents": ["warm up"]}) st.warm_ttft_ms = round((time.time() - t0) * 1000, 1) return {"ok": True, "ttft_ms": st.warm_ttft_ms, "kind": "rerank"} body = {"messages": [{"role": "user", "content": "Say OK."}], "max_tokens": 4, "temperature": 0.0} resp = await chat_impl(body, stream=False) if isinstance(resp, JSONResponse): return resp st.warm_ttft_ms = resp.get("timings", {}).get("ttft_ms") # do not keep the warm-up prompt in the cache st.cache_obj = None st.cache_tokens = [] return {"ok": True, "ttft_ms": st.warm_ttft_ms, "text": resp["choices"][0]["message"]["content"]} async def chat_impl(body: dict, stream: bool): st = STATE assert st if st.status != "ready": return _err(503, f"model not ready ({st.status})", "MODEL_NOT_READY") messages = body.get("messages") if not isinstance(messages, list) or not messages: return _err(400, "messages is required", "INVALID_REQUEST", "invalid_request_error") msgs, images = normalize_messages(messages) tools = body.get("tools") or None tk = dict(body.get("chat_template_kwargs") or {}) # reasoning controls (Qwen3-style enable_thinking) reasoning = body.get("reasoning") if isinstance(reasoning, dict) and "effort" in reasoning: tk.setdefault("enable_thinking", reasoning["effort"] not in ("none", "minimal")) if body.get("enable_thinking") is not None: tk["enable_thinking"] = bool(body["enable_thinking"]) if body.get("reasoning_effort") is not None: tk.setdefault("enable_thinking", body["reasoning_effort"] not in ("none", "minimal")) stops = body.get("stop") or [] if isinstance(stops, str): stops = [stops] rid = new_id("chatcmpl") created = int(time.time()) model_name = body.get("model") or st.args.model_id if not st.gen_lock.acquire(timeout=float(st.args.queue_timeout)): return _err(503, "worker busy", "WORKER_BUSY") st.requests += 1 st.last_used = time.time() try: if st.vision and st.processor is not None: # mlx-vlm handles both text-only and image requests for vision-language models q, cancel, th = await _run_generation(st, _vision_generate_thread, msgs, images, body, tk, tools) think_start, think_end = "", "" thinking_enabled = st.thinks and tk.get("enable_thinking", True) else: if images: return _err(400, "This model does not accept images.", "VISION_UNSUPPORTED", "invalid_request_error") try: prompt = _apply_template(st, msgs, tools, tk) except Exception as e: return _err(400, f"chat template error: {e}", "TEMPLATE_ERROR", "invalid_request_error") if len(prompt) >= st.max_context - 16: return _err(400, f"Prompt has {len(prompt)} tokens but the context window is {st.max_context}.", "CONTEXT_TOO_LARGE", "invalid_request_error") q, cancel, th = await _run_generation(st, _generate_thread, prompt, body, True) think_start = getattr(st.tokenizer, "think_start", None) or "" think_end = getattr(st.tokenizer, "think_end", None) or "" thinking_enabled = (bool(getattr(st.tokenizer, "has_thinking", False)) or st.thinks) and tk.get("enable_thinking", True) markers: list[str] = [] if st.harmony: think_start, think_end = HARMONY_THINK_START, HARMONY_THINK_END thinking_enabled = True markers = HARMONY_MARKERS splitter = ThinkSplitter(think_start, think_end) if thinking_enabled else ThinkSplitter(None, None) stripper = MarkerStripper(markers) stopper = StopMatcher(list(stops)) has_tools = bool(tools) if stream: async def gen(): try: yield sse(chat_chunk(rid, model_name, created, {"role": "assistant", "content": ""})) content_acc = "" reasoning_acc = "" finish = "stop" final: dict = {} while True: ev = await _next_event(q, th) if "error" in ev: yield sse({"error": {"message": ev["error"], "type": "runtime_error", "code": "GENERATION_FAILED"}}) yield sse("[DONE]") return if ev.get("done"): final = ev finish = ev["finish_reason"] break r, c = splitter.feed(ev["text"]) if r: reasoning_acc += r yield sse(chat_chunk(rid, model_name, created, {"reasoning_content": r})) c = stripper.feed(c) if c else c if c: c = stopper.feed(c) if c and not has_tools: content_acc += c yield sse(chat_chunk(rid, model_name, created, {"content": c})) elif c: content_acc += c if stopper.done: cancel.set() finish = "stop" # drain while True: ev2 = await _next_event(q, th) if ev2.get("done") or "error" in ev2: final = ev2 if ev2.get("done") else {} break break r, c = splitter.flush() c = (stripper.feed(c) + stripper.flush()) if not stopper.done else "" c = (stopper.flush() + c) if not stopper.done else "" if r: yield sse(chat_chunk(rid, model_name, created, {"reasoning_content": r})) tool_calls: list[dict] = [] if has_tools: rest, tool_calls = parse_tool_calls(content_acc + c) if tool_calls: finish = "tool_calls" yield sse(chat_chunk(rid, model_name, created, {"tool_calls": [ {"index": i, **tc} for i, tc in enumerate(tool_calls)]})) elif content_acc + c: yield sse(chat_chunk(rid, model_name, created, {"content": content_acc + c})) elif c: yield sse(chat_chunk(rid, model_name, created, {"content": c})) if finish == "cancelled": finish = "stop" if finish not in ("stop", "length", "tool_calls"): finish = "stop" usage = {"prompt_tokens": final.get("prompt_tokens", 0), "completion_tokens": final.get("completion_tokens", 0), "total_tokens": final.get("prompt_tokens", 0) + final.get("completion_tokens", 0)} if final.get("cached_tokens"): usage["prompt_tokens_details"] = {"cached_tokens": final["cached_tokens"]} yield sse(chat_chunk(rid, model_name, created, {}, finish, usage, {"timings": final.get("timings", {})})) yield sse("[DONE]") finally: cancel.set() th.join(timeout=float(st.args.generation_timeout)) st.gen_lock.release() return StreamingResponse(gen(), media_type="text/event-stream", headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}) # non-streaming try: content_acc = "" reasoning_acc = "" final = {} finish = "stop" while True: ev = await _next_event(q, th) if "error" in ev: return _err(500, ev["error"], "GENERATION_FAILED") if ev.get("done"): final = ev finish = ev["finish_reason"] break r, c = splitter.feed(ev["text"]) reasoning_acc += r c = stripper.feed(c) if c else c if c: c = stopper.feed(c) content_acc += c if stopper.done: cancel.set() finish = "stop" while True: ev2 = await _next_event(q, th) if ev2.get("done") or "error" in ev2: final = ev2 if ev2.get("done") else {} break break r, c = splitter.flush() reasoning_acc += r if not stopper.done: content_acc += stopper.flush() + stripper.feed(c) + stripper.flush() content_acc = content_acc.strip("\n") if reasoning_acc else content_acc tool_calls: list[dict] = [] if has_tools: content_acc, tool_calls = parse_tool_calls(content_acc) if tool_calls: finish = "tool_calls" if finish not in ("stop", "length", "tool_calls"): finish = "stop" msg: dict[str, Any] = {"role": "assistant", "content": content_acc if not tool_calls else (content_acc or None)} if reasoning_acc.strip(): msg["reasoning_content"] = reasoning_acc.strip() if tool_calls: msg["tool_calls"] = tool_calls usage = {"prompt_tokens": final.get("prompt_tokens", 0), "completion_tokens": final.get("completion_tokens", 0), "total_tokens": final.get("prompt_tokens", 0) + final.get("completion_tokens", 0)} if final.get("cached_tokens"): usage["prompt_tokens_details"] = {"cached_tokens": final["cached_tokens"]} return {"id": rid, "object": "chat.completion", "created": created, "model": model_name, "choices": [{"index": 0, "message": msg, "logprobs": None, "finish_reason": finish}], "usage": usage, "timings": final.get("timings", {})} finally: cancel.set() th.join(timeout=float(st.args.generation_timeout)) st.gen_lock.release() except Exception: if st.gen_lock.locked(): try: st.gen_lock.release() except RuntimeError: pass raise @app.post("/v1/chat/completions") async def chat_completions(req: Request): body = await req.json() return await chat_impl(body, bool(body.get("stream"))) @app.post("/v1/completions") async def completions(req: Request): st = STATE assert st body = await req.json() if st.status != "ready": return _err(503, f"model not ready ({st.status})", "MODEL_NOT_READY") prompt = body.get("prompt", "") if isinstance(prompt, list): prompt = prompt[0] if prompt and isinstance(prompt[0], str) else "" stream = bool(body.get("stream")) rid = new_id("cmpl") created = int(time.time()) model_name = body.get("model") or st.args.model_id toks = list(st.tokenizer.encode(prompt)) if len(toks) >= st.max_context - 16: return _err(400, f"Prompt has {len(toks)} tokens but the context window is {st.max_context}.", "CONTEXT_TOO_LARGE", "invalid_request_error") stops = body.get("stop") or [] if isinstance(stops, str): stops = [stops] if not st.gen_lock.acquire(timeout=float(st.args.queue_timeout)): return _err(503, "worker busy", "WORKER_BUSY") st.requests += 1 st.last_used = time.time() # completions never reuse the chat cache st.cache_obj = None st.cache_tokens = [] q, cancel, th = await _run_generation(st, _generate_thread, toks, body, False) stopper = StopMatcher(list(stops)) echo = bool(body.get("echo")) def chunk(text: str, finish=None, usage=None, extra=None): d: dict[str, Any] = {"id": rid, "object": "text_completion", "created": created, "model": model_name, "choices": [{"index": 0, "text": text, "logprobs": None, "finish_reason": finish}]} if usage: d["usage"] = usage if extra: d.update(extra) return d if stream: async def gen(): try: if echo: yield sse(chunk(prompt)) finish = "stop" final: dict = {} while True: ev = await _next_event(q, th) if "error" in ev: yield sse({"error": {"message": ev["error"], "type": "runtime_error", "code": "GENERATION_FAILED"}}) break if ev.get("done"): final, finish = ev, ev["finish_reason"] break c = stopper.feed(ev["text"]) if c: yield sse(chunk(c)) if stopper.done: cancel.set() finish = "stop" while True: ev2 = await _next_event(q, th) if ev2.get("done") or "error" in ev2: final = ev2 if ev2.get("done") else {} break break tail = stopper.flush() if not stopper.done else "" if tail: yield sse(chunk(tail)) usage = {"prompt_tokens": final.get("prompt_tokens", 0), "completion_tokens": final.get("completion_tokens", 0), "total_tokens": final.get("prompt_tokens", 0) + final.get("completion_tokens", 0)} yield sse(chunk("", finish if finish in ("stop", "length") else "stop", usage, {"timings": final.get("timings", {})})) yield sse("[DONE]") finally: cancel.set() th.join(timeout=float(st.args.generation_timeout)) st.cache_obj = None st.cache_tokens = [] st.gen_lock.release() return StreamingResponse(gen(), media_type="text/event-stream", headers={"Cache-Control": "no-cache"}) try: text = "" final = {} finish = "stop" while True: ev = await _next_event(q, th) if "error" in ev: return _err(500, ev["error"], "GENERATION_FAILED") if ev.get("done"): final, finish = ev, ev["finish_reason"] break text += stopper.feed(ev["text"]) if stopper.done: cancel.set() finish = "stop" while True: ev2 = await _next_event(q, th) if ev2.get("done") or "error" in ev2: final = ev2 if ev2.get("done") else {} break break if not stopper.done: text += stopper.flush() usage = {"prompt_tokens": final.get("prompt_tokens", 0), "completion_tokens": final.get("completion_tokens", 0), "total_tokens": final.get("prompt_tokens", 0) + final.get("completion_tokens", 0)} return chunk((prompt if echo else "") + text, finish if finish in ("stop", "length") else "stop", usage, {"timings": final.get("timings", {})}) finally: cancel.set() th.join(timeout=float(st.args.generation_timeout)) st.cache_obj = None st.cache_tokens = [] st.gen_lock.release() # --------------------------------------------------------------------------- # Embeddings / rerank (causal-LM style: last-token pooling, Qwen3-Embedding & co.) # --------------------------------------------------------------------------- def _hidden_states(st: WorkerState, tokens: list[int]): mx = _mx() inner = getattr(st.model, "model", None) or getattr(st.model, "language_model", None) x = mx.array([tokens]) if inner is not None and callable(inner): h = inner(x) else: h = st.model(x) mx.eval(h) return h def _embed_one(st: WorkerState, text: str, dims: int | None) -> tuple[list[float], int]: mx = _mx() tok = st.tokenizer ids = list(tok.encode(text)) eos = getattr(tok, "eos_token_id", None) if eos is None: eos_ids = getattr(tok, "eos_token_ids", None) or set() eos = next(iter(eos_ids), None) if eos is not None and (not ids or ids[-1] != eos): ids.append(eos) ids = ids[: st.max_context] h = _hidden_states(st, ids) v = h[0, -1, :].astype(mx.float32) if dims and dims < v.shape[0]: v = v[:dims] norm = mx.sqrt(mx.sum(v * v)) + 1e-12 v = v / norm mx.eval(v) return v.tolist(), len(ids) async def embeddings_impl(body: dict): st = STATE assert st if st.status != "ready": return _err(503, f"model not ready ({st.status})", "MODEL_NOT_READY") inp = body.get("input") if isinstance(inp, str): inputs = [inp] elif isinstance(inp, list): if inp and isinstance(inp[0], list): # token ids inputs = [st.tokenizer.decode(x) for x in inp] else: inputs = [str(x) for x in inp] else: return _err(400, "input must be a string or a list of strings", "INVALID_REQUEST", "invalid_request_error") dims = body.get("dimensions") enc = body.get("encoding_format", "float") if not st.gen_lock.acquire(timeout=float(st.args.queue_timeout)): return _err(503, "worker busy", "WORKER_BUSY") st.requests += 1 st.last_used = time.time() try: t0 = time.time() data = [] total = 0 for i, text in enumerate(inputs): vec, n = await asyncio.wrap_future(MLX_THREAD.submit(_embed_one, st, text, int(dims) if dims else None)) total += n if enc == "base64": import struct b = struct.pack(f"<{len(vec)}f", *vec) data.append({"object": "embedding", "index": i, "embedding": base64.b64encode(b).decode()}) else: data.append({"object": "embedding", "index": i, "embedding": vec}) return {"object": "list", "data": data, "model": body.get("model") or st.args.model_id, "usage": {"prompt_tokens": total, "total_tokens": total}, "timings": {"total_ms": round((time.time() - t0) * 1000, 1)}} except Exception as e: log.exception("embedding failed") return _err(500, f"{type(e).__name__}: {e}", "EMBEDDING_FAILED") finally: try: _mx().clear_cache() except Exception: pass st.gen_lock.release() @app.post("/v1/embeddings") async def embeddings(req: Request): return await embeddings_impl(await req.json()) def _rerank_score(st: WorkerState, query: str, doc: str, instruction: str | None) -> float: """Qwen3-Reranker style: P(yes) vs P(no) at the final position.""" mx = _mx() tok = st.tokenizer instr = instruction or "Given a web search query, retrieve relevant passages that answer the query" prefix = ("<|im_start|>system\nJudge whether the Document meets the requirements based on the Query and the " "Instruct provided. Note that the answer can only be \"yes\" or \"no\".<|im_end|>\n<|im_start|>user\n") suffix = "<|im_end|>\n<|im_start|>assistant\n\n\n\n\n" text = f"{prefix}: {instr}\n: {query}\n: {doc}{suffix}" ids = list(tok.encode(text))[: st.max_context] logits = st.model(mx.array([ids])) last = logits[0, -1, :].astype(mx.float32) yes_id = tok.encode("yes")[-1] no_id = tok.encode("no")[-1] pair = mx.stack([last[no_id], last[yes_id]]) probs = mx.softmax(pair) mx.eval(probs) return float(probs[1].item()) async def rerank_impl(body: dict): st = STATE assert st if st.status != "ready": return _err(503, f"model not ready ({st.status})", "MODEL_NOT_READY") query = body.get("query") docs = body.get("documents") or [] if not query or not isinstance(docs, list): return _err(400, "query and documents are required", "INVALID_REQUEST", "invalid_request_error") texts = [d if isinstance(d, str) else (d.get("text") or json.dumps(d)) for d in docs] top_n = body.get("top_n") or len(texts) if not st.gen_lock.acquire(timeout=float(st.args.queue_timeout)): return _err(503, "worker busy", "WORKER_BUSY") st.requests += 1 st.last_used = time.time() try: t0 = time.time() scores = [] for i, d in enumerate(texts): s = await asyncio.wrap_future(MLX_THREAD.submit(_rerank_score, st, query, d, body.get("instruction"))) scores.append({"index": i, "relevance_score": s, **({"document": {"text": d}} if body.get("return_documents") else {})}) scores.sort(key=lambda x: -x["relevance_score"]) return {"object": "list", "model": body.get("model") or st.args.model_id, "results": scores[: int(top_n)], "usage": {"total_tokens": 0}, "timings": {"total_ms": round((time.time() - t0) * 1000, 1)}} except Exception as e: log.exception("rerank failed") return _err(500, f"{type(e).__name__}: {e}", "RERANK_FAILED") finally: try: _mx().clear_cache() except Exception: pass st.gen_lock.release() @app.post("/v1/rerank") async def rerank(req: Request): return await rerank_impl(await req.json()) # --------------------------------------------------------------------------- # Entrypoint # --------------------------------------------------------------------------- def main(argv: list[str] | None = None) -> None: global STATE ap = argparse.ArgumentParser() ap.add_argument("--model-path", required=True) ap.add_argument("--model-id", required=True) ap.add_argument("--port", type=int, required=True) ap.add_argument("--host", default="127.0.0.1") ap.add_argument("--max-context", type=int, default=16384) ap.add_argument("--default-max-tokens", type=int, default=2048) ap.add_argument("--vision", action="store_true") ap.add_argument("--task", default="text", choices=["text", "embedding", "reranking"]) ap.add_argument("--kv-bits", type=int, default=0) ap.add_argument("--queue-timeout", type=float, default=600) ap.add_argument("--generation-timeout", type=float, default=1800) ap.add_argument("--memory-limit-gb", type=float, default=0) args = ap.parse_args(argv) logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s") STATE = WorkerState(args) if args.memory_limit_gb: try: _mx().set_memory_limit(int(args.memory_limit_gb * GB)) except Exception: pass MLX_THREAD.submit(_load_model, STATE) import uvicorn uvicorn.run(app, host=args.host, port=args.port, log_level="warning", access_log=False, timeout_keep_alive=75) if __name__ == "__main__": main()