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
@@ -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))