mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -07:00
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]>
55 lines
2.1 KiB
Python
55 lines
2.1 KiB
Python
"""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
|