From 5b0ccd9f806b6d6a38673c7a263a2c1a600f1cf5 Mon Sep 17 00:00:00 2001 From: Fill Date: Sat, 31 May 2025 09:01:09 +0900 Subject: [PATCH] added mac support + keep model loaded --- chatterbox_node.py | 280 +++++++++++++++++++++++++++++++++++++++------ 1 file changed, 245 insertions(+), 35 deletions(-) diff --git a/chatterbox_node.py b/chatterbox_node.py index 8660b10..7bf07aa 100644 --- a/chatterbox_node.py +++ b/chatterbox_node.py @@ -9,14 +9,25 @@ from typing import Optional from chatterbox.tts import ChatterboxTTS from chatterbox.vc import ChatterboxVC -# Monkey patch torch.load to always use CPU if needed +from comfy.utils import ProgressBar + +# Monkey patch torch.load to use MPS or CPU if map_location is not specified original_torch_load = torch.load def patched_torch_load(*args, **kwargs): if 'map_location' not in kwargs: - kwargs['map_location'] = torch.device('cpu') + # Determine the appropriate device (MPS for Mac, else CPU) + if torch.backends.mps.is_available(): + device = "mps" + elif torch.cuda.is_available(): + device = "cuda" + else: + device = "cpu" + kwargs['map_location'] = torch.device(device) return original_torch_load(*args, **kwargs) + torch.load = patched_torch_load + class AudioNodeBase: """Base class for audio nodes with common utilities.""" @@ -35,6 +46,8 @@ class FL_ChatterboxTTSNode(AudioNodeBase): """ ComfyUI node for Chatterbox Text-to-Speech functionality. """ + _tts_model = None + _tts_device = None @classmethod def INPUT_TYPES(cls): @@ -48,6 +61,7 @@ class FL_ChatterboxTTSNode(AudioNodeBase): "optional": { "audio_prompt": ("AUDIO",), "use_cpu": ("BOOLEAN", {"default": False}), + "keep_model_loaded": ("BOOLEAN", {"default": False}), } } @@ -56,7 +70,7 @@ class FL_ChatterboxTTSNode(AudioNodeBase): FUNCTION = "generate_speech" CATEGORY = "ChatterBox" - def generate_speech(self, text, exaggeration, cfg_weight, temperature, audio_prompt=None, use_cpu=False): + def generate_speech(self, text, exaggeration, cfg_weight, temperature, audio_prompt=None, use_cpu=False, keep_model_loaded=False): """ Generate speech from text. @@ -67,16 +81,21 @@ class FL_ChatterboxTTSNode(AudioNodeBase): temperature: Controls randomness in generation (0.05-5.0). audio_prompt: AUDIO object containing the reference voice for TTS voice cloning. use_cpu: If True, forces CPU usage even if CUDA is available. + keep_model_loaded: If True, keeps the model loaded in memory after generation. Returns: Tuple of (audio, message) """ # Determine device to use - device = "cpu" if use_cpu else ("cuda" if torch.cuda.is_available() else "cpu") + device = "cpu" if use_cpu else ("mps" if torch.backends.mps.is_available() else ("cuda" if torch.cuda.is_available() else "cpu")) if use_cpu: - message = "Using CPU for inference (CUDA disabled)" + message = "Using CPU for inference (GPU disabled)" + elif torch.backends.mps.is_available() and device == "mps": + message = "Using MPS (Mac GPU) for inference" + elif torch.cuda.is_available() and device == "cuda": + message = "Using CUDA (NVIDIA GPU) for inference" else: - message = f"Using {device} for inference" + message = f"Using {device} for inference" # Should be CPU if no GPU found # Create temporary files for any audio inputs import tempfile @@ -105,16 +124,44 @@ class FL_ChatterboxTTSNode(AudioNodeBase): message += f"\nError creating audio prompt file: {str(e)}" audio_prompt_path = None + tts_model = None + wav = None # Initialize wav to None + audio_data = {"waveform": torch.zeros((1, 2, 1)), "sample_rate": 16000} # Initialize with empty audio + pbar = ProgressBar(100) # Simple progress bar for overall process try: - # Load the TTS model - message += f"\nLoading TTS model on {device}..." - tts_model = ChatterboxTTS.from_pretrained(device=device) - + # Load the TTS model or reuse if loaded and device matches + if FL_ChatterboxTTSNode._tts_model is not None and FL_ChatterboxTTSNode._tts_device == device: + tts_model = FL_ChatterboxTTSNode._tts_model + message += f"\nReusing loaded TTS model on {device}..." + else: + if FL_ChatterboxTTSNode._tts_model is not None: + message += f"\nUnloading previous TTS model (device mismatch or keep_model_loaded is False)..." + FL_ChatterboxTTSNode._tts_model = None + FL_ChatterboxTTSNode._tts_device = None + if torch.cuda.is_available(): + torch.cuda.empty_cache() # Clear CUDA cache if possible + if torch.backends.mps.is_available(): + torch.mps.empty_cache() # Clear MPS cache if possible + + + message += f"\nLoading TTS model on {device}..." + pbar.update_absolute(10) # Indicate model loading started + tts_model = ChatterboxTTS.from_pretrained(device=device) + pbar.update_absolute(50) # Indicate model loading finished + + if keep_model_loaded: + FL_ChatterboxTTSNode._tts_model = tts_model + FL_ChatterboxTTSNode._tts_device = device + message += "\nModel will be kept loaded in memory." + else: + message += "\nModel will be unloaded after use." + # Generate speech message += f"\nGenerating speech for: {text[:50]}..." if len(text) > 50 else f"\nGenerating speech for: {text}" if audio_prompt_path: message += f"\nUsing audio prompt: {audio_prompt_path}" + pbar.update_absolute(60) # Indicate generation started wav = tts_model.generate( text=text, audio_prompt_path=audio_prompt_path, @@ -122,13 +169,46 @@ class FL_ChatterboxTTSNode(AudioNodeBase): cfg_weight=cfg_weight, temperature=temperature, ) + pbar.update_absolute(90) # Indicate generation finished + + audio_data = { + "waveform": wav.unsqueeze(0), # Add batch dimension + "sample_rate": tts_model.sr + } + message += f"\nSpeech generated successfully" + return (audio_data, message) except RuntimeError as e: - if "CUDA" in str(e) and device != "cpu": + # Check for CUDA or MPS errors and attempt fallback to CPU + error_str = str(e) + fallback_to_cpu = False + if "CUDA" in error_str and device == "cuda": message += "\nCUDA error detected during TTS. Falling back to CPU..." - # Try again with CPU + fallback_to_cpu = True + elif "MPS" in error_str and device == "mps": + message += "\nMPS error detected during TTS. Falling back to CPU..." + fallback_to_cpu = True + + if fallback_to_cpu: device = "cpu" + # Unload previous model if it exists + if FL_ChatterboxTTSNode._tts_model is not None: + message += f"\nUnloading previous TTS model..." + FL_ChatterboxTTSNode._tts_model = None + FL_ChatterboxTTSNode._tts_device = None + if torch.cuda.is_available(): + torch.cuda.empty_cache() # Clear CUDA cache if possible + if torch.backends.mps.is_available(): + torch.mps.empty_cache() # Clear MPS cache if possible + + + message += f"\nLoading TTS model on {device}..." + pbar.update_absolute(10) # Indicate model loading started (fallback) tts_model = ChatterboxTTS.from_pretrained(device=device) + pbar.update_absolute(50) # Indicate model loading finished (fallback) + # Note: keep_model_loaded logic is applied after successful generation + # to avoid keeping a failed model loaded. + wav = tts_model.generate( text=text, audio_prompt_path=audio_prompt_path, @@ -136,29 +216,64 @@ class FL_ChatterboxTTSNode(AudioNodeBase): cfg_weight=cfg_weight, temperature=temperature, ) + pbar.update_absolute(90) # Indicate generation finished (fallback) + audio_data = { + "waveform": wav.unsqueeze(0), # Add batch dimension + "sample_rate": tts_model.sr + } + message += f"\nSpeech generated successfully after fallback." + return (audio_data, message) else: - # Re-raise if it's not a CUDA error or we're already on CPU message += f"\nError during TTS: {str(e)}" - # Return empty audio data - empty_audio = {"waveform": torch.zeros((1, 2, 1)), "sample_rate": 16000} - # Clean up any temporary files - for temp_file in temp_files: - if os.path.exists(temp_file): - os.unlink(temp_file) - return (empty_audio, message) + return (audio_data, message) + except Exception as e: + message += f"\nAn unexpected error occurred during TTS: {str(e)}" + return (audio_data, message) finally: # Clean up all temporary files for temp_file in temp_files: if os.path.exists(temp_file): os.unlink(temp_file) - + # If keep_model_loaded is False, ensure model is not stored + # This is done here to ensure model is only kept if generation was successful + if not keep_model_loaded and FL_ChatterboxTTSNode._tts_model is not None: + message += "\nUnloading TTS model as keep_model_loaded is False." + FL_ChatterboxTTSNode._tts_model = None + FL_ChatterboxTTSNode._tts_device = None + if torch.cuda.is_available(): + torch.cuda.empty_cache() # Clear CUDA cache if possible + if torch.backends.mps.is_available(): + torch.mps.empty_cache() # Clear MPS cache if possible + + pbar.update_absolute(100) # Ensure progress bar completes on success or error + return (audio_data, message) # Fallback return, should ideally not be reached + + + # If generation was successful and keep_model_loaded is True, store the model + if keep_model_loaded and tts_model is not None: + FL_ChatterboxTTSNode._tts_model = tts_model + FL_ChatterboxTTSNode._tts_device = device + message += "\nModel will be kept loaded in memory." + elif not keep_model_loaded and FL_ChatterboxTTSNode._tts_model is not None: + # This case handles successful generation when keep_model_loaded was True previously + # but is now False. Ensure the model is unloaded. + message += "\nUnloading TTS model as keep_model_loaded is now False." + FL_ChatterboxTTSNode._tts_model = None + FL_ChatterboxTTSNode._tts_device = None + if torch.cuda.is_available(): + torch.cuda.empty_cache() # Clear CUDA cache if possible + if torch.backends.mps.is_available(): + torch.mps.empty_cache() # Clear MPS cache if possible + + # Create audio data structure for the output audio_data = { "waveform": wav.unsqueeze(0), # Add batch dimension - "sample_rate": tts_model.sr + "sample_rate": tts_model.sr if tts_model else 16000 # Use default sample rate if model loading failed } message += f"\nSpeech generated successfully" + pbar.update_absolute(100) # Ensure progress bar completes on success return (audio_data, message) @@ -167,6 +282,8 @@ class FL_ChatterboxVCNode(AudioNodeBase): """ ComfyUI node for Chatterbox Voice Conversion functionality. """ + _vc_model = None + _vc_device = None @classmethod def INPUT_TYPES(cls): @@ -177,6 +294,7 @@ class FL_ChatterboxVCNode(AudioNodeBase): }, "optional": { "use_cpu": ("BOOLEAN", {"default": False}), + "keep_model_loaded": ("BOOLEAN", {"default": False}), } } @@ -185,7 +303,7 @@ class FL_ChatterboxVCNode(AudioNodeBase): FUNCTION = "convert_voice" CATEGORY = "ChatterBox" - def convert_voice(self, input_audio, target_voice, use_cpu=False): + def convert_voice(self, input_audio, target_voice, use_cpu=False, keep_model_loaded=False): """ Convert the voice in an audio file to match a target voice. @@ -193,16 +311,21 @@ class FL_ChatterboxVCNode(AudioNodeBase): input_audio: AUDIO object containing the audio to convert. target_voice: AUDIO object containing the target voice. use_cpu: If True, forces CPU usage even if CUDA is available. + keep_model_loaded: If True, keeps the model loaded in memory after conversion. Returns: Tuple of (audio, message) """ # Determine device to use - device = "cpu" if use_cpu else ("cuda" if torch.cuda.is_available() else "cpu") + device = "cpu" if use_cpu else ("mps" if torch.backends.mps.is_available() else ("cuda" if torch.cuda.is_available() else "cpu")) if use_cpu: - message = "Using CPU for inference (CUDA disabled)" + message = "Using CPU for inference (GPU disabled)" + elif torch.backends.mps.is_available() and device == "mps": + message = "Using MPS (Mac GPU) for inference" + elif torch.cuda.is_available() and device == "cuda": + message = "Using CUDA (NVIDIA GPU) for inference" else: - message = f"Using {device} for inference" + message = f"Using {device} for inference" # Should be CPU if no GPU found # Create temporary files for the audio inputs import tempfile @@ -226,48 +349,135 @@ class FL_ChatterboxVCNode(AudioNodeBase): target_waveform = target_voice['waveform'].squeeze(0) torchaudio.save(target_voice_path, target_waveform, target_voice['sample_rate']) + vc_model = None + pbar = ProgressBar(100) # Simple progress bar for overall process try: - # Load the VC model - message += f"\nLoading VC model on {device}..." - vc_model = ChatterboxVC.from_pretrained(device=device) - + # Load the VC model or reuse if loaded and device matches + if FL_ChatterboxVCNode._vc_model is not None and FL_ChatterboxVCNode._vc_device == device: + vc_model = FL_ChatterboxVCNode._vc_model + message += f"\nReusing loaded VC model on {device}..." + else: + if FL_ChatterboxVCNode._vc_model is not None: + message += f"\nUnloading previous VC model (device mismatch or keep_model_loaded is False)..." + FL_ChatterboxVCNode._vc_model = None + FL_ChatterboxVCNode._vc_device = None + if torch.cuda.is_available(): + torch.cuda.empty_cache() # Clear CUDA cache if possible + if torch.backends.mps.is_available(): + torch.mps.empty_cache() # Clear MPS cache if possible + + message += f"\nLoading VC model on {device}..." + pbar.update_absolute(10) # Indicate model loading started + vc_model = ChatterboxVC.from_pretrained(device=device) + pbar.update_absolute(50) # Indicate model loading finished + + if keep_model_loaded: + FL_ChatterboxVCNode._vc_model = vc_model + FL_ChatterboxVCNode._vc_device = device + message += "\nModel will be kept loaded in memory." + else: + message += "\nModel will be unloaded after use." + # Convert voice message += f"\nConverting voice to match target voice" + pbar.update_absolute(60) # Indicate conversion started converted_wav = vc_model.generate( audio=input_audio_path, target_voice_path=target_voice_path, ) + pbar.update_absolute(90) # Indicate conversion finished except RuntimeError as e: - if "CUDA" in str(e) and device != "cpu": + # Check for CUDA or MPS errors and attempt fallback to CPU + error_str = str(e) + fallback_to_cpu = False + if "CUDA" in error_str and device == "cuda": message += "\nCUDA error detected during VC. Falling back to CPU..." - # Try again with CPU + fallback_to_cpu = True + elif "MPS" in error_str and device == "mps": + message += "\nMPS error detected during VC. Falling back to CPU..." + fallback_to_cpu = True + + if fallback_to_cpu: device = "cpu" + # Unload previous model if it exists + if FL_ChatterboxVCNode._vc_model is not None: + message += f"\nUnloading previous VC model..." + FL_ChatterboxVCNode._vc_model = None + FL_ChatterboxVCNode._vc_device = None + if torch.cuda.is_available(): + torch.cuda.empty_cache() # Clear CUDA cache if possible + if torch.backends.mps.is_available(): + torch.mps.empty_cache() # Clear MPS cache if possible + + message += f"\nLoading VC model on {device}..." + pbar.update_absolute(10) # Indicate model loading started (fallback) vc_model = ChatterboxVC.from_pretrained(device=device) + pbar.update_absolute(50) # Indicate model loading finished (fallback) + # Note: keep_model_loaded logic is applied after successful generation + # to avoid keeping a failed model loaded. + converted_wav = vc_model.generate( audio=input_audio_path, target_voice_path=target_voice_path, ) + pbar.update_absolute(90) # Indicate conversion finished (fallback) else: - # Re-raise if it's not a CUDA error or we're already on CPU + # Re-raise if it's not a CUDA/MPS error or we're already on CPU message += f"\nError during VC: {str(e)}" # Return the original audio message += f"\nError: {str(e)}" + pbar.update_absolute(100) # Ensure progress bar completes on error return (input_audio, message) + except Exception as e: + message += f"\nAn unexpected error occurred during VC: {str(e)}" + empty_audio = {"waveform": torch.zeros((1, 2, 1)), "sample_rate": 16000} + for temp_file in temp_files: + if os.path.exists(temp_file): + os.unlink(temp_file) + pbar.update_absolute(100) # Ensure progress bar completes on error + return (empty_audio, message) finally: # Clean up all temporary files for temp_file in temp_files: if os.path.exists(temp_file): os.unlink(temp_file) - + # If keep_model_loaded is False, ensure model is not stored + # This is done here to ensure model is only kept if generation was successful + if not keep_model_loaded and FL_ChatterboxVCNode._vc_model is not None: + message += "\nUnloading VC model as keep_model_loaded is False." + FL_ChatterboxVCNode._vc_model = None + FL_ChatterboxVCNode._vc_device = None + if torch.cuda.is_available(): + torch.cuda.empty_cache() # Clear CUDA cache if possible + if torch.backends.mps.is_available(): + torch.mps.empty_cache() # Clear MPS cache if possible + + # If generation was successful and keep_model_loaded is True, store the model + if keep_model_loaded and vc_model is not None: + FL_ChatterboxVCNode._vc_model = vc_model + FL_ChatterboxVCNode._vc_device = device + message += "\nModel will be kept loaded in memory." + elif not keep_model_loaded and FL_ChatterboxVCNode._vc_model is not None: + # This case handles successful generation when keep_model_loaded was True previously + # but is now False. Ensure the model is unloaded. + message += "\nUnloading VC model as keep_model_loaded is now False." + FL_ChatterboxVCNode._vc_model = None + FL_ChatterboxVCNode._vc_device = None + if torch.cuda.is_available(): + torch.cuda.empty_cache() # Clear CUDA cache if possible + if torch.backends.mps.is_available(): + torch.mps.empty_cache() # Clear MPS cache if possible + # Create audio data structure for the output audio_data = { "waveform": converted_wav.unsqueeze(0), # Add batch dimension - "sample_rate": vc_model.sr + "sample_rate": vc_model.sr if vc_model else 16000 # Use default sample rate if model loading failed } message += f"\nVoice converted successfully" + pbar.update_absolute(100) # Ensure progress bar completes on success return (audio_data, message)