fix(backend): prevent unbounded memory accumulation over consecutive TTS generations (#923)

This commit is contained in:
devangkantharia
2026-10-04 00:01:18 +00:00
committed by capy-ai-staging[bot]
parent 63f09455ef
commit 38389db82b
8 changed files with 88 additions and 60 deletions
+9 -3
View File
@@ -205,17 +205,23 @@ def check_cuda_compatibility() -> tuple[bool, str | None]:
def empty_device_cache(device: str) -> None:
"""
Free cached memory on the given device (CUDA or XPU).
Free cached memory and unreferenced tensors on the given device (CUDA, XPU, MPS, CPU).
Backends should call this after unloading models so VRAM is returned
to the OS.
Backends call this after model unloading and post-generation cleanup to return
memory to the OS and prevent process heap accumulation.
"""
import gc
import torch
gc.collect()
if device == "cuda" and torch.cuda.is_available():
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()
def manual_seed(seed: int, device: str) -> None:
+1
View File
@@ -203,6 +203,7 @@ class ChatterboxTTSBackend:
logger.info(f"[Chatterbox] Generating: lang={language}")
with torch.inference_mode():
wav = self.model.generate(
text,
language_id=language,
@@ -184,6 +184,7 @@ class ChatterboxTurboTTSBackend:
logger.info("[Chatterbox Turbo] Generating (English)")
with torch.inference_mode():
wav = self.model.generate(
text,
audio_prompt_path=ref_audio,
+1
View File
@@ -276,6 +276,7 @@ class KokoroTTSBackend:
# Generate all chunks and concatenate
audio_chunks = []
with torch.inference_mode():
for result in pipeline(text, voice=voice_name, speed=1.0):
if result.audio is not None:
chunk = result.audio
+3
View File
@@ -167,6 +167,9 @@ class LuxTTSBackend:
if seed is not None:
manual_seed(seed, self.device)
import torch
with torch.inference_mode():
wav = self.model.generate_speech(
text=text,
encode_dict=voice_prompt,
+1
View File
@@ -231,6 +231,7 @@ class PyTorchTTSBackend:
# See _create_prompt_sync comment — inference runs with the
# process's default HF_HUB_OFFLINE state (issue #462).
with torch.inference_mode():
wavs, sample_rate = self.model.generate_voice_clone(
text=text,
voice_clone_prompt=voice_prompt,
@@ -207,6 +207,7 @@ class QwenCustomVoiceBackend:
# state. Forcing offline here (issue #462) regressed online
# users whose libraries issue legitimate metadata lookups
# during generation.
with torch.inference_mode():
wavs, sample_rate = self.model.generate_custom_voice(**kwargs)
return wavs[0], sample_rate
+14
View File
@@ -156,6 +156,12 @@ 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
def _notify_speak_end(generation_id: str, *, status: str) -> None:
@@ -313,6 +319,7 @@ async def generate_audio_sync(
if crossfade_ms is not None:
gen_kwargs["crossfade_ms"] = crossfade_ms
try:
audio, sample_rate = await generate_chunked(
tts_model, text, voice_prompt, **gen_kwargs
)
@@ -321,6 +328,13 @@ async def generate_audio_sync(
audio = normalize_audio(audio)
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
def _save_regenerate(