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:
|
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:
|
||||||
|
|||||||
@@ -203,6 +203,7 @@ class ChatterboxTTSBackend:
|
|||||||
|
|
||||||
logger.info(f"[Chatterbox] Generating: lang={language}")
|
logger.info(f"[Chatterbox] Generating: lang={language}")
|
||||||
|
|
||||||
|
with torch.inference_mode():
|
||||||
wav = self.model.generate(
|
wav = self.model.generate(
|
||||||
text,
|
text,
|
||||||
language_id=language,
|
language_id=language,
|
||||||
|
|||||||
@@ -184,6 +184,7 @@ class ChatterboxTurboTTSBackend:
|
|||||||
|
|
||||||
logger.info("[Chatterbox Turbo] Generating (English)")
|
logger.info("[Chatterbox Turbo] Generating (English)")
|
||||||
|
|
||||||
|
with torch.inference_mode():
|
||||||
wav = self.model.generate(
|
wav = self.model.generate(
|
||||||
text,
|
text,
|
||||||
audio_prompt_path=ref_audio,
|
audio_prompt_path=ref_audio,
|
||||||
|
|||||||
@@ -276,6 +276,7 @@ class KokoroTTSBackend:
|
|||||||
|
|
||||||
# Generate all chunks and concatenate
|
# Generate all chunks and concatenate
|
||||||
audio_chunks = []
|
audio_chunks = []
|
||||||
|
with torch.inference_mode():
|
||||||
for result in pipeline(text, voice=voice_name, speed=1.0):
|
for result in pipeline(text, voice=voice_name, speed=1.0):
|
||||||
if result.audio is not None:
|
if result.audio is not None:
|
||||||
chunk = result.audio
|
chunk = result.audio
|
||||||
|
|||||||
@@ -167,6 +167,9 @@ class LuxTTSBackend:
|
|||||||
if seed is not None:
|
if seed is not None:
|
||||||
manual_seed(seed, self.device)
|
manual_seed(seed, self.device)
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
with torch.inference_mode():
|
||||||
wav = self.model.generate_speech(
|
wav = self.model.generate_speech(
|
||||||
text=text,
|
text=text,
|
||||||
encode_dict=voice_prompt,
|
encode_dict=voice_prompt,
|
||||||
|
|||||||
@@ -231,6 +231,7 @@ 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).
|
||||||
|
with torch.inference_mode():
|
||||||
wavs, sample_rate = self.model.generate_voice_clone(
|
wavs, sample_rate = self.model.generate_voice_clone(
|
||||||
text=text,
|
text=text,
|
||||||
voice_clone_prompt=voice_prompt,
|
voice_clone_prompt=voice_prompt,
|
||||||
|
|||||||
@@ -207,6 +207,7 @@ 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.
|
||||||
|
with torch.inference_mode():
|
||||||
wavs, sample_rate = self.model.generate_custom_voice(**kwargs)
|
wavs, sample_rate = self.model.generate_custom_voice(**kwargs)
|
||||||
return wavs[0], sample_rate
|
return wavs[0], sample_rate
|
||||||
|
|
||||||
|
|||||||
@@ -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,6 +319,7 @@ 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
|
||||||
|
|
||||||
|
try:
|
||||||
audio, sample_rate = await generate_chunked(
|
audio, sample_rate = await generate_chunked(
|
||||||
tts_model, text, voice_prompt, **gen_kwargs
|
tts_model, text, voice_prompt, **gen_kwargs
|
||||||
)
|
)
|
||||||
@@ -321,6 +328,13 @@ async def generate_audio_sync(
|
|||||||
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(
|
||||||
|
|||||||
Reference in New Issue
Block a user