From b788dc383cbdab01c3cf8139c738843e402838f8 Mon Sep 17 00:00:00 2001 From: jamiepine <32987599+jamiepine@users.noreply.github.com> Date: Sat, 3 Oct 2026 09:47:13 +0000 Subject: [PATCH] fix(generation): make post-generation cleanup best-effort and cover /generate/stream Wrap empty_device_cache in release_generation_memory(), which logs and swallows cleanup failures so a poisoned CUDA context cannot replace a finished generation's result, and call it from the streaming endpoint too, which drives generate_chunked directly. --- backend/routes/generations.py | 29 ++++++++++++++++------------- backend/services/generation.py | 23 +++++++++++++++++++---- 2 files changed, 35 insertions(+), 17 deletions(-) diff --git a/backend/routes/generations.py b/backend/routes/generations.py index fbbeece6..639a2902 100644 --- a/backend/routes/generations.py +++ b/backend/routes/generations.py @@ -12,7 +12,7 @@ from sqlalchemy.orm import Session from .. import config, models from ..services import history, personality, profiles, tts from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db -from ..services.generation import run_generation +from ..services.generation import release_generation_memory, run_generation from ..services.task_queue import cancel_generation as cancel_generation_job, enqueue_generation from ..utils.audio import load_audio from ..utils.tasks import get_task_manager @@ -363,18 +363,21 @@ async def stream_speech( runaway_detector = has_tts_runaway - audio, sample_rate = await generate_chunked( - tts_model, - data.text, - voice_prompt, - language=data.language, - seed=data.seed, - instruct=data.instruct, - max_chunk_chars=data.max_chunk_chars, - crossfade_ms=data.crossfade_ms, - trim_fn=trim_fn, - runaway_detector=runaway_detector, - ) + try: + audio, sample_rate = await generate_chunked( + tts_model, + data.text, + voice_prompt, + language=data.language, + seed=data.seed, + instruct=data.instruct, + max_chunk_chars=data.max_chunk_chars, + crossfade_ms=data.crossfade_ms, + trim_fn=trim_fn, + runaway_detector=runaway_detector, + ) + finally: + release_generation_memory(tts_model) effects_chain_config = None if data.effects_chain is not None: diff --git a/backend/services/generation.py b/backend/services/generation.py index 78b272bd..02e8d2fb 100644 --- a/backend/services/generation.py +++ b/backend/services/generation.py @@ -17,6 +17,7 @@ Mode differences: from __future__ import annotations import asyncio +import logging import traceback from typing import Literal, Optional @@ -26,6 +27,22 @@ from ..database import get_db from ..utils.tasks import get_task_manager +def release_generation_memory(tts_model) -> None: + """Best-effort post-generation memory cleanup. + + Collects garbage and flushes the device allocator cache so the process heap + does not grow across consecutive generations (#923). Never raises: a + cleanup failure (e.g. a poisoned CUDA context) must not replace the + generation's own result or error. + """ + from ..backends.base import empty_device_cache + + try: + empty_device_cache(getattr(tts_model, "device", "cpu")) + except Exception as e: + logging.getLogger(__name__).debug("post-generation cache cleanup failed: %s", e) + + async def run_generation( *, generation_id: str, @@ -54,7 +71,6 @@ 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 @@ -158,7 +174,7 @@ async def run_generation( finally: task_manager.complete_generation(generation_id) bg_db.close() - empty_device_cache(getattr(tts_model, "device", "cpu")) + release_generation_memory(tts_model) def _notify_speak_end(generation_id: str, *, status: str) -> None: @@ -283,7 +299,6 @@ 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 @@ -328,7 +343,7 @@ async def generate_audio_sync( return tts.audio_to_wav_bytes(audio, sample_rate) finally: - empty_device_cache(getattr(tts_model, "device", "cpu")) + release_generation_memory(tts_model) def _save_regenerate(