From f21369e4ac52e8d7b7efbcff7935152b3f6a54a3 Mon Sep 17 00:00:00 2001 From: Vuong Hoang Date: Wed, 30 Sep 2026 09:18:46 -0700 Subject: [PATCH] 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. --- .../intern-decision-serve.contract.md | 14 +++++--- .../src/intern_decision_serve/app.py | 20 +++++++++--- .../src/intern_decision_serve/main.py | 7 ++-- .../intern-decision-serve/tests/test_app.py | 32 +++++++++++++++++++ .../intern-decision-serve/tests/test_main.py | 25 +++++++++++++++ 5 files changed, 87 insertions(+), 11 deletions(-) diff --git a/services/intern-decision-serve/intern-decision-serve.contract.md b/services/intern-decision-serve/intern-decision-serve.contract.md index 31c180e..aa368f6 100644 --- a/services/intern-decision-serve/intern-decision-serve.contract.md +++ b/services/intern-decision-serve/intern-decision-serve.contract.md @@ -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. diff --git a/services/intern-decision-serve/src/intern_decision_serve/app.py b/services/intern-decision-serve/src/intern_decision_serve/app.py index 245e0e6..2d960a1 100644 --- a/services/intern-decision-serve/src/intern_decision_serve/app.py +++ b/services/intern-decision-serve/src/intern_decision_serve/app.py @@ -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: diff --git a/services/intern-decision-serve/src/intern_decision_serve/main.py b/services/intern-decision-serve/src/intern_decision_serve/main.py index cfacac3..63beba5 100644 --- a/services/intern-decision-serve/src/intern_decision_serve/main.py +++ b/services/intern-decision-serve/src/intern_decision_serve/main.py @@ -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) diff --git a/services/intern-decision-serve/tests/test_app.py b/services/intern-decision-serve/tests/test_app.py index ebc8283..cf98f7c 100644 --- a/services/intern-decision-serve/tests/test_app.py +++ b/services/intern-decision-serve/tests/test_app.py @@ -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() diff --git a/services/intern-decision-serve/tests/test_main.py b/services/intern-decision-serve/tests/test_main.py index 5700b1f..a8f7eec 100644 --- a/services/intern-decision-serve/tests/test_main.py +++ b/services/intern-decision-serve/tests/test_main.py @@ -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()