added mac support + keep model loaded

This commit is contained in:
Fill
2025-05-31 09:01:09 +09:00
parent f3c0eb3791
commit 5b0ccd9f80
+245 -35
View File
@@ -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)