mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -07:00
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:
committed by
capy-ai-staging[bot]
parent
17fd1ddd1b
commit
b788dc383c
@@ -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,18 +363,21 @@ async def stream_speech(
|
|||||||
|
|
||||||
runaway_detector = has_tts_runaway
|
runaway_detector = has_tts_runaway
|
||||||
|
|
||||||
audio, sample_rate = await generate_chunked(
|
try:
|
||||||
tts_model,
|
audio, sample_rate = await generate_chunked(
|
||||||
data.text,
|
tts_model,
|
||||||
voice_prompt,
|
data.text,
|
||||||
language=data.language,
|
voice_prompt,
|
||||||
seed=data.seed,
|
language=data.language,
|
||||||
instruct=data.instruct,
|
seed=data.seed,
|
||||||
max_chunk_chars=data.max_chunk_chars,
|
instruct=data.instruct,
|
||||||
crossfade_ms=data.crossfade_ms,
|
max_chunk_chars=data.max_chunk_chars,
|
||||||
trim_fn=trim_fn,
|
crossfade_ms=data.crossfade_ms,
|
||||||
runaway_detector=runaway_detector,
|
trim_fn=trim_fn,
|
||||||
)
|
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:
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
Reference in New Issue
Block a user