mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 09:35:16 -07:00
fix(mlx): drain the MLX pool after an in-flight op when unload landed mid-generation
Review follow-up: unload_model() runs inline on the event loop, so when it lands during a generation it only drops the backend's reference; the worker's local keeps the model alive and mx.clear_cache() finds nothing to free. Once the folded load-and-generate (or transcribe) callable finishes and releases that local, check whether the model was unloaded meanwhile and drain the pool then. Covered by a test that unloads from inside a fake model.generate().
This commit is contained in:
committed by
capy-ai-staging[bot]
parent
b1323ed0a8
commit
cffe24ffd2
@@ -80,3 +80,31 @@ def test_mlx_backends_do_not_clear_cache_when_already_unloaded(monkeypatch):
|
||||
backend.unload_model()
|
||||
|
||||
mock_clear.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mlx_tts_generate_drains_cache_when_unloaded_mid_generation(monkeypatch):
|
||||
"""An unload that lands while generate() runs must still release the pool afterwards."""
|
||||
mock_clear = MagicMock()
|
||||
monkeypatch.setattr(mlx_backend, "empty_mlx_cache", mock_clear)
|
||||
|
||||
backend = mlx_backend.MLXTTSBackend(model_size="0.6B")
|
||||
|
||||
class _FakeResult:
|
||||
audio = [0.0, 0.0]
|
||||
sample_rate = 24000
|
||||
|
||||
class _FakeModel:
|
||||
def generate(self, text, **kwargs):
|
||||
# Simulate /models/unload arriving from the event loop mid-generation.
|
||||
backend.unload_model()
|
||||
yield _FakeResult()
|
||||
|
||||
backend.model = _FakeModel()
|
||||
backend._current_model_size = "0.6B"
|
||||
|
||||
audio, sample_rate = await backend.generate("hello", {}, "en")
|
||||
|
||||
assert len(audio) == 2
|
||||
assert backend.model is None
|
||||
mock_clear.assert_called_once()
|
||||
|
||||
Reference in New Issue
Block a user