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.
This commit is contained in:
jamiepine
2026-10-04 00:01:18 +00:00
committed by capy-ai-staging[bot]
parent 38389db82b
commit 17fd1ddd1b
2 changed files with 8 additions and 15 deletions
+2 -3
View File
@@ -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:
+6 -12
View File
@@ -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(