mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -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