284 lines
11 KiB
Python
284 lines
11 KiB
Python
import os
|
|
import torch
|
|
import torchaudio
|
|
import numpy as np
|
|
from pathlib import Path
|
|
from typing import Optional
|
|
|
|
# Import directly from the chatterbox package
|
|
from chatterbox.tts import ChatterboxTTS
|
|
from chatterbox.vc import ChatterboxVC
|
|
|
|
# Monkey patch torch.load to always use CPU if needed
|
|
original_torch_load = torch.load
|
|
def patched_torch_load(*args, **kwargs):
|
|
if 'map_location' not in kwargs:
|
|
kwargs['map_location'] = torch.device('cpu')
|
|
return original_torch_load(*args, **kwargs)
|
|
torch.load = patched_torch_load
|
|
|
|
class AudioNodeBase:
|
|
"""Base class for audio nodes with common utilities."""
|
|
|
|
@staticmethod
|
|
def create_empty_tensor(audio, frame_rate, height, width, channels=None):
|
|
"""Create an empty tensor with dimensions based on audio duration."""
|
|
audio_duration = audio['waveform'].shape[-1] / audio['sample_rate']
|
|
num_frames = int(audio_duration * frame_rate)
|
|
if channels is None:
|
|
return torch.zeros((num_frames, height, width), dtype=torch.float32)
|
|
else:
|
|
return torch.zeros((num_frames, height, width, channels), dtype=torch.float32)
|
|
|
|
# Text-to-Speech node
|
|
class FL_ChatterboxTTSNode(AudioNodeBase):
|
|
"""
|
|
ComfyUI node for Chatterbox Text-to-Speech functionality.
|
|
"""
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"text": ("STRING", {"multiline": True, "default": "Hello, this is a test."}),
|
|
"exaggeration": ("FLOAT", {"default": 0.5, "min": 0.25, "max": 2.0, "step": 0.05}),
|
|
"cfg_weight": ("FLOAT", {"default": 0.5, "min": 0.2, "max": 1.0, "step": 0.05}),
|
|
"temperature": ("FLOAT", {"default": 0.8, "min": 0.05, "max": 5.0, "step": 0.05}),
|
|
},
|
|
"optional": {
|
|
"audio_prompt": ("AUDIO",),
|
|
"use_cpu": ("BOOLEAN", {"default": False}),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("AUDIO", "STRING")
|
|
RETURN_NAMES = ("audio", "message")
|
|
FUNCTION = "generate_speech"
|
|
CATEGORY = "ChatterBox"
|
|
|
|
def generate_speech(self, text, exaggeration, cfg_weight, temperature, audio_prompt=None, use_cpu=False):
|
|
"""
|
|
Generate speech from text.
|
|
|
|
Args:
|
|
text: The text to convert to speech.
|
|
exaggeration: Controls emotion intensity (0.25-2.0).
|
|
cfg_weight: Controls pace/classifier-free guidance (0.2-1.0).
|
|
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.
|
|
|
|
Returns:
|
|
Tuple of (audio, message)
|
|
"""
|
|
# Determine device to use
|
|
device = "cpu" if use_cpu else ("cuda" if torch.cuda.is_available() else "cpu")
|
|
if use_cpu:
|
|
message = "Using CPU for inference (CUDA disabled)"
|
|
else:
|
|
message = f"Using {device} for inference"
|
|
|
|
# Create temporary files for any audio inputs
|
|
import tempfile
|
|
temp_files = []
|
|
|
|
# Create a temporary file for the audio prompt if provided
|
|
audio_prompt_path = None
|
|
if audio_prompt is not None:
|
|
try:
|
|
with tempfile.NamedTemporaryFile(suffix='.wav', delete=False) as temp_prompt:
|
|
audio_prompt_path = temp_prompt.name
|
|
temp_files.append(audio_prompt_path)
|
|
|
|
# Save the audio prompt to the temporary file
|
|
prompt_waveform = audio_prompt['waveform'].squeeze(0)
|
|
torchaudio.save(audio_prompt_path, prompt_waveform, audio_prompt['sample_rate'])
|
|
message += f"\nUsing provided audio prompt for voice cloning: {audio_prompt_path}"
|
|
|
|
# Debug: Check if the file exists and has content
|
|
if os.path.exists(audio_prompt_path):
|
|
file_size = os.path.getsize(audio_prompt_path)
|
|
message += f"\nAudio prompt file created successfully: {file_size} bytes"
|
|
else:
|
|
message += f"\nWarning: Audio prompt file was not created properly"
|
|
except Exception as e:
|
|
message += f"\nError creating audio prompt file: {str(e)}"
|
|
audio_prompt_path = None
|
|
|
|
try:
|
|
# Load the TTS model
|
|
message += f"\nLoading TTS model on {device}..."
|
|
tts_model = ChatterboxTTS.from_pretrained(device=device)
|
|
|
|
# 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}"
|
|
|
|
wav = tts_model.generate(
|
|
text=text,
|
|
audio_prompt_path=audio_prompt_path,
|
|
exaggeration=exaggeration,
|
|
cfg_weight=cfg_weight,
|
|
temperature=temperature,
|
|
)
|
|
|
|
except RuntimeError as e:
|
|
if "CUDA" in str(e) and device != "cpu":
|
|
message += "\nCUDA error detected during TTS. Falling back to CPU..."
|
|
# Try again with CPU
|
|
device = "cpu"
|
|
tts_model = ChatterboxTTS.from_pretrained(device=device)
|
|
wav = tts_model.generate(
|
|
text=text,
|
|
audio_prompt_path=audio_prompt_path,
|
|
exaggeration=exaggeration,
|
|
cfg_weight=cfg_weight,
|
|
temperature=temperature,
|
|
)
|
|
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)
|
|
finally:
|
|
# Clean up all temporary files
|
|
for temp_file in temp_files:
|
|
if os.path.exists(temp_file):
|
|
os.unlink(temp_file)
|
|
|
|
# Create audio data structure for the output
|
|
audio_data = {
|
|
"waveform": wav.unsqueeze(0), # Add batch dimension
|
|
"sample_rate": tts_model.sr
|
|
}
|
|
|
|
message += f"\nSpeech generated successfully"
|
|
|
|
return (audio_data, message)
|
|
|
|
# Voice Conversion node
|
|
class FL_ChatterboxVCNode(AudioNodeBase):
|
|
"""
|
|
ComfyUI node for Chatterbox Voice Conversion functionality.
|
|
"""
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"input_audio": ("AUDIO",),
|
|
"target_voice": ("AUDIO",),
|
|
},
|
|
"optional": {
|
|
"use_cpu": ("BOOLEAN", {"default": False}),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("AUDIO", "STRING")
|
|
RETURN_NAMES = ("audio", "message")
|
|
FUNCTION = "convert_voice"
|
|
CATEGORY = "ChatterBox"
|
|
|
|
def convert_voice(self, input_audio, target_voice, use_cpu=False):
|
|
"""
|
|
Convert the voice in an audio file to match a target voice.
|
|
|
|
Args:
|
|
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.
|
|
|
|
Returns:
|
|
Tuple of (audio, message)
|
|
"""
|
|
# Determine device to use
|
|
device = "cpu" if use_cpu else ("cuda" if torch.cuda.is_available() else "cpu")
|
|
if use_cpu:
|
|
message = "Using CPU for inference (CUDA disabled)"
|
|
else:
|
|
message = f"Using {device} for inference"
|
|
|
|
# Create temporary files for the audio inputs
|
|
import tempfile
|
|
temp_files = []
|
|
|
|
# Create a temporary file for the input audio
|
|
with tempfile.NamedTemporaryFile(suffix='.wav', delete=False) as temp_input:
|
|
input_audio_path = temp_input.name
|
|
temp_files.append(input_audio_path)
|
|
|
|
# Save the input audio to the temporary file
|
|
input_waveform = input_audio['waveform'].squeeze(0)
|
|
torchaudio.save(input_audio_path, input_waveform, input_audio['sample_rate'])
|
|
|
|
# Create a temporary file for the target voice
|
|
with tempfile.NamedTemporaryFile(suffix='.wav', delete=False) as temp_target:
|
|
target_voice_path = temp_target.name
|
|
temp_files.append(target_voice_path)
|
|
|
|
# Save the target voice to the temporary file
|
|
target_waveform = target_voice['waveform'].squeeze(0)
|
|
torchaudio.save(target_voice_path, target_waveform, target_voice['sample_rate'])
|
|
|
|
try:
|
|
# Load the VC model
|
|
message += f"\nLoading VC model on {device}..."
|
|
vc_model = ChatterboxVC.from_pretrained(device=device)
|
|
|
|
# Convert voice
|
|
message += f"\nConverting voice to match target voice"
|
|
|
|
converted_wav = vc_model.generate(
|
|
audio=input_audio_path,
|
|
target_voice_path=target_voice_path,
|
|
)
|
|
|
|
except RuntimeError as e:
|
|
if "CUDA" in str(e) and device != "cpu":
|
|
message += "\nCUDA error detected during VC. Falling back to CPU..."
|
|
# Try again with CPU
|
|
device = "cpu"
|
|
vc_model = ChatterboxVC.from_pretrained(device=device)
|
|
converted_wav = vc_model.generate(
|
|
audio=input_audio_path,
|
|
target_voice_path=target_voice_path,
|
|
)
|
|
else:
|
|
# Re-raise if it's not a CUDA error or we're already on CPU
|
|
message += f"\nError during VC: {str(e)}"
|
|
# Return the original audio
|
|
message += f"\nError: {str(e)}"
|
|
return (input_audio, message)
|
|
finally:
|
|
# Clean up all temporary files
|
|
for temp_file in temp_files:
|
|
if os.path.exists(temp_file):
|
|
os.unlink(temp_file)
|
|
|
|
# Create audio data structure for the output
|
|
audio_data = {
|
|
"waveform": converted_wav.unsqueeze(0), # Add batch dimension
|
|
"sample_rate": vc_model.sr
|
|
}
|
|
|
|
message += f"\nVoice converted successfully"
|
|
|
|
return (audio_data, message)
|
|
|
|
# Node mappings for ComfyUI
|
|
NODE_CLASS_MAPPINGS = {
|
|
"FL_ChatterboxTTS": FL_ChatterboxTTSNode,
|
|
"FL_ChatterboxVC": FL_ChatterboxVCNode,
|
|
}
|
|
|
|
# Display names for the nodes
|
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
|
"FL_ChatterboxTTS": "FL Chatterbox TTS",
|
|
"FL_ChatterboxVC": "FL Chatterbox VC",
|
|
} |