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:
@@ -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
|
||||
@@ -0,0 +1,87 @@
|
||||
"""Settings for intern-decision-serve. Contract: intern-decision-serve.contract.md § Configuration.
|
||||
|
||||
Every value is validated at startup and a bad one is refused with a ValueError naming the
|
||||
variable: a service that starts and then rejects every request (or runs uncapped) is worse than
|
||||
one that does not start.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
|
||||
PREFIX = "INTERN_DECISION_"
|
||||
MIN_TOKEN_CHARS = 32
|
||||
MODEL_ID = "internlm/Intern-Decision-4B"
|
||||
REVISION = "0e5e6aa7d6d750e2b1504ba11a8136cb58aeb3cd"
|
||||
# The checkpoint's own inference.py is EXECUTED from the (data) mount, so its content is pinned.
|
||||
INFERENCE_PY_SHA256 = "c904e2c67ca0775621a22375ee373d2ba30b52117cda870c6c9ef74143b29863"
|
||||
DEFAULT_CHECKPOINT = f"/hf/hub/models--internlm--Intern-Decision-4B/snapshots/{REVISION}"
|
||||
MAX_QUESTIONS_PER_CALL = 16 # the model's own limit (inference.validate_request)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Settings:
|
||||
api_token: str
|
||||
checkpoint: str = DEFAULT_CHECKPOINT
|
||||
device: str = "cuda"
|
||||
vram_cap_gib: float | None = None
|
||||
max_tokens: int = 8192
|
||||
max_decisions: int = 64
|
||||
max_body_bytes: int = 1024 * 1024
|
||||
max_queue: int = 32
|
||||
release_slack_mib: int = 512
|
||||
keep_vision: bool = False
|
||||
|
||||
@classmethod
|
||||
def from_env(cls, env: Mapping[str, str]) -> "Settings":
|
||||
token = env.get(f"{PREFIX}API_TOKEN", "")
|
||||
# INV-6: visible ASCII only. A CR, LF or NUL can never arrive in a header, so a token
|
||||
# carrying one would lock every caller out while /health still said ok.
|
||||
if len(token) < MIN_TOKEN_CHARS or not all(33 <= ord(c) <= 126 for c in token):
|
||||
raise ValueError(f"{PREFIX}API_TOKEN must be at least {MIN_TOKEN_CHARS} visible ASCII characters")
|
||||
device = env.get(f"{PREFIX}DEVICE", "cuda")
|
||||
if device not in ("cuda", "cpu"):
|
||||
raise ValueError(f"{PREFIX}DEVICE must be cuda or cpu, not {device!r}")
|
||||
keep_vision = env.get(f"{PREFIX}KEEP_VISION", "0")
|
||||
if keep_vision not in ("0", "1"):
|
||||
raise ValueError(f"{PREFIX}KEEP_VISION must be 0 or 1, not {keep_vision!r}")
|
||||
return cls(
|
||||
api_token=token,
|
||||
checkpoint=env.get(f"{PREFIX}CHECKPOINT", DEFAULT_CHECKPOINT),
|
||||
device=device,
|
||||
vram_cap_gib=_positive_float(env, "VRAM_CAP_GIB"),
|
||||
max_tokens=_int(env, "MAX_TOKENS", 8192),
|
||||
max_decisions=_int(env, "MAX_DECISIONS", 64),
|
||||
max_body_bytes=_int(env, "MAX_BODY_BYTES", 1024 * 1024),
|
||||
max_queue=_int(env, "MAX_QUEUE", 32),
|
||||
release_slack_mib=_int(env, "RELEASE_SLACK_MIB", 512, minimum=0),
|
||||
keep_vision=keep_vision == "1",
|
||||
)
|
||||
|
||||
|
||||
def _int(env: Mapping[str, str], name: str, default: int, minimum: int = 1) -> int:
|
||||
raw = env.get(PREFIX + name)
|
||||
if raw is None:
|
||||
return default
|
||||
try:
|
||||
value = int(raw)
|
||||
except ValueError:
|
||||
raise ValueError(f"{PREFIX}{name} must be an integer, got {raw!r}") from None
|
||||
if value < minimum:
|
||||
raise ValueError(f"{PREFIX}{name} must be >= {minimum}, got {value}")
|
||||
return value
|
||||
|
||||
|
||||
def _positive_float(env: Mapping[str, str], name: str) -> float | None:
|
||||
"""Unset or empty means no cap. When set it must be finite and > 0."""
|
||||
raw = env.get(PREFIX + name)
|
||||
if raw is None or raw == "":
|
||||
return None
|
||||
try:
|
||||
value = float(raw)
|
||||
except ValueError:
|
||||
raise ValueError(f"{PREFIX}{name} must be a number, got {raw!r}") from None
|
||||
if not math.isfinite(value) or value <= 0:
|
||||
raise ValueError(f"{PREFIX}{name} must be a finite number > 0, got {raw!r}")
|
||||
return value
|
||||
@@ -0,0 +1,197 @@
|
||||
"""The real engine: Intern-Decision-4B's own DecisionEngine (the checkpoint's inference.py) over one
|
||||
resident model. Needs the `model` extra.
|
||||
|
||||
Contract: intern-decision-serve.contract.md, INV-3 (fail-closed startup), INV-4 (VRAM cap + OOM),
|
||||
INV-5 (offline weights), INV-7 (text-only model), INV-8 (honest prompt hash). load() is exercised on
|
||||
the card at acceptance; its checks and the OOM path are unit-tested against a fake torch and a
|
||||
fake checkpoint (tests/test_engine.py).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import hashlib
|
||||
import importlib.metadata
|
||||
import importlib.util
|
||||
import logging
|
||||
import traceback
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from .config import INFERENCE_PY_SHA256, MODEL_ID, REVISION, Settings
|
||||
from .errors import OutOfMemory, ScoringFailed
|
||||
|
||||
log = logging.getLogger("intern_decision_serve.engine")
|
||||
|
||||
MIB = 2**20
|
||||
WARMUP_REQUEST = {
|
||||
"state": "The deployment completed at 14:02 UTC. Health checks passed in all three zones.",
|
||||
"questions": {"q": {"type": "choice", "instructions": "Is there evidence that the deployment succeeded?",
|
||||
"criteria": {"yes": "The deployment succeeded.", "no": "The deployment did not succeed."}}},
|
||||
}
|
||||
|
||||
|
||||
def _first_line(exc: BaseException) -> str:
|
||||
lines = str(exc).splitlines()
|
||||
return lines[0] if lines else ""
|
||||
|
||||
|
||||
def _version(distribution: str) -> str:
|
||||
try:
|
||||
return importlib.metadata.version(distribution)
|
||||
except importlib.metadata.PackageNotFoundError:
|
||||
return "n/a"
|
||||
|
||||
|
||||
def _sha256(path: Path) -> str:
|
||||
return hashlib.sha256(path.read_bytes()).hexdigest()
|
||||
|
||||
|
||||
def _import_inference(path: Path):
|
||||
"""The checkpoint's own inference.py, imported by file path under a private module name."""
|
||||
spec = importlib.util.spec_from_file_location("intern_decision_inference", path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _vision_stub(torch: Any):
|
||||
class VisionTowerRemoved(torch.nn.Module):
|
||||
"""INV-7: stands where the vision tower was. The service takes no images; if anything ever
|
||||
routes pixels here, fail loudly rather than answer from a missing tower."""
|
||||
def forward(self, *_args, **_kwargs):
|
||||
raise RuntimeError("the vision tower was removed at startup (text-only service, INV-7)")
|
||||
return VisionTowerRemoved()
|
||||
|
||||
|
||||
class TorchEngine:
|
||||
def __init__(self, torch: Any, engine: Any, inference: Any, tokenizer: Any, metadata: dict,
|
||||
settings: Settings, release_above_bytes: int | None = None):
|
||||
self._torch, self._engine, self._inference, self._tokenizer = torch, engine, inference, tokenizer
|
||||
self.metadata, self._settings = metadata, settings
|
||||
self._release_above = release_above_bytes
|
||||
|
||||
@classmethod
|
||||
def load(cls, settings: Settings, *, torch: Any = None,
|
||||
inference_sha256: str = INFERENCE_PY_SHA256) -> "TorchEngine":
|
||||
checkpoint = Path(settings.checkpoint)
|
||||
# INV-3: the pinned snapshot, and the pinned code. Checked before torch touches the card.
|
||||
if checkpoint.name != REVISION or checkpoint.parent.name != "snapshots":
|
||||
raise RuntimeError(f"INTERN_DECISION_CHECKPOINT must be a snapshots/{REVISION} directory, "
|
||||
f"got {checkpoint}")
|
||||
code = checkpoint / "inference.py"
|
||||
found = _sha256(code)
|
||||
if found != inference_sha256:
|
||||
raise RuntimeError(f"{code} has sha256 {found}, not the pinned {inference_sha256}: "
|
||||
"re-check the checkpoint's inference.py before serving it")
|
||||
if torch is None:
|
||||
import torch
|
||||
|
||||
if settings.device == "cuda":
|
||||
if not torch.cuda.is_available():
|
||||
raise RuntimeError("INTERN_DECISION_DEVICE=cuda but torch sees no CUDA device")
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
arch = f"sm_{major}{minor}"
|
||||
if arch not in torch.cuda.get_arch_list(): # INV-3: no silent PTX/CPU fallback
|
||||
raise RuntimeError(f"torch {torch.__version__} has no kernels for {arch}: {torch.cuda.get_arch_list()}")
|
||||
if settings.vram_cap_gib is not None: # INV-4: cap BEFORE the weights land
|
||||
total = torch.cuda.get_device_properties(0).total_memory
|
||||
fraction = settings.vram_cap_gib * 2**30 / total
|
||||
if not 0 < fraction <= 1:
|
||||
raise ValueError(f"INTERN_DECISION_VRAM_CAP_GIB={settings.vram_cap_gib} does not fit a "
|
||||
f"{total / 2**30:.1f} GiB card")
|
||||
torch.cuda.set_per_process_memory_fraction(fraction, 0)
|
||||
|
||||
inference = _import_inference(code)
|
||||
decision_engine = inference.DecisionEngine(checkpoint=str(checkpoint), max_length=settings.max_tokens,
|
||||
device=settings.device, dtype="bfloat16",
|
||||
attn_implementation="sdpa")
|
||||
backend = decision_engine.backend
|
||||
placed = next(backend.model.parameters()).device.type
|
||||
if placed != settings.device: # INV-3
|
||||
raise RuntimeError(f"model landed on {placed}, expected {settings.device}")
|
||||
metadata = {"name": inference.MODEL_NAME, "source": MODEL_ID, "revision": REVISION,
|
||||
"checkpoint": str(checkpoint), "inference_py_sha256": found,
|
||||
"temperature": decision_engine.temperature, "dtype": "bfloat16", "attn_implementation": "sdpa",
|
||||
"device": settings.device, "max_length": settings.max_tokens,
|
||||
"torch_version": torch.__version__, "transformers_version": _version("transformers"),
|
||||
"vision_tower": "loaded" if settings.keep_vision else "removed"}
|
||||
engine = cls(torch, decision_engine, inference, backend.tokenizer, metadata, settings)
|
||||
|
||||
first, _ = engine.predict(WARMUP_REQUEST) # INV-3: one decision must score
|
||||
engine._prove_prompt_hash(first) # INV-8
|
||||
if not settings.keep_vision: # INV-7
|
||||
engine._remove_vision_tower(first)
|
||||
if settings.device == "cuda": # INV-4: the resting footprint
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
engine._release_above = torch.cuda.memory_reserved(0) + settings.release_slack_mib * MIB
|
||||
return engine
|
||||
|
||||
def _prompt_text(self, request: dict) -> str:
|
||||
"""INV-8: the chat-template text, rendered with the same arguments HFBackend.encode uses for a
|
||||
text-only row (the tokenizer is its template when there are no images)."""
|
||||
compiled = self._inference.compile_row(self._inference.validate_request(request))
|
||||
return self._tokenizer.apply_chat_template(compiled.messages, tokenize=False, add_generation_prompt=False,
|
||||
enable_thinking=False, add_vision_id=True)
|
||||
|
||||
def _prove_prompt_hash(self, warmup: dict) -> None:
|
||||
tokens = len(self._tokenizer(self._prompt_text(WARMUP_REQUEST), add_special_tokens=False)["input_ids"])
|
||||
if tokens != warmup["usage"]["input_tokens"]:
|
||||
raise RuntimeError(f"INV-8: the rendered prompt tokenises to {tokens} tokens but the model read "
|
||||
f"{warmup['usage']['input_tokens']}: prompt_sha256 would hash a prompt it never saw")
|
||||
|
||||
def _remove_vision_tower(self, before: dict) -> None:
|
||||
inner = getattr(getattr(self._engine.backend, "model", None), "model", None)
|
||||
if inner is None or not hasattr(inner, "visual"):
|
||||
raise RuntimeError("INV-7: the vision tower is not at backend.model.model.visual; "
|
||||
"set INTERN_DECISION_KEEP_VISION=1 or re-check the model class")
|
||||
inner.visual = _vision_stub(self._torch)
|
||||
gc.collect()
|
||||
after, _ = self.predict(WARMUP_REQUEST)
|
||||
if after["answers"] != before["answers"]:
|
||||
raise RuntimeError(f"INV-7: removing the vision tower changed the warm-up answer "
|
||||
f"({before['answers']} -> {after['answers']})")
|
||||
|
||||
def health(self) -> dict:
|
||||
info = dict(self.metadata)
|
||||
if self._settings.device == "cuda":
|
||||
cuda = self._torch.cuda
|
||||
info["device_name"] = cuda.get_device_name(0)
|
||||
info["allocated_gib"] = round(cuda.memory_allocated(0) / 2**30, 3)
|
||||
info["reserved_gib"] = round(cuda.memory_reserved(0) / 2**30, 3)
|
||||
if hasattr(cuda, "max_memory_reserved"):
|
||||
info["max_reserved_gib"] = round(cuda.max_memory_reserved(0) / 2**30, 3)
|
||||
return info
|
||||
|
||||
def _release_burst(self) -> None:
|
||||
"""INV-4: hand a burst back to the driver so GPU 1's shared headroom (scriberr, the vLLM
|
||||
seats) returns after a big request, instead of sitting in torch's cache."""
|
||||
if self._release_above is not None and self._torch.cuda.memory_reserved(0) > self._release_above:
|
||||
self._torch.cuda.empty_cache()
|
||||
|
||||
def predict(self, request: dict) -> tuple[dict, str]:
|
||||
"""One call: the model's own predict(), then the prompt hash (INV-8). Returns (response, sha)."""
|
||||
try:
|
||||
response = self._engine.predict(request)
|
||||
sha = hashlib.sha256(self._prompt_text(request).encode()).hexdigest()
|
||||
except ValueError:
|
||||
raise # the model's validation: raised before any GPU work
|
||||
except self._torch.cuda.OutOfMemoryError as exc:
|
||||
failure, message = OutOfMemory, _first_line(exc) or "CUDA out of memory"
|
||||
except Exception as exc: # noqa: BLE001 — every other failure is released and reported below
|
||||
message = _first_line(exc)
|
||||
if "out of memory" in message.lower(): # cuBLAS/cuDNN allocation failures
|
||||
failure = OutOfMemory
|
||||
else:
|
||||
failure, message = ScoringFailed, f"{type(exc).__name__}: {message}"
|
||||
# Formatted text, not exc_info: a record that keeps the traceback alive pins the tensors.
|
||||
log.error("predict failed:\n%s", traceback.format_exc())
|
||||
else:
|
||||
self._release_burst()
|
||||
return response, sha
|
||||
# INV-4, outside the except block on purpose: the exception's traceback holds the failed
|
||||
# forward's frames and with them its tensors. Raising inside the block, or `from exc`, would
|
||||
# chain to it and keep them allocated after the response (semif-serve, found on the card).
|
||||
gc.collect()
|
||||
self._torch.cuda.empty_cache()
|
||||
raise failure(message)
|
||||
@@ -0,0 +1,10 @@
|
||||
"""Torch-free exceptions shared by the HTTP layer and the engine."""
|
||||
|
||||
|
||||
class OutOfMemory(RuntimeError):
|
||||
"""The engine ran out of GPU memory during a call and has already released its cache (INV-4)."""
|
||||
|
||||
|
||||
class ScoringFailed(RuntimeError):
|
||||
"""A call failed for a reason other than validation or OOM. Raised unchained, after the failed
|
||||
call's memory has been released; the original traceback is logged, not carried (INV-4)."""
|
||||
@@ -0,0 +1,20 @@
|
||||
"""uvicorn entry point: `uvicorn intern_decision_serve.main:app_from_env --factory --workers 1`."""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
from fastapi import FastAPI
|
||||
|
||||
from .app import create_app
|
||||
from .config import Settings
|
||||
|
||||
|
||||
def app_from_env() -> FastAPI:
|
||||
settings = Settings.from_env(os.environ)
|
||||
# INV-5: never download at runtime, inside the image or out of it. Set before torch /
|
||||
# transformers / huggingface_hub are imported, since they read it at import time.
|
||||
os.environ["HF_HUB_OFFLINE"] = "1"
|
||||
os.environ["TRANSFORMERS_OFFLINE"] = "1"
|
||||
from .engine import TorchEngine # torch loads only inside load(), never in the unit tests
|
||||
|
||||
return create_app(settings, TorchEngine.load(settings))
|
||||
Reference in New Issue
Block a user