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,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)
|
||||
Reference in New Issue
Block a user