laya / server.py
site0230's picture
Upload 24 files
a3699dd verified
Raw History Blame Contribute Delete
20.7 kB
"""Mapika Decider-2B Decision Server
High-performance non-autoregressive decision model API.
100% wire-compatible with Jev, Laya, and SystemOne /v1/systemone endpoints.
Optimized for Hugging Face Free CPU Spaces (2 vCPU cgroup quota, 16 GB RAM).
Design rules enforced:
* Forward pass runs in a dedicated worker thread, never blocking the Uvicorn event loop.
* Asyncio semaphore and queue limits prevent CPU oversubscription on 2 vCPUs.
* Requests are validated before entering the inference pipeline, turning 500s into 400s.
* Memory is bounded: body size, state size, question count, and glibc arena trimming.
* In-flight identical requests are coalesced to avoid duplicate forwards.
"""
from __future__ import annotations
import asyncio
import gc
import json
import logging
import math
import os
import time
from concurrent.futures import ThreadPoolExecutor
from contextlib import asynccontextmanager
from typing import Any, Dict, Optional, Tuple
# --------------------------------------------------------------------------------------
# Logging & CPU runtime configuration (must configure threads before torch import)
# --------------------------------------------------------------------------------------
logging.basicConfig(
level=os.environ.get("DECIDER_LOG_LEVEL", "INFO").upper(),
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
)
logger = logging.getLogger("decider_server")
from decider import cpuopt # noqa: E402
THREAD_INFO = cpuopt.configure_threads()
THREADS = THREAD_INFO.get("threads", 2)
# --------------------------------------------------------------------------------------
try:
import torch
except ImportError:
torch = None
import uvicorn # noqa: E402
from fastapi import FastAPI, HTTPException, Request # noqa: E402
from fastapi.responses import JSONResponse, PlainTextResponse # noqa: E402
if torch is not None and hasattr(torch, "set_float32_matmul_precision"):
try:
torch.set_float32_matmul_precision("high")
except Exception:
pass
# --------------------------------------------------------------------------------------
# Configuration
# --------------------------------------------------------------------------------------
MODEL_ID = os.environ.get("DECIDER_MODEL", "Mapika/decider-2b")
DEVICE = os.environ.get("DECIDER_DEVICE") or ("cuda" if (torch is not None and torch.cuda.is_available()) else "cpu")
PORT = int(os.environ.get("PORT", 7860))
MAX_STATE_TOKENS = int(os.environ.get("DECIDER_MAX_STATE_TOKENS", 8192))
MAX_STATE_CHARS = int(os.environ.get("DECIDER_MAX_STATE_CHARS", 400_000))
MAX_BODY_BYTES = int(os.environ.get("DECIDER_MAX_BODY_BYTES", 4 * 1024 * 1024))
MAX_QUESTIONS = int(os.environ.get("DECIDER_MAX_QUESTIONS", 64))
MAX_SCHEMAS = int(os.environ.get("DECIDER_MAX_SCHEMAS", 16))
MAX_CONCURRENT = int(os.environ.get("DECIDER_MAX_CONCURRENT", 2))
MAX_ROWS = int(os.environ.get("DECIDER_MAX_ROWS", 256))
QUEUE_TIMEOUT_S = float(os.environ.get("DECIDER_QUEUE_TIMEOUT_S", 120.0))
PRECISION = os.environ.get("DECIDER_PRECISION", "auto")
os.environ.setdefault("DECIDER_SHARED_FORK_GB", "1.5")
RESERVED_KEYS = frozenset({"model", "answers", "usage", "latency_ms", "request_id", "coalesced"})
MAX_CHOICE, MAX_LEVELS = 255, 10
_KNOWN_TYPES = {"choice", "score", "noul", "bool"}
# --------------------------------------------------------------------------------------
# Request Validation Layer (Runs on event loop in microseconds -> clean 400s)
# --------------------------------------------------------------------------------------
def _bad(detail: str) -> HTTPException:
return HTTPException(status_code=400, detail=detail)
def _text_len(value: Any) -> int:
if isinstance(value, str):
return len(value)
if isinstance(value, (int, float, bool)) or value is None:
return 8
if isinstance(value, list):
return sum(_text_len(v) for v in value) + 2
if isinstance(value, dict):
return sum(_text_len(k) + _text_len(v) for k, v in value.items()) + 2
return 16
def _instruction_of(spec: dict, where: str) -> str:
raw = spec.get("instructions", spec.get("question", ""))
if raw is None:
raw = ""
if not isinstance(raw, str):
try:
raw = json.dumps(raw, ensure_ascii=False)
except (TypeError, ValueError):
raise _bad(f"{where}: 'instructions' is not JSON-serialisable")
return raw
def _validate_question(qid: str, spec: Any) -> None:
where = f"questions[{qid!r}]"
if not isinstance(spec, dict):
raise _bad(f"{where} is a {type(spec).__name__}, expected an object with 'type' and 'criteria'")
qtype = spec.get("type", "choice")
if not isinstance(qtype, str) or qtype not in _KNOWN_TYPES:
raise _bad(f"{where}: unknown type {qtype!r}; expected one of {sorted(_KNOWN_TYPES)}")
crit = spec.get("criteria", spec.get("options"))
instructions = _instruction_of(spec, where)
if qtype == "choice":
if isinstance(crit, (list, tuple)):
crit = {str(c): None for c in crit}
if not isinstance(crit, dict):
raise _bad(f"{where}: 'criteria' must be an object of 2..{MAX_CHOICE} options (or a list), got {type(crit).__name__}")
if not 2 <= len(crit) <= MAX_CHOICE:
raise _bad(f"{where}: choice needs 2..{MAX_CHOICE} options, got {len(crit)}")
if not instructions:
raise _bad(f"{where}: 'instructions' is required for a choice question")
elif qtype == "score":
if isinstance(crit, dict):
try:
sorted(crit, key=float)
except (TypeError, ValueError):
raise _bad(f"{where}: a criteria map's keys must be numbers, e.g. {{\"0\": \"weak\", \"1\": \"moderate\", \"2\": \"strong\"}}")
elif isinstance(crit, (list, tuple)):
pass
else:
raise _bad(f"{where}: 'criteria' must be an ordered list of 2..{MAX_LEVELS} level descriptions or a map, got {type(crit).__name__}")
if not 2 <= len(crit) <= MAX_LEVELS:
raise _bad(f"{where}: score needs 2..{MAX_LEVELS} levels, got {len(crit)}")
if not instructions:
raise _bad(f"{where}: 'instructions' is required for a score question")
else: # noul / bool
if crit is not None and not isinstance(crit, dict):
raise _bad(f"{where}: 'criteria' must be an object with optional 'true'/'false' descriptions, got {type(crit).__name__}")
described = isinstance(crit, dict) and any(
d not in (None, "") for d in (crit.get("true", crit.get(True)), crit.get("false", crit.get(False)))
)
if not instructions and not described:
raise _bad(f"{where}: a noul question needs 'instructions', or a criteria object describing 'true' or 'false'")
def _validate(body: Any) -> Tuple[Any, Dict[str, dict], bool]:
if not isinstance(body, dict):
raise _bad("Payload must be a JSON object")
if "state" not in body:
raise _bad("'state' is required")
state = body["state"]
if state is None:
raise _bad("'state' must not be null")
if not isinstance(state, (str, dict, list)):
raise _bad(f"'state' must be a string, object or array, got {type(state).__name__}")
if isinstance(state, (dict, list)) and not state:
raise _bad("'state' must not be empty")
if _text_len(state) > MAX_STATE_CHARS:
raise HTTPException(413, f"'state' exceeds character limit of {MAX_STATE_CHARS}")
if isinstance(state, str) and not state.strip():
raise _bad("'state' must not be empty")
questions = body.get("questions")
if questions is None:
raise _bad("'questions' is required")
if isinstance(questions, list):
if not questions:
raise _bad("'questions' must not be empty")
normalized = {}
for i, q in enumerate(questions):
if isinstance(q, dict):
key = q.get("name") or q.get("id") or f"q_{i}"
else:
key = f"q_{i}"
normalized[str(key)] = q
questions = normalized
if not isinstance(questions, dict):
raise _bad(f"'questions' must be an object or a list, got {type(questions).__name__}")
if not questions:
raise _bad("'questions' must not be empty")
if len(questions) > MAX_QUESTIONS:
raise HTTPException(413, f"Too many questions ({len(questions)} > {MAX_QUESTIONS})")
for qid, spec in questions.items():
_validate_question(qid, spec)
independent = body.get("independent", False)
if not isinstance(independent, bool):
raise _bad(f"'independent' must be a boolean, got {type(independent).__name__}")
return state, questions, independent
# --------------------------------------------------------------------------------------
# Serving Engine Manager
# --------------------------------------------------------------------------------------
class ServingEngine:
def __init__(self) -> None:
self.decider = None
self.detail = "uninitialized"
self.ready = False
self.model_version = "v11"
self.precision_note = ""
self.quantized = False
self.layout_mode = "state_first"
self.head_bytes_saved = 0
self.dtype = None
self.worker = ThreadPoolExecutor(max_workers=1, thread_name_prefix="decider-cpu")
self.slot: Optional[asyncio.Semaphore] = None
self.inflight = 0
self.stats = {
"requests": 0, "served": 0, "coalesced": 0, "rejected_busy": 0,
"rejected_invalid": 0, "errors": 0, "truncated_states": 0,
"total_latency_ms": 0, "max_latency_ms": 0, "total_tokens": 0,
"last_latency_ms": None, "last_queue_ms": None,
}
self._inflight_calls: Dict[str, Any] = {}
def load(self) -> None:
from decider.infer import Decider
t0 = time.time()
self.dtype, quantize, self.precision_note = cpuopt.resolve_precision(DEVICE, PRECISION)
if torch.cuda.is_available() and DEVICE.startswith("cuda"):
self.dtype, quantize = torch.bfloat16, False
self.precision_note = f"device={DEVICE}"
logger.info("Loading %s device=%s dtype=%s precision=%s threads=%d",
MODEL_ID, DEVICE, self.dtype, PRECISION, THREADS)
kwargs = dict(device=DEVICE, use_graphs=False, shared_prefix=True,
trim_head=True, quantize="int8" if quantize else "off",
max_schemas=MAX_SCHEMAS)
self.decider = Decider(MODEL_ID, dtype=self.dtype, **kwargs)
self.quantized = bool(getattr(self.decider, "quantized", quantize))
self.layout_mode = ("schema_first" if getattr(self.decider, "schema_first", False) else "state_first")
self.head_bytes_saved = int(getattr(self.decider, "head_bytes_saved", 0))
raw = getattr(self.decider, "name", "decider-2b-v11")
clean = raw.replace("decider-", "")
parts = clean.split("-")
ver = parts[-1] if len(parts) > 1 else clean
self.model_version = ver if ver.startswith("v") else f"v{ver}"
self._warmup()
self.ready = True
self.detail = f"ready (version: {self.model_version}, loaded in {round(time.time() - t0, 1)}s on {DEVICE}, {THREADS} threads, {self.precision_note})"
logger.info("✅ Decider-2B (%s) ready in %.1fs | threads=%d | %s | layout=%s | head_saved=%.0f MB",
self.model_version, time.time() - t0, THREADS, self.precision_note,
self.layout_mode, self.head_bytes_saved / 1e6)
def _warmup(self) -> None:
t0 = time.time()
try:
state = ("Warmup: synthetic market state. Price +1.2%, RSI 58, no material catalyst. " * 4)
qs = {
"action": {"type": "choice", "instructions": "What is the recommended action?",
"criteria": {"BUY": "positive", "SELL": "negative", "HOLD": "neutral"}},
"conviction": {"type": "score", "instructions": "Conviction level?",
"criteria": ["weak", "moderate", "strong"]},
"material": {"type": "noul", "instructions": "Is this a material event?"},
}
out = self.decider.system_one(state, qs, independent=False, max_state_tokens=512)
logger.info("Warmup completed in %.2fs (%d answers)", time.time() - t0, len(out.get("answers", {})))
except Exception as e:
logger.warning("Warmup notice (%s); proceeding to serve", e)
cpuopt.trim()
def decide(self, state, questions, independent: bool) -> Dict[str, Any]:
with torch.inference_mode():
t0 = time.perf_counter()
result = self.decider.system_one(
state, questions, independent=independent,
max_state_tokens=MAX_STATE_TOKENS, max_fwd_tokens=65536,
)
result["latency_ms"] = round((time.perf_counter() - t0) * 1000)
usage = result.get("usage") or {}
self.stats["total_tokens"] += int(usage.get("input_tokens") or 0)
if usage.get("truncated"):
self.stats["truncated_states"] += 1
return result
def shutdown(self) -> None:
self.worker.shutdown(wait=False, cancel_futures=True)
self.decider = None
ENGINE = ServingEngine()
# --------------------------------------------------------------------------------------
# Application Lifespan
# --------------------------------------------------------------------------------------
@asynccontextmanager
async def lifespan(app: FastAPI):
ENGINE.slot = asyncio.Semaphore(1)
loop = asyncio.get_running_loop()
try:
await loop.run_in_executor(ENGINE.worker, ENGINE.load)
except Exception as e:
ENGINE.detail = f"error: {e}"
logger.error("❌ Failed to load Decider-2B: %s", e, exc_info=True)
yield
logger.info("Shutting down Decider Decision Server...")
ENGINE.shutdown()
app = FastAPI(title="Service", docs_url=None, redoc_url=None, openapi_url=None, lifespan=lifespan)
# --------------------------------------------------------------------------------------
# Routes
# --------------------------------------------------------------------------------------
@app.api_route("/", methods=["GET", "HEAD"], response_class=PlainTextResponse)
def index():
return f"Laya is live. Version: {ENGINE.model_version} (Engine: Decider-2B)"
@app.get("/health")
def health():
return {
"status": "healthy" if ENGINE.ready else "degraded",
"service": "laya",
"engine": "Decider-2B",
"version": ENGINE.model_version,
"model": MODEL_ID,
"detail": ENGINE.detail,
"device": DEVICE,
"cuda_available": torch.cuda.is_available(),
"quantized": ENGINE.quantized,
"layout": getattr(ENGINE, "layout_mode", "state_first"),
"threads": THREAD_INFO,
"memory": cpuopt.memory_budget() if ENGINE.ready else None,
"queue_depth": ENGINE.inflight,
}
@app.get("/stats")
def stats():
s = dict(ENGINE.stats)
avg = s["total_latency_ms"] / s["served"] if s["served"] else None
return {
**s,
"avg_latency_ms": round(avg) if avg is not None else None,
"precision": ENGINE.precision_note,
"quantized": ENGINE.quantized,
"layout": getattr(ENGINE, "layout_mode", "state_first"),
"lm_head_bytes_saved": getattr(ENGINE, "head_bytes_saved", 0),
"dtype": str(ENGINE.dtype),
"threads": THREADS,
"device": DEVICE,
"limits": {
"max_state_tokens": MAX_STATE_TOKENS,
"max_state_chars": MAX_STATE_CHARS,
"max_body_bytes": MAX_BODY_BYTES,
"max_questions": MAX_QUESTIONS,
"max_concurrent": MAX_CONCURRENT,
"queue_timeout_s": QUEUE_TIMEOUT_S,
},
"memory": cpuopt.memory_budget(),
}
@app.post("/v1/systemone")
async def handle_systemone(request: Request):
if ENGINE.decider is None or not ENGINE.ready:
raise HTTPException(status_code=503, detail=f"Model is not ready: {ENGINE.detail}")
# Bounded body read
content_length = request.headers.get("content-length")
if content_length:
try:
if int(content_length) > MAX_BODY_BYTES:
raise HTTPException(413, f"Request body is {content_length} bytes, limit is {MAX_BODY_BYTES}")
except ValueError:
raise _bad("Invalid Content-Length header")
raw = await request.body()
if len(raw) > MAX_BODY_BYTES:
raise HTTPException(413, f"Request body is {len(raw)} bytes, limit is {MAX_BODY_BYTES}")
try:
body = json.loads(raw or b"null")
except Exception:
raise _bad("Invalid JSON payload")
# Fast validation on event loop
try:
state, questions, independent = _validate(body)
except HTTPException:
ENGINE.stats["rejected_invalid"] += 1
raise
except Exception as e:
ENGINE.stats["rejected_invalid"] += 1
logger.warning("Validation rejected unexpectedly: %s", e)
raise _bad(f"Invalid request: {e}")
ENGINE.stats["requests"] += 1
key = json.dumps({"s": state, "q": questions, "i": independent}, sort_keys=True, ensure_ascii=False, default=str)
loop = asyncio.get_running_loop()
t_enqueue = time.perf_counter()
# In-flight coalescing for simultaneous identical requests
existing = ENGINE._inflight_calls.get(key)
if existing is not None and not existing.done():
ENGINE.stats["coalesced"] += 1
try:
shared = await asyncio.shield(existing)
except asyncio.CancelledError:
raise
except Exception as e:
raise HTTPException(502, f"In-flight decision failed: {e}")
res = dict(shared)
res["coalesced"] = True
return JSONResponse(content=res)
# Admission control
if ENGINE.inflight >= MAX_CONCURRENT:
ENGINE.stats["rejected_busy"] += 1
raise HTTPException(
status_code=503,
detail=f"Server busy: {ENGINE.inflight} requests queued (limit {MAX_CONCURRENT}); retry shortly",
headers={"Retry-After": "2"},
)
fut: asyncio.Future = loop.create_future()
ENGINE._inflight_calls[key] = fut
ENGINE.inflight += 1
async def _run() -> Dict[str, Any]:
try:
if ENGINE.slot is not None:
await asyncio.wait_for(ENGINE.slot.acquire(), timeout=QUEUE_TIMEOUT_S)
except asyncio.TimeoutError:
raise HTTPException(
status_code=503,
detail=f"Timed out after {QUEUE_TIMEOUT_S:.0f}s waiting for inference slot",
headers={"Retry-After": "2"},
)
try:
return await loop.run_in_executor(ENGINE.worker, ENGINE.decide, state, questions, independent)
finally:
if ENGINE.slot is not None:
ENGINE.slot.release()
try:
try:
result = await _run()
except asyncio.CancelledError:
ENGINE.stats["errors"] += 1
if not fut.done():
fut.cancel()
raise
except HTTPException:
raise
except Exception as e:
ENGINE.stats["errors"] += 1
logger.error("Inference error: %s", e, exc_info=True)
if not fut.done():
fut.set_exception(e)
raise HTTPException(500, f"Inference failed: {e}")
finally:
ENGINE.inflight -= 1
ENGINE._inflight_calls.pop(key, None)
cpuopt.trim()
if not fut.done():
fut.set_result(result)
queue_ms = round((time.perf_counter() - t_enqueue) * 1000)
s = ENGINE.stats
s["served"] += 1
s["total_latency_ms"] += int(result.get("latency_ms") or 0)
s["max_latency_ms"] = max(s["max_latency_ms"], int(result.get("latency_ms") or 0))
s["last_latency_ms"] = result.get("latency_ms")
s["last_queue_ms"] = queue_ms
# Universal compatibility: nested answers + flat mirror (safe against reserved envelope keys)
answers = result.get("answers") or {}
if isinstance(answers, dict):
for k, v in answers.items():
if k not in result and k not in RESERVED_KEYS:
result[k] = v
return JSONResponse(content=result)
finally:
gc.collect()
if __name__ == "__main__":
uvicorn.run(app, host="0.0.0.0", port=PORT, workers=1, log_level="info")