Ines-1 / scripts /serve_jev.py
Endikavi's picture
Ines-1 RC1 (private staging; release commit f198669): repository root without resolving symlinks
bc5b068 verified
Raw History Blame Contribute Delete
12.1 kB
"""Ines-1 (codename mini-v41) behind the typed-decision `/v1/systemone` HTTP interface, on the fast prefill path.
python scripts/serve_jev.py [--checkpoint .] [--port 30040] [--max-rows 64] [--max-tokens 65536] [--graphs 64] [--graph-tokens 6144]
POST /v1/systemone {state, questions, lang?: "es"|"en"} -> {model, answers, usage}
GET /v1/models, GET /health
Every question is one prompt (decisions.Reader, the same prompt the model was trained and
tested with); the answer is the softmax over the option letters at the last position.
One thread owns the GPU. Whatever arrives while it works goes into the next call, so a lone
request is not delayed and concurrent ones share a forward. Each call:
- groups the waiting prompts by length (the oldest one plus its length neighbours while
padding stays <= 30 %, up to --max-rows rows and --max-tokens padded tokens), as serve.Engine
does for admissions;
- a group whose (batch, length) bucket holds <= --graph-tokens padded tokens replays a
PrefillGraphs graph (launch-bound regime); a larger one runs the exact eager `prefill_batch`
(GPU-bound regime, where padding up to a bucket is waste);
- only the logits at each row's last real position are computed (`_prefill_forward`: the
decoder runs over each row's last R positions).
The HTTP side is asyncio (aiohttp). Prompts are built in worker processes (tokenizer only):
in threads they held the GIL the GPU thread needs to launch ~2,000 kernels per forward.
"""
from __future__ import annotations
import argparse
import asyncio
import json
import queue
import sys
import threading
import time
from concurrent.futures import Future, ProcessPoolExecutor
from pathlib import Path
import torch
sys.path.insert(0, str(Path(__file__).absolute().parent.parent)) # the repository root (not resolve(): HF cache symlinks)
from mini_v41_jev.decisions import Reader, options # noqa: E402,F401
class Batcher:
def __init__(self, im, reader, max_rows, max_tokens, graph_rows, graph_tokens=4096, max_padding=0.3):
from mini_v41.fast_decode import PrefillGraphs, prefill_batch
self.model, self.reader = im.model, reader
self.prefill_batch = prefill_batch
self.max_rows, self.max_tokens, self.max_padding = max_rows, max_tokens, max_padding
self.graph_tokens = graph_tokens
self.graphs = {} # batch bucket -> PrefillGraphs over the lengths that keep b * n <= graph_tokens
# A forward is launch-bound below a few thousand padded tokens (eager ~20 ms whatever the
# batch, graph 8.6 ms at 8x128) and GPU-bound above (graph == eager at 16x1024): graphs
# only for those small buckets, one set per batch size so no big bucket is ever captured.
for b in (1, 2, 4, 8, 16, 32, 64):
lengths = tuple(n for n in (32, 64, 96, 128, 192, 256, 384, 512, 1024, 2048) if b * n <= graph_tokens)
if b <= graph_rows and lengths:
with torch.autocast("cuda", dtype=torch.bfloat16): # canonical: bf16 weights + bf16 autocast
self.graphs[b] = PrefillGraphs(im.model, lengths=lengths, batches=(b,))
self.letters = torch.tensor(reader.letter_ids(26), device=im.device)
self.q = queue.Queue()
self.stats = {"calls": 0, "rows": 0, "graph_calls": 0, "busy_s": 0.0, "gpu_s": 0.0, "real_tokens": 0, "padded_tokens": 0}
threading.Thread(target=self._loop, name="gpu", daemon=True).start()
def submit(self, rows):
"""rows: [(ids, n_options)] -> Future of [probabilities over the options]."""
fut = Future()
self.q.put((rows, fut))
return fut
def _groups(self, pending):
"""pending: [(ids, n, slot)] in arrival order -> list of groups, oldest first."""
out = []
left = sorted(pending, key=lambda r: len(r[0]))
while left:
oldest = min(left, key=lambda r: r[2])
i = left.index(oldest)
lo = hi = i
width = len(oldest[0])
def fits(j):
w = max(width, len(left[j][0]))
rows = hi - lo + 2
real = sum(len(left[k][0]) for k in range(lo, hi + 1)) + len(left[j][0])
return rows <= self.max_rows and w * rows <= self.max_tokens and real >= (1 - self.max_padding) * w * rows
while True:
cand = [j for j in (lo - 1, hi + 1) if 0 <= j < len(left)]
cand.sort(key=lambda j: abs(len(left[j][0]) - len(oldest[0])))
j = next((j for j in cand if fits(j)), None)
if j is None:
break
lo, hi = min(lo, j), max(hi, j)
width = max(width, len(left[j][0]))
out.append(left[lo:hi + 1])
del left[lo:hi + 1]
return out
@torch.no_grad()
def _run(self, group):
prompts = [r[0] for r in group]
b = next((x for x in sorted(self.graphs) if x >= len(prompts)), None)
width = max(map(len, prompts))
n = next((x for x in self.graphs[b].lengths if x >= width), None) if b else None
with torch.autocast("cuda", dtype=torch.bfloat16): # both paths: bf16 weights + bf16 autocast
if n is not None:
logits, _, _ = self.graphs[b](prompts)
self.stats["graph_calls"] += 1
else:
logits, _, _ = self.prefill_batch(self.model, prompts)
z = logits[:, self.letters].float() # [N, 26]
n = torch.tensor([r[1] for r in group], device=z.device)
z = z.masked_fill(torch.arange(26, device=z.device)[None] >= n[:, None], float("-inf"))
return z.softmax(-1).cpu().tolist()
def _loop(self):
while True:
items = [self.q.get()]
while True:
try:
items.append(self.q.get_nowait())
except queue.Empty:
break
t0 = time.perf_counter()
pending, slot = [], 0
for k, (rows, _) in enumerate(items):
for j, (ids, n) in enumerate(rows):
pending.append((ids, n, slot, k, j))
slot += 1
results = [[None] * len(rows) for rows, _ in items]
try:
for g in self._groups([(p[0], p[1], p[2]) for p in pending]):
t1 = time.perf_counter()
probs = self._run(g)
self.stats["gpu_s"] += time.perf_counter() - t1
self.stats["real_tokens"] += sum(len(r[0]) for r in g)
self.stats["padded_tokens"] += len(g) * max(len(r[0]) for r in g)
for (ids, n, s), p in zip(g, probs):
_, _, _, k, j = pending[s]
results[k][j] = p[:n]
self.stats["calls"] += 1
self.stats["rows"] += len(g)
for (rows, fut), r in zip(items, results):
fut.set_result(r)
except Exception as e: # noqa: BLE001 - one bad batch must not kill the server
for _, fut in items:
if not fut.done():
fut.set_exception(e)
self.stats["busy_s"] += time.perf_counter() - t0
_W = {}
def _init_worker(checkpoint, max_len, lang):
from types import SimpleNamespace
from mini_v41.tokenizer import Tokenizer
tok = Tokenizer(Path(checkpoint) / "tokenizer") # the checkpoint's own tokenizer, as InferenceModel
_W["reader"] = Reader(SimpleNamespace(tokenizer=tok, max_sequence_length=max_len, device="cpu"))
_W["lang"] = lang
def _ping(_):
return 0
def _build(body):
"""Request -> ([(ids, n options)], [(qid, question, keys)]), in a worker process."""
reader = _W["reader"]
if "state" not in body:
raise ValueError("the body needs `state` (text, object or list)")
qs = body.get("questions") or {}
if not isinstance(qs, dict) or not qs:
raise ValueError("`questions` must be a non-empty object")
lang = body.get("lang") or _W["lang"]
rows, meta = [], []
for qid, q in qs.items():
if q.get("type") not in ("choice", "score", "noul"):
raise ValueError("question %r: type must be choice, score or noul" % qid)
ids, keys = reader.prompt(body["state"], q, lang)
rows.append((ids, len(keys)))
meta.append((qid, q, keys))
return rows, meta
def answer(q, keys, p):
probs = dict(zip(keys, p))
if q["type"] == "noul":
return {"type": "noul", "noul": float(probs["true"])}
r = {"type": q["type"], "probabilities": probs, "confidence": float(max(p))}
if q["type"] == "choice":
r["choice"] = max(probs, key=probs.get)
else:
r["score"] = float(sum(i * x for i, x in enumerate(p)))
r["legend"] = {str(i): t for i, t in enumerate(q["criteria"])}
return r
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--checkpoint", type=Path, default=Path(__file__).absolute().parent.parent,
help="the repository directory (config.json, model.safetensors, tokenizer/)")
ap.add_argument("--name", default="Ines-1")
ap.add_argument("--port", type=int, default=30040)
ap.add_argument("--max-rows", type=int, default=64)
ap.add_argument("--max-tokens", type=int, default=65536, help="padded tokens per forward")
ap.add_argument("--graphs", type=int, default=64, help="groups up to this many rows replay a prefill graph (0: never)")
ap.add_argument("--lang", default="es")
ap.add_argument("--prompt-procs", type=int, default=16, help="processes building prompts (tokenizer only)")
ap.add_argument("--graph-tokens", type=int, default=6144)
a = ap.parse_args()
from aiohttp import web
from mini_v41_jev.decider import load
im = load(a.checkpoint, device="cuda:0")
reader = Reader(im)
t0 = time.time()
batcher = Batcher(im, reader, a.max_rows, a.max_tokens, a.graphs, a.graph_tokens)
print("graphs captured in %.0fs" % (time.time() - t0), flush=True)
import multiprocessing as mp
pool = ProcessPoolExecutor(a.prompt_procs, mp_context=mp.get_context("spawn"), initializer=_init_worker,
initargs=(str(a.checkpoint), im.max_sequence_length, a.lang))
list(pool.map(_ping, range(a.prompt_procs * 2)))
async def systemone(req):
try:
body = await req.json()
except Exception: # noqa: BLE001
return web.json_response({"error": "invalid JSON"}, status=400)
loop = asyncio.get_running_loop()
try:
rows, meta = await loop.run_in_executor(pool, _build, body)
except (ValueError, KeyError, TypeError) as e:
return web.json_response({"error": str(e)}, status=422)
probs = await asyncio.wrap_future(batcher.submit(rows))
answers = {qid: answer(q, keys, p) for (qid, q, keys), p in zip(meta, probs)}
return web.json_response({"model": a.name, "answers": answers,
"usage": {"input_tokens": sum(len(r[0]) for r in rows), "output_tokens": 0}})
async def models(_):
return web.json_response({"models": [{"name": a.name, "description": "Ines-1: 1.6B MoE typed-decision model, one prompt per question"}]})
async def health(_):
s = dict(batcher.stats)
s["rows_per_call"] = round(s["rows"] / s["calls"], 2) if s["calls"] else 0
return web.json_response({"status": "ok", "model": a.name, "checkpoint": str(a.checkpoint), "batching": s})
# warm-up: one prompt per length bucket, through the batcher
for n in (50, 300, 900, 1900):
batcher.submit([(list(range(10, 10 + n)), 2)]).result()
app = web.Application(client_max_size=8 * 2**20)
app.add_routes([web.post("/v1/systemone", systemone), web.get("/v1/models", models), web.get("/health", health)])
web.run_app(app, host="0.0.0.0", port=a.port, access_log=None, backlog=4096)
if __name__ == "__main__":
main()