From 99d05b917c897cc241020588e8a2584ff327f586 Mon Sep 17 00:00:00 2001 From: David Strouk Date: Tue, 23 Jun 2026 16:13:39 +0300 Subject: [PATCH] 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) --- backend/tests/test_mlx_thread_affinity.py | 54 +++++++++++++++++++++++ 1 file changed, 54 insertions(+) create mode 100644 backend/tests/test_mlx_thread_affinity.py diff --git a/backend/tests/test_mlx_thread_affinity.py b/backend/tests/test_mlx_thread_affinity.py new file mode 100644 index 00000000..968c17f2 --- /dev/null +++ b/backend/tests/test_mlx_thread_affinity.py @@ -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