1"""Standalone MLX inference worker: one process = one model.23Started by the Model Manager:4 python -m llm_api.worker.mlx_worker --model-path P --port N --model-id ID [--vision] [--task text|embedding|reranking]56Exposes on 127.0.0.1:7 GET /health loading|ready|error + memory8 POST /v1/chat/completions OpenAI compatible (stream or not)9 POST /v1/completions10 POST /v1/embeddings11 POST /v1/rerank12 POST /tokenize13 POST /warmup14 POST /clear-cache15 POST /shutdown1617Killing the process is the memory-release mechanism: everything (weights, KV cache, Metal heaps)18goes away with it.19"""2021from __future__ import annotations2223import argparse24import asyncio25import base6426import json27import logging28import math29import os30import queue31import sys32import tempfile33import threading34import time35from pathlib import Path36from typing import Any3738import psutil39from fastapi import FastAPI, Request40from fastapi.responses import JSONResponse, StreamingResponse4142from .openai_types import (HARMONY_MARKERS, HARMONY_THINK_END, HARMONY_THINK_START, MarkerStripper, StopMatcher,43 ThinkSplitter, chat_chunk, new_id, normalize_messages, parse_tool_calls, sse)4445log = logging.getLogger("mlx_worker")46GB = 1024**34748# ---------------------------------------------------------------------------49# State50# ---------------------------------------------------------------------------515253class WorkerState:54 def __init__(self, args: argparse.Namespace):55 self.args = args56 self.status = "loading"57 self.error: str | None = None58 self.started = time.time()59 self.loaded_at: float | None = None60 self.load_ms: float | None = None61 self.model = None62 self.tokenizer = None63 self.processor = None # mlx_vlm64 self.config: dict = {}65 self.gen_config: dict = {}66 self.vision = bool(args.vision)67 self.task = args.task68 self.max_context = args.max_context69 self.gen_lock = threading.Lock()70 self.requests = 071 self.tokens_generated = 072 self.last_used = time.time()73 # prompt cache (single conversation)74 self.cache_tokens: list[int] = []75 self.cache_obj = None76 self.warm_ttft_ms: float | None = None77 self.template_text = ""78 self.harmony = False79 self.thinks = False808182class MLXThread:83 """All MLX work (import, load, generation, embeddings) runs on this single thread.84 MLX streams/command buffers are thread-affine and MLX is not thread-safe."""8586 def __init__(self):87 self._q: "queue.Queue[tuple]" = queue.Queue()88 self._th = threading.Thread(target=self._run, name="mlx", daemon=True)89 self._th.start()9091 def _run(self):92 while True:93 fn, args, fut = self._q.get()94 if fn is None:95 return96 try:97 res = fn(*args)98 if fut is not None:99 fut.set_result(res)100 except BaseException as e: # noqa: BLE001101 if fut is not None:102 fut.set_exception(e)103 else:104 log.exception("mlx thread task failed")105106 def submit(self, fn, *args):107 import concurrent.futures108 fut: concurrent.futures.Future = concurrent.futures.Future()109 self._q.put((fn, args, fut))110 return fut111112 def alive(self) -> bool:113 return self._th.is_alive()114115116MLX_THREAD = MLXThread()117STATE: WorkerState | None = None118app = FastAPI(title="llm-api mlx worker")119120121def _mx():122 import mlx.core as mx123 return mx124125126def memory_info() -> dict:127 out: dict[str, Any] = {"rss_gb": round(psutil.Process().memory_info().rss / GB, 3)}128 try:129 mx = _mx()130 out.update({131 "active_gb": round(mx.get_active_memory() / GB, 3),132 "peak_gb": round(mx.get_peak_memory() / GB, 3),133 "cache_gb": round(mx.get_cache_memory() / GB, 3),134 })135 except Exception:136 pass137 return out138139140# ---------------------------------------------------------------------------141# Loading142# ---------------------------------------------------------------------------143144145def _load_model(st: WorkerState) -> None:146 t0 = time.time()147 try:148 mx = _mx()149 path = st.args.model_path150 cfg_p = Path(path) / "config.json"151 if cfg_p.exists():152 st.config = json.loads(cfg_p.read_text())153 gc_p = Path(path) / "generation_config.json"154 if gc_p.exists():155 try:156 st.gen_config = json.loads(gc_p.read_text())157 except Exception:158 st.gen_config = {}159 tpl_p = Path(path) / "chat_template.jinja"160 if tpl_p.exists():161 st.template_text = tpl_p.read_text(errors="replace")162 else:163 tc_p = Path(path) / "tokenizer_config.json"164 if tc_p.exists():165 try:166 t = json.loads(tc_p.read_text()).get("chat_template") or ""167 st.template_text = t if isinstance(t, str) else json.dumps(t)168 except Exception:169 pass170 st.harmony = "<|channel|>" in st.template_text171 st.thinks = st.harmony or "<think>" in st.template_text or "enable_thinking" in st.template_text172 if st.vision:173 try:174 import mlx_vlm175 from mlx_vlm.utils import load_config176 st.model, st.processor = mlx_vlm.load(path)177 st.config = load_config(path)178 tok = getattr(st.processor, "tokenizer", st.processor)179 st.tokenizer = tok180 log.info("loaded with mlx_vlm")181 except Exception as e: # fall back to text-only182 log.warning("mlx_vlm load failed (%s); falling back to mlx_lm text-only", e)183 st.vision = False184 if not st.vision:185 from mlx_lm import load186 st.model, st.tokenizer = load(path)187 # Touch the weights so lazily-loaded parameters are materialized188 try:189 from mlx.utils import tree_flatten190 params = tree_flatten(st.model.parameters())191 mx.eval([p for _, p in params])192 except Exception:193 pass194 st.load_ms = (time.time() - t0) * 1000195 st.loaded_at = time.time()196 st.status = "ready"197 log.info("model ready in %.0f ms (%s)", st.load_ms, memory_info())198 except Exception as e:199 st.status = "error"200 st.error = f"{type(e).__name__}: {e}"201 log.exception("model load failed")202203204# ---------------------------------------------------------------------------205# Prompt building206# ---------------------------------------------------------------------------207208209def _apply_template(st: WorkerState, messages: list[dict], tools: list | None, kwargs: dict) -> list[int]:210 tok = st.tokenizer211 if st.vision and st.processor is not None:212 # handled by vision path213 raise RuntimeError("use vision path")214 has_tpl = getattr(tok, "has_chat_template", True)215 if has_tpl:216 try:217 out = tok.apply_chat_template(messages, tools=tools, add_generation_prompt=True, tokenize=True, **kwargs)218 except TypeError:219 out = tok.apply_chat_template(messages, add_generation_prompt=True, tokenize=True)220 if isinstance(out, dict):221 out = out.get("input_ids")222 if hasattr(out, "tolist"):223 out = out.tolist()224 if out and isinstance(out[0], list):225 out = out[0]226 return list(out)227 # No chat template: simple fallback228 text = ""229 for m in messages:230 text += f"{m.get('role','user').capitalize()}: {m.get('content','')}\n"231 text += "Assistant:"232 return list(tok.encode(text))233234235def _sampler_and_processors(body: dict, st: WorkerState):236 from mlx_lm.sample_utils import make_logits_processors, make_sampler237 gc = st.gen_config or {}238 temp = body.get("temperature")239 if temp is None:240 temp = gc.get("temperature", 0.7)241 top_p = body.get("top_p")242 if top_p is None:243 top_p = gc.get("top_p", 0.95)244 top_k = body.get("top_k")245 if top_k is None:246 top_k = gc.get("top_k", 0) or 0247 min_p = body.get("min_p", 0.0) or 0.0248 sampler = make_sampler(temp=float(temp), top_p=float(top_p) if top_p and top_p < 1.0 else 0.0,249 min_p=float(min_p), top_k=int(top_k))250 logit_bias = body.get("logit_bias")251 lb = None252 if isinstance(logit_bias, dict) and logit_bias:253 lb = {}254 for k, v in logit_bias.items():255 try:256 lb[int(k)] = float(v)257 except (TypeError, ValueError):258 pass259 rep = body.get("repetition_penalty")260 pres = body.get("presence_penalty") or None261 freq = body.get("frequency_penalty") or None262 procs = make_logits_processors(logit_bias=lb, repetition_penalty=float(rep) if rep else None,263 presence_penalty=float(pres) if pres else None,264 frequency_penalty=float(freq) if freq else None)265 return sampler, procs266267268def _max_tokens(body: dict, prompt_len: int, st: WorkerState) -> int:269 mt = body.get("max_completion_tokens") or body.get("max_tokens")270 if mt is None:271 mt = int(st.args.default_max_tokens)272 ctx = st.max_context273 room = ctx - prompt_len274 if room < 16:275 raise ValueError(f"prompt has {prompt_len} tokens; the loaded context window is {ctx} tokens.")276 return max(1, min(int(mt), room))277278279def _prepare_cache(st: WorkerState, prompt: list[int]):280 """Reuse KV cache for the common prefix of the previous conversation."""281 from mlx_lm.models.cache import can_trim_prompt_cache, make_prompt_cache, trim_prompt_cache282 if st.cache_obj is not None and st.cache_tokens and can_trim_prompt_cache(st.cache_obj):283 n = 0284 for a, b in zip(st.cache_tokens, prompt):285 if a != b:286 break287 n += 1288 # always leave at least one token to process289 n = min(n, len(prompt) - 1)290 if n > 0:291 to_trim = len(st.cache_tokens) - n292 if to_trim > 0:293 trim_prompt_cache(st.cache_obj, to_trim)294 st.cache_tokens = prompt[:n]295 return st.cache_obj, prompt[n:], n296 st.cache_obj = make_prompt_cache(st.model)297 st.cache_tokens = []298 return st.cache_obj, prompt, 0299300301# ---------------------------------------------------------------------------302# Generation (runs in a thread; yields events into a queue)303# ---------------------------------------------------------------------------304305306def _generate_thread(st: WorkerState, prompt: list[int], body: dict, chat: bool, out: "queue.Queue[dict]",307 cancel: threading.Event) -> None:308 mx = _mx()309 try:310 from mlx_lm import stream_generate311 if body.get("seed") is not None:312 mx.random.seed(int(body["seed"]))313 sampler, procs = _sampler_and_processors(body, st)314 max_tokens = _max_tokens(body, len(prompt), st)315 cache, rest, cached = _prepare_cache(st, prompt)316 t0 = time.time()317 first = None318 n_gen = 0319 gen_tokens: list[int] = []320 prompt_tps = 0.0321 gen_tps = 0.0322 peak = 0.0323 finish = "length"324 eos_ids = getattr(st.tokenizer, "eos_token_ids", None) or set()325 kwargs: dict[str, Any] = {}326 if st.args.kv_bits:327 kwargs["kv_bits"] = int(st.args.kv_bits)328 kwargs["quantized_kv_start"] = 4096329 for g in stream_generate(st.model, st.tokenizer, rest, max_tokens=max_tokens, sampler=sampler,330 logits_processors=procs, prompt_cache=cache, prefill_step_size=2048, **kwargs):331 if first is None:332 first = time.time()333 n_gen += 1334 gen_tokens.append(g.token)335 prompt_tps = g.prompt_tps336 gen_tps = g.generation_tps337 peak = g.peak_memory338 if g.finish_reason:339 finish = g.finish_reason340 out.put({"text": g.text, "token": g.token})341 if cancel.is_set():342 finish = "cancelled"343 break344 if g.finish_reason:345 break346 t1 = time.time()347 # Remember exactly the tokens the KV cache holds (the last sampled token is never fed back)348 all_tokens = list(prompt) + gen_tokens349 try:350 off = int(cache[0].offset)351 except Exception:352 off = len(all_tokens) - 1353 st.cache_tokens = all_tokens[:max(0, min(off, len(all_tokens)))]354 st.tokens_generated += n_gen355 out.put({356 "done": True, "finish_reason": finish, "prompt_tokens": len(prompt), "completion_tokens": n_gen,357 "cached_tokens": cached,358 "timings": {359 "ttft_ms": round(((first or t1) - t0) * 1000, 1),360 "prompt_ms": round(((first or t1) - t0) * 1000, 1),361 "generation_ms": round((t1 - (first or t1)) * 1000, 1),362 "total_ms": round((t1 - t0) * 1000, 1),363 "prompt_tps": round(prompt_tps, 1), "generation_tps": round(gen_tps, 2),364 "peak_memory_gb": round(peak, 3),365 },366 })367 except Exception as e:368 log.exception("generation failed")369 # A failed generation may leave the cache inconsistent370 st.cache_obj = None371 st.cache_tokens = []372 out.put({"error": f"{type(e).__name__}: {e}"})373 finally:374 try:375 mx.clear_cache()376 except Exception:377 pass378379380def _vision_generate_thread(st: WorkerState, messages: list[dict], images: list[dict], body: dict, tk: dict,381 tools: list | None, out: "queue.Queue[dict]", cancel: threading.Event) -> None:382 mx = _mx()383 tmpfiles: list[str] = []384 try:385 from mlx_vlm import stream_generate as vlm_stream386 from mlx_vlm.prompt_utils import apply_chat_template387 paths: list[str] = []388 for part in images:389 url = part.get("image_url", {}).get("url") if isinstance(part.get("image_url"), dict) else part.get("image_url") or part.get("url")390 if not url:391 continue392 if url.startswith("data:"):393 header, b64 = url.split(",", 1)394 ext = ".png" if "png" in header else ".jpg"395 fd, p = tempfile.mkstemp(suffix=ext)396 with os.fdopen(fd, "wb") as f:397 f.write(base64.b64decode(b64))398 tmpfiles.append(p)399 paths.append(p)400 else:401 paths.append(url)402 tpl_kwargs = dict(tk or {})403 if tools:404 tpl_kwargs["tools"] = tools405 try:406 prompt = apply_chat_template(st.processor, st.config, messages, num_images=len(paths), **tpl_kwargs)407 except TypeError:408 prompt = apply_chat_template(st.processor, st.config, messages, num_images=len(paths))409 if body.get("seed") is not None:410 mx.random.seed(int(body["seed"]))411 mt = body.get("max_completion_tokens") or body.get("max_tokens") or int(st.args.default_max_tokens)412 gc = st.gen_config or {}413 temp = body.get("temperature")414 if temp is None:415 temp = gc.get("temperature", 0.7)416 top_p = body.get("top_p")417 if top_p is None:418 top_p = gc.get("top_p", 0.95)419 gen_kwargs: dict[str, Any] = {"max_tokens": int(mt), "temperature": float(temp), "top_p": float(top_p)}420 top_k = body.get("top_k", gc.get("top_k"))421 if top_k:422 gen_kwargs["top_k"] = int(top_k)423 if body.get("min_p"):424 gen_kwargs["min_p"] = float(body["min_p"])425 for k in ("repetition_penalty", "presence_penalty", "frequency_penalty"):426 if body.get(k):427 gen_kwargs[k] = float(body[k])428 if isinstance(body.get("logit_bias"), dict) and body["logit_bias"]:429 try:430 gen_kwargs["logit_bias"] = {int(k): float(v) for k, v in body["logit_bias"].items()}431 except (TypeError, ValueError):432 pass433 if st.args.kv_bits:434 gen_kwargs["kv_bits"] = int(st.args.kv_bits)435 t0 = time.time()436 first = None437 n = 0438 ptoks = 0439 ptps = gtps = peak = 0.0440 finish = "length"441 for g in vlm_stream(st.model, st.processor, prompt, image=paths or None, **gen_kwargs):442 if first is None:443 first = time.time()444 n += 1445 ptoks = getattr(g, "prompt_tokens", ptoks)446 ptps = getattr(g, "prompt_tps", ptps)447 gtps = getattr(g, "generation_tps", gtps)448 peak = getattr(g, "peak_memory", peak)449 fr = getattr(g, "finish_reason", None)450 if fr:451 finish = fr452 out.put({"text": g.text, "token": getattr(g, "token", 0)})453 if cancel.is_set():454 finish = "cancelled"455 break456 if fr:457 break458 t1 = time.time()459 st.tokens_generated += n460 out.put({"done": True, "finish_reason": finish, "prompt_tokens": int(ptoks), "completion_tokens": n,461 "cached_tokens": 0,462 "timings": {"ttft_ms": round(((first or t1) - t0) * 1000, 1), "total_ms": round((t1 - t0) * 1000, 1),463 "prompt_tps": round(ptps, 1), "generation_tps": round(gtps, 2),464 "peak_memory_gb": round(peak, 3)}})465 except Exception as e:466 log.exception("vision generation failed")467 out.put({"error": f"{type(e).__name__}: {e}"})468 finally:469 for p in tmpfiles:470 try:471 os.unlink(p)472 except OSError:473 pass474 try:475 mx.clear_cache()476 except Exception:477 pass478479480class _Task:481 """Thread-like wrapper around a Future running on the MLX thread."""482483 def __init__(self, fut):484 self.fut = fut485486 def is_alive(self) -> bool:487 return not self.fut.done()488489 def join(self, timeout: float | None = None) -> None:490 try:491 self.fut.result(timeout=timeout)492 except Exception:493 pass494495496async def _run_generation(st: WorkerState, target, *args) -> tuple["queue.Queue[dict]", threading.Event, _Task]:497 out: queue.Queue[dict] = queue.Queue()498 cancel = threading.Event()499 fut = MLX_THREAD.submit(target, st, *args, out, cancel)500 return out, cancel, _Task(fut)501502503async def _next_event(q: "queue.Queue[dict]", th: "_Task | None" = None) -> dict:504 """Next event from the generation thread; detects a dead thread instead of waiting forever."""505 loop = asyncio.get_running_loop()506507 def _get():508 while True:509 try:510 return q.get(timeout=1.0)511 except queue.Empty:512 if th is not None and not th.is_alive():513 try:514 return q.get_nowait()515 except queue.Empty:516 return {"error": "generation thread exited unexpectedly"}517518 return await loop.run_in_executor(None, _get)519520521# ---------------------------------------------------------------------------522# Endpoints523# ---------------------------------------------------------------------------524525526def _err(status: int, message: str, code: str = "WORKER_ERROR", etype: str = "runtime_error") -> JSONResponse:527 return JSONResponse(status_code=status, content={"error": {"message": message, "type": etype, "code": code}})528529530@app.get("/health")531async def health():532 st = STATE533 assert st534 return {535 "status": st.status, "error": st.error, "model_id": st.args.model_id, "runtime": "mlx",536 "vision": st.vision, "task": st.task, "elapsed_seconds": round(time.time() - st.started, 1),537 "load_ms": st.load_ms, "memory": memory_info(), "requests": st.requests, "tokens_generated": st.tokens_generated,538 "max_context": st.max_context, "pid": os.getpid(), "busy": st.gen_lock.locked(),539 "warm_ttft_ms": st.warm_ttft_ms, "mlx_thread_alive": MLX_THREAD.alive(),540 }541542543@app.post("/clear-cache")544async def clear_cache():545 st = STATE546 assert st547 st.cache_obj = None548 st.cache_tokens = []549 try:550 _mx().clear_cache()551 except Exception:552 pass553 return memory_info()554555556@app.post("/shutdown")557async def shutdown():558 async def _exit():559 await asyncio.sleep(0.2)560 os._exit(0)561 asyncio.create_task(_exit())562 return {"ok": True}563564565@app.post("/tokenize")566async def tokenize(req: Request):567 st = STATE568 assert st569 body = await req.json()570 if st.status != "ready":571 return _err(503, "model not ready", "MODEL_NOT_READY")572 if "messages" in body:573 msgs, _ = normalize_messages(body["messages"])574 try:575 toks = _apply_template(st, msgs, body.get("tools"), body.get("chat_template_kwargs") or {})576 except RuntimeError:577 toks = list(st.tokenizer.encode(" ".join(m.get("content", "") for m in msgs)))578 else:579 toks = list(st.tokenizer.encode(body.get("text", "") or body.get("prompt", "")))580 return {"tokens": len(toks), "max_context": st.max_context}581582583@app.post("/warmup")584async def warmup():585 st = STATE586 assert st587 if st.status != "ready":588 return _err(503, "model not ready", "MODEL_NOT_READY")589 t0 = time.time()590 if st.task in ("embedding",):591 r = await embeddings_impl({"input": "warm up"})592 st.warm_ttft_ms = round((time.time() - t0) * 1000, 1)593 return {"ok": True, "ttft_ms": st.warm_ttft_ms, "kind": "embedding", "dims": len(r["data"][0]["embedding"])}594 if st.task == "reranking":595 r = await rerank_impl({"query": "warm", "documents": ["warm up"]})596 st.warm_ttft_ms = round((time.time() - t0) * 1000, 1)597 return {"ok": True, "ttft_ms": st.warm_ttft_ms, "kind": "rerank"}598 body = {"messages": [{"role": "user", "content": "Say OK."}], "max_tokens": 4, "temperature": 0.0}599 resp = await chat_impl(body, stream=False)600 if isinstance(resp, JSONResponse):601 return resp602 st.warm_ttft_ms = resp.get("timings", {}).get("ttft_ms")603 # do not keep the warm-up prompt in the cache604 st.cache_obj = None605 st.cache_tokens = []606 return {"ok": True, "ttft_ms": st.warm_ttft_ms, "text": resp["choices"][0]["message"]["content"]}607608609async def chat_impl(body: dict, stream: bool):610 st = STATE611 assert st612 if st.status != "ready":613 return _err(503, f"model not ready ({st.status})", "MODEL_NOT_READY")614 messages = body.get("messages")615 if not isinstance(messages, list) or not messages:616 return _err(400, "messages is required", "INVALID_REQUEST", "invalid_request_error")617 msgs, images = normalize_messages(messages)618 tools = body.get("tools") or None619 tk = dict(body.get("chat_template_kwargs") or {})620 # reasoning controls (Qwen3-style enable_thinking)621 reasoning = body.get("reasoning")622 if isinstance(reasoning, dict) and "effort" in reasoning:623 tk.setdefault("enable_thinking", reasoning["effort"] not in ("none", "minimal"))624 if body.get("enable_thinking") is not None:625 tk["enable_thinking"] = bool(body["enable_thinking"])626 if body.get("reasoning_effort") is not None:627 tk.setdefault("enable_thinking", body["reasoning_effort"] not in ("none", "minimal"))628 stops = body.get("stop") or []629 if isinstance(stops, str):630 stops = [stops]631 rid = new_id("chatcmpl")632 created = int(time.time())633 model_name = body.get("model") or st.args.model_id634635 if not st.gen_lock.acquire(timeout=float(st.args.queue_timeout)):636 return _err(503, "worker busy", "WORKER_BUSY")637 st.requests += 1638 st.last_used = time.time()639 try:640 if st.vision and st.processor is not None:641 # mlx-vlm handles both text-only and image requests for vision-language models642 q, cancel, th = await _run_generation(st, _vision_generate_thread, msgs, images, body, tk, tools)643 think_start, think_end = "<think>", "</think>"644 thinking_enabled = st.thinks and tk.get("enable_thinking", True)645 else:646 if images:647 return _err(400, "This model does not accept images.", "VISION_UNSUPPORTED", "invalid_request_error")648 try:649 prompt = _apply_template(st, msgs, tools, tk)650 except Exception as e:651 return _err(400, f"chat template error: {e}", "TEMPLATE_ERROR", "invalid_request_error")652 if len(prompt) >= st.max_context - 16:653 return _err(400, f"Prompt has {len(prompt)} tokens but the context window is {st.max_context}.",654 "CONTEXT_TOO_LARGE", "invalid_request_error")655 q, cancel, th = await _run_generation(st, _generate_thread, prompt, body, True)656 think_start = getattr(st.tokenizer, "think_start", None) or "<think>"657 think_end = getattr(st.tokenizer, "think_end", None) or "</think>"658 thinking_enabled = (bool(getattr(st.tokenizer, "has_thinking", False)) or st.thinks) and tk.get("enable_thinking", True)659 markers: list[str] = []660 if st.harmony:661 think_start, think_end = HARMONY_THINK_START, HARMONY_THINK_END662 thinking_enabled = True663 markers = HARMONY_MARKERS664 splitter = ThinkSplitter(think_start, think_end) if thinking_enabled else ThinkSplitter(None, None)665 stripper = MarkerStripper(markers)666 stopper = StopMatcher(list(stops))667 has_tools = bool(tools)668669 if stream:670 async def gen():671 try:672 yield sse(chat_chunk(rid, model_name, created, {"role": "assistant", "content": ""}))673 content_acc = ""674 reasoning_acc = ""675 finish = "stop"676 final: dict = {}677 while True:678 ev = await _next_event(q, th)679 if "error" in ev:680 yield sse({"error": {"message": ev["error"], "type": "runtime_error", "code": "GENERATION_FAILED"}})681 yield sse("[DONE]")682 return683 if ev.get("done"):684 final = ev685 finish = ev["finish_reason"]686 break687 r, c = splitter.feed(ev["text"])688 if r:689 reasoning_acc += r690 yield sse(chat_chunk(rid, model_name, created, {"reasoning_content": r}))691 c = stripper.feed(c) if c else c692 if c:693 c = stopper.feed(c)694 if c and not has_tools:695 content_acc += c696 yield sse(chat_chunk(rid, model_name, created, {"content": c}))697 elif c:698 content_acc += c699 if stopper.done:700 cancel.set()701 finish = "stop"702 # drain703 while True:704 ev2 = await _next_event(q, th)705 if ev2.get("done") or "error" in ev2:706 final = ev2 if ev2.get("done") else {}707 break708 break709 r, c = splitter.flush()710 c = (stripper.feed(c) + stripper.flush()) if not stopper.done else ""711 c = (stopper.flush() + c) if not stopper.done else ""712 if r:713 yield sse(chat_chunk(rid, model_name, created, {"reasoning_content": r}))714 tool_calls: list[dict] = []715 if has_tools:716 rest, tool_calls = parse_tool_calls(content_acc + c)717 if tool_calls:718 finish = "tool_calls"719 yield sse(chat_chunk(rid, model_name, created, {"tool_calls": [720 {"index": i, **tc} for i, tc in enumerate(tool_calls)]}))721 elif content_acc + c:722 yield sse(chat_chunk(rid, model_name, created, {"content": content_acc + c}))723 elif c:724 yield sse(chat_chunk(rid, model_name, created, {"content": c}))725 if finish == "cancelled":726 finish = "stop"727 if finish not in ("stop", "length", "tool_calls"):728 finish = "stop"729 usage = {"prompt_tokens": final.get("prompt_tokens", 0), "completion_tokens": final.get("completion_tokens", 0),730 "total_tokens": final.get("prompt_tokens", 0) + final.get("completion_tokens", 0)}731 if final.get("cached_tokens"):732 usage["prompt_tokens_details"] = {"cached_tokens": final["cached_tokens"]}733 yield sse(chat_chunk(rid, model_name, created, {}, finish, usage, {"timings": final.get("timings", {})}))734 yield sse("[DONE]")735 finally:736 cancel.set()737 th.join(timeout=float(st.args.generation_timeout))738 st.gen_lock.release()739 return StreamingResponse(gen(), media_type="text/event-stream",740 headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"})741742 # non-streaming743 try:744 content_acc = ""745 reasoning_acc = ""746 final = {}747 finish = "stop"748 while True:749 ev = await _next_event(q, th)750 if "error" in ev:751 return _err(500, ev["error"], "GENERATION_FAILED")752 if ev.get("done"):753 final = ev754 finish = ev["finish_reason"]755 break756 r, c = splitter.feed(ev["text"])757 reasoning_acc += r758 c = stripper.feed(c) if c else c759 if c:760 c = stopper.feed(c)761 content_acc += c762 if stopper.done:763 cancel.set()764 finish = "stop"765 while True:766 ev2 = await _next_event(q, th)767 if ev2.get("done") or "error" in ev2:768 final = ev2 if ev2.get("done") else {}769 break770 break771 r, c = splitter.flush()772 reasoning_acc += r773 if not stopper.done:774 content_acc += stopper.flush() + stripper.feed(c) + stripper.flush()775 content_acc = content_acc.strip("\n") if reasoning_acc else content_acc776 tool_calls: list[dict] = []777 if has_tools:778 content_acc, tool_calls = parse_tool_calls(content_acc)779 if tool_calls:780 finish = "tool_calls"781 if finish not in ("stop", "length", "tool_calls"):782 finish = "stop"783 msg: dict[str, Any] = {"role": "assistant", "content": content_acc if not tool_calls else (content_acc or None)}784 if reasoning_acc.strip():785 msg["reasoning_content"] = reasoning_acc.strip()786 if tool_calls:787 msg["tool_calls"] = tool_calls788 usage = {"prompt_tokens": final.get("prompt_tokens", 0), "completion_tokens": final.get("completion_tokens", 0),789 "total_tokens": final.get("prompt_tokens", 0) + final.get("completion_tokens", 0)}790 if final.get("cached_tokens"):791 usage["prompt_tokens_details"] = {"cached_tokens": final["cached_tokens"]}792 return {"id": rid, "object": "chat.completion", "created": created, "model": model_name,793 "choices": [{"index": 0, "message": msg, "logprobs": None, "finish_reason": finish}],794 "usage": usage, "timings": final.get("timings", {})}795 finally:796 cancel.set()797 th.join(timeout=float(st.args.generation_timeout))798 st.gen_lock.release()799 except Exception:800 if st.gen_lock.locked():801 try:802 st.gen_lock.release()803 except RuntimeError:804 pass805 raise806807808@app.post("/v1/chat/completions")809async def chat_completions(req: Request):810 body = await req.json()811 return await chat_impl(body, bool(body.get("stream")))812813814@app.post("/v1/completions")815async def completions(req: Request):816 st = STATE817 assert st818 body = await req.json()819 if st.status != "ready":820 return _err(503, f"model not ready ({st.status})", "MODEL_NOT_READY")821 prompt = body.get("prompt", "")822 if isinstance(prompt, list):823 prompt = prompt[0] if prompt and isinstance(prompt[0], str) else ""824 stream = bool(body.get("stream"))825 rid = new_id("cmpl")826 created = int(time.time())827 model_name = body.get("model") or st.args.model_id828 toks = list(st.tokenizer.encode(prompt))829 if len(toks) >= st.max_context - 16:830 return _err(400, f"Prompt has {len(toks)} tokens but the context window is {st.max_context}.",831 "CONTEXT_TOO_LARGE", "invalid_request_error")832 stops = body.get("stop") or []833 if isinstance(stops, str):834 stops = [stops]835 if not st.gen_lock.acquire(timeout=float(st.args.queue_timeout)):836 return _err(503, "worker busy", "WORKER_BUSY")837 st.requests += 1838 st.last_used = time.time()839 # completions never reuse the chat cache840 st.cache_obj = None841 st.cache_tokens = []842 q, cancel, th = await _run_generation(st, _generate_thread, toks, body, False)843 stopper = StopMatcher(list(stops))844 echo = bool(body.get("echo"))845846 def chunk(text: str, finish=None, usage=None, extra=None):847 d: dict[str, Any] = {"id": rid, "object": "text_completion", "created": created, "model": model_name,848 "choices": [{"index": 0, "text": text, "logprobs": None, "finish_reason": finish}]}849 if usage:850 d["usage"] = usage851 if extra:852 d.update(extra)853 return d854855 if stream:856 async def gen():857 try:858 if echo:859 yield sse(chunk(prompt))860 finish = "stop"861 final: dict = {}862 while True:863 ev = await _next_event(q, th)864 if "error" in ev:865 yield sse({"error": {"message": ev["error"], "type": "runtime_error", "code": "GENERATION_FAILED"}})866 break867 if ev.get("done"):868 final, finish = ev, ev["finish_reason"]869 break870 c = stopper.feed(ev["text"])871 if c:872 yield sse(chunk(c))873 if stopper.done:874 cancel.set()875 finish = "stop"876 while True:877 ev2 = await _next_event(q, th)878 if ev2.get("done") or "error" in ev2:879 final = ev2 if ev2.get("done") else {}880 break881 break882 tail = stopper.flush() if not stopper.done else ""883 if tail:884 yield sse(chunk(tail))885 usage = {"prompt_tokens": final.get("prompt_tokens", 0), "completion_tokens": final.get("completion_tokens", 0),886 "total_tokens": final.get("prompt_tokens", 0) + final.get("completion_tokens", 0)}887 yield sse(chunk("", finish if finish in ("stop", "length") else "stop", usage, {"timings": final.get("timings", {})}))888 yield sse("[DONE]")889 finally:890 cancel.set()891 th.join(timeout=float(st.args.generation_timeout))892 st.cache_obj = None893 st.cache_tokens = []894 st.gen_lock.release()895 return StreamingResponse(gen(), media_type="text/event-stream", headers={"Cache-Control": "no-cache"})896 try:897 text = ""898 final = {}899 finish = "stop"900 while True:901 ev = await _next_event(q, th)902 if "error" in ev:903 return _err(500, ev["error"], "GENERATION_FAILED")904 if ev.get("done"):905 final, finish = ev, ev["finish_reason"]906 break907 text += stopper.feed(ev["text"])908 if stopper.done:909 cancel.set()910 finish = "stop"911 while True:912 ev2 = await _next_event(q, th)913 if ev2.get("done") or "error" in ev2:914 final = ev2 if ev2.get("done") else {}915 break916 break917 if not stopper.done:918 text += stopper.flush()919 usage = {"prompt_tokens": final.get("prompt_tokens", 0), "completion_tokens": final.get("completion_tokens", 0),920 "total_tokens": final.get("prompt_tokens", 0) + final.get("completion_tokens", 0)}921 return chunk((prompt if echo else "") + text, finish if finish in ("stop", "length") else "stop", usage,922 {"timings": final.get("timings", {})})923 finally:924 cancel.set()925 th.join(timeout=float(st.args.generation_timeout))926 st.cache_obj = None927 st.cache_tokens = []928 st.gen_lock.release()929930931# ---------------------------------------------------------------------------932# Embeddings / rerank (causal-LM style: last-token pooling, Qwen3-Embedding & co.)933# ---------------------------------------------------------------------------934935936def _hidden_states(st: WorkerState, tokens: list[int]):937 mx = _mx()938 inner = getattr(st.model, "model", None) or getattr(st.model, "language_model", None)939 x = mx.array([tokens])940 if inner is not None and callable(inner):941 h = inner(x)942 else:943 h = st.model(x)944 mx.eval(h)945 return h946947948def _embed_one(st: WorkerState, text: str, dims: int | None) -> tuple[list[float], int]:949 mx = _mx()950 tok = st.tokenizer951 ids = list(tok.encode(text))952 eos = getattr(tok, "eos_token_id", None)953 if eos is None:954 eos_ids = getattr(tok, "eos_token_ids", None) or set()955 eos = next(iter(eos_ids), None)956 if eos is not None and (not ids or ids[-1] != eos):957 ids.append(eos)958 ids = ids[: st.max_context]959 h = _hidden_states(st, ids)960 v = h[0, -1, :].astype(mx.float32)961 if dims and dims < v.shape[0]:962 v = v[:dims]963 norm = mx.sqrt(mx.sum(v * v)) + 1e-12964 v = v / norm965 mx.eval(v)966 return v.tolist(), len(ids)967968969async def embeddings_impl(body: dict):970 st = STATE971 assert st972 if st.status != "ready":973 return _err(503, f"model not ready ({st.status})", "MODEL_NOT_READY")974 inp = body.get("input")975 if isinstance(inp, str):976 inputs = [inp]977 elif isinstance(inp, list):978 if inp and isinstance(inp[0], list): # token ids979 inputs = [st.tokenizer.decode(x) for x in inp]980 else:981 inputs = [str(x) for x in inp]982 else:983 return _err(400, "input must be a string or a list of strings", "INVALID_REQUEST", "invalid_request_error")984 dims = body.get("dimensions")985 enc = body.get("encoding_format", "float")986 if not st.gen_lock.acquire(timeout=float(st.args.queue_timeout)):987 return _err(503, "worker busy", "WORKER_BUSY")988 st.requests += 1989 st.last_used = time.time()990 try:991 t0 = time.time()992 data = []993 total = 0994 for i, text in enumerate(inputs):995 vec, n = await asyncio.wrap_future(MLX_THREAD.submit(_embed_one, st, text, int(dims) if dims else None))996 total += n997 if enc == "base64":998 import struct999 b = struct.pack(f"<{len(vec)}f", *vec)1000 data.append({"object": "embedding", "index": i, "embedding": base64.b64encode(b).decode()})1001 else:1002 data.append({"object": "embedding", "index": i, "embedding": vec})1003 return {"object": "list", "data": data, "model": body.get("model") or st.args.model_id,1004 "usage": {"prompt_tokens": total, "total_tokens": total},1005 "timings": {"total_ms": round((time.time() - t0) * 1000, 1)}}1006 except Exception as e:1007 log.exception("embedding failed")1008 return _err(500, f"{type(e).__name__}: {e}", "EMBEDDING_FAILED")1009 finally:1010 try:1011 _mx().clear_cache()1012 except Exception:1013 pass1014 st.gen_lock.release()101510161017@app.post("/v1/embeddings")1018async def embeddings(req: Request):1019 return await embeddings_impl(await req.json())102010211022def _rerank_score(st: WorkerState, query: str, doc: str, instruction: str | None) -> float:1023 """Qwen3-Reranker style: P(yes) vs P(no) at the final position."""1024 mx = _mx()1025 tok = st.tokenizer1026 instr = instruction or "Given a web search query, retrieve relevant passages that answer the query"1027 prefix = ("<|im_start|>system\nJudge whether the Document meets the requirements based on the Query and the "1028 "Instruct provided. Note that the answer can only be \"yes\" or \"no\".<|im_end|>\n<|im_start|>user\n")1029 suffix = "<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n"1030 text = f"{prefix}<Instruct>: {instr}\n<Query>: {query}\n<Document>: {doc}{suffix}"1031 ids = list(tok.encode(text))[: st.max_context]1032 logits = st.model(mx.array([ids]))1033 last = logits[0, -1, :].astype(mx.float32)1034 yes_id = tok.encode("yes")[-1]1035 no_id = tok.encode("no")[-1]1036 pair = mx.stack([last[no_id], last[yes_id]])1037 probs = mx.softmax(pair)1038 mx.eval(probs)1039 return float(probs[1].item())104010411042async def rerank_impl(body: dict):1043 st = STATE1044 assert st1045 if st.status != "ready":1046 return _err(503, f"model not ready ({st.status})", "MODEL_NOT_READY")1047 query = body.get("query")1048 docs = body.get("documents") or []1049 if not query or not isinstance(docs, list):1050 return _err(400, "query and documents are required", "INVALID_REQUEST", "invalid_request_error")1051 texts = [d if isinstance(d, str) else (d.get("text") or json.dumps(d)) for d in docs]1052 top_n = body.get("top_n") or len(texts)1053 if not st.gen_lock.acquire(timeout=float(st.args.queue_timeout)):1054 return _err(503, "worker busy", "WORKER_BUSY")1055 st.requests += 11056 st.last_used = time.time()1057 try:1058 t0 = time.time()1059 scores = []1060 for i, d in enumerate(texts):1061 s = await asyncio.wrap_future(MLX_THREAD.submit(_rerank_score, st, query, d, body.get("instruction")))1062 scores.append({"index": i, "relevance_score": s, **({"document": {"text": d}} if body.get("return_documents") else {})})1063 scores.sort(key=lambda x: -x["relevance_score"])1064 return {"object": "list", "model": body.get("model") or st.args.model_id, "results": scores[: int(top_n)],1065 "usage": {"total_tokens": 0}, "timings": {"total_ms": round((time.time() - t0) * 1000, 1)}}1066 except Exception as e:1067 log.exception("rerank failed")1068 return _err(500, f"{type(e).__name__}: {e}", "RERANK_FAILED")1069 finally:1070 try:1071 _mx().clear_cache()1072 except Exception:1073 pass1074 st.gen_lock.release()107510761077@app.post("/v1/rerank")1078async def rerank(req: Request):1079 return await rerank_impl(await req.json())108010811082# ---------------------------------------------------------------------------1083# Entrypoint1084# ---------------------------------------------------------------------------108510861087def main(argv: list[str] | None = None) -> None:1088 global STATE1089 ap = argparse.ArgumentParser()1090 ap.add_argument("--model-path", required=True)1091 ap.add_argument("--model-id", required=True)1092 ap.add_argument("--port", type=int, required=True)1093 ap.add_argument("--host", default="127.0.0.1")1094 ap.add_argument("--max-context", type=int, default=16384)1095 ap.add_argument("--default-max-tokens", type=int, default=2048)1096 ap.add_argument("--vision", action="store_true")1097 ap.add_argument("--task", default="text", choices=["text", "embedding", "reranking"])1098 ap.add_argument("--kv-bits", type=int, default=0)1099 ap.add_argument("--queue-timeout", type=float, default=600)1100 ap.add_argument("--generation-timeout", type=float, default=1800)1101 ap.add_argument("--memory-limit-gb", type=float, default=0)1102 args = ap.parse_args(argv)1103 logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s")1104 STATE = WorkerState(args)1105 if args.memory_limit_gb:1106 try:1107 _mx().set_memory_limit(int(args.memory_limit_gb * GB))1108 except Exception:1109 pass1110 MLX_THREAD.submit(_load_model, STATE)1111 import uvicorn1112 uvicorn.run(app, host=args.host, port=args.port, log_level="warning", access_log=False,1113 timeout_keep_alive=75)111411151116if __name__ == "__main__":1117 main()1118