mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 09:35:16 -07:00
fix(backends): run MLX load and inference on one thread
MLX streams are thread-local and mlx-audio caches one on the model at load time. Model load and inference were each dispatched through asyncio.to_thread(), which uses a multi-worker pool, so load and generate/transcribe could land on different OS threads -- the inference thread then has no Stream(gpu, N) and MLX aborts with "There is no Stream(gpu, 1) in current thread." Route every MLX call (load, generate, transcribe; TTS and STT) through a single dedicated worker thread so a model and its stream always share a thread. max_workers=1 also serialises the single local GPU. Adds a regression test covering the thread-affinity invariant. Fixes #699. Also addresses #675. Co-Authored-By: Claude Opus 4.8 (1M context) <[email protected]>
This commit is contained in:
committed by
capy-ai-staging[bot]
co-authored by
Claude Opus 4.8
parent
2693c983f2
commit
99d05b917c
@@ -0,0 +1,54 @@
|
||||
"""Regression tests for MLX worker-thread affinity (issues #675, #699)."""
|
||||
|
||||
import asyncio
|
||||
import threading
|
||||
|
||||
import pytest
|
||||
|
||||
from backend.backends.mlx_backend import _run_on_mlx_thread
|
||||
|
||||
|
||||
class TestMLXThreadAffinity:
|
||||
"""MLX model load and inference must stay pinned to one OS thread.
|
||||
|
||||
MLX streams are thread-local and mlx-audio caches one at load time, so a
|
||||
model loaded on one thread cannot be used from another -- inference aborts
|
||||
with "There is no Stream(gpu, 1) in current thread". These tests lock in the
|
||||
single-dedicated-thread guarantee that prevents that.
|
||||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_all_calls_run_on_one_thread(self):
|
||||
"""Concurrently dispatched MLX calls all land on the same OS thread.
|
||||
|
||||
Dispatching concurrently (not sequentially) is what makes this prove a
|
||||
single worker: with more than one worker the gather would fan out across
|
||||
threads and the id set would have more than one entry.
|
||||
"""
|
||||
thread_ids = await asyncio.gather(*[_run_on_mlx_thread(threading.get_ident) for _ in range(25)])
|
||||
|
||||
assert len(set(thread_ids)) == 1, f"MLX work spread across threads: {set(thread_ids)}"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_thread_local_stream_survives_across_calls(self):
|
||||
"""A stream cached at load time stays visible to later inference calls.
|
||||
|
||||
Mirrors mlx-audio: the stream is created on the load thread and reused
|
||||
on every generate. On a different thread the second step would raise,
|
||||
reproducing the user-facing error verbatim.
|
||||
"""
|
||||
stream_registry = threading.local()
|
||||
|
||||
def _load_model():
|
||||
stream_registry.stream = "Stream(gpu, 1)"
|
||||
|
||||
def _generate():
|
||||
stream = getattr(stream_registry, "stream", None)
|
||||
if stream is None:
|
||||
raise RuntimeError("There is no Stream(gpu, 1) in current thread.")
|
||||
return stream
|
||||
|
||||
await _run_on_mlx_thread(_load_model)
|
||||
results = [await _run_on_mlx_thread(_generate) for _ in range(10)]
|
||||
|
||||
assert results == ["Stream(gpu, 1)"] * 10
|
||||
Reference in New Issue
Block a user