mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 09:35:16 -07:00
fix(backend): prevent unbounded memory accumulation over consecutive TTS generations (#923)
This commit is contained in:
committed by
capy-ai-staging[bot]
parent
63f09455ef
commit
38389db82b
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user