feat(intern-decision-serve): Intern-Decision-4B behind semif-serve's HTTP surface

Contract, service and tests (fake engine, no GPU). Scores through the checkpoint's own
inference.py (DecisionEngine.predict, sha256-pinned); maps semif decisions onto Jev choice
questions, packs /decide/shared into calls of at most 16, runs orderings in waves, and keeps
semif's error mapping, admission, body limit and hard VRAM cap. Deltas from semif-serve are
listed in the contract.
This commit is contained in:
vh
2026-09-30 09:04:39 -07:00
parent bb806e3596
commit 5bbf0aaeba
18 changed files with 3196 additions and 0 deletions
@@ -0,0 +1,331 @@
"""intern-decision-serve HTTP layer. Contract: intern-decision-serve.contract.md.
The request surface is semif-serve's. Each semif decision becomes one Jev `choice` question; the
questions of a request are asked through the engine (the checkpoint's own DecisionEngine.predict)
in calls of at most 16, and the answers are re-keyed into semif's result shape.
"""
from __future__ import annotations
import hmac
import itertools
import json
import math
import threading
import time
from dataclasses import dataclass
from typing import Any, Literal
from fastapi import FastAPI, Request
from fastapi.concurrency import run_in_threadpool
from fastapi.responses import JSONResponse
from pydantic import BaseModel, ConfigDict, ValidationError
from .config import INFERENCE_PY_SHA256, MAX_QUESTIONS_PER_CALL, Settings
from .errors import OutOfMemory, ScoringFailed
OPEN_PATHS = frozenset({"/health"})
State = str | dict[str, Any] | list[Any]
MIN_OPTIONS, MAX_OPTIONS = 2, 16 # SemIf's rule (its 16 answer letters), kept for the surface
MAX_OPTIONS_FOR_ALL = 4 # "orderings": "all" asks n! orderings
LOG_FLOOR = 1e-300 # a probability of exactly 0 before the log (combine)
PROMPT_VERSION = f"intern-decision-jev/inference.py@{INFERENCE_PY_SHA256[:12]}"
READOUT = ("logits at the position before each <decision> marker, softmax over the field's answer "
"symbols (the checkpoint's own inference.py, DecisionEngine.predict)")
PROBABILITY_STATUS = ("temperature-scaled by the checkpoint's own shipped calibration (see calibration); "
"vendor-fitted, not fitted on our workloads")
CHUNKING = (f"/decide/shared questions are packed greedily, in request order, into calls of at most "
f"{MAX_QUESTIONS_PER_CALL} (1-{MAX_QUESTIONS_PER_CALL}, {MAX_QUESTIONS_PER_CALL + 1}-"
f"{2 * MAX_QUESTIONS_PER_CALL}, ...); each call is one prompt, so the questions in a call are "
f"asked together. With orderings, ordering k of every decision forms wave k, packed the same way.")
class Option(BaseModel):
model_config = ConfigDict(extra="ignore")
id: str
description: str
class Decision(BaseModel):
model_config = ConfigDict(extra="ignore")
id: str
question: str
options: list[Option]
orderings: Literal["none", "rotations", "all"] = "none"
class DecideBody(Decision):
state: State
workload: str | None = None
class SharedBody(BaseModel):
model_config = ConfigDict(extra="ignore")
state: State
decisions: list[Decision]
workload: str | None = None
class ApiError(Exception):
def __init__(self, status: int, code: str, message: str):
super().__init__(message)
self.status, self.code, self.message = status, code, message
def error(status: int, code: str, message: str) -> JSONResponse:
return JSONResponse(status_code=status, content={"error": {"code": code, "message": message}})
def _first_error(exc: ValidationError) -> str:
first = exc.errors()[0]
where = ".".join(str(p) for p in first.get("loc", ())) or "body"
return f"{where}: {first.get('msg', 'invalid')}"
async def read_limited(stream, declared: str | None, limit: int) -> bytes:
"""Read a request body, refusing it once it would exceed `limit` bytes. The check runs BEFORE a
chunk is kept, and nothing after the crossing chunk is read. A declared length is trusted only
as ASCII digits: `"²".isdigit()` is True but `int("²")` raises."""
too_large = ApiError(413, "request_too_large", f"request body exceeds {limit} bytes")
if declared is not None and declared.isascii() and declared.isdigit() and int(declared) > limit:
raise too_large
body = bytearray()
async for chunk in stream:
if len(body) + len(chunk) > limit:
raise too_large
body.extend(chunk)
return bytes(body)
def check_request(state: State, decisions: list[Decision], workload: str | None) -> None:
"""SemIf's row validation, re-stated because SemIf is gone; a violation is a 422."""
def bad(message: str) -> ApiError:
return ApiError(422, "invalid_request", message)
if workload is not None:
raise bad(f"unknown workload {workload!r}: no per-workload calibration is configured; "
"the model's own calibration always applies")
if not state:
raise bad("state must be a nonempty string, object, or array")
try:
json.dumps(state, ensure_ascii=False, allow_nan=False)
except (TypeError, ValueError):
raise bad("state must be finite JSON-compatible data") from None
if len({d.id for d in decisions}) != len(decisions):
raise bad("decision ids must be unique")
for d in decisions:
if not d.id or not d.question:
raise bad("id and question must be nonempty strings")
if not MIN_OPTIONS <= len(d.options) <= MAX_OPTIONS:
raise bad(f"decision {d.id!r}: options must contain {MIN_OPTIONS}-{MAX_OPTIONS} entries")
if len({o.id for o in d.options}) != len(d.options):
raise bad(f"decision {d.id!r}: option ids must be unique")
def ordering_perms(d: Decision) -> list[tuple[int, ...]]:
"""Index permutations of the caller's options, the caller's own order first."""
n = len(d.options)
if d.orderings == "none":
return [tuple(range(n))]
if d.orderings == "rotations":
return [tuple((start + k) % n for k in range(n)) for start in range(n)]
if n > MAX_OPTIONS_FOR_ALL:
raise ApiError(422, "invalid_request", f"decision {d.id!r}: orderings 'all' allows at most "
f"{MAX_OPTIONS_FOR_ALL} options ({n} given); use 'rotations'")
return list(itertools.permutations(range(n)))
@dataclass(frozen=True)
class Slot:
"""One question to ask: ordering `k` of decision number `decision`."""
decision: int
k: int
row_id: str
option_ids: list[str]
question: dict
def plan_waves(decisions: list[Decision]) -> list[list[Slot]]:
"""Wave k holds ordering k of every decision that has more than k orderings, in request order.
Wave 0 is the request as written; no wave holds two orderings of one decision."""
perms = [ordering_perms(d) for d in decisions]
waves = []
for k in range(max(map(len, perms), default=0)):
wave = []
for i, (d, ps) in enumerate(zip(decisions, perms)):
if k < len(ps):
options = [d.options[j] for j in ps[k]]
wave.append(Slot(i, k, d.id if d.orderings == "none" else f"{d.id}#o{k}", [o.id for o in options],
{"type": "choice", "instructions": d.question,
"criteria": {o.id: o.description for o in options}}))
waves.append(wave)
return waves
def pack(slots: list[Slot]) -> list[list[Slot]]:
"""Greedy, in request order: at most MAX_QUESTIONS_PER_CALL questions per call."""
return [slots[k:k + MAX_QUESTIONS_PER_CALL] for k in range(0, len(slots), MAX_QUESTIONS_PER_CALL)]
def field_names(count: int) -> list[str]:
"""Positional: `q` for a one-question call, `q1`..`qN` otherwise. Decision ids never reach the prompt."""
return ["q"] if count == 1 else [f"q{i + 1}" for i in range(count)]
def row_result(slot: Slot, response: dict, sha: str, field: str, questions: int, index: int, model: dict) -> dict:
"""INV-1: re-key the model's answer for `field`; never fill in a number it did not return."""
answer = response["answers"][field]
probs = answer["probabilities"]
missing = [i for i in slot.option_ids if i not in probs]
if missing:
raise ScoringFailed(f"the model's answer for field {field!r} lacks option ids {missing}")
return {"id": slot.row_id, "option_ids": slot.option_ids, "probabilities": [probs[i] for i in slot.option_ids],
"top": answer["choice"], "confidence": answer["confidence"],
"calibration": response["calibration"], "native": answer,
"input_tokens": response["usage"]["input_tokens"], "prompt_sha256": sha,
"prompt_version": PROMPT_VERSION, "model": model, "readout": READOUT,
"probability_status": PROBABILITY_STATUS,
"call": {"index": index, "field": field, "questions": questions}}
def combine(d: Decision, results: list[dict]) -> dict:
"""semif-serve's averaging over log p (the model returns probabilities, not logits): the mean
per option id, renormalised. The per-ordering results ride along unchanged."""
option_ids = [o.id for o in d.options]
logp: dict[str, list[float]] = {i: [] for i in option_ids}
probs: dict[str, list[float]] = {i: [] for i in option_ids}
for result in results:
for oid, p in zip(result["option_ids"], result["probabilities"]):
logp[oid].append(math.log(max(p, LOG_FLOOR)))
probs[oid].append(p)
means = [sum(logp[i]) / len(logp[i]) for i in option_ids]
peak = max(means)
weights = [math.exp(m - peak) for m in means]
combined = [w / sum(weights) for w in weights]
winner = option_ids[combined.index(max(combined))]
tops = [r["top"] for r in results]
return {"id": d.id, "option_ids": option_ids,
"combined": {"method": d.orderings, "orderings": len(results), "probabilities": combined,
"top": winner, "agreement": tops.count(winner) / len(tops),
"spread": {i: [min(probs[i]), max(probs[i])] for i in option_ids}},
"orderings": results}
def create_app(settings: Settings, engine: Any) -> FastAPI:
app = FastAPI(title="intern-decision-serve")
expected = f"Bearer {settings.api_token}".encode()
inference = threading.Lock() # INV-2: one request's calls at a time, off the event loop
in_progress = 0 # POSTs admitted and not yet answered
def admit():
nonlocal in_progress
if in_progress >= settings.max_queue:
raise ApiError(429, "busy", f"{in_progress} requests already in progress (limit {settings.max_queue})")
in_progress += 1
def leave():
nonlocal in_progress
in_progress -= 1
@app.middleware("http")
async def require_bearer(request: Request, call_next):
if request.url.path not in OPEN_PATHS:
supplied = request.headers.get("authorization", "").encode()
if not hmac.compare_digest(supplied, expected): # INV-6
return error(401, "unauthorized", "missing or wrong bearer token")
return await call_next(request)
@app.exception_handler(ApiError)
async def api_error(_request: Request, exc: ApiError):
return error(exc.status, exc.code, exc.message)
async def parse(request: Request, model: type[BaseModel]):
try:
return model.model_validate_json(
await read_limited(request.stream(), request.headers.get("content-length"), settings.max_body_bytes))
except ValidationError as exc:
raise ApiError(422, "invalid_request", _first_error(exc)) from exc
def planned(state: State, decisions: list[Decision], workload: str | None) -> list[list[Slot]]:
check_request(state, decisions, workload)
waves = plan_waves(decisions)
rows = sum(map(len, waves))
if not 1 <= rows <= settings.max_decisions:
raise ApiError(422, "invalid_request",
f"this request scores {rows} rows; the limit is 1..{settings.max_decisions}")
return waves
def run(state: State, decisions: list[Decision], waves: list[list[Slot]]) -> tuple[list[dict], dict]:
"""Every call of one request, back to back under the lock (INV-2); results in request order."""
with inference:
return _run(state, decisions, waves)
def _run(state: State, decisions: list[Decision], waves: list[list[Slot]]) -> tuple[list[dict], dict]:
started = time.perf_counter()
model = engine.metadata # static: results never carry live memory numbers
by_slot: dict[tuple[int, int], dict] = {}
sizes, tokens, inference_ms = [], [], 0.0
for chunk in (chunk for wave in waves for chunk in pack(wave)):
fields = field_names(len(chunk))
response, sha = engine.predict({"state": state,
"questions": {f: s.question for f, s in zip(fields, chunk)}})
index = len(sizes)
sizes.append(len(chunk))
tokens.append(response["usage"]["input_tokens"])
inference_ms += response["timing"]["inference_ms"]
for f, s in zip(fields, chunk):
by_slot[(s.decision, s.k)] = row_result(s, response, sha, f, len(chunk), index, model)
out = []
for i, d in enumerate(decisions):
if d.orderings == "none":
out.append(by_slot[(i, 0)])
else:
out.append(combine(d, [by_slot[(i, k)] for k in range(len(ordering_perms(d)))]))
timing = {"total_seconds": time.perf_counter() - started, "batch_size": len(by_slot),
"calls": len(sizes), "questions_per_call": sizes, "input_tokens": tokens,
"inference_seconds": inference_ms / 1000}
return out, timing
async def score(state: State, decisions: list[Decision], waves: list[list[Slot]]):
"""Run one request's calls in a worker thread; map their failures to contract codes."""
try:
return await run_in_threadpool(run, state, decisions, waves)
except ValueError as exc: # the model's own validation, token limit, ...
raise ApiError(422, "invalid_request", str(exc)) from exc
except OutOfMemory as exc:
raise ApiError(503, "out_of_memory", str(exc)) from exc
except Exception as exc: # noqa: BLE001 — any other failure, building the response included
raise ApiError(500, "scoring_failed", f"{type(exc).__name__}: {exc}") from exc
@app.get("/health")
async def health():
return {"status": "ok", "model": engine.health(), "vram_cap_gib": settings.vram_cap_gib,
"max_tokens": settings.max_tokens, "max_decisions": settings.max_decisions,
"max_questions_per_call": MAX_QUESTIONS_PER_CALL, "chunking": CHUNKING, "workloads": []}
@app.post("/decide")
async def decide(request: Request):
admit() # before the body is read
try:
body = await parse(request, DecideBody)
results, timing = await score(body.state, [body], planned(body.state, [body], body.workload))
if body.orderings != "none":
return results[0]
return {**results[0], "total_seconds": timing["total_seconds"],
"forward_seconds": timing["inference_seconds"]}
finally:
leave()
@app.post("/decide/shared")
async def decide_shared(request: Request):
admit()
try:
body = await parse(request, SharedBody)
results, timing = await score(body.state, body.decisions,
planned(body.state, body.decisions, body.workload))
return {"results": results, "timing": timing}
finally:
leave()
return app