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: 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 Backends call this after model unloading and post-generation cleanup to return
to the OS. memory to the OS and prevent process heap accumulation.
""" """
import gc
import torch import torch
gc.collect()
if device == "cuda" and torch.cuda.is_available(): if device == "cuda" and torch.cuda.is_available():
torch.cuda.empty_cache() torch.cuda.empty_cache()
elif device == "xpu" and hasattr(torch, "xpu"): elif device == "xpu" and hasattr(torch, "xpu"):
torch.xpu.empty_cache() 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: def manual_seed(seed: int, device: str) -> None:
+15 -14
View File
@@ -203,21 +203,22 @@ class ChatterboxTTSBackend:
logger.info(f"[Chatterbox] Generating: lang={language}") logger.info(f"[Chatterbox] Generating: lang={language}")
wav = self.model.generate( with torch.inference_mode():
text, wav = self.model.generate(
language_id=language, text,
audio_prompt_path=ref_audio, language_id=language,
exaggeration=lang_defaults["exaggeration"], audio_prompt_path=ref_audio,
cfg_weight=lang_defaults["cfg_weight"], exaggeration=lang_defaults["exaggeration"],
temperature=lang_defaults["temperature"], cfg_weight=lang_defaults["cfg_weight"],
repetition_penalty=lang_defaults["repetition_penalty"], temperature=lang_defaults["temperature"],
) repetition_penalty=lang_defaults["repetition_penalty"],
)
# Convert tensor -> numpy # Convert tensor -> numpy
if isinstance(wav, torch.Tensor): if isinstance(wav, torch.Tensor):
audio = wav.squeeze().cpu().numpy().astype(np.float32) audio = wav.squeeze().cpu().numpy().astype(np.float32)
else: else:
audio = np.asarray(wav, dtype=np.float32) audio = np.asarray(wav, dtype=np.float32)
sample_rate = getattr(self.model, "sr", None) or getattr(self.model, "sample_rate", 24000) sample_rate = getattr(self.model, "sr", None) or getattr(self.model, "sample_rate", 24000)
+14 -13
View File
@@ -184,20 +184,21 @@ class ChatterboxTurboTTSBackend:
logger.info("[Chatterbox Turbo] Generating (English)") logger.info("[Chatterbox Turbo] Generating (English)")
wav = self.model.generate( with torch.inference_mode():
text, wav = self.model.generate(
audio_prompt_path=ref_audio, text,
temperature=0.8, audio_prompt_path=ref_audio,
top_k=1000, temperature=0.8,
top_p=0.95, top_k=1000,
repetition_penalty=1.2, top_p=0.95,
) repetition_penalty=1.2,
)
# Convert tensor -> numpy # Convert tensor -> numpy
if isinstance(wav, torch.Tensor): if isinstance(wav, torch.Tensor):
audio = wav.squeeze().cpu().numpy().astype(np.float32) audio = wav.squeeze().cpu().numpy().astype(np.float32)
else: else:
audio = np.asarray(wav, dtype=np.float32) audio = np.asarray(wav, dtype=np.float32)
sample_rate = getattr(self.model, "sr", None) or getattr(self.model, "sample_rate", 24000) sample_rate = getattr(self.model, "sr", None) or getattr(self.model, "sample_rate", 24000)
+7 -6
View File
@@ -276,12 +276,13 @@ class KokoroTTSBackend:
# Generate all chunks and concatenate # Generate all chunks and concatenate
audio_chunks = [] audio_chunks = []
for result in pipeline(text, voice=voice_name, speed=1.0): with torch.inference_mode():
if result.audio is not None: for result in pipeline(text, voice=voice_name, speed=1.0):
chunk = result.audio if result.audio is not None:
if isinstance(chunk, torch.Tensor): chunk = result.audio
chunk = chunk.detach().cpu().numpy() if isinstance(chunk, torch.Tensor):
audio_chunks.append(chunk.squeeze()) chunk = chunk.detach().cpu().numpy()
audio_chunks.append(chunk.squeeze())
if not audio_chunks: if not audio_chunks:
# Return 1 second of silence as fallback # Return 1 second of silence as fallback
+14 -11
View File
@@ -167,18 +167,21 @@ class LuxTTSBackend:
if seed is not None: if seed is not None:
manual_seed(seed, self.device) manual_seed(seed, self.device)
wav = self.model.generate_speech( import torch
text=text,
encode_dict=voice_prompt,
num_steps=4,
guidance_scale=3.0,
t_shift=0.5,
speed=1.0,
return_smooth=False, # 48kHz output
)
# LuxTTS returns a tensor (may be on GPU/MPS), move to CPU first with torch.inference_mode():
audio = wav.detach().cpu().numpy().squeeze() wav = self.model.generate_speech(
text=text,
encode_dict=voice_prompt,
num_steps=4,
guidance_scale=3.0,
t_shift=0.5,
speed=1.0,
return_smooth=False, # 48kHz output
)
# LuxTTS returns a tensor (may be on GPU/MPS), move to CPU first
audio = wav.detach().cpu().numpy().squeeze()
return audio, 48000 return audio, 48000
return await asyncio.to_thread(_generate_sync) return await asyncio.to_thread(_generate_sync)
+7 -6
View File
@@ -231,12 +231,13 @@ class PyTorchTTSBackend:
# See _create_prompt_sync comment — inference runs with the # See _create_prompt_sync comment — inference runs with the
# process's default HF_HUB_OFFLINE state (issue #462). # process's default HF_HUB_OFFLINE state (issue #462).
wavs, sample_rate = self.model.generate_voice_clone( with torch.inference_mode():
text=text, wavs, sample_rate = self.model.generate_voice_clone(
voice_clone_prompt=voice_prompt, text=text,
language=LANGUAGE_CODE_TO_NAME.get(language, "auto"), voice_clone_prompt=voice_prompt,
instruct=instruct, language=LANGUAGE_CODE_TO_NAME.get(language, "auto"),
) instruct=instruct,
)
return wavs[0], sample_rate return wavs[0], sample_rate
# Run blocking inference in thread pool to avoid blocking event loop # Run blocking inference in thread pool to avoid blocking event loop
@@ -207,7 +207,8 @@ class QwenCustomVoiceBackend:
# state. Forcing offline here (issue #462) regressed online # state. Forcing offline here (issue #462) regressed online
# users whose libraries issue legitimate metadata lookups # users whose libraries issue legitimate metadata lookups
# during generation. # during generation.
wavs, sample_rate = self.model.generate_custom_voice(**kwargs) with torch.inference_mode():
wavs, sample_rate = self.model.generate_custom_voice(**kwargs)
return wavs[0], sample_rate return wavs[0], sample_rate
audio, sample_rate = await asyncio.to_thread(_generate_sync) audio, sample_rate = await asyncio.to_thread(_generate_sync)
+20 -6
View File
@@ -156,6 +156,12 @@ async def run_generation(
finally: finally:
task_manager.complete_generation(generation_id) task_manager.complete_generation(generation_id)
bg_db.close() 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: def _notify_speak_end(generation_id: str, *, status: str) -> None:
@@ -313,14 +319,22 @@ async def generate_audio_sync(
if crossfade_ms is not None: if crossfade_ms is not None:
gen_kwargs["crossfade_ms"] = crossfade_ms gen_kwargs["crossfade_ms"] = crossfade_ms
audio, sample_rate = await generate_chunked( try:
tts_model, text, voice_prompt, **gen_kwargs audio, sample_rate = await generate_chunked(
) tts_model, text, voice_prompt, **gen_kwargs
)
if normalize: if normalize:
audio = normalize_audio(audio) audio = normalize_audio(audio)
return tts.audio_to_wav_bytes(audio, sample_rate) 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( def _save_regenerate(