ver 1.2.0
This commit is contained in:
@@ -5,7 +5,7 @@
|
||||
<br />
|
||||
<div align="center">
|
||||
<a href="https://github.com/wildminder/ComfyUI-Chatterbox">
|
||||
<img src="./assets/preview.png" alt="Chatterbox Nodes in ComfyUI">
|
||||
<img src="./workflow-examples/ChatterboxTTS-workflow.png" alt="Chatterbox Nodes in ComfyUI">
|
||||
</a>
|
||||
|
||||
<h3 align="center">ComfyUI Chatterbox</h3>
|
||||
|
||||
+62
-14
@@ -1,7 +1,59 @@
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
import logging
|
||||
import folder_paths
|
||||
from importlib import metadata as importlib_metadata
|
||||
|
||||
# Configure a logger for the entire custom node package
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.setLevel(logging.WARNING)
|
||||
|
||||
# Add a handler if none exist to avoid duplicate logs
|
||||
if not logger.hasHandlers():
|
||||
# Use stdout for info, stderr for errors
|
||||
handler = logging.StreamHandler(sys.stdout)
|
||||
formatter = logging.Formatter(f"[%(name)s] %(message)s")
|
||||
handler.setFormatter(formatter)
|
||||
logger.addHandler(handler)
|
||||
|
||||
|
||||
# Monkey-Patch for 'chatterbox-tts' version
|
||||
# original_importlib_version_func = importlib_metadata.version
|
||||
# def patched_version_lookup(package_name):
|
||||
# if package_name == "chatterbox-tts":
|
||||
# return "local-vendored"
|
||||
# return original_importlib_version_func(package_name)
|
||||
# importlib_metadata.version = patched_version_lookup
|
||||
|
||||
|
||||
# easily check if the real 'perth' is being used.
|
||||
try:
|
||||
import perth
|
||||
# A simple check to ensure it's not our mock from a previous run
|
||||
if not hasattr(perth, '_is_mock'):
|
||||
logger.info("Found and using 'resemble-perth' library for watermarking.")
|
||||
except ImportError:
|
||||
logger.warning("'resemble-perth' not found. Watermarking will be unavailable.")
|
||||
class DummyPerthImplicitWatermarker:
|
||||
def apply_watermark(self, wav, sample_rate):
|
||||
logger.warning("Watermarking skipped: 'resemble-perth' is not installed.")
|
||||
return wav
|
||||
perth_mock = types.ModuleType('perth')
|
||||
perth_mock.PerthImplicitWatermarker = DummyPerthImplicitWatermarker
|
||||
# Flag to identify our mock module
|
||||
perth_mock._is_mock = True
|
||||
sys.modules['perth'] = perth_mock
|
||||
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
src_dir = os.path.join(current_dir, "src")
|
||||
if src_dir not in sys.path:
|
||||
sys.path.insert(0, src_dir)
|
||||
|
||||
|
||||
from .nodes import ChatterboxTTSNode, ChatterboxVCNode
|
||||
from .modules.chatterbox_handler import CHATTERBOX_MODEL_SUBDIR, DEFAULT_MODEL_PACK_NAME
|
||||
from .modules.chatterbox_handler import CHATTERBOX_MODEL_SUBDIR
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"ChatterboxTTS": ChatterboxTTSNode,
|
||||
@@ -13,25 +65,21 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ChatterboxVC": "Chatterbox Voice Conversion 🗣️",
|
||||
}
|
||||
|
||||
# WEB_DIRECTORY = "./js"
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
|
||||
# Model path setup for ComfyUI
|
||||
chatterbox_models_full_path = os.path.join(folder_paths.models_dir, CHATTERBOX_MODEL_SUBDIR)
|
||||
if not os.path.exists(chatterbox_models_full_path):
|
||||
try:
|
||||
os.makedirs(chatterbox_models_full_path, exist_ok=True)
|
||||
#print(f"ChatterboxTTS/VC Init: Created models directory at {chatterbox_models_full_path}")
|
||||
except OSError as e:
|
||||
print(f"ChatterboxTTS/VC Init: Error creating models directory {chatterbox_models_full_path}: {e}")
|
||||
logger.error(f"Error creating models directory {chatterbox_models_full_path}: {e}")
|
||||
|
||||
if CHATTERBOX_MODEL_SUBDIR not in folder_paths.folder_names_and_paths:
|
||||
folder_paths.folder_names_and_paths[CHATTERBOX_MODEL_SUBDIR] = (
|
||||
[chatterbox_models_full_path],
|
||||
folder_paths.supported_pt_extensions
|
||||
)
|
||||
#print(f"ChatterboxTTS/VC Init: Registered '{CHATTERBOX_MODEL_SUBDIR}' with ComfyUI model paths: {chatterbox_models_full_path}")
|
||||
# Register the tts/chatterbox path with ComfyUI
|
||||
tts_chatterbox_path = os.path.join(folder_paths.models_dir, "tts")
|
||||
if "tts" not in folder_paths.folder_names_and_paths:
|
||||
supported_exts = folder_paths.supported_pt_extensions.union({".safetensors"})
|
||||
folder_paths.folder_names_and_paths["tts"] = ([tts_chatterbox_path], supported_exts)
|
||||
else:
|
||||
if chatterbox_models_full_path not in folder_paths.folder_names_and_paths[CHATTERBOX_MODEL_SUBDIR][0]:
|
||||
folder_paths.folder_names_and_paths[CHATTERBOX_MODEL_SUBDIR][0].append(chatterbox_models_full_path)
|
||||
#print(f"ChatterboxTTS/VC Init: Appended path for '{CHATTERBOX_MODEL_SUBDIR}': {chatterbox_models_full_path}")
|
||||
|
||||
if tts_chatterbox_path not in folder_paths.folder_names_and_paths["tts"][0]:
|
||||
folder_paths.folder_names_and_paths["tts"][0].append(tts_chatterbox_path)
|
||||
+50
-174
@@ -1,222 +1,98 @@
|
||||
import os
|
||||
import sys
|
||||
import logging
|
||||
import torch
|
||||
import random
|
||||
import numpy as np
|
||||
import folder_paths
|
||||
import hashlib
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
current_script_path = os.path.dirname(os.path.abspath(__file__))
|
||||
chatterbox_src_path = os.path.abspath(os.path.join(current_script_path, '..', 'src'))
|
||||
if chatterbox_src_path not in sys.path:
|
||||
sys.path.insert(0, chatterbox_src_path)
|
||||
|
||||
|
||||
from chatterbox.tts import ChatterboxTTS
|
||||
from chatterbox.vc import ChatterboxVC # <--- Added VC import
|
||||
from chatterbox.vc import ChatterboxVC
|
||||
|
||||
TTS_MODEL_CACHE = {}
|
||||
VC_MODEL_CACHE = {}
|
||||
CHATTERBOX_MODEL_SUBDIR = "chatterbox_tts"
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
CHATTERBOX_MODEL_SUBDIR = os.path.join("tts", "chatterbox")
|
||||
CHATTERBOX_REPO_ID = "ResembleAI/chatterbox"
|
||||
|
||||
CHATTERBOX_FILES_TO_DOWNLOAD = ["ve.pt", "t3_cfg.pt", "s3gen.pt", "tokenizer.json", "conds.pt"]
|
||||
CHATTERBOX_FILES_TO_DOWNLOAD = ["ve.safetensors", "t3_cfg.safetensors", "s3gen.safetensors", "tokenizer.json", "conds.pt"]
|
||||
DEFAULT_MODEL_PACK_NAME = "resembleai_default_voice"
|
||||
|
||||
def get_chatterbox_model_pack_names():
|
||||
"""Returns a list of available Chatterbox model pack names (subdirectories)."""
|
||||
chatterbox_models_base_path = os.path.join(folder_paths.models_dir, CHATTERBOX_MODEL_SUBDIR)
|
||||
if not os.path.isdir(chatterbox_models_base_path):
|
||||
os.makedirs(chatterbox_models_base_path, exist_ok=True)
|
||||
#print(f"ChatterboxTTS/VC: Created models directory at {chatterbox_models_base_path}")
|
||||
return []
|
||||
|
||||
return [DEFAULT_MODEL_PACK_NAME]
|
||||
packs = [d for d in os.listdir(chatterbox_models_base_path) if os.path.isdir(os.path.join(chatterbox_models_base_path, d))]
|
||||
if not packs:
|
||||
return []
|
||||
return packs
|
||||
# Ensure default is first if it exists
|
||||
if DEFAULT_MODEL_PACK_NAME in packs:
|
||||
packs.insert(0, packs.pop(packs.index(DEFAULT_MODEL_PACK_NAME)))
|
||||
|
||||
# Return default even if folder doesn't exist yet, to prompt the user to download it.
|
||||
return packs if packs else [DEFAULT_MODEL_PACK_NAME]
|
||||
|
||||
def get_model_pack_path(model_pack_name):
|
||||
"""Gets the full path to a specific model pack."""
|
||||
if not model_pack_name:
|
||||
# Added check for None or empty string to prevent errors
|
||||
if not model_pack_name:
|
||||
return None
|
||||
return os.path.join(folder_paths.models_dir, CHATTERBOX_MODEL_SUBDIR, model_pack_name)
|
||||
|
||||
def _download_file_from_hf(repo_id, filename, local_dir, local_filename=None):
|
||||
if local_filename is None:
|
||||
local_filename = filename
|
||||
destination = os.path.join(local_dir, local_filename)
|
||||
|
||||
def _download_file_from_hf(repo_id, filename, local_dir):
|
||||
destination = os.path.join(local_dir, filename)
|
||||
if not os.path.exists(destination):
|
||||
print(f"ChatterboxTTS/VC: Downloading '{filename}' from '{repo_id}' to '{destination}'...")
|
||||
logger.info(f"Downloading '{filename}' from '{repo_id}'...")
|
||||
try:
|
||||
hf_hub_download(
|
||||
repo_id=repo_id,
|
||||
filename=filename,
|
||||
local_dir=local_dir,
|
||||
)
|
||||
if filename != local_filename and os.path.exists(os.path.join(local_dir, filename)) and not os.path.exists(destination):
|
||||
os.rename(os.path.join(local_dir, filename), destination)
|
||||
print(f"ChatterboxTTS/VC: Successfully downloaded '{local_filename}'.")
|
||||
hf_hub_download(repo_id=repo_id, filename=filename, local_dir=local_dir, local_dir_use_symlinks=False, resume_download=True)
|
||||
logger.info(f"Successfully downloaded '{filename}'.")
|
||||
return True
|
||||
except Exception as e:
|
||||
print(f"ChatterboxTTS/VC: Failed to download '{filename}' from '{repo_id}'. Error: {e}")
|
||||
if os.path.exists(destination + ".incomplete"):
|
||||
os.remove(destination + ".incomplete")
|
||||
logger.error(f"Failed to download '{filename}': {e}")
|
||||
if os.path.exists(destination + ".incomplete"): os.remove(destination + ".incomplete")
|
||||
return False
|
||||
else:
|
||||
return True
|
||||
return True
|
||||
|
||||
def download_chatterbox_model_pack_if_missing(model_pack_name):
|
||||
"""Downloads all necessary model files for a given pack name if they don't exist."""
|
||||
ckpt_dir = get_model_pack_path(model_pack_name)
|
||||
if not ckpt_dir:
|
||||
print(f"ChatterboxTTS/VC: Invalid model pack name '{model_pack_name}', cannot download.")
|
||||
logger.warning(f"Invalid model pack name '{model_pack_name}', cannot download.")
|
||||
return False
|
||||
|
||||
os.makedirs(ckpt_dir, exist_ok=True)
|
||||
|
||||
all_files_successfully_managed = True
|
||||
all_files_ok = all(_download_file_from_hf(CHATTERBOX_REPO_ID, f, ckpt_dir) for f in CHATTERBOX_FILES_TO_DOWNLOAD)
|
||||
if not all_files_ok:
|
||||
logger.error(f"Some files failed to download for model pack '{model_pack_name}'. Check logs.")
|
||||
return all_files_ok
|
||||
|
||||
vc_required_files = ["s3gen.pt", "conds.pt"]
|
||||
|
||||
files_to_check = CHATTERBOX_FILES_TO_DOWNLOAD
|
||||
|
||||
for file_name in files_to_check:
|
||||
if not _download_file_from_hf(CHATTERBOX_REPO_ID, file_name, ckpt_dir):
|
||||
all_files_successfully_managed = False
|
||||
if file_name in vc_required_files:
|
||||
print(f"ChatterboxTTS/VC: Critical file '{file_name}' for VC failed to download for pack '{model_pack_name}'.")
|
||||
# return False
|
||||
|
||||
if all_files_successfully_managed:
|
||||
pass #print(f"ChatterboxTTS/VC: Verified/Downloaded all files for model pack '{model_pack_name}' into '{ckpt_dir}'.")
|
||||
else:
|
||||
print(f"ChatterboxTTS/VC: Some files failed to download for model pack '{model_pack_name}'. Check logs.")
|
||||
return all_files_successfully_managed
|
||||
|
||||
|
||||
def load_chatterbox_tts_model(model_pack_name, device_str="cuda"):
|
||||
def load_chatterbox_models(model_pack_name, device):
|
||||
"""Loads both TTS and VC models for a given pack onto a specified device."""
|
||||
ckpt_dir = get_model_pack_path(model_pack_name)
|
||||
if not ckpt_dir:
|
||||
raise ValueError(f"ChatterboxTTS: Invalid model_pack_name: {model_pack_name}")
|
||||
#print(f"ChatterboxTTS: Attempting to load TTS model pack '{model_pack_name}' from '{ckpt_dir}'.")
|
||||
raise ValueError(f"Invalid model_pack_name: {model_pack_name}")
|
||||
|
||||
if not download_chatterbox_model_pack_if_missing(model_pack_name):
|
||||
print(f"ChatterboxTTS: Warning - Not all TTS model files could be verified/downloaded for '{model_pack_name}'. Loading may fail.")
|
||||
if not os.path.isdir(ckpt_dir):
|
||||
raise FileNotFoundError(f"ChatterboxTTS: Model pack directory '{model_pack_name}' not found and could not be created at '{ckpt_dir}'.")
|
||||
#print(f"ChatterboxTTS: Loading ChatterboxTTS model from local directory: {ckpt_dir} onto device: {device_str}")
|
||||
try:
|
||||
model = ChatterboxTTS.from_local(ckpt_dir, device=device_str)
|
||||
except Exception as e:
|
||||
print(f"ChatterboxTTS: Error during ChatterboxTTS.from_local('{ckpt_dir}', device='{device_str}'): {e}")
|
||||
print(f"ChatterboxTTS: Please ensure all required model files ({', '.join(CHATTERBOX_FILES_TO_DOWNLOAD)}) are present in {ckpt_dir} or can be downloaded from {CHATTERBOX_REPO_ID}.")
|
||||
raise
|
||||
return model
|
||||
|
||||
def get_cached_chatterbox_tts_model(model_pack_name, device_str="cuda"):
|
||||
"""Loads and caches the ChatterboxTTS model."""
|
||||
if not model_pack_name:
|
||||
available_packs = get_chatterbox_model_pack_names()
|
||||
model_pack_name = available_packs[0] if available_packs else DEFAULT_MODEL_PACK_NAME
|
||||
print(f"ChatterboxTTS: No model pack specified for TTS, using '{model_pack_name}'.")
|
||||
|
||||
cache_key = (model_pack_name, device_str, "tts")
|
||||
current_model = TTS_MODEL_CACHE.get(cache_key)
|
||||
model_device_correct = False
|
||||
if current_model is not None and hasattr(current_model, 'device'):
|
||||
try:
|
||||
if str(current_model.device) == device_str:
|
||||
model_device_correct = True
|
||||
except Exception as e:
|
||||
print(f"ChatterboxTTS: Error checking cached TTS model device: {e}. Will reload.")
|
||||
|
||||
if current_model is not None and model_device_correct:
|
||||
return current_model
|
||||
else:
|
||||
if current_model is not None and not model_device_correct:
|
||||
print(f"ChatterboxTTS: Device mismatch for cached TTS model '{model_pack_name}'. Reloading.")
|
||||
else:
|
||||
print(f"ChatterboxTTS: TTS Model for '{model_pack_name}' on '{device_str}' not in cache. Loading...")
|
||||
logger.warning(f"Not all model files could be verified for '{model_pack_name}'. Loading may fail.")
|
||||
|
||||
TTS_MODEL_CACHE[cache_key] = load_chatterbox_tts_model(model_pack_name, device_str)
|
||||
return TTS_MODEL_CACHE[cache_key]
|
||||
|
||||
|
||||
def load_chatterbox_vc_model(model_pack_name, device_str="cuda"):
|
||||
"""Loads the ChatterboxVC model from a specified pack."""
|
||||
ckpt_dir = get_model_pack_path(model_pack_name)
|
||||
if not ckpt_dir:
|
||||
raise ValueError(f"ChatterboxVC: Invalid model_pack_name: {model_pack_name}")
|
||||
|
||||
#print(f"ChatterboxVC: Attempting to load VC model pack '{model_pack_name}' from '{ckpt_dir}'.")
|
||||
|
||||
if not download_chatterbox_model_pack_if_missing(model_pack_name):
|
||||
vc_min_files = ["s3gen.pt", "conds.pt"]
|
||||
for f_name in vc_min_files:
|
||||
if not os.path.exists(os.path.join(ckpt_dir, f_name)):
|
||||
print(f"ChatterboxVC: Critical file '{f_name}' for VC is missing from pack '{model_pack_name}' after download attempt. Loading will likely fail.")
|
||||
break
|
||||
else:
|
||||
print(f"ChatterboxVC: Warning - Not all model files could be verified/downloaded for '{model_pack_name}', but attempting to load VC with available files.")
|
||||
|
||||
|
||||
if not os.path.isdir(ckpt_dir):
|
||||
raise FileNotFoundError(f"ChatterboxVC: Model pack directory '{model_pack_name}' not found and could not be created at '{ckpt_dir}'.")
|
||||
|
||||
# print(f"ChatterboxVC: Loading ChatterboxVC model from local directory: {ckpt_dir} onto device: {device_str}")
|
||||
raise FileNotFoundError(f"Model pack directory '{model_pack_name}' not found at '{ckpt_dir}'.")
|
||||
|
||||
try:
|
||||
model = ChatterboxVC.from_local(ckpt_dir, device=device_str)
|
||||
logger.info(f"Loading Chatterbox TTS model from {ckpt_dir} onto {device}")
|
||||
tts_model = ChatterboxTTS.from_local(ckpt_dir, device=device)
|
||||
except Exception as e:
|
||||
print(f"ChatterboxVC: Error during ChatterboxVC.from_local('{ckpt_dir}', device='{device_str}'): {e}")
|
||||
print(f"ChatterboxVC: Please ensure at least 's3gen.pt' and optionally 'conds.pt' (for default voice) are present in {ckpt_dir} or can be downloaded from {CHATTERBOX_REPO_ID}.")
|
||||
logger.error(f"Error loading ChatterboxTTS from '{ckpt_dir}': {e}", exc_info=True)
|
||||
raise
|
||||
|
||||
try:
|
||||
logger.info(f"Loading Chatterbox VC model from {ckpt_dir} onto {device}")
|
||||
vc_model = ChatterboxVC.from_local(ckpt_dir, device=device)
|
||||
except Exception as e:
|
||||
logger.error(f"Error loading ChatterboxVC from '{ckpt_dir}': {e}", exc_info=True)
|
||||
raise
|
||||
return model
|
||||
|
||||
def get_cached_chatterbox_vc_model(model_pack_name, device_str="cuda"):
|
||||
"""Loads and caches the ChatterboxVC model."""
|
||||
if not model_pack_name:
|
||||
available_packs = get_chatterbox_model_pack_names()
|
||||
model_pack_name = available_packs[0] if available_packs else DEFAULT_MODEL_PACK_NAME
|
||||
print(f"ChatterboxVC: No model pack specified for VC, using '{model_pack_name}'.")
|
||||
|
||||
cache_key = (model_pack_name, device_str, "vc")
|
||||
current_model = VC_MODEL_CACHE.get(cache_key)
|
||||
model_device_correct = False
|
||||
if current_model is not None and hasattr(current_model, 'device'):
|
||||
try:
|
||||
if str(current_model.device) == device_str:
|
||||
model_device_correct = True
|
||||
except Exception as e:
|
||||
print(f"ChatterboxVC: Error checking cached VC model device: {e}. Will reload.")
|
||||
|
||||
if current_model is not None and model_device_correct:
|
||||
return current_model
|
||||
else:
|
||||
if current_model is not None and not model_device_correct:
|
||||
print(f"ChatterboxVC: Device mismatch for cached VC model '{model_pack_name}'. Reloading.")
|
||||
else:
|
||||
print(f"ChatterboxVC: VC Model for '{model_pack_name}' on '{device_str}' not in cache. Loading...")
|
||||
VC_MODEL_CACHE[cache_key] = load_chatterbox_vc_model(model_pack_name, device_str)
|
||||
return VC_MODEL_CACHE[cache_key]
|
||||
|
||||
return tts_model, vc_model
|
||||
|
||||
def set_chatterbox_seed(seed: int):
|
||||
"""Sets the seed for Chatterbox TTS/VC."""
|
||||
MAX_NUMPY_SEED = 2**32 - 1
|
||||
if seed == 0:
|
||||
actual_seed_for_torch_random = random.randint(1, 0xffffffffffffffff)
|
||||
actual_seed_for_numpy = random.randint(1, MAX_NUMPY_SEED)
|
||||
else:
|
||||
actual_seed_for_torch_random = seed
|
||||
# fast fix range for seed, sorry
|
||||
actual_seed_for_numpy = seed % MAX_NUMPY_SEED
|
||||
|
||||
actual_seed_for_torch_random = random.randint(1, 0xffffffffffffffff) if seed == 0 else seed
|
||||
actual_seed_for_numpy = random.randint(1, MAX_NUMPY_SEED) if seed == 0 else (seed % MAX_NUMPY_SEED)
|
||||
torch.manual_seed(actual_seed_for_torch_random)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(actual_seed_for_torch_random)
|
||||
torch.cuda.manual_seed_all(actual_seed_for_torch_random)
|
||||
if torch.cuda.is_available(): torch.cuda.manual_seed_all(actual_seed_for_torch_random)
|
||||
random.seed(actual_seed_for_torch_random)
|
||||
np.random.seed(actual_seed_for_numpy)
|
||||
|
||||
np.random.seed(actual_seed_for_numpy)
|
||||
@@ -1,231 +1,325 @@
|
||||
import os
|
||||
import torch
|
||||
import folder_paths
|
||||
import tempfile
|
||||
import soundfile as sf
|
||||
import numpy as np
|
||||
import logging
|
||||
import perth
|
||||
|
||||
import comfy.model_management as mm
|
||||
import comfy.model_patcher
|
||||
from comfy.utils import ProgressBar
|
||||
|
||||
from .modules.chatterbox_handler import (
|
||||
get_chatterbox_model_pack_names,
|
||||
get_cached_chatterbox_tts_model,
|
||||
get_cached_chatterbox_vc_model,
|
||||
load_chatterbox_models,
|
||||
set_chatterbox_seed,
|
||||
CHATTERBOX_MODEL_SUBDIR,
|
||||
DEFAULT_MODEL_PACK_NAME
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
CHATTERBOX_PATCHER_CACHE = {}
|
||||
|
||||
class ChatterboxModelWrapper(torch.nn.Module):
|
||||
"""
|
||||
A simple torch.nn.Module wrapper for the Chatterbox models.
|
||||
This allows ComfyUI's model management to treat our custom models like any other
|
||||
torch module, enabling device placement (.to()) and other standard operations.
|
||||
"""
|
||||
def __init__(self, model_pack_name):
|
||||
super().__init__()
|
||||
self.model_pack_name = model_pack_name
|
||||
self.tts_model = None
|
||||
self.vc_model = None
|
||||
|
||||
def load_model(self, device):
|
||||
self.tts_model, self.vc_model = load_chatterbox_models(self.model_pack_name, device)
|
||||
|
||||
class ChatterboxPatcher(comfy.model_patcher.ModelPatcher):
|
||||
"""
|
||||
Custom ModelPatcher for Chatterbox. This class hooks into ComfyUI's
|
||||
model management system (loading, offloading) to handle our non-standard models.
|
||||
"""
|
||||
def __init__(self, model, *args, **kwargs):
|
||||
super().__init__(model, *args, **kwargs)
|
||||
|
||||
def patch_model(self, device_to=None, *args, **kwargs):
|
||||
"""
|
||||
This method is called by ComfyUI's model manager when it's time to load
|
||||
the model onto the target device (usually the GPU). Our responsibility here
|
||||
is to ensure the model weights are loaded from disk if they haven't been already.
|
||||
"""
|
||||
target_device = self.load_device
|
||||
|
||||
# The core loading logic: If the model isn't in memory, load it from disk.
|
||||
if self.model.tts_model is None:
|
||||
logger.info(f"Loading Chatterbox models for '{self.model.model_pack_name}' to {target_device}...")
|
||||
self.model.load_model(target_device)
|
||||
self.model.model_loaded_weight_memory = self.size
|
||||
else:
|
||||
logger.info(f"Chatterbox models for '{self.model.model_pack_name}' already in memory.")
|
||||
|
||||
return super().patch_model(device_to=target_device, *args, **kwargs)
|
||||
|
||||
def unpatch_model(self, device_to=None, unpatch_weights=True, *args, **kwargs):
|
||||
"""
|
||||
This method is called by ComfyUI's model manager to offload the model
|
||||
(usually to the CPU) and free up VRAM.
|
||||
"""
|
||||
if unpatch_weights:
|
||||
logger.info(f"Offloading Chatterbox models for '{self.model.model_pack_name}' to {device_to}...")
|
||||
self.model.tts_model = None
|
||||
self.vc_model = None
|
||||
# Reset memory footprint
|
||||
self.model.model_loaded_weight_memory = 0
|
||||
mm.soft_empty_cache()
|
||||
return super().unpatch_model(device_to, unpatch_weights, *args, **kwargs)
|
||||
|
||||
class ChatterboxTTSNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
available_model_packs = get_chatterbox_model_pack_names()
|
||||
displayed_packs = [DEFAULT_MODEL_PACK_NAME] + [p for p in available_model_packs if p != DEFAULT_MODEL_PACK_NAME]
|
||||
if not displayed_packs:
|
||||
displayed_packs = [DEFAULT_MODEL_PACK_NAME]
|
||||
return {"required": {
|
||||
"model_pack_name": (get_chatterbox_model_pack_names(), {
|
||||
"default": DEFAULT_MODEL_PACK_NAME,
|
||||
"tooltip": "Select the Chatterbox voice model pack to use for generation."
|
||||
}),
|
||||
"text": ("STRING", {
|
||||
"multiline": True,
|
||||
"default": "Hello, this is a test of Chatterbox TTS in ComfyUI.",
|
||||
"tooltip": "Text to be synthesized into speech."
|
||||
}),
|
||||
"max_new_tokens": ("INT", {
|
||||
"default": 1000, "min": 16, "max": 4000, "step": 8,
|
||||
"tooltip": "Maximum number of audio tokens to generate. 25 tokens ≈ 1 second. The hard limit is 4096 tokens (≈ 163 seconds)."
|
||||
}),
|
||||
"flow_cfg_scale": ("FLOAT", {"default": 0.7, "min": 0.0, "max": 2.0, "step": 0.05, "tooltip": "CFG scale for the mel spectrogram decoder (flow matching). Higher values increase adherence to content and timbre but may reduce naturalness."}),
|
||||
"exaggeration": ("FLOAT", {
|
||||
"default": 0.5, "min": 0.25, "max": 2.0, "step": 0.05,
|
||||
"tooltip": "Controls the expressiveness and emotional intensity. Higher values lead to more exaggerated prosody."
|
||||
}),
|
||||
"temperature": ("FLOAT", {
|
||||
"default": 0.8, "min": 0.05, "max": 5.0, "step": 0.05,
|
||||
"tooltip": "Controls the randomness of the sampling process. Higher values produce more diverse speech, while lower values are more deterministic."
|
||||
}),
|
||||
"cfg_weight": ("FLOAT", {
|
||||
"default": 0.5, "min": 0.2, "max": 1.0, "step": 0.05,
|
||||
"tooltip": "Classifier-Free Guidance (CFG) weight. Controls how strongly the model adheres to the text prompt. Higher values may reduce naturalness."
|
||||
}),
|
||||
"repetition_penalty": ("FLOAT", {
|
||||
"default": 1.2, "min": 1.0, "max": 2.0, "step": 0.1,
|
||||
"tooltip": "Penalizes repeated tokens to discourage monotonous or repetitive speech. A value of 1.0 means no penalty."
|
||||
}),
|
||||
"min_p": ("FLOAT", {
|
||||
"default": 0.05, "min": 0.0, "max": 1.0, "step": 0.01,
|
||||
"tooltip": "Sets a minimum probability threshold for nucleus sampling (Min-P). Filters out tokens with very low probability."
|
||||
}),
|
||||
"top_p": ("FLOAT", {
|
||||
"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01,
|
||||
"tooltip": "Nucleus sampling (Top-P) parameter. The model samples from the smallest set of tokens whose cumulative probability exceeds this value."
|
||||
}),
|
||||
"seed": ("INT", {
|
||||
"default": 0, "min": 0, "max": 0xffffffffffffffff, "control_after_generate": True,
|
||||
"tooltip": "Seed for random number generation. A value of 0 will use a random seed."
|
||||
}),
|
||||
"use_watermark": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Enable or disable the audio watermark. Requires 'resemble-perth' to be installed."
|
||||
}),
|
||||
}, "optional": {"audio_prompt": ("AUDIO",),}}
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"model_pack_name": (displayed_packs, {"default": DEFAULT_MODEL_PACK_NAME if DEFAULT_MODEL_PACK_NAME in displayed_packs else (displayed_packs[0] if displayed_packs else DEFAULT_MODEL_PACK_NAME)}),
|
||||
"text": ("STRING", {"multiline": True, "default": "Hello, this is a test of Chatterbox TTS in ComfyUI."}),
|
||||
"exaggeration": ("FLOAT", {"default": 0.5, "min": 0.25, "max": 2.0, "step": 0.05}),
|
||||
"temperature": ("FLOAT", {"default": 0.8, "min": 0.05, "max": 5.0, "step": 0.05}),
|
||||
"cfg_weight": ("FLOAT", {"default": 0.5, "min": 0.2, "max": 1.0, "step": 0.05}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "control_after_generate": True}),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda" if torch.cuda.is_available() else "cpu"}),
|
||||
},
|
||||
"optional": {
|
||||
"audio_prompt": ("AUDIO",),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("AUDIO",); RETURN_NAMES = ("audio",); FUNCTION = "synthesize"; CATEGORY = "audio/generation"; OUTPUT_NODE = True
|
||||
|
||||
RETURN_TYPES = ("AUDIO",)
|
||||
RETURN_NAMES = ("audio",)
|
||||
FUNCTION = "synthesize"
|
||||
CATEGORY = "audio/generation"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def synthesize(self, model_pack_name, text, exaggeration, temperature, cfg_weight, seed, device, audio_prompt=None):
|
||||
def synthesize(self, model_pack_name, text, max_new_tokens, flow_cfg_scale, exaggeration, temperature, cfg_weight, repetition_penalty, min_p, top_p, seed, use_watermark, audio_prompt=None):
|
||||
if not text.strip():
|
||||
#print("Chatterbox TTS: Empty text provided, returning silent audio.")
|
||||
dummy_sr = 24000
|
||||
silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
|
||||
logger.info("Empty text provided, returning silent audio.")
|
||||
dummy_sr = 24000; silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
|
||||
return ({"waveform": silent_waveform.unsqueeze(0), "sample_rate": dummy_sr},)
|
||||
|
||||
try:
|
||||
chatterbox_model = get_cached_chatterbox_tts_model(model_pack_name, device_str=device)
|
||||
except Exception as e:
|
||||
print(f"ChatterboxTTS: Error loading/downloading TTS model pack '{model_pack_name}': {e}")
|
||||
dummy_sr = 24000
|
||||
silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
|
||||
cache_key = model_pack_name
|
||||
if cache_key not in CHATTERBOX_PATCHER_CACHE:
|
||||
load_device = mm.get_torch_device()
|
||||
logger.info(f"Creating Chatterbox ModelPatcher for {model_pack_name} on device {load_device}")
|
||||
model_wrapper = ChatterboxModelWrapper(model_pack_name)
|
||||
patcher = ChatterboxPatcher(
|
||||
model=model_wrapper,
|
||||
load_device=load_device,
|
||||
offload_device=mm.unet_offload_device(),
|
||||
size=int(1.5 * 1024**3)
|
||||
)
|
||||
CHATTERBOX_PATCHER_CACHE[cache_key] = patcher
|
||||
|
||||
patcher = CHATTERBOX_PATCHER_CACHE[cache_key]
|
||||
|
||||
mm.load_model_gpu(patcher)
|
||||
tts_model = patcher.model.tts_model
|
||||
|
||||
if tts_model is None:
|
||||
logger.error("TTS model failed to load. Please check logs for download or loading errors.")
|
||||
dummy_sr = 24000; silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
|
||||
return ({"waveform": silent_waveform.unsqueeze(0), "sample_rate": dummy_sr},)
|
||||
|
||||
set_chatterbox_seed(seed)
|
||||
|
||||
is_perth_installed = not getattr(perth, '_is_mock', False)
|
||||
if use_watermark and not is_perth_installed:
|
||||
logger.warning("Watermarking is enabled, but 'resemble-perth' is not installed. Output will not be watermarked.")
|
||||
|
||||
original_watermarker = tts_model.watermarker
|
||||
if not use_watermark:
|
||||
class TmpDummyWatermarker:
|
||||
def apply_watermark(self, wav, sample_rate): return wav
|
||||
tts_model.watermarker = TmpDummyWatermarker()
|
||||
if is_perth_installed: logger.info("Watermarking disabled by user.")
|
||||
|
||||
wav_tensor_chatterbox = None; audio_prompt_path_temp = None
|
||||
|
||||
audio_prompt_path_temp = None
|
||||
if audio_prompt is not None and \
|
||||
audio_prompt.get("waveform") is not None and \
|
||||
audio_prompt["waveform"].numel() > 0:
|
||||
|
||||
waveform_in = audio_prompt["waveform"]
|
||||
sample_rate_in = audio_prompt["sample_rate"]
|
||||
waveform_cpu = waveform_in.cpu()
|
||||
if waveform_cpu.shape[0] > 1:
|
||||
print(f"ChatterboxTTS: Audio prompt has batch size {waveform_cpu.shape[0]}, using first item.")
|
||||
current_waveform = waveform_cpu[0]
|
||||
if current_waveform.shape[0] > 1:
|
||||
current_waveform = torch.mean(current_waveform, dim=0)
|
||||
else:
|
||||
current_waveform = current_waveform.squeeze(0)
|
||||
|
||||
processed_audio_prompt = current_waveform.numpy().astype(np.float32)
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp_wav:
|
||||
audio_prompt_path_temp = tmp_wav.name
|
||||
sf.write(audio_prompt_path_temp, processed_audio_prompt, sample_rate_in)
|
||||
#print(f"ChatterboxTTS: Using audio prompt from temp file: {audio_prompt_path_temp}")
|
||||
except Exception as e:
|
||||
print(f"ChatterboxTTS: Error writing temp audio prompt file: {e}")
|
||||
audio_prompt_path_temp = None
|
||||
pbar = ProgressBar(max_new_tokens)
|
||||
|
||||
try:
|
||||
wav_tensor_chatterbox = chatterbox_model.generate(
|
||||
if audio_prompt and audio_prompt.get("waveform") is not None and audio_prompt["waveform"].numel() > 0:
|
||||
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp_wav:
|
||||
audio_prompt_path_temp = tmp_wav.name
|
||||
waveform_in = audio_prompt["waveform"]; sample_rate_in = audio_prompt["sample_rate"]
|
||||
waveform_cpu = waveform_in.cpu()[0]
|
||||
current_waveform = torch.mean(waveform_cpu, dim=0) if waveform_cpu.shape[0] > 1 else waveform_cpu.squeeze(0)
|
||||
sf.write(audio_prompt_path_temp, current_waveform.numpy().astype(np.float32), sample_rate_in)
|
||||
|
||||
wav_tensor_chatterbox = tts_model.generate(
|
||||
text,
|
||||
audio_prompt_path=audio_prompt_path_temp,
|
||||
exaggeration=exaggeration,
|
||||
temperature=temperature,
|
||||
cfg_weight=cfg_weight
|
||||
)
|
||||
cfg_weight=cfg_weight,
|
||||
repetition_penalty=repetition_penalty,
|
||||
min_p=min_p,
|
||||
top_p=top_p,
|
||||
pbar=pbar,
|
||||
max_new_tokens=max_new_tokens,
|
||||
flow_cfg_scale=flow_cfg_scale
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"ChatterboxTTS: Error during TTS generation: {e}")
|
||||
dummy_sr = 24000
|
||||
silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
|
||||
logger.error(f"Error during TTS generation: {e}", exc_info=True)
|
||||
dummy_sr = 24000; silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
|
||||
return ({"waveform": silent_waveform.unsqueeze(0), "sample_rate": dummy_sr},)
|
||||
finally:
|
||||
tts_model.watermarker = original_watermarker
|
||||
if audio_prompt_path_temp and os.path.exists(audio_prompt_path_temp):
|
||||
try:
|
||||
os.remove(audio_prompt_path_temp)
|
||||
except Exception as e:
|
||||
print(f"ChatterboxTTS: Error removing temp audio prompt file {audio_prompt_path_temp}: {e}")
|
||||
|
||||
wav_tensor_comfy = wav_tensor_chatterbox.cpu().unsqueeze(0)
|
||||
return ({"waveform": wav_tensor_comfy, "sample_rate": chatterbox_model.sr},)
|
||||
try: os.remove(audio_prompt_path_temp)
|
||||
except Exception as e: logger.error(f"Error removing temp audio prompt file: {e}")
|
||||
|
||||
wav_tensor_comfy = wav_tensor_chatterbox.cpu().unsqueeze(0)
|
||||
return ({"waveform": wav_tensor_comfy, "sample_rate": tts_model.sr},)
|
||||
|
||||
class ChatterboxVCNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
available_model_packs = get_chatterbox_model_pack_names()
|
||||
displayed_packs = [DEFAULT_MODEL_PACK_NAME] + [p for p in available_model_packs if p != DEFAULT_MODEL_PACK_NAME]
|
||||
if not displayed_packs:
|
||||
displayed_packs = [DEFAULT_MODEL_PACK_NAME]
|
||||
return {"required": {
|
||||
"model_pack_name": (get_chatterbox_model_pack_names(), {
|
||||
"default": DEFAULT_MODEL_PACK_NAME,
|
||||
"tooltip": "Select the Chatterbox voice model pack to use for conversion."
|
||||
}),
|
||||
"source_audio": ("AUDIO", {
|
||||
"tooltip": "The audio containing the speech content to be converted."
|
||||
}),
|
||||
"n_timesteps": ("INT", {
|
||||
"default": 10, "min": 2, "max": 50, "step": 1,
|
||||
"tooltip": "Number of diffusion steps for the flow matching process. Higher values may improve quality at the cost of speed."
|
||||
}),
|
||||
"temperature": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.05,
|
||||
"tooltip": "Controls the randomness of the initial noise. 1.0 is standard. Lower values are more deterministic."}),
|
||||
"flow_cfg_scale": ("FLOAT", {"default": 0.7, "min": 0.0, "max": 2.0, "step": 0.05,
|
||||
"tooltip": "CFG scale for the mel spectrogram decoder. Higher values increase adherence to the target voice but may reduce naturalness."}),
|
||||
"use_watermark": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Enable or disable the audio watermark. Requires 'resemble-perth' to be installed."
|
||||
}),
|
||||
},
|
||||
"optional": {"target_voice_audio": ("AUDIO", {
|
||||
"tooltip": "The audio file containing the target voice timbre. If not provided, the default voice from the model pack will be used."
|
||||
}), }}
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"model_pack_name": (displayed_packs, {"default": DEFAULT_MODEL_PACK_NAME if DEFAULT_MODEL_PACK_NAME in displayed_packs else (displayed_packs[0] if displayed_packs else DEFAULT_MODEL_PACK_NAME)}),
|
||||
"source_audio": ("AUDIO",),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda" if torch.cuda.is_available() else "cpu"}),
|
||||
},
|
||||
"optional": {
|
||||
"target_voice_audio": ("AUDIO",), # Optional: if not provided, uses default voice from conds.pt
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("AUDIO",)
|
||||
RETURN_NAMES = ("converted_audio",)
|
||||
FUNCTION = "convert_voice"
|
||||
CATEGORY = "audio/generation"
|
||||
OUTPUT_NODE = True
|
||||
RETURN_TYPES = ("AUDIO",); RETURN_NAMES = ("converted_audio",); FUNCTION = "convert_voice"; CATEGORY = "audio/generation"; OUTPUT_NODE = True
|
||||
|
||||
def _save_audio_to_temp_file(self, audio_data, prefix=""):
|
||||
"""Helper to save ComfyUI AUDIO dict to a temporary WAV file."""
|
||||
if audio_data is None or \
|
||||
audio_data.get("waveform") is None or \
|
||||
audio_data["waveform"].numel() == 0:
|
||||
return None
|
||||
if audio_data is None or audio_data.get("waveform") is None or audio_data["waveform"].numel() == 0: return None
|
||||
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp_wav:
|
||||
try:
|
||||
waveform_in = audio_data["waveform"]; sample_rate_in = audio_data["sample_rate"]
|
||||
waveform_cpu = waveform_in.cpu()[0]
|
||||
current_waveform = torch.mean(waveform_cpu, dim=0) if waveform_cpu.shape[0] > 1 else waveform_cpu.squeeze(0)
|
||||
sf.write(tmp_wav.name, current_waveform.numpy().astype(np.float32), sample_rate_in)
|
||||
return tmp_wav.name
|
||||
except Exception as e:
|
||||
logger.error(f"Error writing temp {prefix}audio file: {e}", exc_info=True)
|
||||
return None
|
||||
|
||||
waveform_in = audio_data["waveform"]
|
||||
sample_rate_in = audio_data["sample_rate"]
|
||||
waveform_cpu = waveform_in.cpu()
|
||||
|
||||
if waveform_cpu.shape[0] > 1:
|
||||
print(f"ChatterboxVC: {prefix}Audio has batch size {waveform_cpu.shape[0]}, using first item.")
|
||||
current_waveform = waveform_cpu[0]
|
||||
|
||||
if current_waveform.shape[0] > 1: # If C > 1 (stereo or more)
|
||||
print(f"ChatterboxVC: {prefix}Audio has {current_waveform.shape[0]} channels, converting to mono by averaging.")
|
||||
current_waveform = torch.mean(current_waveform, dim=0)
|
||||
else: # If C == 1
|
||||
current_waveform = current_waveform.squeeze(0)
|
||||
|
||||
processed_audio = current_waveform.numpy().astype(np.float32)
|
||||
|
||||
temp_file_path = None
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp_wav:
|
||||
temp_file_path = tmp_wav.name
|
||||
sf.write(temp_file_path, processed_audio, sample_rate_in)
|
||||
#print(f"ChatterboxVC: Saved {prefix}audio to temp file: {temp_file_path}")
|
||||
return temp_file_path
|
||||
except Exception as e:
|
||||
print(f"ChatterboxVC: Error writing temp {prefix}audio file: {e}")
|
||||
if temp_file_path and os.path.exists(temp_file_path):
|
||||
os.remove(temp_file_path)
|
||||
return None
|
||||
|
||||
def convert_voice(self, model_pack_name, source_audio, device, target_voice_audio=None):
|
||||
def convert_voice(self, model_pack_name, source_audio, n_timesteps, temperature, flow_cfg_scale, use_watermark, target_voice_audio=None):
|
||||
if source_audio is None or source_audio.get("waveform") is None or source_audio["waveform"].numel() == 0:
|
||||
print("ChatterboxVC: No source audio provided, returning silent audio.")
|
||||
dummy_sr = 24000
|
||||
silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
|
||||
logger.warning("No source audio provided, returning silent audio.")
|
||||
dummy_sr = 24000; silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
|
||||
return ({"waveform": silent_waveform.unsqueeze(0), "sample_rate": dummy_sr},)
|
||||
|
||||
try:
|
||||
vc_model = get_cached_chatterbox_vc_model(model_pack_name, device_str=device)
|
||||
except Exception as e:
|
||||
print(f"ChatterboxVC: Error loading/downloading VC model pack '{model_pack_name}': {e}")
|
||||
dummy_sr = 24000
|
||||
silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
|
||||
cache_key = model_pack_name
|
||||
if cache_key not in CHATTERBOX_PATCHER_CACHE:
|
||||
load_device = mm.get_torch_device()
|
||||
logger.info(f"Creating Chatterbox ModelPatcher for {model_pack_name} on device {load_device}")
|
||||
model_wrapper = ChatterboxModelWrapper(model_pack_name)
|
||||
patcher = ChatterboxPatcher(
|
||||
model=model_wrapper,
|
||||
load_device=load_device,
|
||||
offload_device=mm.unet_offload_device(),
|
||||
size=int(1.5 * 1024**3)
|
||||
)
|
||||
CHATTERBOX_PATCHER_CACHE[cache_key] = patcher
|
||||
|
||||
patcher = CHATTERBOX_PATCHER_CACHE[cache_key]
|
||||
|
||||
mm.load_model_gpu(patcher)
|
||||
vc_model = patcher.model.vc_model
|
||||
|
||||
if vc_model is None:
|
||||
logger.error("VC model failed to load. Please check logs for download or loading errors.")
|
||||
dummy_sr = 24000; silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
|
||||
return ({"waveform": silent_waveform.unsqueeze(0), "sample_rate": dummy_sr},)
|
||||
|
||||
source_audio_path_temp = None
|
||||
target_voice_path_temp = None
|
||||
is_perth_installed = not getattr(perth, '_is_mock', False)
|
||||
if use_watermark and not is_perth_installed:
|
||||
logger.warning("Watermarking is enabled, but 'resemble-perth' is not installed. Output will not be watermarked.")
|
||||
|
||||
original_watermarker = vc_model.watermarker
|
||||
if not use_watermark:
|
||||
class TmpDummyWatermarker:
|
||||
def apply_watermark(self, wav, sample_rate): return wav
|
||||
vc_model.watermarker = TmpDummyWatermarker()
|
||||
if is_perth_installed: logger.info("Watermarking disabled by user.")
|
||||
|
||||
source_audio_path_temp = None; target_voice_path_temp = None
|
||||
|
||||
pbar = ProgressBar(n_timesteps)
|
||||
|
||||
try:
|
||||
source_audio_path_temp = self._save_audio_to_temp_file(source_audio, prefix="Source ")
|
||||
if not source_audio_path_temp:
|
||||
raise ValueError("Failed to process source audio.")
|
||||
|
||||
if target_voice_audio is not None and \
|
||||
target_voice_audio.get("waveform") is not None and \
|
||||
target_voice_audio["waveform"].numel() > 0:
|
||||
if not source_audio_path_temp: raise ValueError("Failed to process source audio.")
|
||||
if target_voice_audio and target_voice_audio.get("waveform") is not None and target_voice_audio["waveform"].numel() > 0:
|
||||
target_voice_path_temp = self._save_audio_to_temp_file(target_voice_audio, prefix="Target ")
|
||||
#print(f"ChatterboxVC: Using target voice from temp file: {target_voice_path_temp}")
|
||||
else:
|
||||
print("ChatterboxVC: No target voice audio provided or it's empty. Using default reference from model pack if available.")
|
||||
|
||||
# ChatterboxVC.generate expects file paths
|
||||
converted_wav_tensor = vc_model.generate(
|
||||
audio=source_audio_path_temp,
|
||||
target_voice_path=target_voice_path_temp # This will be None if no target_voice_audio was provided or saving it failed
|
||||
) # Expected output: (1, num_samples)
|
||||
|
||||
target_voice_path=target_voice_path_temp,
|
||||
n_timesteps=n_timesteps,
|
||||
pbar=pbar,
|
||||
temperature=temperature,
|
||||
flow_cfg_scale=flow_cfg_scale
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"ChatterboxVC: Error during voice conversion: {e}")
|
||||
dummy_sr = 24000
|
||||
silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
|
||||
logger.error(f"Error during voice conversion: {e}", exc_info=True)
|
||||
dummy_sr = 24000; silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
|
||||
return ({"waveform": silent_waveform.unsqueeze(0), "sample_rate": dummy_sr},)
|
||||
finally:
|
||||
if source_audio_path_temp and os.path.exists(source_audio_path_temp):
|
||||
try:
|
||||
os.remove(source_audio_path_temp)
|
||||
except Exception as e:
|
||||
print(f"ChatterboxVC: Error removing temp source audio file {source_audio_path_temp}: {e}")
|
||||
if target_voice_path_temp and os.path.exists(target_voice_path_temp):
|
||||
try:
|
||||
os.remove(target_voice_path_temp)
|
||||
except Exception as e:
|
||||
print(f"ChatterboxVC: Error removing temp target audio file {target_voice_path_temp}: {e}")
|
||||
vc_model.watermarker = original_watermarker
|
||||
if source_audio_path_temp and os.path.exists(source_audio_path_temp): os.remove(source_audio_path_temp)
|
||||
if target_voice_path_temp and os.path.exists(target_voice_path_temp): os.remove(target_voice_path_temp)
|
||||
|
||||
# ComfyUI AUDIO format: {"waveform": tensor (B, C, T), "sample_rate": int}
|
||||
# ChatterboxVC output: tensor (1, T)
|
||||
vc_wav_tensor_comfy = converted_wav_tensor.cpu().unsqueeze(0) # (1, T) -> (1, 1, T)
|
||||
return ({"waveform": vc_wav_tensor_comfy, "sample_rate": vc_model.sr},)
|
||||
vc_wav_tensor_comfy = converted_wav_tensor.cpu().unsqueeze(0)
|
||||
return ({"waveform": vc_wav_tensor_comfy, "sample_rate": vc_model.sr},)
|
||||
|
||||
@@ -1,2 +1,9 @@
|
||||
try:
|
||||
from importlib.metadata import version
|
||||
except ImportError:
|
||||
from importlib_metadata import version # For Python <3.8
|
||||
|
||||
# __version__ = version("chatterbox-tts")
|
||||
|
||||
from .tts import ChatterboxTTS
|
||||
from .vc import ChatterboxVC
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
from ..utils import AttrDict
|
||||
|
||||
CFM_PARAMS = AttrDict({
|
||||
"sigma_min": 1e-06,
|
||||
"solver": "euler",
|
||||
"t_scheduler": "cosine",
|
||||
"training_cfg_rate": 0.2,
|
||||
"inference_cfg_rate": 0.7,
|
||||
"reg_loss_type": "l1"
|
||||
})
|
||||
@@ -14,32 +14,54 @@
|
||||
import logging
|
||||
import random
|
||||
from typing import Dict, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn import functional as F
|
||||
from omegaconf import DictConfig
|
||||
from .utils.mask import make_pad_mask
|
||||
from .configs import CFM_PARAMS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
class MaskedDiffWithXvec(torch.nn.Module):
|
||||
def __init__(self,
|
||||
input_size: int = 512,
|
||||
output_size: int = 80,
|
||||
spk_embed_dim: int = 192,
|
||||
output_type: str = "mel",
|
||||
vocab_size: int = 4096,
|
||||
input_frame_rate: int = 50,
|
||||
only_mask_loss: bool = True,
|
||||
encoder: torch.nn.Module = None,
|
||||
length_regulator: torch.nn.Module = None,
|
||||
decoder: torch.nn.Module = None,
|
||||
decoder_conf: Dict = {'in_channels': 240, 'out_channel': 80, 'spk_emb_dim': 80, 'n_spks': 1,
|
||||
'cfm_params': DictConfig({'sigma_min': 1e-06, 'solver': 'euler', 't_scheduler': 'cosine',
|
||||
'training_cfg_rate': 0.2, 'inference_cfg_rate': 0.7, 'reg_loss_type': 'l1'}),
|
||||
'decoder_params': {'channels': [256, 256], 'dropout': 0.0, 'attention_head_dim': 64,
|
||||
'n_blocks': 4, 'num_mid_blocks': 12, 'num_heads': 8, 'act_fn': 'gelu'}},
|
||||
mel_feat_conf: Dict = {'n_fft': 1024, 'num_mels': 80, 'sampling_rate': 22050,
|
||||
'hop_size': 256, 'win_size': 1024, 'fmin': 0, 'fmax': 8000}):
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int = 512,
|
||||
output_size: int = 80,
|
||||
spk_embed_dim: int = 192,
|
||||
output_type: str = "mel",
|
||||
vocab_size: int = 4096,
|
||||
input_frame_rate: int = 50,
|
||||
only_mask_loss: bool = True,
|
||||
encoder: torch.nn.Module = None,
|
||||
length_regulator: torch.nn.Module = None,
|
||||
decoder: torch.nn.Module = None,
|
||||
decoder_conf: Dict = {
|
||||
'in_channels': 240,
|
||||
'out_channel': 80,
|
||||
'spk_emb_dim': 80,
|
||||
'n_spks': 1,
|
||||
'cfm_params': CFM_PARAMS,
|
||||
'decoder_params': {
|
||||
'channels': [256, 256],
|
||||
'dropout': 0.0,
|
||||
'attention_head_dim': 64,
|
||||
'n_blocks': 4,
|
||||
'num_mid_blocks': 12,
|
||||
'num_heads': 8,
|
||||
'act_fn': 'gelu',
|
||||
}
|
||||
},
|
||||
mel_feat_conf: Dict = {
|
||||
'n_fft': 1024,
|
||||
'num_mels': 80,
|
||||
'sampling_rate': 22050,
|
||||
'hop_size': 256,
|
||||
'win_size': 1024,
|
||||
'fmin': 0,
|
||||
'fmax': 8000
|
||||
}
|
||||
):
|
||||
super().__init__()
|
||||
self.input_size = input_size
|
||||
self.output_size = output_size
|
||||
@@ -48,7 +70,7 @@ class MaskedDiffWithXvec(torch.nn.Module):
|
||||
self.vocab_size = vocab_size
|
||||
self.output_type = output_type
|
||||
self.input_frame_rate = input_frame_rate
|
||||
logging.info(f"input frame rate={self.input_frame_rate}")
|
||||
# logger.info(f"input frame rate={self.input_frame_rate}")
|
||||
self.input_embedding = nn.Embedding(vocab_size, input_size)
|
||||
self.spk_embed_affine_layer = torch.nn.Linear(spk_embed_dim, output_size)
|
||||
self.encoder = encoder
|
||||
@@ -110,8 +132,13 @@ class MaskedDiffWithXvec(torch.nn.Module):
|
||||
prompt_feat,
|
||||
prompt_feat_len,
|
||||
embedding,
|
||||
flow_cache):
|
||||
if self.fp16 is True:
|
||||
flow_cache,
|
||||
n_timesteps=10,
|
||||
pbar=None,
|
||||
temperature=1.0,
|
||||
flow_cfg_scale=0.7
|
||||
):
|
||||
if hasattr(self, 'fp16') and self.fp16 is True:
|
||||
prompt_feat = prompt_feat.half()
|
||||
embedding = embedding.half()
|
||||
|
||||
@@ -143,9 +170,12 @@ class MaskedDiffWithXvec(torch.nn.Module):
|
||||
mask=mask.unsqueeze(1),
|
||||
spks=embedding,
|
||||
cond=conds,
|
||||
n_timesteps=10,
|
||||
n_timesteps=n_timesteps,
|
||||
prompt_len=mel_len1,
|
||||
flow_cache=flow_cache
|
||||
flow_cache=flow_cache,
|
||||
pbar=pbar,
|
||||
temperature=temperature,
|
||||
flow_cfg_scale=flow_cfg_scale
|
||||
)
|
||||
feat = feat[:, :, mel_len1:]
|
||||
assert feat.shape[2] == mel_len2
|
||||
@@ -153,25 +183,45 @@ class MaskedDiffWithXvec(torch.nn.Module):
|
||||
|
||||
|
||||
class CausalMaskedDiffWithXvec(torch.nn.Module):
|
||||
def __init__(self,
|
||||
input_size: int = 512,
|
||||
output_size: int = 80,
|
||||
spk_embed_dim: int = 192,
|
||||
output_type: str = "mel",
|
||||
vocab_size: int = 6561,
|
||||
input_frame_rate: int = 25,
|
||||
only_mask_loss: bool = True,
|
||||
token_mel_ratio: int = 2,
|
||||
pre_lookahead_len: int = 3,
|
||||
encoder: torch.nn.Module = None,
|
||||
decoder: torch.nn.Module = None,
|
||||
decoder_conf: Dict = {'in_channels': 240, 'out_channel': 80, 'spk_emb_dim': 80, 'n_spks': 1,
|
||||
'cfm_params': DictConfig({'sigma_min': 1e-06, 'solver': 'euler', 't_scheduler': 'cosine',
|
||||
'training_cfg_rate': 0.2, 'inference_cfg_rate': 0.7, 'reg_loss_type': 'l1'}),
|
||||
'decoder_params': {'channels': [256, 256], 'dropout': 0.0, 'attention_head_dim': 64,
|
||||
'n_blocks': 4, 'num_mid_blocks': 12, 'num_heads': 8, 'act_fn': 'gelu'}},
|
||||
mel_feat_conf: Dict = {'n_fft': 1024, 'num_mels': 80, 'sampling_rate': 22050,
|
||||
'hop_size': 256, 'win_size': 1024, 'fmin': 0, 'fmax': 8000}):
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int = 512,
|
||||
output_size: int = 80,
|
||||
spk_embed_dim: int = 192,
|
||||
output_type: str = "mel",
|
||||
vocab_size: int = 6561,
|
||||
input_frame_rate: int = 25,
|
||||
only_mask_loss: bool = True,
|
||||
token_mel_ratio: int = 2,
|
||||
pre_lookahead_len: int = 3,
|
||||
encoder: torch.nn.Module = None,
|
||||
decoder: torch.nn.Module = None,
|
||||
decoder_conf: Dict = {
|
||||
'in_channels': 240,
|
||||
'out_channel': 80,
|
||||
'spk_emb_dim': 80,
|
||||
'n_spks': 1,
|
||||
'cfm_params': CFM_PARAMS,
|
||||
'decoder_params': {
|
||||
'channels': [256, 256],
|
||||
'dropout': 0.0,
|
||||
'attention_head_dim': 64,
|
||||
'n_blocks': 4,
|
||||
'num_mid_blocks': 12,
|
||||
'num_heads': 8,
|
||||
'act_fn': 'gelu',
|
||||
}
|
||||
},
|
||||
mel_feat_conf: Dict = {
|
||||
'n_fft': 1024,
|
||||
'num_mels': 80,
|
||||
'sampling_rate': 22050,
|
||||
'hop_size': 256,
|
||||
'win_size': 1024,
|
||||
'fmin': 0,
|
||||
'fmax': 8000
|
||||
}
|
||||
):
|
||||
super().__init__()
|
||||
self.input_size = input_size
|
||||
self.output_size = output_size
|
||||
@@ -180,7 +230,7 @@ class CausalMaskedDiffWithXvec(torch.nn.Module):
|
||||
self.vocab_size = vocab_size
|
||||
self.output_type = output_type
|
||||
self.input_frame_rate = input_frame_rate
|
||||
logging.info(f"input frame rate={self.input_frame_rate}")
|
||||
# logger.info(f"input frame rate={self.input_frame_rate}")
|
||||
self.input_embedding = nn.Embedding(vocab_size, input_size)
|
||||
self.spk_embed_affine_layer = torch.nn.Linear(spk_embed_dim, output_size)
|
||||
self.encoder = encoder
|
||||
@@ -202,8 +252,13 @@ class CausalMaskedDiffWithXvec(torch.nn.Module):
|
||||
prompt_feat,
|
||||
prompt_feat_len,
|
||||
embedding,
|
||||
finalize):
|
||||
if self.fp16 is True:
|
||||
finalize,
|
||||
n_timesteps=10,
|
||||
pbar=None,
|
||||
temperature=1.0,
|
||||
flow_cfg_scale=0.7
|
||||
):
|
||||
if hasattr(self, 'fp16') and self.fp16 is True:
|
||||
prompt_feat = prompt_feat.half()
|
||||
embedding = embedding.half()
|
||||
|
||||
@@ -235,7 +290,10 @@ class CausalMaskedDiffWithXvec(torch.nn.Module):
|
||||
mask=mask.unsqueeze(1),
|
||||
spks=embedding,
|
||||
cond=conds,
|
||||
n_timesteps=10
|
||||
n_timesteps=n_timesteps,
|
||||
pbar=pbar,
|
||||
temperature=temperature,
|
||||
flow_cfg_scale=flow_cfg_scale,
|
||||
)
|
||||
feat = feat[:, :, mel_len1:]
|
||||
assert feat.shape[2] == mel_len2
|
||||
|
||||
@@ -14,18 +14,9 @@
|
||||
import threading
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from tqdm import tqdm
|
||||
from .matcha.flow_matching import BASECFM
|
||||
from omegaconf import OmegaConf
|
||||
|
||||
|
||||
CFM_PARAMS = OmegaConf.create({
|
||||
"sigma_min": 1e-06,
|
||||
"solver": "euler",
|
||||
"t_scheduler": "cosine",
|
||||
"training_cfg_rate": 0.2,
|
||||
"inference_cfg_rate": 0.7,
|
||||
"reg_loss_type": "l1"
|
||||
})
|
||||
from .configs import CFM_PARAMS
|
||||
|
||||
|
||||
class ConditionalCFM(BASECFM):
|
||||
@@ -45,7 +36,7 @@ class ConditionalCFM(BASECFM):
|
||||
self.lock = threading.Lock()
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(self, mu, mask, n_timesteps, temperature=1.0, spks=None, cond=None, prompt_len=0, flow_cache=torch.zeros(1, 80, 0, 2)):
|
||||
def forward(self, mu, mask, n_timesteps, temperature=1.0, spks=None, cond=None, prompt_len=0, flow_cache=torch.zeros(1, 80, 0, 2), pbar=None, flow_cfg_scale=None):
|
||||
"""Forward diffusion
|
||||
|
||||
Args:
|
||||
@@ -58,11 +49,15 @@ class ConditionalCFM(BASECFM):
|
||||
spks (torch.Tensor, optional): speaker ids. Defaults to None.
|
||||
shape: (batch_size, spk_emb_dim)
|
||||
cond: Not used but kept for future purposes
|
||||
pbar: ComfyUI ProgressBar instance
|
||||
|
||||
Returns:
|
||||
sample: generated mel-spectrogram
|
||||
shape: (batch_size, n_feats, mel_timesteps)
|
||||
"""
|
||||
|
||||
if flow_cfg_scale is not None:
|
||||
self.inference_cfg_rate = flow_cfg_scale
|
||||
|
||||
z = torch.randn_like(mu).to(mu.device).to(mu.dtype) * temperature
|
||||
cache_size = flow_cache.shape[2]
|
||||
@@ -77,9 +72,10 @@ class ConditionalCFM(BASECFM):
|
||||
t_span = torch.linspace(0, 1, n_timesteps + 1, device=mu.device, dtype=mu.dtype)
|
||||
if self.t_scheduler == 'cosine':
|
||||
t_span = 1 - torch.cos(t_span * 0.5 * torch.pi)
|
||||
return self.solve_euler(z, t_span=t_span, mu=mu, mask=mask, spks=spks, cond=cond), flow_cache
|
||||
|
||||
return self.solve_euler(z, t_span=t_span, mu=mu, mask=mask, spks=spks, cond=cond, pbar=pbar), flow_cache
|
||||
|
||||
def solve_euler(self, x, t_span, mu, mask, spks, cond):
|
||||
def solve_euler(self, x, t_span, mu, mask, spks, cond, pbar=None):
|
||||
"""
|
||||
Fixed euler solver for ODEs.
|
||||
Args:
|
||||
@@ -93,6 +89,7 @@ class ConditionalCFM(BASECFM):
|
||||
spks (torch.Tensor, optional): speaker ids. Defaults to None.
|
||||
shape: (batch_size, spk_emb_dim)
|
||||
cond: Not used but kept for future purposes
|
||||
pbar: ComfyUI ProgressBar instance
|
||||
"""
|
||||
t, _, dt = t_span[0], t_span[-1], t_span[1] - t_span[0]
|
||||
t = t.unsqueeze(dim=0)
|
||||
@@ -100,7 +97,8 @@ class ConditionalCFM(BASECFM):
|
||||
# I am storing this because I can later plot it by putting a debugger here and saving it to a file
|
||||
# Or in future might add like a return_all_steps flag
|
||||
sol = []
|
||||
|
||||
iterator = tqdm(range(1, len(t_span)), desc="Voice Converting", dynamic_ncols=True)
|
||||
|
||||
# Do not use concat, it may cause memory format changed and trt infer with wrong results!
|
||||
x_in = torch.zeros([2, 80, x.size(2)], device=x.device, dtype=x.dtype)
|
||||
mask_in = torch.zeros([2, 1, x.size(2)], device=x.device, dtype=x.dtype)
|
||||
@@ -108,7 +106,7 @@ class ConditionalCFM(BASECFM):
|
||||
t_in = torch.zeros([2], device=x.device, dtype=x.dtype)
|
||||
spks_in = torch.zeros([2, 80], device=x.device, dtype=x.dtype)
|
||||
cond_in = torch.zeros([2, 80, x.size(2)], device=x.device, dtype=x.dtype)
|
||||
for step in range(1, len(t_span)):
|
||||
for step in iterator:
|
||||
# Classifier-Free Guidance inference introduced in VoiceBox
|
||||
x_in[:] = x
|
||||
mask_in[:] = mask
|
||||
@@ -129,6 +127,9 @@ class ConditionalCFM(BASECFM):
|
||||
sol.append(x)
|
||||
if step < len(t_span) - 1:
|
||||
dt = t_span[step + 1] - t
|
||||
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
|
||||
return sol[-1].float()
|
||||
|
||||
@@ -201,7 +202,7 @@ class CausalConditionalCFM(ConditionalCFM):
|
||||
self.rand_noise = torch.randn([1, 80, 50 * 300])
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(self, mu, mask, n_timesteps, temperature=1.0, spks=None, cond=None):
|
||||
def forward(self, mu, mask, n_timesteps, temperature=1.0, spks=None, cond=None, pbar=None, flow_cfg_scale=None, **kwargs):
|
||||
"""Forward diffusion
|
||||
|
||||
Args:
|
||||
@@ -220,9 +221,12 @@ class CausalConditionalCFM(ConditionalCFM):
|
||||
shape: (batch_size, n_feats, mel_timesteps)
|
||||
"""
|
||||
|
||||
if flow_cfg_scale is not None:
|
||||
self.inference_cfg_rate = flow_cfg_scale
|
||||
|
||||
z = self.rand_noise[:, :, :mu.size(2)].to(mu.device).to(mu.dtype) * temperature
|
||||
# fix prompt and overlap part mu and z
|
||||
t_span = torch.linspace(0, 1, n_timesteps + 1, device=mu.device, dtype=mu.dtype)
|
||||
if self.t_scheduler == 'cosine':
|
||||
t_span = 1 - torch.cos(t_span * 0.5 * torch.pi)
|
||||
return self.solve_euler(z, t_span=t_span, mu=mu, mask=mask, spks=spks, cond=cond), None
|
||||
return self.solve_euler(z, t_span=t_span, mu=mu, mask=mask, spks=spks, cond=cond, pbar=pbar), None
|
||||
|
||||
@@ -2,15 +2,21 @@ from typing import Any, Dict, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from diffusers.models.attention import (
|
||||
from diffusers.models.activations import (
|
||||
GEGLU,
|
||||
GELU,
|
||||
AdaLayerNorm,
|
||||
AdaLayerNormZero,
|
||||
ApproximateGELU,
|
||||
)
|
||||
from diffusers.models.normalization import (
|
||||
AdaLayerNorm,
|
||||
AdaLayerNormZero,
|
||||
)
|
||||
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from diffusers.models.lora import LoRACompatibleLinear
|
||||
|
||||
# deprecated LoRACompatibleLinear
|
||||
# from diffusers.models.lora import LoRACompatibleLinear
|
||||
|
||||
from diffusers.utils.torch_utils import maybe_allow_in_graph
|
||||
|
||||
|
||||
@@ -45,7 +51,9 @@ class SnakeBeta(nn.Module):
|
||||
"""
|
||||
super().__init__()
|
||||
self.in_features = out_features if isinstance(out_features, list) else [out_features]
|
||||
self.proj = LoRACompatibleLinear(in_features, out_features)
|
||||
|
||||
# Switched to standard torch.nn.Linear
|
||||
self.proj = nn.Linear(in_features, out_features)
|
||||
|
||||
# initialize alpha
|
||||
self.alpha_logscale = alpha_logscale
|
||||
@@ -123,7 +131,9 @@ class FeedForward(nn.Module):
|
||||
# project dropout
|
||||
self.net.append(nn.Dropout(dropout))
|
||||
# project out
|
||||
self.net.append(LoRACompatibleLinear(inner_dim, dim_out))
|
||||
|
||||
# Switched to standard torch.nn.Linear
|
||||
self.net.append(nn.Linear(inner_dim, dim_out))
|
||||
# FF as used in Vision Transformer, MLP-Mixer, etc. have a final dropout
|
||||
if final_dropout:
|
||||
self.net.append(nn.Dropout(dropout))
|
||||
|
||||
@@ -19,7 +19,6 @@ import torch
|
||||
import torchaudio as ta
|
||||
from functools import lru_cache
|
||||
from typing import Optional
|
||||
from omegaconf import DictConfig
|
||||
|
||||
from ..s3tokenizer import S3_SR, SPEECH_VOCAB_SIZE, S3Tokenizer
|
||||
from .const import S3GEN_SR
|
||||
@@ -31,7 +30,9 @@ from .hifigan import HiFTGenerator
|
||||
from .transformer.upsample_encoder import UpsampleConformerEncoder
|
||||
from .flow_matching import CausalConditionalCFM
|
||||
from .decoder import ConditionalDecoder
|
||||
from .configs import CFM_PARAMS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
def drop_invalid_tokens(x):
|
||||
assert len(x.shape) <= 2 and x.shape[0] == 1, "only batch size of one allowed for now"
|
||||
@@ -85,14 +86,7 @@ class S3Token2Mel(torch.nn.Module):
|
||||
num_heads=8,
|
||||
act_fn='gelu',
|
||||
)
|
||||
cfm_params = DictConfig({
|
||||
"sigma_min": 1e-06,
|
||||
"solver": 'euler',
|
||||
"t_scheduler": 'cosine',
|
||||
"training_cfg_rate": 0.2,
|
||||
"inference_cfg_rate": 0.7,
|
||||
"reg_loss_type": 'l1',
|
||||
})
|
||||
cfm_params = CFM_PARAMS
|
||||
decoder = CausalConditionalCFM(
|
||||
spk_emb_dim=80,
|
||||
cfm_params=cfm_params,
|
||||
@@ -129,7 +123,7 @@ class S3Token2Mel(torch.nn.Module):
|
||||
ref_wav = ref_wav.unsqueeze(0) # (B, L)
|
||||
|
||||
if ref_wav.size(1) > 10 * ref_sr:
|
||||
print("WARNING: cosydec received ref longer than 10s")
|
||||
logger.warning("WARNING: cosydec received ref longer than 10s")
|
||||
|
||||
ref_wav_24 = ref_wav
|
||||
if ref_sr != S3GEN_SR:
|
||||
@@ -149,7 +143,7 @@ class S3Token2Mel(torch.nn.Module):
|
||||
|
||||
# Make sure mel_len = 2 * stoken_len (happens when the input is not padded to multiple of 40ms)
|
||||
if ref_mels_24.shape[1] != 2 * ref_speech_tokens.shape[1]:
|
||||
logging.warning(
|
||||
logger.warning(
|
||||
"Reference mel length is not equal to 2 * reference token length.\n"
|
||||
)
|
||||
ref_speech_tokens = ref_speech_tokens[:, :ref_mels_24.shape[1] // 2]
|
||||
@@ -172,6 +166,8 @@ class S3Token2Mel(torch.nn.Module):
|
||||
# pre-computed ref embedding (prod API)
|
||||
ref_dict: Optional[dict] = None,
|
||||
finalize: bool = False,
|
||||
n_timesteps: int = 10,
|
||||
pbar=None
|
||||
):
|
||||
"""
|
||||
Generate waveforms from S3 speech tokens and a reference waveform, which the speaker timbre is inferred from.
|
||||
@@ -211,11 +207,15 @@ class S3Token2Mel(torch.nn.Module):
|
||||
token=speech_tokens,
|
||||
token_len=speech_token_lens,
|
||||
finalize=finalize,
|
||||
n_timesteps=n_timesteps,
|
||||
pbar=pbar,
|
||||
**ref_dict,
|
||||
)
|
||||
return output_mels
|
||||
|
||||
|
||||
|
||||
|
||||
class S3Token2Wav(S3Token2Mel):
|
||||
"""
|
||||
The decoder of CosyVoice2 is a concat of token-to-mel (CFM) and a mel-to-waveform (HiFiGAN) modules.
|
||||
@@ -250,9 +250,11 @@ class S3Token2Wav(S3Token2Mel):
|
||||
ref_sr: Optional[int],
|
||||
# pre-computed ref embedding (prod API)
|
||||
ref_dict: Optional[dict] = None,
|
||||
finalize: bool = False
|
||||
finalize: bool = False,
|
||||
n_timesteps: int = 10,
|
||||
pbar=None
|
||||
):
|
||||
output_mels = super().forward(speech_tokens, ref_wav=ref_wav, ref_sr=ref_sr, ref_dict=ref_dict, finalize=finalize)
|
||||
output_mels = super().forward(speech_tokens, ref_wav=ref_wav, ref_sr=ref_sr, ref_dict=ref_dict, finalize=finalize, n_timesteps=n_timesteps, pbar=pbar)
|
||||
|
||||
# TODO jrm: ignoring the speed control (mel interpolation) and the HiFTGAN caching mechanisms for now.
|
||||
hift_cache_source = torch.zeros(1, 1, 0).to(self.device)
|
||||
@@ -275,8 +277,41 @@ class S3Token2Wav(S3Token2Mel):
|
||||
# pre-computed ref embedding (prod API)
|
||||
ref_dict: Optional[dict] = None,
|
||||
finalize: bool = False,
|
||||
n_timesteps: int = 10,
|
||||
pbar=None,
|
||||
temperature: float = 1.0,
|
||||
flow_cfg_scale: float = 0.7
|
||||
):
|
||||
return super().forward(speech_tokens, ref_wav=ref_wav, ref_sr=ref_sr, ref_dict=ref_dict, finalize=finalize)
|
||||
# This method in the base class now needs to accept and pass the new params
|
||||
# (The base class `S3Token2Mel` needs this change)
|
||||
|
||||
assert (ref_wav is None) ^ (ref_dict is None), f"Must provide exactly one of ref_wav or ref_dict (got {ref_wav} and {ref_dict})"
|
||||
|
||||
if ref_dict is None:
|
||||
ref_dict = self.embed_ref(ref_wav, ref_sr)
|
||||
else:
|
||||
for rk in list(ref_dict):
|
||||
if isinstance(ref_dict[rk], np.ndarray):
|
||||
ref_dict[rk] = torch.from_numpy(ref_dict[rk])
|
||||
if torch.is_tensor(ref_dict[rk]):
|
||||
ref_dict[rk] = ref_dict[rk].to(self.device)
|
||||
|
||||
if len(speech_tokens.shape) == 1:
|
||||
speech_tokens = speech_tokens.unsqueeze(0)
|
||||
|
||||
speech_token_lens = torch.LongTensor([speech_tokens.size(1)]).to(self.device)
|
||||
|
||||
output_mels, _ = self.flow.inference(
|
||||
token=speech_tokens,
|
||||
token_len=speech_token_lens,
|
||||
finalize=finalize,
|
||||
n_timesteps=n_timesteps,
|
||||
pbar=pbar,
|
||||
temperature=temperature,
|
||||
flow_cfg_scale=flow_cfg_scale,
|
||||
**ref_dict,
|
||||
)
|
||||
return output_mels
|
||||
|
||||
@torch.inference_mode()
|
||||
def hift_inference(self, speech_feat, cache_source: torch.Tensor = None):
|
||||
@@ -295,8 +330,17 @@ class S3Token2Wav(S3Token2Mel):
|
||||
ref_dict: Optional[dict] = None,
|
||||
cache_source: torch.Tensor = None, # NOTE: this arg is for streaming, it can probably be removed here
|
||||
finalize: bool = True,
|
||||
n_timesteps: int = 10,
|
||||
pbar=None,
|
||||
temperature: float = 1.0,
|
||||
flow_cfg_scale: float = 0.7
|
||||
):
|
||||
output_mels = self.flow_inference(speech_tokens, ref_wav=ref_wav, ref_sr=ref_sr, ref_dict=ref_dict, finalize=finalize)
|
||||
output_mels = self.flow_inference(
|
||||
speech_tokens,
|
||||
ref_wav=ref_wav, ref_sr=ref_sr, ref_dict=ref_dict,
|
||||
finalize=finalize, n_timesteps=n_timesteps, pbar=pbar,
|
||||
temperature=temperature, flow_cfg_scale=flow_cfg_scale
|
||||
)
|
||||
output_wavs, output_sources = self.hift_inference(output_mels, cache_source)
|
||||
|
||||
# NOTE: ad-hoc method to reduce "spillover" from the reference clip.
|
||||
|
||||
@@ -1,17 +1,19 @@
|
||||
# FILE: src/chatterbox/models/t3/llama_configs.py
|
||||
LLAMA_520M_CONFIG_DICT = dict(
|
||||
# Arbitrary small number that won't cause problems when loading.
|
||||
# These param are unused due to custom input layers.
|
||||
vocab_size=8,
|
||||
# This defines the maximum sequence length the RoPE mechanism is initially configured for
|
||||
# if rope_scaling is None. 8192 is a common value for Llama models.
|
||||
max_position_embeddings=8192,
|
||||
# default params needed for loading most pretrained 1B weights
|
||||
max_position_embeddings=131072,
|
||||
hidden_size=1024,
|
||||
intermediate_size=4096,
|
||||
num_hidden_layers=30,
|
||||
num_attention_heads=16,
|
||||
attn_implementation="sdpa", # Use "eager" if "sdpa" causes issues with older torch/transformers
|
||||
head_dim=64, # hidden_size // num_attention_heads
|
||||
# This is required because AlignmentStreamAnalyzer needs to inspect attention weights,
|
||||
# which is not supported by optimized backends like FlashAttention (SDPA).
|
||||
# This informs transformers that the fallback to the slower path is intentional, silencing the warning.
|
||||
attn_implementation="eager",
|
||||
#attn_implementation="sdpa",
|
||||
head_dim=64,
|
||||
tie_word_embeddings=False,
|
||||
hidden_act="silu",
|
||||
attention_bias=False,
|
||||
@@ -19,15 +21,21 @@ LLAMA_520M_CONFIG_DICT = dict(
|
||||
initializer_range=0.02,
|
||||
mlp_bias=False,
|
||||
model_type="llama",
|
||||
num_key_value_heads=16, # For Llama GQA/MQA. For standard MHA, num_key_value_heads = num_attention_heads
|
||||
num_key_value_heads=16,
|
||||
pretraining_tp=1,
|
||||
rms_norm_eps=1e-05,
|
||||
rope_scaling=None, # MODIFICATION: Explicitly set to None
|
||||
rope_theta=500000.0, # This is crucial for Llama 3 style RoPE
|
||||
torch_dtype="bfloat16", # Consider "float16" or "float32" if bf16 is not supported or causes issues
|
||||
rope_scaling=dict(
|
||||
factor=8.0,
|
||||
high_freq_factor=4.0,
|
||||
low_freq_factor=1.0,
|
||||
original_max_position_embeddings=8192,
|
||||
rope_type="llama3"
|
||||
),
|
||||
rope_theta=500000.0,
|
||||
torch_dtype="bfloat16",
|
||||
use_cache=True,
|
||||
)
|
||||
|
||||
LLAMA_CONFIGS = {
|
||||
"Llama_520M": LLAMA_520M_CONFIG_DICT,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,7 +7,15 @@ import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from torch.nn.attention import SDPBackend, sdpa_kernel
|
||||
|
||||
try:
|
||||
from torch.nn.attention import SDPBackend, sdpa_kernel
|
||||
# Flag to indicate that the modern SDPA API is available.
|
||||
SUPPORTS_SDPA_KERNEL = True
|
||||
except (ImportError, AttributeError):
|
||||
# Fallback for older PyTorch versions.
|
||||
SUPPORTS_SDPA_KERNEL = False
|
||||
|
||||
|
||||
class RelativePositionBias(nn.Module):
|
||||
def __init__(self, scale, causal=False, num_buckets=32, max_distance=128, heads=8):
|
||||
@@ -66,6 +74,7 @@ class AttentionQKV(nn.Module):
|
||||
def setup_flash_config(self):
|
||||
# Setup flash attention configuration
|
||||
flash_config = {
|
||||
'enable_flash': True,
|
||||
'enable_math': True,
|
||||
'enable_mem_efficient': True
|
||||
}
|
||||
@@ -89,14 +98,32 @@ class AttentionQKV(nn.Module):
|
||||
return torch.einsum("bhts,bhls->bhlt", attn, v)
|
||||
|
||||
def flash_attention(self, q, k, v, mask=None):
|
||||
config = self.flash_config if self.flash_config else {}
|
||||
with sdpa_kernel(backends=[SDPBackend.EFFICIENT_ATTENTION]):
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
attn_mask=mask,
|
||||
dropout_p=self.dropout_rate if self.training else 0.
|
||||
)
|
||||
if SUPPORTS_SDPA_KERNEL:
|
||||
# Modern PyTorch (>= 2.2) API for enabling specific backends.
|
||||
# This replaces the deprecated `torch.backends.cuda.sdp_kernel`.
|
||||
backends = []
|
||||
if self.flash_config.get('enable_flash', True):
|
||||
backends.append(SDPBackend.FLASH_ATTENTION)
|
||||
if self.flash_config.get('enable_mem_efficient', True):
|
||||
backends.append(SDPBackend.EFFICIENT_ATTENTION)
|
||||
if self.flash_config.get('enable_math', True):
|
||||
backends.append(SDPBackend.MATH)
|
||||
|
||||
with sdpa_kernel(backends=backends):
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
attn_mask=mask,
|
||||
dropout_p=self.dropout_rate if self.training else 0.
|
||||
)
|
||||
else:
|
||||
# Fallback for older PyTorch versions.
|
||||
config = self.flash_config if self.flash_config else {}
|
||||
with torch.backends.cuda.sdp_kernel(**config):
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
attn_mask=mask,
|
||||
dropout_p=self.dropout_rate if self.training else 0.
|
||||
)
|
||||
return out
|
||||
|
||||
def split_heads(self, x):
|
||||
|
||||
@@ -8,7 +8,7 @@ import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn, Tensor
|
||||
from transformers import LlamaModel, LlamaConfig
|
||||
from transformers.generation.logits_process import TopPLogitsWarper, RepetitionPenaltyLogitsProcessor
|
||||
from transformers.generation.logits_process import MinPLogitsWarper, RepetitionPenaltyLogitsProcessor, TopPLogitsWarper
|
||||
|
||||
from .modules.learned_pos_emb import LearnedPositionEmbeddings
|
||||
|
||||
@@ -16,18 +16,12 @@ from .modules.cond_enc import T3CondEnc, T3Cond
|
||||
from .modules.t3_config import T3Config
|
||||
from .llama_configs import LLAMA_CONFIGS
|
||||
from .inference.t3_hf_backend import T3HuggingfaceBackend
|
||||
from .inference.alignment_stream_analyzer import AlignmentStreamAnalyzer
|
||||
from ..utils import AttrDict
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AttrDict(dict):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super(AttrDict, self).__init__(*args, **kwargs)
|
||||
self.__dict__ = self
|
||||
|
||||
|
||||
def _ensure_BOT_EOT(text_tokens: Tensor, hp):
|
||||
B = text_tokens.size(0)
|
||||
assert (text_tokens == hp.start_text_token).int().sum() >= B, "missing start_text_token"
|
||||
@@ -89,11 +83,13 @@ class T3(nn.Module):
|
||||
t3_cond: T3Cond,
|
||||
text_tokens: torch.LongTensor,
|
||||
speech_tokens: torch.LongTensor,
|
||||
cfg_weight: float = 0.0,
|
||||
):
|
||||
# prepare input embeddings (skip backbone tranformer embeddings)
|
||||
cond_emb = self.prepare_conditioning(t3_cond) # (B, len_cond, dim)
|
||||
text_emb = self.text_emb(text_tokens) # (B, len_text, dim)
|
||||
text_emb[1].zero_() # CFG uncond
|
||||
if cfg_weight > 0.0:
|
||||
text_emb[1].zero_() # CFG uncond
|
||||
|
||||
speech_emb = self.speech_emb(speech_tokens) # (B, len_speech, dim)
|
||||
if self.hp.input_pos_emb == "learned":
|
||||
@@ -217,14 +213,16 @@ class T3(nn.Module):
|
||||
|
||||
# HF generate args
|
||||
num_return_sequences=1,
|
||||
max_new_tokens=None,
|
||||
max_new_tokens,
|
||||
stop_on_eos=True,
|
||||
do_sample=True,
|
||||
temperature=0.8,
|
||||
top_p=0.8,
|
||||
min_p=0.05,
|
||||
top_p=1.0,
|
||||
length_penalty=1.0,
|
||||
repetition_penalty=2.0,
|
||||
repetition_penalty=1.2,
|
||||
cfg_weight=0,
|
||||
pbar=None
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
@@ -235,6 +233,9 @@ class T3(nn.Module):
|
||||
_ensure_BOT_EOT(text_tokens, self.hp)
|
||||
text_tokens = torch.atleast_2d(text_tokens).to(dtype=torch.long, device=self.device)
|
||||
|
||||
if max_new_tokens is None:
|
||||
max_new_tokens = self.hp.max_speech_tokens
|
||||
|
||||
# Default initial speech to a single start-of-speech token
|
||||
if initial_speech_tokens is None:
|
||||
initial_speech_tokens = self.hp.start_speech_token * torch.ones_like(text_tokens[:, :1])
|
||||
@@ -244,6 +245,7 @@ class T3(nn.Module):
|
||||
t3_cond=t3_cond,
|
||||
text_tokens=text_tokens,
|
||||
speech_tokens=initial_speech_tokens,
|
||||
cfg_weight=cfg_weight,
|
||||
)
|
||||
|
||||
# In order to use the standard HF generate method, we need to extend some methods to inject our custom logic
|
||||
@@ -254,19 +256,12 @@ class T3(nn.Module):
|
||||
# TODO? synchronize the expensive compile function
|
||||
# with self.compile_lock:
|
||||
if not self.compiled:
|
||||
alignment_stream_analyzer = AlignmentStreamAnalyzer(
|
||||
self.tfmr,
|
||||
None,
|
||||
text_tokens_slice=(len_cond, len_cond + text_tokens.size(-1)),
|
||||
alignment_layer_idx=9, # TODO: hparam or something?
|
||||
eos_idx=self.hp.stop_speech_token,
|
||||
)
|
||||
patched_model = T3HuggingfaceBackend(
|
||||
config=self.cfg,
|
||||
llama=self.tfmr,
|
||||
speech_enc=self.speech_emb,
|
||||
speech_head=self.speech_head,
|
||||
alignment_stream_analyzer=alignment_stream_analyzer,
|
||||
alignment_stream_analyzer=None,
|
||||
)
|
||||
self.patched_model = patched_model
|
||||
self.compiled = True
|
||||
@@ -281,7 +276,7 @@ class T3(nn.Module):
|
||||
# max_new_tokens=max_new_tokens or self.hp.max_speech_tokens,
|
||||
# num_return_sequences=num_return_sequences,
|
||||
# temperature=temperature,
|
||||
# top_p=top_p,
|
||||
# min_p=min_p,
|
||||
# length_penalty=length_penalty,
|
||||
# repetition_penalty=repetition_penalty,
|
||||
# do_sample=do_sample,
|
||||
@@ -297,18 +292,22 @@ class T3(nn.Module):
|
||||
# batch_size=2 for CFG
|
||||
bos_embed = torch.cat([bos_embed, bos_embed])
|
||||
|
||||
# Combine condition and BOS token for the initial input
|
||||
inputs_embeds = torch.cat([embeds, bos_embed], dim=1)
|
||||
# Combine condition and BOS token for the initial input if cfg_weight > 0
|
||||
if cfg_weight > 0:
|
||||
inputs_embeds = torch.cat([embeds, bos_embed], dim=1)
|
||||
else:
|
||||
inputs_embeds = embeds
|
||||
|
||||
# Track generated token ids; start with the BOS token.
|
||||
generated_ids = bos_token.clone()
|
||||
predicted = [] # To store the predicted tokens
|
||||
|
||||
# Instantiate the logits processors.
|
||||
min_p_warper = MinPLogitsWarper(min_p=min_p)
|
||||
top_p_warper = TopPLogitsWarper(top_p=top_p)
|
||||
repetition_penalty_processor = RepetitionPenaltyLogitsProcessor(penalty=repetition_penalty)
|
||||
repetition_penalty_processor = RepetitionPenaltyLogitsProcessor(penalty=float(repetition_penalty))
|
||||
|
||||
# ---- Initial Forward Pass (no kv_cache yet) ----
|
||||
# Initial Forward Pass (no kv_cache yet)
|
||||
output = self.patched_model(
|
||||
inputs_embeds=inputs_embeds,
|
||||
past_key_values=None,
|
||||
@@ -320,14 +319,17 @@ class T3(nn.Module):
|
||||
# Initialize kv_cache with the full context.
|
||||
past = output.past_key_values
|
||||
|
||||
# ---- Generation Loop using kv_cache ----
|
||||
for i in tqdm(range(max_new_tokens), desc="Sampling", dynamic_ncols=True):
|
||||
iterator = tqdm(range(max_new_tokens), desc="Chatterbox Sampling", dynamic_ncols=True)
|
||||
|
||||
for i in iterator:
|
||||
logits = output.logits[:, -1, :]
|
||||
|
||||
# CFG
|
||||
logits_cond = logits[0:1]
|
||||
logits_uncond = logits[1:2]
|
||||
logits = logits_cond + cfg_weight * (logits_cond - logits_uncond)
|
||||
if cfg_weight > 0.0:
|
||||
logits_cond = logits[0:1]
|
||||
logits_uncond = logits[1:2]
|
||||
logits = logits_cond + cfg_weight * (logits_cond - logits_uncond)
|
||||
|
||||
logits = logits.squeeze(1)
|
||||
|
||||
# Apply temperature scaling.
|
||||
@@ -336,6 +338,7 @@ class T3(nn.Module):
|
||||
|
||||
# Apply repetition penalty and top‑p filtering.
|
||||
logits = repetition_penalty_processor(generated_ids, logits)
|
||||
logits = min_p_warper(None, logits)
|
||||
logits = top_p_warper(None, logits)
|
||||
|
||||
# Convert logits to probabilities and sample the next token.
|
||||
@@ -345,8 +348,7 @@ class T3(nn.Module):
|
||||
predicted.append(next_token)
|
||||
generated_ids = torch.cat([generated_ids, next_token], dim=1)
|
||||
|
||||
# Check for EOS token.
|
||||
if next_token.view(-1) == self.hp.stop_speech_token:
|
||||
if stop_on_eos and next_token.view(-1) == self.hp.stop_speech_token:
|
||||
break
|
||||
|
||||
# Get embedding for the new token.
|
||||
@@ -354,7 +356,8 @@ class T3(nn.Module):
|
||||
next_token_embed = next_token_embed + self.speech_pos_emb.get_fixed_embedding(i + 1)
|
||||
|
||||
# For CFG
|
||||
next_token_embed = torch.cat([next_token_embed, next_token_embed])
|
||||
if cfg_weight > 0.0:
|
||||
next_token_embed = torch.cat([next_token_embed, next_token_embed])
|
||||
|
||||
# Forward pass with only the new token and the cached past.
|
||||
output = self.patched_model(
|
||||
@@ -367,6 +370,9 @@ class T3(nn.Module):
|
||||
# Update the kv_cache.
|
||||
past = output.past_key_values
|
||||
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
|
||||
# Concatenate all predicted tokens along the sequence dimension.
|
||||
predicted_tokens = torch.cat(predicted, dim=1) # shape: (B, num_tokens)
|
||||
return predicted_tokens
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
class AttrDict(dict):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super(AttrDict, self).__init__(*args, **kwargs)
|
||||
self.__dict__ = self
|
||||
@@ -259,7 +259,7 @@ class VoiceEncoder(nn.Module):
|
||||
"""
|
||||
if sample_rate != self.hp.sample_rate:
|
||||
wavs = [
|
||||
librosa.resample(wav, orig_sr=sample_rate, target_sr=self.hp.sample_rate, res_type="kaiser_fast")
|
||||
librosa.resample(wav, orig_sr=sample_rate, target_sr=self.hp.sample_rate, res_type="kaiser_best")
|
||||
for wav in wavs
|
||||
]
|
||||
|
||||
|
||||
+69
-11
@@ -1,3 +1,4 @@
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
@@ -6,14 +7,16 @@ import torch
|
||||
import perth
|
||||
import torch.nn.functional as F
|
||||
from huggingface_hub import hf_hub_download
|
||||
from safetensors.torch import load_file
|
||||
|
||||
from .models.t3 import T3
|
||||
from .models.s3tokenizer import S3_SR, drop_invalid_tokens
|
||||
from .models.s3tokenizer import S3_SR, S3_TOKEN_RATE, drop_invalid_tokens
|
||||
from .models.s3gen import S3GEN_SR, S3Gen
|
||||
from .models.tokenizers import EnTokenizer
|
||||
from .models.voice_encoder import VoiceEncoder
|
||||
from .models.t3.modules.cond_enc import T3Cond
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
REPO_ID = "ResembleAI/chatterbox"
|
||||
|
||||
@@ -96,9 +99,27 @@ class Conditionals:
|
||||
|
||||
@classmethod
|
||||
def load(cls, fpath, map_location="cpu"):
|
||||
if isinstance(map_location, str):
|
||||
map_location = torch.device(map_location)
|
||||
kwargs = torch.load(fpath, map_location=map_location, weights_only=True)
|
||||
return cls(T3Cond(**kwargs['t3']), kwargs['gen'])
|
||||
|
||||
# helper function for padding
|
||||
def _pad_wav_to_40ms_multiple(wav: torch.Tensor, sr: int) -> torch.Tensor:
|
||||
"""
|
||||
Pads a waveform to be a multiple of 40ms to prevent rounding errors between
|
||||
the mel spectrogram (20ms hop) and the speech tokenizer (40ms hop).
|
||||
"""
|
||||
S3_TOKEN_DURATION_S = 1 / S3_TOKEN_RATE # 0.04 seconds
|
||||
samples_per_token = int(sr * S3_TOKEN_DURATION_S)
|
||||
current_samples = wav.shape[-1]
|
||||
remainder = current_samples % samples_per_token
|
||||
if remainder != 0:
|
||||
padding_needed = samples_per_token - remainder
|
||||
padded_wav = F.pad(wav, (0, padding_needed))
|
||||
return padded_wav
|
||||
return wav
|
||||
|
||||
|
||||
class ChatterboxTTS:
|
||||
ENC_COND_LEN = 6 * S3_SR
|
||||
@@ -126,14 +147,20 @@ class ChatterboxTTS:
|
||||
def from_local(cls, ckpt_dir, device) -> 'ChatterboxTTS':
|
||||
ckpt_dir = Path(ckpt_dir)
|
||||
|
||||
# Always load to CPU first for non-CUDA devices to handle CUDA-saved models
|
||||
if device in ["cpu", "mps"]:
|
||||
map_location = torch.device('cpu')
|
||||
else:
|
||||
map_location = None
|
||||
|
||||
ve = VoiceEncoder()
|
||||
ve.load_state_dict(
|
||||
torch.load(ckpt_dir / "ve.pt")
|
||||
load_file(ckpt_dir / "ve.safetensors")
|
||||
)
|
||||
ve.to(device).eval()
|
||||
|
||||
t3 = T3()
|
||||
t3_state = torch.load(ckpt_dir / "t3_cfg.pt")
|
||||
t3_state = load_file(ckpt_dir / "t3_cfg.safetensors")
|
||||
if "model" in t3_state.keys():
|
||||
t3_state = t3_state["model"][0]
|
||||
t3.load_state_dict(t3_state)
|
||||
@@ -141,7 +168,7 @@ class ChatterboxTTS:
|
||||
|
||||
s3gen = S3Gen()
|
||||
s3gen.load_state_dict(
|
||||
torch.load(ckpt_dir / "s3gen.pt")
|
||||
load_file(ckpt_dir / "s3gen.safetensors"), strict=False
|
||||
)
|
||||
s3gen.to(device).eval()
|
||||
|
||||
@@ -151,13 +178,21 @@ class ChatterboxTTS:
|
||||
|
||||
conds = None
|
||||
if (builtin_voice := ckpt_dir / "conds.pt").exists():
|
||||
conds = Conditionals.load(builtin_voice).to(device)
|
||||
conds = Conditionals.load(builtin_voice, map_location=map_location).to(device)
|
||||
|
||||
return cls(t3, s3gen, ve, tokenizer, device, conds=conds)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, device) -> 'ChatterboxTTS':
|
||||
for fpath in ["ve.pt", "t3_cfg.pt", "s3gen.pt", "tokenizer.json", "conds.pt"]:
|
||||
# Check if MPS is available on macOS
|
||||
if device == "mps" and not torch.backends.mps.is_available():
|
||||
if not torch.backends.mps.is_built():
|
||||
logger.warning("MPS not available because the current PyTorch install was not built with MPS enabled.")
|
||||
else:
|
||||
logger.warning("MPS not available because the current MacOS version is not 12.3+ and/or you do not have an MPS-enabled device on this machine.")
|
||||
device = "cpu"
|
||||
|
||||
for fpath in ["ve.safetensors", "t3_cfg.safetensors", "s3gen.safetensors", "tokenizer.json", "conds.pt"]:
|
||||
local_path = hf_hub_download(repo_id=REPO_ID, filename=fpath)
|
||||
|
||||
return cls.from_local(Path(local_path).parent, device)
|
||||
@@ -165,11 +200,18 @@ class ChatterboxTTS:
|
||||
def prepare_conditionals(self, wav_fpath, exaggeration=0.5):
|
||||
## Load reference wav
|
||||
s3gen_ref_wav, _sr = librosa.load(wav_fpath, sr=S3GEN_SR)
|
||||
|
||||
# Convert to tensor and pad to a 40ms boundary
|
||||
s3gen_ref_wav = torch.from_numpy(s3gen_ref_wav).float().unsqueeze(0)
|
||||
s3gen_ref_wav = _pad_wav_to_40ms_multiple(s3gen_ref_wav, S3GEN_SR)
|
||||
|
||||
# Now convert back to numpy for librosa, or use torch audio for resampling
|
||||
s3gen_ref_wav_np = s3gen_ref_wav.squeeze(0).numpy()
|
||||
|
||||
ref_16k_wav = librosa.resample(s3gen_ref_wav, orig_sr=S3GEN_SR, target_sr=S3_SR)
|
||||
ref_16k_wav = librosa.resample(s3gen_ref_wav_np, orig_sr=S3GEN_SR, target_sr=S3_SR)
|
||||
|
||||
s3gen_ref_wav = s3gen_ref_wav[:self.DEC_COND_LEN]
|
||||
s3gen_ref_dict = self.s3gen.embed_ref(s3gen_ref_wav, S3GEN_SR, device=self.device)
|
||||
s3gen_ref_wav_np = s3gen_ref_wav_np[:self.DEC_COND_LEN]
|
||||
s3gen_ref_dict = self.s3gen.embed_ref(s3gen_ref_wav_np, S3GEN_SR, device=self.device)
|
||||
|
||||
# Speech cond prompt tokens
|
||||
if plen := self.t3.hp.speech_cond_prompt_len:
|
||||
@@ -191,10 +233,16 @@ class ChatterboxTTS:
|
||||
def generate(
|
||||
self,
|
||||
text,
|
||||
repetition_penalty=1.2,
|
||||
min_p=0.05,
|
||||
top_p=1.0,
|
||||
audio_prompt_path=None,
|
||||
exaggeration=0.5,
|
||||
cfg_weight=0.5,
|
||||
temperature=0.8,
|
||||
pbar=None,
|
||||
max_new_tokens=1000,
|
||||
flow_cfg_scale=0.7
|
||||
):
|
||||
if audio_prompt_path:
|
||||
self.prepare_conditionals(audio_prompt_path, exaggeration=exaggeration)
|
||||
@@ -213,7 +261,9 @@ class ChatterboxTTS:
|
||||
# Norm and tokenize text
|
||||
text = punc_norm(text)
|
||||
text_tokens = self.tokenizer.text_to_tokens(text).to(self.device)
|
||||
text_tokens = torch.cat([text_tokens, text_tokens], dim=0) # Need two seqs for CFG
|
||||
|
||||
if cfg_weight > 0.0:
|
||||
text_tokens = torch.cat([text_tokens, text_tokens], dim=0) # Need two seqs for CFG
|
||||
|
||||
sot = self.t3.hp.start_text_token
|
||||
eot = self.t3.hp.stop_text_token
|
||||
@@ -224,20 +274,28 @@ class ChatterboxTTS:
|
||||
speech_tokens = self.t3.inference(
|
||||
t3_cond=self.conds.t3,
|
||||
text_tokens=text_tokens,
|
||||
max_new_tokens=1000, # TODO: use the value in config
|
||||
max_new_tokens=max_new_tokens,
|
||||
temperature=temperature,
|
||||
cfg_weight=cfg_weight,
|
||||
repetition_penalty=repetition_penalty,
|
||||
min_p=min_p,
|
||||
top_p=top_p,
|
||||
pbar=pbar
|
||||
)
|
||||
# Extract only the conditional batch.
|
||||
speech_tokens = speech_tokens[0]
|
||||
|
||||
# TODO: output becomes 1D
|
||||
speech_tokens = drop_invalid_tokens(speech_tokens)
|
||||
|
||||
speech_tokens = speech_tokens[speech_tokens < 6561]
|
||||
|
||||
speech_tokens = speech_tokens.to(self.device)
|
||||
|
||||
wav, _ = self.s3gen.inference(
|
||||
speech_tokens=speech_tokens,
|
||||
ref_dict=self.conds.gen,
|
||||
flow_cfg_scale=flow_cfg_scale
|
||||
)
|
||||
wav = wav.squeeze(0).detach().cpu().numpy()
|
||||
watermarked_wav = self.watermarker.apply_watermark(wav, sample_rate=self.sr)
|
||||
|
||||
+53
-7
@@ -1,16 +1,34 @@
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
import librosa
|
||||
import torch
|
||||
import perth
|
||||
import torch.nn.functional as F
|
||||
from huggingface_hub import hf_hub_download
|
||||
from safetensors.torch import load_file
|
||||
|
||||
from .models.s3tokenizer import S3_SR
|
||||
from .models.s3tokenizer import S3_SR, S3_TOKEN_RATE
|
||||
from .models.s3gen import S3GEN_SR, S3Gen
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
REPO_ID = "ResembleAI/chatterbox"
|
||||
|
||||
def _pad_wav_to_40ms_multiple(wav: torch.Tensor, sr: int) -> torch.Tensor:
|
||||
"""
|
||||
Pads a waveform to be a multiple of 40ms to prevent rounding errors between
|
||||
the mel spectrogram (20ms hop) and the speech tokenizer (40ms hop).
|
||||
"""
|
||||
S3_TOKEN_DURATION_S = 1 / S3_TOKEN_RATE # 0.04 seconds
|
||||
samples_per_token = int(sr * S3_TOKEN_DURATION_S)
|
||||
current_samples = wav.shape[-1]
|
||||
remainder = current_samples % samples_per_token
|
||||
if remainder != 0:
|
||||
padding_needed = samples_per_token - remainder
|
||||
padded_wav = F.pad(wav, (0, padding_needed))
|
||||
return padded_wav
|
||||
return wav
|
||||
|
||||
class ChatterboxVC:
|
||||
ENC_COND_LEN = 6 * S3_SR
|
||||
@@ -37,14 +55,21 @@ class ChatterboxVC:
|
||||
@classmethod
|
||||
def from_local(cls, ckpt_dir, device) -> 'ChatterboxVC':
|
||||
ckpt_dir = Path(ckpt_dir)
|
||||
|
||||
# Always load to CPU first for non-CUDA devices to handle CUDA-saved models
|
||||
if device in ["cpu", "mps"]:
|
||||
map_location = torch.device('cpu')
|
||||
else:
|
||||
map_location = None
|
||||
|
||||
ref_dict = None
|
||||
if (builtin_voice := ckpt_dir / "conds.pt").exists():
|
||||
states = torch.load(builtin_voice)
|
||||
states = torch.load(builtin_voice, map_location=map_location)
|
||||
ref_dict = states['gen']
|
||||
|
||||
s3gen = S3Gen()
|
||||
s3gen.load_state_dict(
|
||||
torch.load(ckpt_dir / "s3gen.pt")
|
||||
load_file(ckpt_dir / "s3gen.safetensors"), strict=False
|
||||
)
|
||||
s3gen.to(device).eval()
|
||||
|
||||
@@ -52,7 +77,15 @@ class ChatterboxVC:
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, device) -> 'ChatterboxVC':
|
||||
for fpath in ["s3gen.pt", "conds.pt"]:
|
||||
# Check if MPS is available on macOS
|
||||
if device == "mps" and not torch.backends.mps.is_available():
|
||||
if not torch.backends.mps.is_built():
|
||||
logger.warning("MPS not available because the current PyTorch install was not built with MPS enabled.")
|
||||
else:
|
||||
logger.warning("MPS not available because the current MacOS version is not 12.3+ and/or you do not have an MPS-enabled device on this machine.")
|
||||
device = "cpu"
|
||||
|
||||
for fpath in ["s3gen.safetensors", "conds.pt"]:
|
||||
local_path = hf_hub_download(repo_id=REPO_ID, filename=fpath)
|
||||
|
||||
return cls.from_local(Path(local_path).parent, device)
|
||||
@@ -61,18 +94,27 @@ class ChatterboxVC:
|
||||
## Load reference wav
|
||||
s3gen_ref_wav, _sr = librosa.load(wav_fpath, sr=S3GEN_SR)
|
||||
|
||||
s3gen_ref_wav = s3gen_ref_wav[:self.DEC_COND_LEN]
|
||||
self.ref_dict = self.s3gen.embed_ref(s3gen_ref_wav, S3GEN_SR, device=self.device)
|
||||
# Convert to tensor and pad to a 40ms boundary
|
||||
s3gen_ref_wav = torch.from_numpy(s3gen_ref_wav).float().unsqueeze(0)
|
||||
s3gen_ref_wav = _pad_wav_to_40ms_multiple(s3gen_ref_wav, S3GEN_SR)
|
||||
s3gen_ref_wav_np = s3gen_ref_wav.squeeze(0).numpy()
|
||||
|
||||
s3gen_ref_wav_np = s3gen_ref_wav_np[:self.DEC_COND_LEN]
|
||||
self.ref_dict = self.s3gen.embed_ref(s3gen_ref_wav_np, S3GEN_SR, device=self.device)
|
||||
|
||||
def generate(
|
||||
self,
|
||||
audio,
|
||||
target_voice_path=None,
|
||||
n_timesteps=10,
|
||||
pbar=None,
|
||||
temperature=1.0,
|
||||
flow_cfg_scale=0.7
|
||||
):
|
||||
if target_voice_path:
|
||||
self.set_target_voice(target_voice_path)
|
||||
else:
|
||||
assert self.ref_dict is not None, "Please `prepare_conditionals` first or specify `target_voice_path`"
|
||||
assert self.ref_dict is not None, "Please `set_target_voice` first or specify `target_voice_path`"
|
||||
|
||||
with torch.inference_mode():
|
||||
audio_16, _ = librosa.load(audio, sr=S3_SR)
|
||||
@@ -82,6 +124,10 @@ class ChatterboxVC:
|
||||
wav, _ = self.s3gen.inference(
|
||||
speech_tokens=s3_tokens,
|
||||
ref_dict=self.ref_dict,
|
||||
n_timesteps=n_timesteps,
|
||||
pbar=pbar,
|
||||
temperature=temperature,
|
||||
flow_cfg_scale=flow_cfg_scale
|
||||
)
|
||||
wav = wav.squeeze(0).detach().cpu().numpy()
|
||||
watermarked_wav = self.watermarker.apply_watermark(wav, sample_rate=self.sr)
|
||||
|
||||
Reference in New Issue
Block a user