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