added mac support + keep model loaded
This commit is contained in:
+245
-35
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user