diff --git a/README.md b/README.md
index 494a80a..5c5f6d5 100644
--- a/README.md
+++ b/README.md
@@ -5,7 +5,7 @@
-
+
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)