mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -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,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)
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
Reference in New Issue
Block a user