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.
49 lines
2.1 KiB
Python
49 lines
2.1 KiB
Python
"""Entry point: INV-5, offline before torch/transformers load. Contract: intern-decision-serve.contract.md."""
|
|
import os
|
|
|
|
from fake_engine import FakeEngine
|
|
|
|
|
|
def test_app_from_env_goes_offline_before_the_engine_loads(monkeypatch):
|
|
from intern_decision_serve import engine as engine_module
|
|
from intern_decision_serve import main
|
|
seen = {}
|
|
|
|
def fake_load(settings, **_kw):
|
|
seen["offline"] = (os.environ.get("HF_HUB_OFFLINE"), os.environ.get("TRANSFORMERS_OFFLINE"))
|
|
seen["token"] = settings.api_token
|
|
return FakeEngine()
|
|
|
|
monkeypatch.delenv("HF_HUB_OFFLINE", raising=False)
|
|
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising=False)
|
|
monkeypatch.setenv("INTERN_DECISION_API_TOKEN", "k" * 40)
|
|
monkeypatch.setattr(engine_module.TorchEngine, "load", staticmethod(fake_load))
|
|
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()
|