Spaces:
Running on Zero
fix: Keep CUDA calls off the dedicated inference thread -- again
Browse files@spaces.GPU calls failed in ZeroGPU's own worker with "RuntimeError: No CUDA
GPUs are available", while the Space genuinely had one attached
(hardware.current: zero-a10g, account is PRO). spaces intercepts CUDA only
inside the context it manages, and app.py's ThreadPoolExecutor is outside it, so
the GPU handoff never happened.
This repo already found and fixed this, in 4de16e3: "Don't route CUDA/ZeroGPU
calls through a manually-created thread." The MLX rewrite made the executor
unconditional again -- reasonably enough, since the CUDA path had been deleted at
the time -- and reintroduced the bug the moment the CUDA path came back.
USE_INFERENCE_THREAD now gates it on the backend, restoring that commit's shape:
MLX keeps its single pinned thread, because MLX genuinely needs one; torch gets
direct calls under the plain lock, and /api/chat/stream hands Starlette a sync
generator to iterate on its own threadpool instead of the queue-and-executor
relay.
Verified both ways with a TestClient: on torch neither endpoint touches a
controlai-prefixed thread, and on MLX both still run on controlai_0.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
|
@@ -222,6 +222,16 @@ generations sharing one KV cache, and splitting them across separate calls would
|
|
| 222 |
state on the far side of a process boundary each time. `_collect()` exists so `/api/chat` honours the
|
| 223 |
hook too; `ControlAgent.run()` consumes `self.stream` directly and would bypass it.
|
| 224 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 225 |
**ZeroGPU platform gotchas, each learned by having the Space fail:**
|
| 226 |
- A `@spaces.GPU` function is only detected if wired to a real Gradio event handler. One called
|
| 227 |
solely from a FastAPI route fails startup with "No @spaces.GPU function detected". Hence the
|
|
|
|
| 222 |
state on the far side of a process boundary each time. `_collect()` exists so `/api/chat` honours the
|
| 223 |
hook too; `ControlAgent.run()` consumes `self.stream` directly and would bypass it.
|
| 224 |
|
| 225 |
+
**The dedicated inference thread must never carry a CUDA call.** `spaces` intercepts CUDA only
|
| 226 |
+
inside the context it manages, and `app.py`'s `ThreadPoolExecutor` is outside it — a `@spaces.GPU`
|
| 227 |
+
call made from there fails in its own worker with `RuntimeError: No CUDA GPUs are available`, even
|
| 228 |
+
with a GPU genuinely attached (`hardware.current: zero-a10g`). `USE_INFERENCE_THREAD` gates the
|
| 229 |
+
executor on the backend: MLX keeps its single pinned thread, torch gets direct calls under the lock
|
| 230 |
+
and hands `/api/chat/stream` a plain sync generator for Starlette's own threadpool. **This was found
|
| 231 |
+
and fixed once before, in 4de16e3, and the MLX rewrite reintroduced it** — the executor was made
|
| 232 |
+
unconditional because the CUDA path had been deleted. Read that commit before touching threading
|
| 233 |
+
here.
|
| 234 |
+
|
| 235 |
**ZeroGPU platform gotchas, each learned by having the Space fail:**
|
| 236 |
- A `@spaces.GPU` function is only detected if wired to a real Gradio event handler. One called
|
| 237 |
solely from a FastAPI route fails startup with "No @spaces.GPU function detected". Hence the
|
|
@@ -49,7 +49,27 @@ for directory in (STATIC_DIR, PLOTS_DIR, UPLOADS_DIR):
|
|
| 49 |
directory.mkdir(parents=True, exist_ok=True)
|
| 50 |
|
| 51 |
# One thread, for the lifetime of the process: see the module docstring.
|
| 52 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 53 |
# The agent holds a single KV cache that every turn mutates, so turns must not
|
| 54 |
# interleave even though they all land on the same thread.
|
| 55 |
inference_lock = threading.Lock()
|
|
@@ -93,6 +113,10 @@ def get_agent() -> ControlAgent:
|
|
| 93 |
|
| 94 |
|
| 95 |
async def _on_inference_thread(fn, *args):
|
|
|
|
|
|
|
|
|
|
|
|
|
| 96 |
return await asyncio.get_running_loop().run_in_executor(inference_executor, fn, *args)
|
| 97 |
|
| 98 |
|
|
@@ -291,6 +315,29 @@ async def chat_stream(req: ChatRequest) -> StreamingResponse:
|
|
| 291 |
if not message:
|
| 292 |
raise HTTPException(status_code=400, detail="Message cannot be empty")
|
| 293 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 294 |
events: queue.Queue = queue.Queue()
|
| 295 |
sentinel = object()
|
| 296 |
|
|
@@ -314,7 +361,7 @@ async def chat_stream(req: ChatRequest) -> StreamingResponse:
|
|
| 314 |
event = await loop.run_in_executor(None, events.get)
|
| 315 |
if event is sentinel:
|
| 316 |
break
|
| 317 |
-
yield
|
| 318 |
|
| 319 |
return StreamingResponse(
|
| 320 |
relay(),
|
|
|
|
| 49 |
directory.mkdir(parents=True, exist_ok=True)
|
| 50 |
|
| 51 |
# One thread, for the lifetime of the process: see the module docstring.
|
| 52 |
+
def _uses_mlx() -> bool:
|
| 53 |
+
return os.environ.get("CONTROLAI_BACKEND", "mlx").strip().lower() not in (
|
| 54 |
+
"torch", "pytorch", "cuda",
|
| 55 |
+
)
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
# MLX keeps its compute stream in thread-local state, so every call must land on
|
| 59 |
+
# one consistent OS thread for the process lifetime.
|
| 60 |
+
#
|
| 61 |
+
# **This must never be used for the CUDA/ZeroGPU path.** `spaces` only intercepts
|
| 62 |
+
# CUDA inside the context it manages; a manually-created thread is outside it, and
|
| 63 |
+
# a @spaces.GPU call made from one fails in its own worker with
|
| 64 |
+
# "RuntimeError: No CUDA GPUs are available" even when the Space genuinely has a
|
| 65 |
+
# GPU attached. This was found once before and fixed the same way (4de16e3); the
|
| 66 |
+
# MLX rewrite reintroduced the unconditional executor and reintroduced the bug.
|
| 67 |
+
USE_INFERENCE_THREAD = _uses_mlx()
|
| 68 |
+
inference_executor = (
|
| 69 |
+
ThreadPoolExecutor(max_workers=1, thread_name_prefix="controlai")
|
| 70 |
+
if USE_INFERENCE_THREAD
|
| 71 |
+
else None
|
| 72 |
+
)
|
| 73 |
# The agent holds a single KV cache that every turn mutates, so turns must not
|
| 74 |
# interleave even though they all land on the same thread.
|
| 75 |
inference_lock = threading.Lock()
|
|
|
|
| 113 |
|
| 114 |
|
| 115 |
async def _on_inference_thread(fn, *args):
|
| 116 |
+
if not USE_INFERENCE_THREAD:
|
| 117 |
+
# Deliberately blocking: on ZeroGPU the call has to stay in the context
|
| 118 |
+
# `spaces` manages, and a demo serving one turn at a time is fine.
|
| 119 |
+
return fn(*args)
|
| 120 |
return await asyncio.get_running_loop().run_in_executor(inference_executor, fn, *args)
|
| 121 |
|
| 122 |
|
|
|
|
| 315 |
if not message:
|
| 316 |
raise HTTPException(status_code=400, detail="Message cannot be empty")
|
| 317 |
|
| 318 |
+
def _sse(event: dict) -> str:
|
| 319 |
+
return f"data: {json.dumps(event, ensure_ascii=False, default=str)}\n\n"
|
| 320 |
+
|
| 321 |
+
if not USE_INFERENCE_THREAD:
|
| 322 |
+
# ZeroGPU: hand Starlette a plain sync generator and let it iterate on
|
| 323 |
+
# its own threadpool. The queue-and-custom-executor relay below would
|
| 324 |
+
# put the @spaces.GPU call on a thread `spaces` does not manage.
|
| 325 |
+
def sync_relay():
|
| 326 |
+
try:
|
| 327 |
+
with inference_lock:
|
| 328 |
+
for event in _to_wire_events(message, req.history):
|
| 329 |
+
yield _sse(event)
|
| 330 |
+
except Exception as exc:
|
| 331 |
+
print(f"[chat] {type(exc).__name__}: {exc}")
|
| 332 |
+
yield _sse({"type": "error", "error": f"{type(exc).__name__}: {exc}"})
|
| 333 |
+
|
| 334 |
+
return StreamingResponse(
|
| 335 |
+
sync_relay(),
|
| 336 |
+
media_type="text/event-stream",
|
| 337 |
+
headers={"Cache-Control": "no-cache", "Connection": "keep-alive",
|
| 338 |
+
"X-Accel-Buffering": "no"},
|
| 339 |
+
)
|
| 340 |
+
|
| 341 |
events: queue.Queue = queue.Queue()
|
| 342 |
sentinel = object()
|
| 343 |
|
|
|
|
| 361 |
event = await loop.run_in_executor(None, events.get)
|
| 362 |
if event is sentinel:
|
| 363 |
break
|
| 364 |
+
yield _sse(event)
|
| 365 |
|
| 366 |
return StreamingResponse(
|
| 367 |
relay(),
|