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
+4 -1
View File
@@ -12,7 +12,7 @@ from sqlalchemy.orm import Session
from .. import config, models from .. import config, models
from ..services import history, personality, profiles, tts from ..services import history, personality, profiles, tts
from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db 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 ..services.task_queue import cancel_generation as cancel_generation_job, enqueue_generation
from ..utils.audio import load_audio from ..utils.audio import load_audio
from ..utils.tasks import get_task_manager from ..utils.tasks import get_task_manager
@@ -363,6 +363,7 @@ async def stream_speech(
runaway_detector = has_tts_runaway runaway_detector = has_tts_runaway
try:
audio, sample_rate = await generate_chunked( audio, sample_rate = await generate_chunked(
tts_model, tts_model,
data.text, data.text,
@@ -375,6 +376,8 @@ async def stream_speech(
trim_fn=trim_fn, trim_fn=trim_fn,
runaway_detector=runaway_detector, runaway_detector=runaway_detector,
) )
finally:
release_generation_memory(tts_model)
effects_chain_config = None effects_chain_config = None
if data.effects_chain is not None: if data.effects_chain is not None:
+19 -4
View File
@@ -17,6 +17,7 @@ Mode differences:
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import logging
import traceback import traceback
from typing import Literal, Optional from typing import Literal, Optional
@@ -26,6 +27,22 @@ from ..database import get_db
from ..utils.tasks import get_task_manager 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( async def run_generation(
*, *,
generation_id: str, generation_id: str,
@@ -54,7 +71,6 @@ async def run_generation(
get_tts_backend_for_engine, get_tts_backend_for_engine,
load_engine_model, load_engine_model,
) )
from ..backends.base import empty_device_cache
from ..utils.chunked_tts import generate_chunked from ..utils.chunked_tts import generate_chunked
from ..utils.audio import has_tts_runaway, normalize_audio, save_audio, trim_tts_output from ..utils.audio import has_tts_runaway, normalize_audio, save_audio, trim_tts_output
@@ -158,7 +174,7 @@ async def run_generation(
finally: finally:
task_manager.complete_generation(generation_id) task_manager.complete_generation(generation_id)
bg_db.close() 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: def _notify_speak_end(generation_id: str, *, status: str) -> None:
@@ -283,7 +299,6 @@ async def generate_audio_sync(
get_tts_backend_for_engine, get_tts_backend_for_engine,
load_engine_model, load_engine_model,
) )
from ..backends.base import empty_device_cache
from ..utils.chunked_tts import generate_chunked from ..utils.chunked_tts import generate_chunked
from ..utils.audio import has_tts_runaway, normalize_audio, trim_tts_output from ..utils.audio import has_tts_runaway, normalize_audio, trim_tts_output
from . import tts from . import tts
@@ -328,7 +343,7 @@ async def generate_audio_sync(
return tts.audio_to_wav_bytes(audio, sample_rate) return tts.audio_to_wav_bytes(audio, sample_rate)
finally: finally:
empty_device_cache(getattr(tts_model, "device", "cpu")) release_generation_memory(tts_model)
def _save_regenerate( def _save_regenerate(