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