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:
+15 -14
View File
@@ -203,21 +203,22 @@ class ChatterboxTTSBackend:
logger.info(f"[Chatterbox] Generating: lang={language}")
wav = self.model.generate(
text,
language_id=language,
audio_prompt_path=ref_audio,
exaggeration=lang_defaults["exaggeration"],
cfg_weight=lang_defaults["cfg_weight"],
temperature=lang_defaults["temperature"],
repetition_penalty=lang_defaults["repetition_penalty"],
)
with torch.inference_mode():
wav = self.model.generate(
text,
language_id=language,
audio_prompt_path=ref_audio,
exaggeration=lang_defaults["exaggeration"],
cfg_weight=lang_defaults["cfg_weight"],
temperature=lang_defaults["temperature"],
repetition_penalty=lang_defaults["repetition_penalty"],
)
# Convert tensor -> numpy
if isinstance(wav, torch.Tensor):
audio = wav.squeeze().cpu().numpy().astype(np.float32)
else:
audio = np.asarray(wav, dtype=np.float32)
# Convert tensor -> numpy
if isinstance(wav, torch.Tensor):
audio = wav.squeeze().cpu().numpy().astype(np.float32)
else:
audio = np.asarray(wav, dtype=np.float32)
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)")
wav = self.model.generate(
text,
audio_prompt_path=ref_audio,
temperature=0.8,
top_k=1000,
top_p=0.95,
repetition_penalty=1.2,
)
with torch.inference_mode():
wav = self.model.generate(
text,
audio_prompt_path=ref_audio,
temperature=0.8,
top_k=1000,
top_p=0.95,
repetition_penalty=1.2,
)
# Convert tensor -> numpy
if isinstance(wav, torch.Tensor):
audio = wav.squeeze().cpu().numpy().astype(np.float32)
else:
audio = np.asarray(wav, dtype=np.float32)
# Convert tensor -> numpy
if isinstance(wav, torch.Tensor):
audio = wav.squeeze().cpu().numpy().astype(np.float32)
else:
audio = np.asarray(wav, dtype=np.float32)
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
audio_chunks = []
for result in pipeline(text, voice=voice_name, speed=1.0):
if result.audio is not None:
chunk = result.audio
if isinstance(chunk, torch.Tensor):
chunk = chunk.detach().cpu().numpy()
audio_chunks.append(chunk.squeeze())
with torch.inference_mode():
for result in pipeline(text, voice=voice_name, speed=1.0):
if result.audio is not None:
chunk = result.audio
if isinstance(chunk, torch.Tensor):
chunk = chunk.detach().cpu().numpy()
audio_chunks.append(chunk.squeeze())
if not audio_chunks:
# Return 1 second of silence as fallback
+14 -11
View File
@@ -167,18 +167,21 @@ class LuxTTSBackend:
if seed is not None:
manual_seed(seed, self.device)
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
)
import torch
# LuxTTS returns a tensor (may be on GPU/MPS), move to CPU first
audio = wav.detach().cpu().numpy().squeeze()
with torch.inference_mode():
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 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
# process's default HF_HUB_OFFLINE state (issue #462).
wavs, sample_rate = self.model.generate_voice_clone(
text=text,
voice_clone_prompt=voice_prompt,
language=LANGUAGE_CODE_TO_NAME.get(language, "auto"),
instruct=instruct,
)
with torch.inference_mode():
wavs, sample_rate = self.model.generate_voice_clone(
text=text,
voice_clone_prompt=voice_prompt,
language=LANGUAGE_CODE_TO_NAME.get(language, "auto"),
instruct=instruct,
)
return wavs[0], sample_rate
# 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
# users whose libraries issue legitimate metadata lookups
# 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
audio, sample_rate = await asyncio.to_thread(_generate_sync)
+20 -6
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,14 +319,22 @@ async def generate_audio_sync(
if crossfade_ms is not None:
gen_kwargs["crossfade_ms"] = crossfade_ms
audio, sample_rate = await generate_chunked(
tts_model, text, voice_prompt, **gen_kwargs
)
try:
audio, sample_rate = await generate_chunked(
tts_model, text, voice_prompt, **gen_kwargs
)
if normalize:
audio = normalize_audio(audio)
if normalize:
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(