mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -07:00
Review follow-ups: clear_voice_prompt_memory_cache() now runs before backend.unload_model() at all three call sites so the device-resident prompt tensors are already released when empty_device_cache() / empty_mlx_cache() runs, instead of going back into the caching allocator afterwards. empty_mlx_cache's docstring states why it may run off the MLX worker thread.
37 lines
863 B
Python
37 lines
863 B
Python
"""
|
|
TTS inference module - delegates to backend abstraction layer.
|
|
"""
|
|
|
|
from typing import Optional
|
|
import numpy as np
|
|
import io
|
|
import soundfile as sf
|
|
|
|
from ..backends import get_tts_backend, TTSBackend
|
|
from ..utils.cache import clear_voice_prompt_memory_cache
|
|
|
|
|
|
def get_tts_model() -> TTSBackend:
|
|
"""
|
|
Get TTS backend instance (MLX or PyTorch based on platform).
|
|
|
|
Returns:
|
|
TTS backend instance
|
|
"""
|
|
return get_tts_backend()
|
|
|
|
|
|
def unload_tts_model():
|
|
"""Unload TTS model to free memory."""
|
|
backend = get_tts_backend()
|
|
clear_voice_prompt_memory_cache()
|
|
backend.unload_model()
|
|
|
|
|
|
def audio_to_wav_bytes(audio: np.ndarray, sample_rate: int) -> bytes:
|
|
"""Convert audio array to WAV bytes."""
|
|
buffer = io.BytesIO()
|
|
sf.write(buffer, audio, sample_rate, format="WAV")
|
|
buffer.seek(0)
|
|
return buffer.read()
|