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.
|
`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
|
If an answer lacks one of the decision's option ids, that is a 500 `scoring_failed`, never a
|
||||||
guess.
|
guess.
|
||||||
- **INV-2 one model, one inference at a time.** The model loads at startup, and a process-wide
|
- **INV-2 one model, one inference thread.** The model loads at startup on a dedicated
|
||||||
lock serialises every request's calls (all of a request's chunks run inside one hold). The app
|
single-thread executor. The warm-ups and every later call run on that **same host thread**,
|
||||||
runs one worker. Calls run off the event loop, so `/health` answers during one.
|
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:
|
- **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
|
- `inference.py` in the checkpoint hashes to the pinned sha256 (it is executed code, loaded
|
||||||
from a data mount);
|
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.
|
- **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
|
- **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
|
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.**
|
- **Engine against a fake torch.**
|
||||||
- An OOM is re-raised unchained, and `empty_cache` runs only after the failed call's tensors
|
- An OOM is re-raised unchained, and `empty_cache` runs only after the failed call's tensors
|
||||||
are freed.
|
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
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import hmac
|
import hmac
|
||||||
import itertools
|
import itertools
|
||||||
import json
|
import json
|
||||||
import math
|
import math
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Literal
|
from typing import Any, Literal
|
||||||
|
|
||||||
from fastapi import FastAPI, Request
|
from fastapi import FastAPI, Request
|
||||||
from fastapi.concurrency import run_in_threadpool
|
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
from pydantic import BaseModel, ConfigDict, ValidationError
|
from pydantic import BaseModel, ConfigDict, ValidationError
|
||||||
|
|
||||||
@@ -212,10 +213,19 @@ def combine(d: Decision, results: list[dict]) -> dict:
|
|||||||
"orderings": results}
|
"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")
|
app = FastAPI(title="intern-decision-serve")
|
||||||
expected = f"Bearer {settings.api_token}".encode()
|
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
|
in_progress = 0 # POSTs admitted and not yet answered
|
||||||
|
|
||||||
def admit():
|
def admit():
|
||||||
@@ -288,9 +298,9 @@ def create_app(settings: Settings, engine: Any) -> FastAPI:
|
|||||||
return out, timing
|
return out, timing
|
||||||
|
|
||||||
async def score(state: State, decisions: list[Decision], waves: list[list[Slot]]):
|
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:
|
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, ...
|
except ValueError as exc: # the model's own validation, token limit, ...
|
||||||
raise ApiError(422, "invalid_request", str(exc)) from exc
|
raise ApiError(422, "invalid_request", str(exc)) from exc
|
||||||
except OutOfMemory as exc:
|
except OutOfMemory as exc:
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ import os
|
|||||||
|
|
||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
|
|
||||||
from .app import create_app
|
from .app import create_app, inference_thread
|
||||||
from .config import Settings
|
from .config import Settings
|
||||||
|
|
||||||
|
|
||||||
@@ -17,4 +17,7 @@ def app_from_env() -> FastAPI:
|
|||||||
os.environ["TRANSFORMERS_OFFLINE"] = "1"
|
os.environ["TRANSFORMERS_OFFLINE"] = "1"
|
||||||
from .engine import TorchEngine # torch loads only inside load(), never in the unit tests
|
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()
|
engine.release.set()
|
||||||
assert [f.result().status_code for f in held] == [200, 200]
|
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
|
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()
|
app = main.app_from_env()
|
||||||
assert seen == {"offline": ("1", "1"), "token": "k" * 40}
|
assert seen == {"offline": ("1", "1"), "token": "k" * 40}
|
||||||
assert app.title == "intern-decision-serve"
|
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