From 8b718b576e99faae77e93d79fd882be142f70a35 Mon Sep 17 00:00:00 2001 From: WildAi <2853742+wildminder@users.noreply.github.com> Date: Mon, 21 Jul 2025 17:12:37 +0300 Subject: [PATCH] ver 1.2.0 --- README.md | 2 +- __init__.py | 76 ++- modules/chatterbox_handler.py | 224 ++------- nodes.py | 438 +++++++++++------- src/chatterbox/__init__.py | 7 + src/chatterbox/models/__init__.py | 0 src/chatterbox/models/s3gen/configs.py | 10 + src/chatterbox/models/s3gen/flow.py | 152 ++++-- src/chatterbox/models/s3gen/flow_matching.py | 40 +- .../models/s3gen/matcha/transformer.py | 22 +- src/chatterbox/models/s3gen/s3gen.py | 74 ++- src/chatterbox/models/t3/llama_configs.py | 30 +- src/chatterbox/models/t3/modules/perceiver.py | 43 +- src/chatterbox/models/t3/t3.py | 72 +-- src/chatterbox/models/utils.py | 4 + .../models/voice_encoder/voice_encoder.py | 2 +- src/chatterbox/tts.py | 80 +++- src/chatterbox/vc.py | 60 ++- 18 files changed, 818 insertions(+), 518 deletions(-) create mode 100644 src/chatterbox/models/__init__.py create mode 100644 src/chatterbox/models/s3gen/configs.py create mode 100644 src/chatterbox/models/utils.py diff --git a/README.md b/README.md index 494a80a..5c5f6d5 100644 --- a/README.md +++ b/README.md @@ -5,7 +5,7 @@
- Chatterbox Nodes in ComfyUI + Chatterbox Nodes in ComfyUI

ComfyUI Chatterbox

diff --git a/__init__.py b/__init__.py index 6181aae..709d2c2 100644 --- a/__init__.py +++ b/__init__.py @@ -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) \ No newline at end of file diff --git a/modules/chatterbox_handler.py b/modules/chatterbox_handler.py index 26b6c13..092c569 100644 --- a/modules/chatterbox_handler.py +++ b/modules/chatterbox_handler.py @@ -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) - \ No newline at end of file + np.random.seed(actual_seed_for_numpy) \ No newline at end of file diff --git a/nodes.py b/nodes.py index 05833e8..99f3f52 100644 --- a/nodes.py +++ b/nodes.py @@ -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},) \ No newline at end of file + vc_wav_tensor_comfy = converted_wav_tensor.cpu().unsqueeze(0) + return ({"waveform": vc_wav_tensor_comfy, "sample_rate": vc_model.sr},) diff --git a/src/chatterbox/__init__.py b/src/chatterbox/__init__.py index 20b4391..6ae2ab1 100644 --- a/src/chatterbox/__init__.py +++ b/src/chatterbox/__init__.py @@ -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 diff --git a/src/chatterbox/models/__init__.py b/src/chatterbox/models/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/chatterbox/models/s3gen/configs.py b/src/chatterbox/models/s3gen/configs.py new file mode 100644 index 0000000..b09b2e5 --- /dev/null +++ b/src/chatterbox/models/s3gen/configs.py @@ -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" +}) diff --git a/src/chatterbox/models/s3gen/flow.py b/src/chatterbox/models/s3gen/flow.py index fee2ec9..41b90f5 100644 --- a/src/chatterbox/models/s3gen/flow.py +++ b/src/chatterbox/models/s3gen/flow.py @@ -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 diff --git a/src/chatterbox/models/s3gen/flow_matching.py b/src/chatterbox/models/s3gen/flow_matching.py index 74fc66f..458c17b 100644 --- a/src/chatterbox/models/s3gen/flow_matching.py +++ b/src/chatterbox/models/s3gen/flow_matching.py @@ -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 diff --git a/src/chatterbox/models/s3gen/matcha/transformer.py b/src/chatterbox/models/s3gen/matcha/transformer.py index 4f6762b..5517e1f 100644 --- a/src/chatterbox/models/s3gen/matcha/transformer.py +++ b/src/chatterbox/models/s3gen/matcha/transformer.py @@ -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)) diff --git a/src/chatterbox/models/s3gen/s3gen.py b/src/chatterbox/models/s3gen/s3gen.py index c611344..fb352e9 100644 --- a/src/chatterbox/models/s3gen/s3gen.py +++ b/src/chatterbox/models/s3gen/s3gen.py @@ -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. diff --git a/src/chatterbox/models/t3/llama_configs.py b/src/chatterbox/models/t3/llama_configs.py index 6247ead..c8484cd 100644 --- a/src/chatterbox/models/t3/llama_configs.py +++ b/src/chatterbox/models/t3/llama_configs.py @@ -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, -} \ No newline at end of file +} diff --git a/src/chatterbox/models/t3/modules/perceiver.py b/src/chatterbox/models/t3/modules/perceiver.py index 8bed6d6..a47e813 100644 --- a/src/chatterbox/models/t3/modules/perceiver.py +++ b/src/chatterbox/models/t3/modules/perceiver.py @@ -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): diff --git a/src/chatterbox/models/t3/t3.py b/src/chatterbox/models/t3/t3.py index 32fc6bc..70b21bd 100644 --- a/src/chatterbox/models/t3/t3.py +++ b/src/chatterbox/models/t3/t3.py @@ -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 diff --git a/src/chatterbox/models/utils.py b/src/chatterbox/models/utils.py new file mode 100644 index 0000000..a4abce5 --- /dev/null +++ b/src/chatterbox/models/utils.py @@ -0,0 +1,4 @@ +class AttrDict(dict): + def __init__(self, *args, **kwargs): + super(AttrDict, self).__init__(*args, **kwargs) + self.__dict__ = self diff --git a/src/chatterbox/models/voice_encoder/voice_encoder.py b/src/chatterbox/models/voice_encoder/voice_encoder.py index b0ed2df..7fc00ca 100644 --- a/src/chatterbox/models/voice_encoder/voice_encoder.py +++ b/src/chatterbox/models/voice_encoder/voice_encoder.py @@ -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 ] diff --git a/src/chatterbox/tts.py b/src/chatterbox/tts.py index 3371dae..2adbaca 100644 --- a/src/chatterbox/tts.py +++ b/src/chatterbox/tts.py @@ -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) diff --git a/src/chatterbox/vc.py b/src/chatterbox/vc.py index 629b686..ef05b10 100644 --- a/src/chatterbox/vc.py +++ b/src/chatterbox/vc.py @@ -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)