fix(intern-decision-serve): load and score on one dedicated inference thread

torch keeps CUDA state per host thread (cuBLAS handles and workspaces), partly outside the
per-process VRAM cap. Scoring on anyio's threadpool let 40 threads each create it: measured on
fv-ml1 GPU 3, +252 MiB outside the cap and +326 MiB inside, which pushed the process past the
10,300 MiB GPU 1 budget. Load, warm-up and every call now run on the same single thread.
This commit is contained in:
vh
2026-09-30 09:18:46 -07:00
parent f7415db5c9
commit f21369e4ac
5 changed files with 87 additions and 11 deletions
@@ -88,9 +88,15 @@ request. A violation is a 422.
`calibration` and `input_tokens` is what `predict()` returned. The wrapper only re-keys it.
If an answer lacks one of the decision's option ids, that is a 500 `scoring_failed`, never a
guess.
- **INV-2 one model, one inference at a time.** The model loads at startup, and a process-wide
lock serialises every request's calls (all of a request's chunks run inside one hold). The app
runs one worker. Calls run off the event loop, so `/health` answers during one.
- **INV-2 one model, one inference thread.** The model loads at startup on a dedicated
single-thread executor. The warm-ups and every later call run on that **same host thread**,
never on the event loop's threadpool. A lock also serialises each request's calls, so all of a
request's chunks run inside one hold. The app runs one worker, and `/health` answers during a
call.
- **Why one thread:** torch keeps CUDA state per host thread (cuBLAS handles and workspaces),
and part of it sits outside the VRAM cap.
- **Measured on 2026-09-30, fv-ml1 GPU 3:** anyio's 40 worker threads added 252 MiB outside the
cap and 326 MiB inside it. That pushed the footprint past the GPU 1 budget.
- **INV-3 fail-closed startup.** Before the service serves, all of these must hold:
- `inference.py` in the checkpoint hashes to the pinned sha256 (it is executed code, loaded
from a data mount);
@@ -237,7 +243,7 @@ on the host, the cap is the single knob `VRAM_CAP_GIB`.
- **Engine failures.** An engine `OutOfMemory` is 503, and any other failure is 500.
- **Concurrency.** Requests are serialised: two never overlap inside the engine, and the
chunks of one request are not interleaved with another's. `/health` answers while a call is
blocked.
blocked. Every call runs on one thread, which is the thread the engine was loaded on.
- **Engine against a fake torch.**
- An OOM is re-raised unchained, and `empty_cache` runs only after the failed call's tensors
are freed.
@@ -6,17 +6,18 @@ in calls of at most 16, and the answers are re-keyed into semif's result shape.
"""
from __future__ import annotations
import asyncio
import hmac
import itertools
import json
import math
import threading
import time
from concurrent.futures import ThreadPoolExecutor
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
@@ -212,10 +213,19 @@ def combine(d: Decision, results: list[dict]) -> dict:
"orderings": results}
def create_app(settings: Settings, engine: Any) -> FastAPI:
def inference_thread() -> ThreadPoolExecutor:
"""INV-2: the ONE host thread that ever touches the model. torch keeps CUDA state per host thread
(cuBLAS handles and workspaces), partly outside the VRAM cap: measured 2026-09-30, 40 worker
threads added 252 MiB outside the cap and 326 MiB inside it."""
return ThreadPoolExecutor(max_workers=1, thread_name_prefix="inference")
def create_app(settings: Settings, engine: Any, executor: ThreadPoolExecutor | None = None) -> FastAPI:
"""`executor` must be the single thread the engine was loaded on (main.py); tests get a fresh one."""
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
executor = executor or inference_thread()
inference = threading.Lock() # INV-2: one request's calls at a time (the executor has one thread too)
in_progress = 0 # POSTs admitted and not yet answered
def admit():
@@ -288,9 +298,9 @@ def create_app(settings: Settings, engine: Any) -> FastAPI:
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."""
"""Run one request's calls on the inference thread; map their failures to contract codes."""
try:
return await run_in_threadpool(run, state, decisions, waves)
return await asyncio.get_running_loop().run_in_executor(executor, 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:
@@ -5,7 +5,7 @@ import os
from fastapi import FastAPI
from .app import create_app
from .app import create_app, inference_thread
from .config import Settings
@@ -17,4 +17,7 @@ def app_from_env() -> FastAPI:
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))
# INV-2: load, warm-up and every later call run on this one thread (per-thread CUDA state).
executor = inference_thread()
engine = executor.submit(TorchEngine.load, settings).result()
return create_app(settings, engine, executor)
@@ -432,3 +432,35 @@ def test_more_than_max_queue_requests_in_progress_get_429_busy_before_the_body_i
engine.release.set()
assert [f.result().status_code for f in held] == [200, 200]
assert make_client(FakeEngine(), max_queue=2).post("/decide", json=ROW, headers=AUTH).status_code == 200
class ThreadRecordingEngine(FakeEngine):
def __init__(self):
super().__init__()
self.threads = set()
def predict(self, request):
self.threads.add(threading.get_ident())
return super().predict(request)
def test_every_call_runs_on_one_dedicated_inference_thread():
"""torch keeps CUDA state per host thread (cuBLAS handles + workspaces), outside the VRAM cap;
measured 2026-09-30: 40 worker threads added 252 MiB outside the cap. So: one thread, always."""
engine = ThreadRecordingEngine()
with make_client(engine) as client, ThreadPoolExecutor(12) as pool:
futures = [pool.submit(client.post, "/decide/shared", json={"state": f"s{i}", "decisions": decisions(20)},
headers=AUTH) for i in range(24)]
assert all(f.result().status_code == 200 for f in futures)
assert len(engine.threads) == 1
def test_the_app_uses_the_executor_it_is_given():
from intern_decision_serve.app import create_app as build
engine = ThreadRecordingEngine()
executor = ThreadPoolExecutor(1)
home = executor.submit(threading.get_ident).result()
client = TestClient(build(Settings(api_token=TOKEN), engine, executor=executor))
assert client.post("/decide", json=ROW, headers=AUTH).status_code == 200
assert engine.threads == {home}
executor.shutdown()
@@ -21,3 +21,28 @@ def test_app_from_env_goes_offline_before_the_engine_loads(monkeypatch):
app = main.app_from_env()
assert seen == {"offline": ("1", "1"), "token": "k" * 40}
assert app.title == "intern-decision-serve"
def test_the_engine_loads_on_the_same_thread_every_call_later_runs_on(monkeypatch):
import threading
from intern_decision_serve import engine as engine_module
from intern_decision_serve import main
from fastapi.testclient import TestClient
seen = {}
class Recording(FakeEngine):
def predict(self, request):
seen.setdefault("calls", set()).add(threading.get_ident())
return super().predict(request)
def fake_load(settings, **_kw):
seen["load"] = threading.get_ident()
return Recording()
monkeypatch.setenv("INTERN_DECISION_API_TOKEN", "k" * 40)
monkeypatch.setattr(engine_module.TorchEngine, "load", staticmethod(fake_load))
client = TestClient(main.app_from_env())
body = {"id": "r", "state": "s", "question": "q?", "options": [{"id": "a", "description": "A"},
{"id": "b", "description": "B"}]}
assert client.post("/decide", json=body, headers={"Authorization": "Bearer " + "k" * 40}).status_code == 200
assert seen["calls"] == {seen["load"]} and seen["load"] != threading.get_ident()