Download runtime/jevlike/server.py from Cem13/kodama-core: direct link, hf CLI and curl.
- Browser
- Download file 21.4 kB
-
https://huggingface.co/Cem13/kodama-core/resolve/main/runtime/jevlike/server.py
- Command line
-
hf download hf://Cem13/kodama-core/runtime/jevlike/server.py
-
curl -L -o server.py https://huggingface.co/Cem13/kodama-core/resolve/main/runtime/jevlike/server.py
21.4 kB
| """Jev-compatible HTTP API for a SystemOne checkpoint. | |
| python -m jevlike.server --checkpoint checkpoints/jevlike-large-nf --port 8077 [--max-len 4096 --long-state chunk] | |
| Endpoints | |
| --------- | |
| POST /v1/systemone {"state": str | object | array, | |
| "questions": {key: {"type": "choice"|"score"|"noul", | |
| "instructions": str, | |
| "criteria": {k: desc} | [k, ...]}}} | |
| -> {"answers": {key: <SPEC answer>}, "model": str, "latency_ms": float} | |
| POST /v1/systemone/batch {"questions": {...}?, # shared default for every item | |
| "items": [{"state": ..., "questions": {...}?}, ...]} | |
| -> {"results": [<same as above>, ...], "model": str, "latency_ms": float} | |
| POST /v1/systemone/file multipart/form-data: file=<PDF | image | HTML | text>, questions=<JSON of the | |
| `questions` object above>, optional kind=pdf|image|html|text, ocr=tesseract|rapidocr, | |
| return_text=true. The file is converted with jevlike.ingest.to_text (text layer, | |
| OCR for scans/photos, tables as rows); text longer than the model's context is | |
| read in windows (long_state="chunk"). | |
| -> {"answers": ..., "model": ..., "latency_ms": ..., "ingest": {kind, n_pages, | |
| method_per_page, seconds, warnings, chars, ocr_confidence, long_state}, | |
| "text": str (only with return_text)} | |
| GET /health -> {"status": "ok", "model": ..., "checkpoint": ..., ...} | |
| State serialization (decided here, documented in README) | |
| -------------------------------------------------------- | |
| * a JSON string is passed to the model verbatim; | |
| * an object or array is rendered as pretty-printed JSON: ``json.dumps(state, indent=2, | |
| ensure_ascii=False)``. Key order is preserved as sent (not sorted): long states are cut in the | |
| *middle* by the serializer (it keeps ~25% of the token budget from the head and ~75% from the | |
| tail), so callers control what survives by where they put fields. | |
| ``--sort-keys`` switches to canonical sorted order (same value -> same string regardless of | |
| key order). Either way the rendering is deterministic. | |
| Validation errors (e.g. a score with 11 levels) come back as HTTP 422 with FastAPI's usual | |
| ``{"detail": [{"loc": [...], "msg": "...", "type": ...}]}`` body. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import inspect | |
| import json | |
| import logging | |
| import os | |
| import threading | |
| import time | |
| from contextlib import asynccontextmanager | |
| from pathlib import Path | |
| from typing import Annotated, Any, Literal, Optional, Union | |
| from fastapi import FastAPI, File, Form, HTTPException, UploadFile | |
| from fastapi.responses import JSONResponse | |
| from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError, field_validator, model_validator | |
| from jevlike.types import Choice, Noul, Score | |
| log = logging.getLogger("jevlike.server") | |
| MAX_QUESTIONS = 128 # per state | |
| MAX_BATCH_ITEMS = 256 # per /batch request | |
| MAX_INSTRUCTIONS_CHARS = 4000 | |
| MAX_STATE_CHARS = 200_000 # the model truncates anyway; this only guards the server | |
| MAX_FILE_BYTES = 50 * 1024 * 1024 # /v1/systemone/file upload limit | |
| MAX_FILE_PAGES = 200 # pages converted per uploaded file | |
| JSONValue = Union[str, dict[str, Any], list[Any]] | |
| # ---------------------------------------------------------------- request schema | |
| def _check_keys(keys: list[str], what: str) -> None: | |
| if any(not str(k).strip() for k in keys): | |
| raise ValueError(f"{what} keys/levels must be non-empty strings") | |
| class ChoiceQ(BaseModel): | |
| model_config = ConfigDict(extra="forbid") | |
| type: Literal["choice"] | |
| instructions: str = Field(min_length=1, max_length=MAX_INSTRUCTIONS_CHARS) | |
| criteria: Union[dict[str, str], list[str]] | |
| def _criteria(cls, v): | |
| keys = list(v.keys()) if isinstance(v, dict) else list(v) | |
| if not 2 <= len(keys) <= 255: | |
| raise ValueError(f"choice needs 2-255 options, got {len(keys)}") | |
| _check_keys(keys, "choice") | |
| dups = sorted({k for k in keys if keys.count(k) > 1}) | |
| if dups: | |
| raise ValueError(f"choice option keys must be unique, duplicated: {dups}") | |
| return v | |
| class ScoreQ(BaseModel): | |
| model_config = ConfigDict(extra="forbid") | |
| type: Literal["score"] | |
| instructions: str = Field(min_length=1, max_length=MAX_INSTRUCTIONS_CHARS) | |
| criteria: Union[list[str], dict[str, str]] # levels lowest -> highest | |
| def _criteria(cls, v): | |
| n = len(v) | |
| if not 2 <= n <= 10: | |
| raise ValueError(f"score needs 2-10 levels (lowest first), got {n}") | |
| _check_keys(list(v), "score") | |
| return v | |
| class NoulQ(BaseModel): | |
| model_config = ConfigDict(extra="forbid") | |
| type: Literal["noul"] | |
| instructions: str = Field(min_length=1, max_length=MAX_INSTRUCTIONS_CHARS) | |
| criteria: Optional[Any] = None | |
| def _criteria(cls, v): | |
| if v not in (None, [], {}): | |
| raise ValueError("noul questions take no criteria; put the statement to test in `instructions`") | |
| return None | |
| QuestionIn = Annotated[Union[ChoiceQ, ScoreQ, NoulQ], Field(discriminator="type")] | |
| Questions = Annotated[dict[str, QuestionIn], Field(min_length=1, max_length=MAX_QUESTIONS)] | |
| def _check_state(v): | |
| if isinstance(v, str) and not v.strip(): | |
| raise ValueError("state must be a non-empty string, object or array") | |
| if isinstance(v, (dict, list)) and not v: | |
| raise ValueError("state must not be an empty object/array") | |
| return v | |
| class PredictRequest(BaseModel): | |
| model_config = ConfigDict(extra="forbid") | |
| state: JSONValue | |
| questions: Questions | |
| def _state_ok(cls, v): | |
| return _check_state(v) | |
| class BatchItem(BaseModel): | |
| model_config = ConfigDict(extra="forbid") | |
| state: JSONValue | |
| questions: Optional[Questions] = None | |
| def _state_ok(cls, v): | |
| return _check_state(v) | |
| class BatchRequest(BaseModel): | |
| model_config = ConfigDict(extra="forbid") | |
| questions: Optional[Questions] = None | |
| items: list[BatchItem] = Field(min_length=1, max_length=MAX_BATCH_ITEMS) | |
| def _every_item_has_questions(self): | |
| if self.questions is None: | |
| missing = [i for i, it in enumerate(self.items) if it.questions is None] | |
| if missing: | |
| raise ValueError(f"items {missing[:10]} have no questions and no top-level `questions` default was given") | |
| return self | |
| # ---------------------------------------------------------------- wire <-> SystemOne | |
| def serialize_state(state: JSONValue, sort_keys: bool = False) -> str: | |
| """Strings verbatim; objects/arrays as 2-space pretty JSON (key order kept unless sort_keys).""" | |
| if isinstance(state, str): | |
| return state | |
| return json.dumps(state, indent=2, ensure_ascii=False, sort_keys=sort_keys) | |
| def to_public(q: Union[ChoiceQ, ScoreQ, NoulQ]): | |
| """Validated wire question -> (public jevlike question, score level keys or None).""" | |
| if isinstance(q, ChoiceQ): | |
| crit = dict(q.criteria) if isinstance(q.criteria, dict) else list(q.criteria) | |
| return Choice(q.instructions, crit), None | |
| if isinstance(q, ScoreQ): | |
| if isinstance(q.criteria, dict): | |
| # {key: description}: the model reads "key: description"; answers report the key. | |
| keys = list(q.criteria) | |
| levels = [f"{k}: {d}" if d else k for k, d in q.criteria.items()] | |
| return Score(q.instructions, levels), keys | |
| return Score(q.instructions, list(q.criteria)), None | |
| return Noul(q.instructions), None | |
| def _jsonable(x): | |
| """Make model output JSON-safe (numpy / torch scalars and arrays -> builtins).""" | |
| if isinstance(x, dict): | |
| return {str(k): _jsonable(v) for k, v in x.items()} | |
| if isinstance(x, (list, tuple)): | |
| return [_jsonable(v) for v in x] | |
| if isinstance(x, (str, bool, int, float)) or x is None: | |
| return x | |
| if hasattr(x, "tolist"): # numpy / torch | |
| return _jsonable(x.tolist()) | |
| if hasattr(x, "item"): | |
| return x.item() | |
| return x | |
| def _chunk_by_questions(sizes: list[int], budget: int) -> list[list[int]]: | |
| """Group consecutive item indices so each group has <= budget questions (min 1 item).""" | |
| groups, cur, n = [], [], 0 | |
| for i, s in enumerate(sizes): | |
| if cur and n + s > budget: | |
| groups.append(cur) | |
| cur, n = [], 0 | |
| cur.append(i) | |
| n += s | |
| if cur: | |
| groups.append(cur) | |
| return groups | |
| def load_model(checkpoint: str, device: Optional[str] = None, **opts): | |
| """Lazy import so the API module (and its tests) never need torch. ``opts``: SystemOne.load | |
| overrides (max_len, long_state, chunk_agg); None values are dropped.""" | |
| from jevlike.predict import SystemOne | |
| kwargs = {k: v for k, v in opts.items() if v is not None} | |
| if device and "device" in inspect.signature(SystemOne.load).parameters: | |
| kwargs["device"] = device | |
| return SystemOne.load(checkpoint, **kwargs) | |
| # ---------------------------------------------------------------- app | |
| class _Runtime: | |
| def __init__(self, model, checkpoint, device, sort_keys, max_batch_questions, model_name, load_opts=None): | |
| self.model = model | |
| self.load_opts = load_opts or {} | |
| self.checkpoint = checkpoint | |
| self.device = device | |
| self.sort_keys = sort_keys | |
| self.max_batch_questions = max_batch_questions | |
| self.model_name = model_name | |
| self.error: Optional[str] = None | |
| self.lock = threading.Lock() # one forward pass at a time (single GPU / CPU pool) | |
| self.started = time.time() | |
| def name(self) -> str: | |
| if self.model_name: | |
| return self.model_name | |
| for attr in ("name", "model_name"): | |
| v = getattr(self.model, attr, None) | |
| if isinstance(v, str) and v: | |
| return v | |
| return Path(self.checkpoint).name if self.checkpoint else "jevlike" | |
| def create_app(model=None, checkpoint: Optional[str] = None, device: Optional[str] = None, | |
| sort_keys: bool = False, max_batch_questions: int = 64, | |
| model_name: Optional[str] = None, max_len: Optional[int] = None, | |
| long_state: Optional[str] = None, chunk_agg: Optional[str] = None) -> FastAPI: | |
| """Build the app. Pass `model` (anything with SystemOne's predict/predict_batch) or a | |
| `checkpoint` directory to load once at startup (with optional `max_len` / `long_state` / | |
| `chunk_agg` overrides, see jevlike/predict.py).""" | |
| checkpoint = checkpoint or os.environ.get("JEVLIKE_CHECKPOINT") | |
| opts = {k: v for k, v in (("max_len", max_len), ("long_state", long_state), ("chunk_agg", chunk_agg)) | |
| if v is not None} | |
| rt = _Runtime(model, checkpoint, device, sort_keys, max_batch_questions, model_name, opts) | |
| async def lifespan(app: FastAPI): | |
| if rt.model is None: | |
| if not rt.checkpoint: | |
| rt.error = "no model: pass --checkpoint (or set JEVLIKE_CHECKPOINT)" | |
| log.error(rt.error) | |
| else: | |
| t0 = time.perf_counter() | |
| try: | |
| rt.model = load_model(rt.checkpoint, rt.device, **rt.load_opts) | |
| log.info("loaded %s in %.1fs", rt.checkpoint, time.perf_counter() - t0) | |
| except Exception as e: # keep serving /health with the reason | |
| rt.error = f"failed to load {rt.checkpoint}: {type(e).__name__}: {e}" | |
| log.exception(rt.error) | |
| yield | |
| app = FastAPI(title="jevlike System One", version="0.1.0", lifespan=lifespan, | |
| description="State in, calibrated answers to typed questions out.") | |
| app.state.runtime = rt | |
| def _model(): | |
| if rt.model is None: | |
| raise HTTPException(503, rt.error or "model is not loaded yet") | |
| return rt.model | |
| def _prepare(state, questions: dict): | |
| text = serialize_state(state, rt.sort_keys) | |
| if len(text) > MAX_STATE_CHARS: | |
| raise HTTPException(413, f"state is {len(text)} chars after serialization; limit is {MAX_STATE_CHARS}") | |
| public, level_keys = {}, {} | |
| for key, q in questions.items(): | |
| public[key], lk = to_public(q) | |
| if lk is not None: | |
| level_keys[key] = lk | |
| return text, public, level_keys | |
| def _finish(res: dict, level_keys: dict) -> dict: | |
| res = _jsonable(res) | |
| answers = res.get("answers", {}) | |
| for key, keys in level_keys.items(): | |
| a = answers.get(key) | |
| if isinstance(a, dict) and isinstance(a.get("level"), int) and 0 <= a["level"] < len(keys): | |
| a["label"] = keys[a["level"]] | |
| return {"answers": answers, "model": res.get("model") or rt.name, | |
| **({"latency_ms": res["latency_ms"]} if "latency_ms" in res else {})} | |
| def _call(fn, *args): | |
| try: | |
| with rt.lock: | |
| return fn(*args) | |
| except (ValueError, AssertionError, TypeError) as e: # model-side validation | |
| raise HTTPException(422, f"model rejected the request: {e}") from e | |
| except Exception as e: | |
| log.exception("predict failed") | |
| raise HTTPException(500, f"prediction failed: {type(e).__name__}") from e | |
| def health(): | |
| body = {"status": "ok" if rt.model is not None else ("error" if rt.error else "loading"), | |
| "model": rt.name, "checkpoint": rt.checkpoint, | |
| "device": str(getattr(rt.model, "device", rt.device)) if rt.model is not None else rt.device, | |
| "uptime_s": round(time.time() - rt.started, 1)} | |
| for attr in ("max_len", "long_state"): | |
| v = getattr(rt.model, attr, None) if rt.model is not None else rt.load_opts.get(attr) | |
| if isinstance(v, (int, str)) and not isinstance(v, bool): | |
| body[attr] = v | |
| if rt.error: | |
| body["error"] = rt.error | |
| return JSONResponse(body, status_code=200 if rt.model is not None else 503) | |
| def systemone(req: PredictRequest): | |
| model = _model() | |
| t0 = time.perf_counter() | |
| text, public, level_keys = _prepare(req.state, req.questions) | |
| out = _finish(_call(model.predict, text, public), level_keys) | |
| out["latency_ms"] = round((time.perf_counter() - t0) * 1000, 3) # server-side wall time | |
| return out | |
| def systemone_batch(req: BatchRequest): | |
| model = _model() | |
| t0 = time.perf_counter() | |
| prepared = [_prepare(it.state, it.questions or req.questions) for it in req.items] | |
| results: list[dict] = [None] * len(prepared) # type: ignore[list-item] | |
| for group in _chunk_by_questions([len(p[1]) for p in prepared], rt.max_batch_questions): | |
| outs = _call(model.predict_batch, [(prepared[i][0], prepared[i][1]) for i in group]) | |
| if len(outs) != len(group): | |
| raise HTTPException(500, "model returned the wrong number of results") | |
| for i, res in zip(group, outs): | |
| results[i] = _finish(res, prepared[i][2]) | |
| return {"results": results, "model": rt.name, | |
| "latency_ms": round((time.perf_counter() - t0) * 1000, 3)} | |
| if _multipart_available(): | |
| _add_file_endpoint(app, rt, _model, _prepare, _finish, _call) | |
| else: # pragma: no cover | |
| log.warning("python-multipart is not installed: POST /v1/systemone/file is disabled") | |
| return app | |
| def _multipart_available() -> bool: | |
| try: | |
| import python_multipart # noqa: F401 | |
| return True | |
| except ImportError: | |
| try: | |
| import multipart # noqa: F401 | |
| return True | |
| except ImportError: | |
| return False | |
| _QUESTIONS_ADAPTER = TypeAdapter(Questions) | |
| def _add_file_endpoint(app: FastAPI, rt: "_Runtime", _model, _prepare, _finish, _call) -> None: | |
| from functools import partial | |
| def systemone_file(file: UploadFile = File(...), questions: str = Form(...), | |
| kind: Optional[str] = Form(None), ocr: str = Form("tesseract"), | |
| return_text: bool = Form(False)): | |
| from jevlike import ingest | |
| model = _model() | |
| t0 = time.perf_counter() | |
| try: | |
| raw = json.loads(questions) | |
| if isinstance(raw, dict) and set(raw) == {"questions"}: | |
| raw = raw["questions"] | |
| qs = _QUESTIONS_ADAPTER.validate_python(raw) | |
| except json.JSONDecodeError as e: | |
| raise HTTPException(422, f"questions is not valid JSON: {e}") from None | |
| except ValidationError as e: | |
| raise HTTPException(422, json.loads(e.json(include_url=False))) from None | |
| if kind is not None and kind not in ingest.KINDS: | |
| raise HTTPException(422, f"kind must be one of {list(ingest.KINDS)}") | |
| if ocr not in ingest.OCR_ENGINES: | |
| raise HTTPException(422, f"ocr must be one of {list(ingest.OCR_ENGINES)}") | |
| data = file.file.read(MAX_FILE_BYTES + 1) | |
| if len(data) > MAX_FILE_BYTES: | |
| raise HTTPException(413, f"file is larger than {MAX_FILE_BYTES} bytes") | |
| if not data: | |
| raise HTTPException(422, "empty file") | |
| kind = kind or ingest.detect_kind(data, file.filename) | |
| try: | |
| ing = ingest.to_text(data, kind=kind, ocr=ocr, max_pages=MAX_FILE_PAGES) | |
| except ValueError as e: | |
| raise HTTPException(422, f"could not read the file: {e}") from None | |
| except Exception as e: | |
| log.exception("ingest failed") | |
| raise HTTPException(500, f"document conversion failed: {type(e).__name__}: {e}") from None | |
| if not ing.text.strip(): | |
| raise HTTPException(422, {"msg": "no text could be extracted from the file", | |
| "ingest": ing.to_dict(with_text=False)}) | |
| text, public, level_keys = _prepare(ing.text[:MAX_STATE_CHARS], qs) | |
| fn = model.predict | |
| long_state = None | |
| if "long_state" in ingest._params(fn) and ingest._too_long(model, text): | |
| long_state = "chunk" | |
| fn = partial(model.predict, long_state="chunk") | |
| out = _finish(_call(fn, text, public), level_keys) | |
| info = ing.to_dict(with_text=False) | |
| info["long_state"] = long_state or getattr(model, "long_state", None) | |
| if len(ing.text) > MAX_STATE_CHARS: | |
| info["warnings"] = info["warnings"] + [f"text cut to {MAX_STATE_CHARS} chars"] | |
| out["ingest"] = _jsonable(info) | |
| if return_text: | |
| out["text"] = ing.text | |
| out["latency_ms"] = round((time.perf_counter() - t0) * 1000, 3) | |
| return out | |
| def main(argv=None): | |
| ap = argparse.ArgumentParser(description="Serve a jevlike SystemOne checkpoint over HTTP (Jev-compatible).") | |
| ap.add_argument("--checkpoint", required=True, help="checkpoint dir, e.g. checkpoints/jevlike-large-nf") | |
| ap.add_argument("--host", default="127.0.0.1") | |
| ap.add_argument("--port", type=int, default=8077) | |
| ap.add_argument("--device", default=None, help="cpu / cuda / cuda:0 (default: SystemOne auto)") | |
| ap.add_argument("--sort-keys", action="store_true", help="serialize object states with sorted keys") | |
| ap.add_argument("--max-batch-questions", type=int, default=64, | |
| help="max questions per forward pass for /v1/systemone/batch") | |
| ap.add_argument("--model-name", default=None, help="name reported in responses (default: checkpoint dir name)") | |
| ap.add_argument("--log-level", default="info") | |
| ap.add_argument("--max-len", type=int, default=None, | |
| help="override the checkpoint's max_len (tokens, <= 8192; default: checkpoint value)") | |
| ap.add_argument("--long-state", choices=["truncate", "chunk"], default=None, | |
| help="states longer than max_len: cut in the middle (default) or score overlapping " | |
| "windows and pool them") | |
| ap.add_argument("--chunk-agg", default=None, | |
| help="pooling rule for --long-state chunk: auto (choice/score mean of log-probs, " | |
| "noul max), mean, max, linear, noisy_or") | |
| args = ap.parse_args(argv) | |
| import uvicorn | |
| logging.basicConfig(level=args.log_level.upper(), format="%(asctime)s %(name)s %(levelname)s %(message)s") | |
| app = create_app(checkpoint=args.checkpoint, device=args.device, sort_keys=args.sort_keys, | |
| max_batch_questions=args.max_batch_questions, model_name=args.model_name, | |
| max_len=args.max_len, long_state=args.long_state, chunk_agg=args.chunk_agg) | |
| uvicorn.run(app, host=args.host, port=args.port, log_level=args.log_level) | |
| if __name__ == "__main__": | |
| main() | |