From 17fd1ddd1bf2aa3a5b75c96acebea69188b050c9 Mon Sep 17 00:00:00 2001 From: jamiepine <32987599+jamiepine@users.noreply.github.com> Date: Sat, 3 Oct 2026 09:32:27 +0000 Subject: [PATCH] refactor(generation): simplify post-generation cache cleanup Initialise tts_model before the try so the finally can read its device without a locals() probe, import empty_device_cache alongside the other backend imports instead of inside a bare try/except, and flatten the MPS branch in empty_device_cache. --- backend/backends/base.py | 5 ++--- backend/services/generation.py | 18 ++++++------------ 2 files changed, 8 insertions(+), 15 deletions(-) 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(