diff --git a/backend/backends/base.py b/backend/backends/base.py index 70ec11ef3..d5fd8532f 100644 --- a/backend/backends/base.py +++ b/backend/backends/base.py @@ -171,17 +171,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: diff --git a/backend/backends/chatterbox_backend.py b/backend/backends/chatterbox_backend.py index e7a025b38..8ea2172dd 100644 --- a/backend/backends/chatterbox_backend.py +++ b/backend/backends/chatterbox_backend.py @@ -203,21 +203,22 @@ def _generate_sync(): 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"], - ) - - # Convert tensor -> numpy - if isinstance(wav, torch.Tensor): - audio = wav.squeeze().cpu().numpy().astype(np.float32) - else: - audio = np.asarray(wav, dtype=np.float32) + 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) sample_rate = getattr(self.model, "sr", None) or getattr(self.model, "sample_rate", 24000) diff --git a/backend/backends/chatterbox_turbo_backend.py b/backend/backends/chatterbox_turbo_backend.py index 6f7d6b94f..29be69e8d 100644 --- a/backend/backends/chatterbox_turbo_backend.py +++ b/backend/backends/chatterbox_turbo_backend.py @@ -184,20 +184,21 @@ def _generate_sync(): 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, - ) - - # Convert tensor -> numpy - if isinstance(wav, torch.Tensor): - audio = wav.squeeze().cpu().numpy().astype(np.float32) - else: - audio = np.asarray(wav, dtype=np.float32) + 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) sample_rate = getattr(self.model, "sr", None) or getattr(self.model, "sample_rate", 24000) diff --git a/backend/backends/kokoro_backend.py b/backend/backends/kokoro_backend.py index 59005f0a7..0f6ff39ec 100644 --- a/backend/backends/kokoro_backend.py +++ b/backend/backends/kokoro_backend.py @@ -276,12 +276,13 @@ def _generate_sync(): # 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 diff --git a/backend/backends/luxtts_backend.py b/backend/backends/luxtts_backend.py index 7f15686af..55be18b24 100644 --- a/backend/backends/luxtts_backend.py +++ b/backend/backends/luxtts_backend.py @@ -167,18 +167,21 @@ def _generate_sync(): 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 + + 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() + # 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) diff --git a/backend/backends/pytorch_backend.py b/backend/backends/pytorch_backend.py index f8ae79b86..faca0d6ab 100644 --- a/backend/backends/pytorch_backend.py +++ b/backend/backends/pytorch_backend.py @@ -231,12 +231,13 @@ def _generate_sync(): # 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 diff --git a/backend/backends/qwen_custom_voice_backend.py b/backend/backends/qwen_custom_voice_backend.py index 74f739bba..791f800b0 100644 --- a/backend/backends/qwen_custom_voice_backend.py +++ b/backend/backends/qwen_custom_voice_backend.py @@ -207,7 +207,8 @@ def _generate_sync(): # 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) diff --git a/backend/services/generation.py b/backend/services/generation.py index a4b2e8a3f..e1893d8d2 100644 --- a/backend/services/generation.py +++ b/backend/services/generation.py @@ -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(