diff --git a/backend/backends/base.py b/backend/backends/base.py index 71b2e011..41ffef12 100644 --- a/backend/backends/base.py +++ b/backend/backends/base.py @@ -219,9 +219,8 @@ def empty_device_cache(device: str) -> None: torch.cuda.empty_cache() elif device == "xpu" and hasattr(torch, "xpu"): torch.xpu.empty_cache() - elif device == "mps" and hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): - if hasattr(torch.mps, "empty_cache"): - torch.mps.empty_cache() + elif device == "mps" and torch.backends.mps.is_available() and hasattr(torch.mps, "empty_cache"): + torch.mps.empty_cache() def manual_seed(seed: int, device: str) -> None: diff --git a/backend/services/generation.py b/backend/services/generation.py index e1893d8d..78b272bd 100644 --- a/backend/services/generation.py +++ b/backend/services/generation.py @@ -54,11 +54,13 @@ async def run_generation( get_tts_backend_for_engine, load_engine_model, ) + from ..backends.base import empty_device_cache from ..utils.chunked_tts import generate_chunked from ..utils.audio import has_tts_runaway, normalize_audio, save_audio, trim_tts_output task_manager = get_task_manager() bg_db = next(get_db()) + tts_model = None try: tts_model = get_tts_backend_for_engine(engine) @@ -156,12 +158,7 @@ async def run_generation( finally: task_manager.complete_generation(generation_id) bg_db.close() - try: - from ..backends.base import empty_device_cache - device = getattr(tts_model, "device", "cpu") if "tts_model" in locals() else "cpu" - empty_device_cache(device) - except Exception: - pass + empty_device_cache(getattr(tts_model, "device", "cpu")) def _notify_speak_end(generation_id: str, *, status: str) -> None: @@ -286,11 +283,13 @@ async def generate_audio_sync( get_tts_backend_for_engine, load_engine_model, ) + from ..backends.base import empty_device_cache from ..utils.chunked_tts import generate_chunked from ..utils.audio import has_tts_runaway, normalize_audio, trim_tts_output from . import tts bg_db = next(get_db()) + tts_model = None try: tts_model = get_tts_backend_for_engine(engine) await load_engine_model(engine, model_size) @@ -329,12 +328,7 @@ async def generate_audio_sync( return tts.audio_to_wav_bytes(audio, sample_rate) finally: - try: - from ..backends.base import empty_device_cache - device = getattr(tts_model, "device", "cpu") if "tts_model" in locals() else "cpu" - empty_device_cache(device) - except Exception: - pass + empty_device_cache(getattr(tts_model, "device", "cpu")) def _save_regenerate(