SPB Git forge

spb/llm-api

Public
0commits 0branches 0releases
0 Bsize
maindefault branch
—last push
44.6 KB · 1,118 lines python
Raw Blame History
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