ver 1.2.0

This commit is contained in:
WildAi
2025-07-21 17:12:37 +03:00
parent 613980ab3c
commit 8b718b576e
18 changed files with 818 additions and 518 deletions
+1 -1
View File
@@ -5,7 +5,7 @@
<br />
<div align="center">
<a href="https://github.com/wildminder/ComfyUI-Chatterbox">
<img src="./assets/preview.png" alt="Chatterbox Nodes in ComfyUI">
<img src="./workflow-examples/ChatterboxTTS-workflow.png" alt="Chatterbox Nodes in ComfyUI">
</a>
<h3 align="center">ComfyUI Chatterbox</h3>
+62 -14
View File
@@ -1,7 +1,59 @@
import os
import sys
import types
import logging
import folder_paths
from importlib import metadata as importlib_metadata
# Configure a logger for the entire custom node package
logger = logging.getLogger(__name__)
logger.setLevel(logging.WARNING)
# Add a handler if none exist to avoid duplicate logs
if not logger.hasHandlers():
# Use stdout for info, stderr for errors
handler = logging.StreamHandler(sys.stdout)
formatter = logging.Formatter(f"[%(name)s] %(message)s")
handler.setFormatter(formatter)
logger.addHandler(handler)
# Monkey-Patch for 'chatterbox-tts' version
# original_importlib_version_func = importlib_metadata.version
# def patched_version_lookup(package_name):
# if package_name == "chatterbox-tts":
# return "local-vendored"
# return original_importlib_version_func(package_name)
# importlib_metadata.version = patched_version_lookup
# easily check if the real 'perth' is being used.
try:
import perth
# A simple check to ensure it's not our mock from a previous run
if not hasattr(perth, '_is_mock'):
logger.info("Found and using 'resemble-perth' library for watermarking.")
except ImportError:
logger.warning("'resemble-perth' not found. Watermarking will be unavailable.")
class DummyPerthImplicitWatermarker:
def apply_watermark(self, wav, sample_rate):
logger.warning("Watermarking skipped: 'resemble-perth' is not installed.")
return wav
perth_mock = types.ModuleType('perth')
perth_mock.PerthImplicitWatermarker = DummyPerthImplicitWatermarker
# Flag to identify our mock module
perth_mock._is_mock = True
sys.modules['perth'] = perth_mock
current_dir = os.path.dirname(os.path.abspath(__file__))
src_dir = os.path.join(current_dir, "src")
if src_dir not in sys.path:
sys.path.insert(0, src_dir)
from .nodes import ChatterboxTTSNode, ChatterboxVCNode
from .modules.chatterbox_handler import CHATTERBOX_MODEL_SUBDIR, DEFAULT_MODEL_PACK_NAME
from .modules.chatterbox_handler import CHATTERBOX_MODEL_SUBDIR
NODE_CLASS_MAPPINGS = {
"ChatterboxTTS": ChatterboxTTSNode,
@@ -13,25 +65,21 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ChatterboxVC": "Chatterbox Voice Conversion 🗣️",
}
# WEB_DIRECTORY = "./js"
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
# Model path setup for ComfyUI
chatterbox_models_full_path = os.path.join(folder_paths.models_dir, CHATTERBOX_MODEL_SUBDIR)
if not os.path.exists(chatterbox_models_full_path):
try:
os.makedirs(chatterbox_models_full_path, exist_ok=True)
#print(f"ChatterboxTTS/VC Init: Created models directory at {chatterbox_models_full_path}")
except OSError as e:
print(f"ChatterboxTTS/VC Init: Error creating models directory {chatterbox_models_full_path}: {e}")
logger.error(f"Error creating models directory {chatterbox_models_full_path}: {e}")
if CHATTERBOX_MODEL_SUBDIR not in folder_paths.folder_names_and_paths:
folder_paths.folder_names_and_paths[CHATTERBOX_MODEL_SUBDIR] = (
[chatterbox_models_full_path],
folder_paths.supported_pt_extensions
)
#print(f"ChatterboxTTS/VC Init: Registered '{CHATTERBOX_MODEL_SUBDIR}' with ComfyUI model paths: {chatterbox_models_full_path}")
# Register the tts/chatterbox path with ComfyUI
tts_chatterbox_path = os.path.join(folder_paths.models_dir, "tts")
if "tts" not in folder_paths.folder_names_and_paths:
supported_exts = folder_paths.supported_pt_extensions.union({".safetensors"})
folder_paths.folder_names_and_paths["tts"] = ([tts_chatterbox_path], supported_exts)
else:
if chatterbox_models_full_path not in folder_paths.folder_names_and_paths[CHATTERBOX_MODEL_SUBDIR][0]:
folder_paths.folder_names_and_paths[CHATTERBOX_MODEL_SUBDIR][0].append(chatterbox_models_full_path)
#print(f"ChatterboxTTS/VC Init: Appended path for '{CHATTERBOX_MODEL_SUBDIR}': {chatterbox_models_full_path}")
if tts_chatterbox_path not in folder_paths.folder_names_and_paths["tts"][0]:
folder_paths.folder_names_and_paths["tts"][0].append(tts_chatterbox_path)
+50 -174
View File
@@ -1,222 +1,98 @@
import os
import sys
import logging
import torch
import random
import numpy as np
import folder_paths
import hashlib
from huggingface_hub import hf_hub_download
current_script_path = os.path.dirname(os.path.abspath(__file__))
chatterbox_src_path = os.path.abspath(os.path.join(current_script_path, '..', 'src'))
if chatterbox_src_path not in sys.path:
sys.path.insert(0, chatterbox_src_path)
from chatterbox.tts import ChatterboxTTS
from chatterbox.vc import ChatterboxVC # <--- Added VC import
from chatterbox.vc import ChatterboxVC
TTS_MODEL_CACHE = {}
VC_MODEL_CACHE = {}
CHATTERBOX_MODEL_SUBDIR = "chatterbox_tts"
logger = logging.getLogger(__name__)
CHATTERBOX_MODEL_SUBDIR = os.path.join("tts", "chatterbox")
CHATTERBOX_REPO_ID = "ResembleAI/chatterbox"
CHATTERBOX_FILES_TO_DOWNLOAD = ["ve.pt", "t3_cfg.pt", "s3gen.pt", "tokenizer.json", "conds.pt"]
CHATTERBOX_FILES_TO_DOWNLOAD = ["ve.safetensors", "t3_cfg.safetensors", "s3gen.safetensors", "tokenizer.json", "conds.pt"]
DEFAULT_MODEL_PACK_NAME = "resembleai_default_voice"
def get_chatterbox_model_pack_names():
"""Returns a list of available Chatterbox model pack names (subdirectories)."""
chatterbox_models_base_path = os.path.join(folder_paths.models_dir, CHATTERBOX_MODEL_SUBDIR)
if not os.path.isdir(chatterbox_models_base_path):
os.makedirs(chatterbox_models_base_path, exist_ok=True)
#print(f"ChatterboxTTS/VC: Created models directory at {chatterbox_models_base_path}")
return []
return [DEFAULT_MODEL_PACK_NAME]
packs = [d for d in os.listdir(chatterbox_models_base_path) if os.path.isdir(os.path.join(chatterbox_models_base_path, d))]
if not packs:
return []
return packs
# Ensure default is first if it exists
if DEFAULT_MODEL_PACK_NAME in packs:
packs.insert(0, packs.pop(packs.index(DEFAULT_MODEL_PACK_NAME)))
# Return default even if folder doesn't exist yet, to prompt the user to download it.
return packs if packs else [DEFAULT_MODEL_PACK_NAME]
def get_model_pack_path(model_pack_name):
"""Gets the full path to a specific model pack."""
if not model_pack_name:
# Added check for None or empty string to prevent errors
if not model_pack_name:
return None
return os.path.join(folder_paths.models_dir, CHATTERBOX_MODEL_SUBDIR, model_pack_name)
def _download_file_from_hf(repo_id, filename, local_dir, local_filename=None):
if local_filename is None:
local_filename = filename
destination = os.path.join(local_dir, local_filename)
def _download_file_from_hf(repo_id, filename, local_dir):
destination = os.path.join(local_dir, filename)
if not os.path.exists(destination):
print(f"ChatterboxTTS/VC: Downloading '{filename}' from '{repo_id}' to '{destination}'...")
logger.info(f"Downloading '{filename}' from '{repo_id}'...")
try:
hf_hub_download(
repo_id=repo_id,
filename=filename,
local_dir=local_dir,
)
if filename != local_filename and os.path.exists(os.path.join(local_dir, filename)) and not os.path.exists(destination):
os.rename(os.path.join(local_dir, filename), destination)
print(f"ChatterboxTTS/VC: Successfully downloaded '{local_filename}'.")
hf_hub_download(repo_id=repo_id, filename=filename, local_dir=local_dir, local_dir_use_symlinks=False, resume_download=True)
logger.info(f"Successfully downloaded '{filename}'.")
return True
except Exception as e:
print(f"ChatterboxTTS/VC: Failed to download '{filename}' from '{repo_id}'. Error: {e}")
if os.path.exists(destination + ".incomplete"):
os.remove(destination + ".incomplete")
logger.error(f"Failed to download '{filename}': {e}")
if os.path.exists(destination + ".incomplete"): os.remove(destination + ".incomplete")
return False
else:
return True
return True
def download_chatterbox_model_pack_if_missing(model_pack_name):
"""Downloads all necessary model files for a given pack name if they don't exist."""
ckpt_dir = get_model_pack_path(model_pack_name)
if not ckpt_dir:
print(f"ChatterboxTTS/VC: Invalid model pack name '{model_pack_name}', cannot download.")
logger.warning(f"Invalid model pack name '{model_pack_name}', cannot download.")
return False
os.makedirs(ckpt_dir, exist_ok=True)
all_files_successfully_managed = True
all_files_ok = all(_download_file_from_hf(CHATTERBOX_REPO_ID, f, ckpt_dir) for f in CHATTERBOX_FILES_TO_DOWNLOAD)
if not all_files_ok:
logger.error(f"Some files failed to download for model pack '{model_pack_name}'. Check logs.")
return all_files_ok
vc_required_files = ["s3gen.pt", "conds.pt"]
files_to_check = CHATTERBOX_FILES_TO_DOWNLOAD
for file_name in files_to_check:
if not _download_file_from_hf(CHATTERBOX_REPO_ID, file_name, ckpt_dir):
all_files_successfully_managed = False
if file_name in vc_required_files:
print(f"ChatterboxTTS/VC: Critical file '{file_name}' for VC failed to download for pack '{model_pack_name}'.")
# return False
if all_files_successfully_managed:
pass #print(f"ChatterboxTTS/VC: Verified/Downloaded all files for model pack '{model_pack_name}' into '{ckpt_dir}'.")
else:
print(f"ChatterboxTTS/VC: Some files failed to download for model pack '{model_pack_name}'. Check logs.")
return all_files_successfully_managed
def load_chatterbox_tts_model(model_pack_name, device_str="cuda"):
def load_chatterbox_models(model_pack_name, device):
"""Loads both TTS and VC models for a given pack onto a specified device."""
ckpt_dir = get_model_pack_path(model_pack_name)
if not ckpt_dir:
raise ValueError(f"ChatterboxTTS: Invalid model_pack_name: {model_pack_name}")
#print(f"ChatterboxTTS: Attempting to load TTS model pack '{model_pack_name}' from '{ckpt_dir}'.")
raise ValueError(f"Invalid model_pack_name: {model_pack_name}")
if not download_chatterbox_model_pack_if_missing(model_pack_name):
print(f"ChatterboxTTS: Warning - Not all TTS model files could be verified/downloaded for '{model_pack_name}'. Loading may fail.")
if not os.path.isdir(ckpt_dir):
raise FileNotFoundError(f"ChatterboxTTS: Model pack directory '{model_pack_name}' not found and could not be created at '{ckpt_dir}'.")
#print(f"ChatterboxTTS: Loading ChatterboxTTS model from local directory: {ckpt_dir} onto device: {device_str}")
try:
model = ChatterboxTTS.from_local(ckpt_dir, device=device_str)
except Exception as e:
print(f"ChatterboxTTS: Error during ChatterboxTTS.from_local('{ckpt_dir}', device='{device_str}'): {e}")
print(f"ChatterboxTTS: Please ensure all required model files ({', '.join(CHATTERBOX_FILES_TO_DOWNLOAD)}) are present in {ckpt_dir} or can be downloaded from {CHATTERBOX_REPO_ID}.")
raise
return model
def get_cached_chatterbox_tts_model(model_pack_name, device_str="cuda"):
"""Loads and caches the ChatterboxTTS model."""
if not model_pack_name:
available_packs = get_chatterbox_model_pack_names()
model_pack_name = available_packs[0] if available_packs else DEFAULT_MODEL_PACK_NAME
print(f"ChatterboxTTS: No model pack specified for TTS, using '{model_pack_name}'.")
cache_key = (model_pack_name, device_str, "tts")
current_model = TTS_MODEL_CACHE.get(cache_key)
model_device_correct = False
if current_model is not None and hasattr(current_model, 'device'):
try:
if str(current_model.device) == device_str:
model_device_correct = True
except Exception as e:
print(f"ChatterboxTTS: Error checking cached TTS model device: {e}. Will reload.")
if current_model is not None and model_device_correct:
return current_model
else:
if current_model is not None and not model_device_correct:
print(f"ChatterboxTTS: Device mismatch for cached TTS model '{model_pack_name}'. Reloading.")
else:
print(f"ChatterboxTTS: TTS Model for '{model_pack_name}' on '{device_str}' not in cache. Loading...")
logger.warning(f"Not all model files could be verified for '{model_pack_name}'. Loading may fail.")
TTS_MODEL_CACHE[cache_key] = load_chatterbox_tts_model(model_pack_name, device_str)
return TTS_MODEL_CACHE[cache_key]
def load_chatterbox_vc_model(model_pack_name, device_str="cuda"):
"""Loads the ChatterboxVC model from a specified pack."""
ckpt_dir = get_model_pack_path(model_pack_name)
if not ckpt_dir:
raise ValueError(f"ChatterboxVC: Invalid model_pack_name: {model_pack_name}")
#print(f"ChatterboxVC: Attempting to load VC model pack '{model_pack_name}' from '{ckpt_dir}'.")
if not download_chatterbox_model_pack_if_missing(model_pack_name):
vc_min_files = ["s3gen.pt", "conds.pt"]
for f_name in vc_min_files:
if not os.path.exists(os.path.join(ckpt_dir, f_name)):
print(f"ChatterboxVC: Critical file '{f_name}' for VC is missing from pack '{model_pack_name}' after download attempt. Loading will likely fail.")
break
else:
print(f"ChatterboxVC: Warning - Not all model files could be verified/downloaded for '{model_pack_name}', but attempting to load VC with available files.")
if not os.path.isdir(ckpt_dir):
raise FileNotFoundError(f"ChatterboxVC: Model pack directory '{model_pack_name}' not found and could not be created at '{ckpt_dir}'.")
# print(f"ChatterboxVC: Loading ChatterboxVC model from local directory: {ckpt_dir} onto device: {device_str}")
raise FileNotFoundError(f"Model pack directory '{model_pack_name}' not found at '{ckpt_dir}'.")
try:
model = ChatterboxVC.from_local(ckpt_dir, device=device_str)
logger.info(f"Loading Chatterbox TTS model from {ckpt_dir} onto {device}")
tts_model = ChatterboxTTS.from_local(ckpt_dir, device=device)
except Exception as e:
print(f"ChatterboxVC: Error during ChatterboxVC.from_local('{ckpt_dir}', device='{device_str}'): {e}")
print(f"ChatterboxVC: Please ensure at least 's3gen.pt' and optionally 'conds.pt' (for default voice) are present in {ckpt_dir} or can be downloaded from {CHATTERBOX_REPO_ID}.")
logger.error(f"Error loading ChatterboxTTS from '{ckpt_dir}': {e}", exc_info=True)
raise
try:
logger.info(f"Loading Chatterbox VC model from {ckpt_dir} onto {device}")
vc_model = ChatterboxVC.from_local(ckpt_dir, device=device)
except Exception as e:
logger.error(f"Error loading ChatterboxVC from '{ckpt_dir}': {e}", exc_info=True)
raise
return model
def get_cached_chatterbox_vc_model(model_pack_name, device_str="cuda"):
"""Loads and caches the ChatterboxVC model."""
if not model_pack_name:
available_packs = get_chatterbox_model_pack_names()
model_pack_name = available_packs[0] if available_packs else DEFAULT_MODEL_PACK_NAME
print(f"ChatterboxVC: No model pack specified for VC, using '{model_pack_name}'.")
cache_key = (model_pack_name, device_str, "vc")
current_model = VC_MODEL_CACHE.get(cache_key)
model_device_correct = False
if current_model is not None and hasattr(current_model, 'device'):
try:
if str(current_model.device) == device_str:
model_device_correct = True
except Exception as e:
print(f"ChatterboxVC: Error checking cached VC model device: {e}. Will reload.")
if current_model is not None and model_device_correct:
return current_model
else:
if current_model is not None and not model_device_correct:
print(f"ChatterboxVC: Device mismatch for cached VC model '{model_pack_name}'. Reloading.")
else:
print(f"ChatterboxVC: VC Model for '{model_pack_name}' on '{device_str}' not in cache. Loading...")
VC_MODEL_CACHE[cache_key] = load_chatterbox_vc_model(model_pack_name, device_str)
return VC_MODEL_CACHE[cache_key]
return tts_model, vc_model
def set_chatterbox_seed(seed: int):
"""Sets the seed for Chatterbox TTS/VC."""
MAX_NUMPY_SEED = 2**32 - 1
if seed == 0:
actual_seed_for_torch_random = random.randint(1, 0xffffffffffffffff)
actual_seed_for_numpy = random.randint(1, MAX_NUMPY_SEED)
else:
actual_seed_for_torch_random = seed
# fast fix range for seed, sorry
actual_seed_for_numpy = seed % MAX_NUMPY_SEED
actual_seed_for_torch_random = random.randint(1, 0xffffffffffffffff) if seed == 0 else seed
actual_seed_for_numpy = random.randint(1, MAX_NUMPY_SEED) if seed == 0 else (seed % MAX_NUMPY_SEED)
torch.manual_seed(actual_seed_for_torch_random)
if torch.cuda.is_available():
torch.cuda.manual_seed(actual_seed_for_torch_random)
torch.cuda.manual_seed_all(actual_seed_for_torch_random)
if torch.cuda.is_available(): torch.cuda.manual_seed_all(actual_seed_for_torch_random)
random.seed(actual_seed_for_torch_random)
np.random.seed(actual_seed_for_numpy)
np.random.seed(actual_seed_for_numpy)
+266 -172
View File
@@ -1,231 +1,325 @@
import os
import torch
import folder_paths
import tempfile
import soundfile as sf
import numpy as np
import logging
import perth
import comfy.model_management as mm
import comfy.model_patcher
from comfy.utils import ProgressBar
from .modules.chatterbox_handler import (
get_chatterbox_model_pack_names,
get_cached_chatterbox_tts_model,
get_cached_chatterbox_vc_model,
load_chatterbox_models,
set_chatterbox_seed,
CHATTERBOX_MODEL_SUBDIR,
DEFAULT_MODEL_PACK_NAME
)
logger = logging.getLogger(__name__)
CHATTERBOX_PATCHER_CACHE = {}
class ChatterboxModelWrapper(torch.nn.Module):
"""
A simple torch.nn.Module wrapper for the Chatterbox models.
This allows ComfyUI's model management to treat our custom models like any other
torch module, enabling device placement (.to()) and other standard operations.
"""
def __init__(self, model_pack_name):
super().__init__()
self.model_pack_name = model_pack_name
self.tts_model = None
self.vc_model = None
def load_model(self, device):
self.tts_model, self.vc_model = load_chatterbox_models(self.model_pack_name, device)
class ChatterboxPatcher(comfy.model_patcher.ModelPatcher):
"""
Custom ModelPatcher for Chatterbox. This class hooks into ComfyUI's
model management system (loading, offloading) to handle our non-standard models.
"""
def __init__(self, model, *args, **kwargs):
super().__init__(model, *args, **kwargs)
def patch_model(self, device_to=None, *args, **kwargs):
"""
This method is called by ComfyUI's model manager when it's time to load
the model onto the target device (usually the GPU). Our responsibility here
is to ensure the model weights are loaded from disk if they haven't been already.
"""
target_device = self.load_device
# The core loading logic: If the model isn't in memory, load it from disk.
if self.model.tts_model is None:
logger.info(f"Loading Chatterbox models for '{self.model.model_pack_name}' to {target_device}...")
self.model.load_model(target_device)
self.model.model_loaded_weight_memory = self.size
else:
logger.info(f"Chatterbox models for '{self.model.model_pack_name}' already in memory.")
return super().patch_model(device_to=target_device, *args, **kwargs)
def unpatch_model(self, device_to=None, unpatch_weights=True, *args, **kwargs):
"""
This method is called by ComfyUI's model manager to offload the model
(usually to the CPU) and free up VRAM.
"""
if unpatch_weights:
logger.info(f"Offloading Chatterbox models for '{self.model.model_pack_name}' to {device_to}...")
self.model.tts_model = None
self.vc_model = None
# Reset memory footprint
self.model.model_loaded_weight_memory = 0
mm.soft_empty_cache()
return super().unpatch_model(device_to, unpatch_weights, *args, **kwargs)
class ChatterboxTTSNode:
@classmethod
def INPUT_TYPES(cls):
available_model_packs = get_chatterbox_model_pack_names()
displayed_packs = [DEFAULT_MODEL_PACK_NAME] + [p for p in available_model_packs if p != DEFAULT_MODEL_PACK_NAME]
if not displayed_packs:
displayed_packs = [DEFAULT_MODEL_PACK_NAME]
return {"required": {
"model_pack_name": (get_chatterbox_model_pack_names(), {
"default": DEFAULT_MODEL_PACK_NAME,
"tooltip": "Select the Chatterbox voice model pack to use for generation."
}),
"text": ("STRING", {
"multiline": True,
"default": "Hello, this is a test of Chatterbox TTS in ComfyUI.",
"tooltip": "Text to be synthesized into speech."
}),
"max_new_tokens": ("INT", {
"default": 1000, "min": 16, "max": 4000, "step": 8,
"tooltip": "Maximum number of audio tokens to generate. 25 tokens ≈ 1 second. The hard limit is 4096 tokens (≈ 163 seconds)."
}),
"flow_cfg_scale": ("FLOAT", {"default": 0.7, "min": 0.0, "max": 2.0, "step": 0.05, "tooltip": "CFG scale for the mel spectrogram decoder (flow matching). Higher values increase adherence to content and timbre but may reduce naturalness."}),
"exaggeration": ("FLOAT", {
"default": 0.5, "min": 0.25, "max": 2.0, "step": 0.05,
"tooltip": "Controls the expressiveness and emotional intensity. Higher values lead to more exaggerated prosody."
}),
"temperature": ("FLOAT", {
"default": 0.8, "min": 0.05, "max": 5.0, "step": 0.05,
"tooltip": "Controls the randomness of the sampling process. Higher values produce more diverse speech, while lower values are more deterministic."
}),
"cfg_weight": ("FLOAT", {
"default": 0.5, "min": 0.2, "max": 1.0, "step": 0.05,
"tooltip": "Classifier-Free Guidance (CFG) weight. Controls how strongly the model adheres to the text prompt. Higher values may reduce naturalness."
}),
"repetition_penalty": ("FLOAT", {
"default": 1.2, "min": 1.0, "max": 2.0, "step": 0.1,
"tooltip": "Penalizes repeated tokens to discourage monotonous or repetitive speech. A value of 1.0 means no penalty."
}),
"min_p": ("FLOAT", {
"default": 0.05, "min": 0.0, "max": 1.0, "step": 0.01,
"tooltip": "Sets a minimum probability threshold for nucleus sampling (Min-P). Filters out tokens with very low probability."
}),
"top_p": ("FLOAT", {
"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01,
"tooltip": "Nucleus sampling (Top-P) parameter. The model samples from the smallest set of tokens whose cumulative probability exceeds this value."
}),
"seed": ("INT", {
"default": 0, "min": 0, "max": 0xffffffffffffffff, "control_after_generate": True,
"tooltip": "Seed for random number generation. A value of 0 will use a random seed."
}),
"use_watermark": ("BOOLEAN", {
"default": False,
"tooltip": "Enable or disable the audio watermark. Requires 'resemble-perth' to be installed."
}),
}, "optional": {"audio_prompt": ("AUDIO",),}}
return {
"required": {
"model_pack_name": (displayed_packs, {"default": DEFAULT_MODEL_PACK_NAME if DEFAULT_MODEL_PACK_NAME in displayed_packs else (displayed_packs[0] if displayed_packs else DEFAULT_MODEL_PACK_NAME)}),
"text": ("STRING", {"multiline": True, "default": "Hello, this is a test of Chatterbox TTS in ComfyUI."}),
"exaggeration": ("FLOAT", {"default": 0.5, "min": 0.25, "max": 2.0, "step": 0.05}),
"temperature": ("FLOAT", {"default": 0.8, "min": 0.05, "max": 5.0, "step": 0.05}),
"cfg_weight": ("FLOAT", {"default": 0.5, "min": 0.2, "max": 1.0, "step": 0.05}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "control_after_generate": True}),
"device": (["cuda", "cpu"], {"default": "cuda" if torch.cuda.is_available() else "cpu"}),
},
"optional": {
"audio_prompt": ("AUDIO",),
}
}
RETURN_TYPES = ("AUDIO",); RETURN_NAMES = ("audio",); FUNCTION = "synthesize"; CATEGORY = "audio/generation"; OUTPUT_NODE = True
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("audio",)
FUNCTION = "synthesize"
CATEGORY = "audio/generation"
OUTPUT_NODE = True
def synthesize(self, model_pack_name, text, exaggeration, temperature, cfg_weight, seed, device, audio_prompt=None):
def synthesize(self, model_pack_name, text, max_new_tokens, flow_cfg_scale, exaggeration, temperature, cfg_weight, repetition_penalty, min_p, top_p, seed, use_watermark, audio_prompt=None):
if not text.strip():
#print("Chatterbox TTS: Empty text provided, returning silent audio.")
dummy_sr = 24000
silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
logger.info("Empty text provided, returning silent audio.")
dummy_sr = 24000; silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
return ({"waveform": silent_waveform.unsqueeze(0), "sample_rate": dummy_sr},)
try:
chatterbox_model = get_cached_chatterbox_tts_model(model_pack_name, device_str=device)
except Exception as e:
print(f"ChatterboxTTS: Error loading/downloading TTS model pack '{model_pack_name}': {e}")
dummy_sr = 24000
silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
cache_key = model_pack_name
if cache_key not in CHATTERBOX_PATCHER_CACHE:
load_device = mm.get_torch_device()
logger.info(f"Creating Chatterbox ModelPatcher for {model_pack_name} on device {load_device}")
model_wrapper = ChatterboxModelWrapper(model_pack_name)
patcher = ChatterboxPatcher(
model=model_wrapper,
load_device=load_device,
offload_device=mm.unet_offload_device(),
size=int(1.5 * 1024**3)
)
CHATTERBOX_PATCHER_CACHE[cache_key] = patcher
patcher = CHATTERBOX_PATCHER_CACHE[cache_key]
mm.load_model_gpu(patcher)
tts_model = patcher.model.tts_model
if tts_model is None:
logger.error("TTS model failed to load. Please check logs for download or loading errors.")
dummy_sr = 24000; silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
return ({"waveform": silent_waveform.unsqueeze(0), "sample_rate": dummy_sr},)
set_chatterbox_seed(seed)
is_perth_installed = not getattr(perth, '_is_mock', False)
if use_watermark and not is_perth_installed:
logger.warning("Watermarking is enabled, but 'resemble-perth' is not installed. Output will not be watermarked.")
original_watermarker = tts_model.watermarker
if not use_watermark:
class TmpDummyWatermarker:
def apply_watermark(self, wav, sample_rate): return wav
tts_model.watermarker = TmpDummyWatermarker()
if is_perth_installed: logger.info("Watermarking disabled by user.")
wav_tensor_chatterbox = None; audio_prompt_path_temp = None
audio_prompt_path_temp = None
if audio_prompt is not None and \
audio_prompt.get("waveform") is not None and \
audio_prompt["waveform"].numel() > 0:
waveform_in = audio_prompt["waveform"]
sample_rate_in = audio_prompt["sample_rate"]
waveform_cpu = waveform_in.cpu()
if waveform_cpu.shape[0] > 1:
print(f"ChatterboxTTS: Audio prompt has batch size {waveform_cpu.shape[0]}, using first item.")
current_waveform = waveform_cpu[0]
if current_waveform.shape[0] > 1:
current_waveform = torch.mean(current_waveform, dim=0)
else:
current_waveform = current_waveform.squeeze(0)
processed_audio_prompt = current_waveform.numpy().astype(np.float32)
try:
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp_wav:
audio_prompt_path_temp = tmp_wav.name
sf.write(audio_prompt_path_temp, processed_audio_prompt, sample_rate_in)
#print(f"ChatterboxTTS: Using audio prompt from temp file: {audio_prompt_path_temp}")
except Exception as e:
print(f"ChatterboxTTS: Error writing temp audio prompt file: {e}")
audio_prompt_path_temp = None
pbar = ProgressBar(max_new_tokens)
try:
wav_tensor_chatterbox = chatterbox_model.generate(
if audio_prompt and audio_prompt.get("waveform") is not None and audio_prompt["waveform"].numel() > 0:
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp_wav:
audio_prompt_path_temp = tmp_wav.name
waveform_in = audio_prompt["waveform"]; sample_rate_in = audio_prompt["sample_rate"]
waveform_cpu = waveform_in.cpu()[0]
current_waveform = torch.mean(waveform_cpu, dim=0) if waveform_cpu.shape[0] > 1 else waveform_cpu.squeeze(0)
sf.write(audio_prompt_path_temp, current_waveform.numpy().astype(np.float32), sample_rate_in)
wav_tensor_chatterbox = tts_model.generate(
text,
audio_prompt_path=audio_prompt_path_temp,
exaggeration=exaggeration,
temperature=temperature,
cfg_weight=cfg_weight
)
cfg_weight=cfg_weight,
repetition_penalty=repetition_penalty,
min_p=min_p,
top_p=top_p,
pbar=pbar,
max_new_tokens=max_new_tokens,
flow_cfg_scale=flow_cfg_scale
)
except Exception as e:
print(f"ChatterboxTTS: Error during TTS generation: {e}")
dummy_sr = 24000
silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
logger.error(f"Error during TTS generation: {e}", exc_info=True)
dummy_sr = 24000; silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
return ({"waveform": silent_waveform.unsqueeze(0), "sample_rate": dummy_sr},)
finally:
tts_model.watermarker = original_watermarker
if audio_prompt_path_temp and os.path.exists(audio_prompt_path_temp):
try:
os.remove(audio_prompt_path_temp)
except Exception as e:
print(f"ChatterboxTTS: Error removing temp audio prompt file {audio_prompt_path_temp}: {e}")
wav_tensor_comfy = wav_tensor_chatterbox.cpu().unsqueeze(0)
return ({"waveform": wav_tensor_comfy, "sample_rate": chatterbox_model.sr},)
try: os.remove(audio_prompt_path_temp)
except Exception as e: logger.error(f"Error removing temp audio prompt file: {e}")
wav_tensor_comfy = wav_tensor_chatterbox.cpu().unsqueeze(0)
return ({"waveform": wav_tensor_comfy, "sample_rate": tts_model.sr},)
class ChatterboxVCNode:
@classmethod
def INPUT_TYPES(cls):
available_model_packs = get_chatterbox_model_pack_names()
displayed_packs = [DEFAULT_MODEL_PACK_NAME] + [p for p in available_model_packs if p != DEFAULT_MODEL_PACK_NAME]
if not displayed_packs:
displayed_packs = [DEFAULT_MODEL_PACK_NAME]
return {"required": {
"model_pack_name": (get_chatterbox_model_pack_names(), {
"default": DEFAULT_MODEL_PACK_NAME,
"tooltip": "Select the Chatterbox voice model pack to use for conversion."
}),
"source_audio": ("AUDIO", {
"tooltip": "The audio containing the speech content to be converted."
}),
"n_timesteps": ("INT", {
"default": 10, "min": 2, "max": 50, "step": 1,
"tooltip": "Number of diffusion steps for the flow matching process. Higher values may improve quality at the cost of speed."
}),
"temperature": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.05,
"tooltip": "Controls the randomness of the initial noise. 1.0 is standard. Lower values are more deterministic."}),
"flow_cfg_scale": ("FLOAT", {"default": 0.7, "min": 0.0, "max": 2.0, "step": 0.05,
"tooltip": "CFG scale for the mel spectrogram decoder. Higher values increase adherence to the target voice but may reduce naturalness."}),
"use_watermark": ("BOOLEAN", {
"default": False,
"tooltip": "Enable or disable the audio watermark. Requires 'resemble-perth' to be installed."
}),
},
"optional": {"target_voice_audio": ("AUDIO", {
"tooltip": "The audio file containing the target voice timbre. If not provided, the default voice from the model pack will be used."
}), }}
return {
"required": {
"model_pack_name": (displayed_packs, {"default": DEFAULT_MODEL_PACK_NAME if DEFAULT_MODEL_PACK_NAME in displayed_packs else (displayed_packs[0] if displayed_packs else DEFAULT_MODEL_PACK_NAME)}),
"source_audio": ("AUDIO",),
"device": (["cuda", "cpu"], {"default": "cuda" if torch.cuda.is_available() else "cpu"}),
},
"optional": {
"target_voice_audio": ("AUDIO",), # Optional: if not provided, uses default voice from conds.pt
}
}
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("converted_audio",)
FUNCTION = "convert_voice"
CATEGORY = "audio/generation"
OUTPUT_NODE = True
RETURN_TYPES = ("AUDIO",); RETURN_NAMES = ("converted_audio",); FUNCTION = "convert_voice"; CATEGORY = "audio/generation"; OUTPUT_NODE = True
def _save_audio_to_temp_file(self, audio_data, prefix=""):
"""Helper to save ComfyUI AUDIO dict to a temporary WAV file."""
if audio_data is None or \
audio_data.get("waveform") is None or \
audio_data["waveform"].numel() == 0:
return None
if audio_data is None or audio_data.get("waveform") is None or audio_data["waveform"].numel() == 0: return None
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp_wav:
try:
waveform_in = audio_data["waveform"]; sample_rate_in = audio_data["sample_rate"]
waveform_cpu = waveform_in.cpu()[0]
current_waveform = torch.mean(waveform_cpu, dim=0) if waveform_cpu.shape[0] > 1 else waveform_cpu.squeeze(0)
sf.write(tmp_wav.name, current_waveform.numpy().astype(np.float32), sample_rate_in)
return tmp_wav.name
except Exception as e:
logger.error(f"Error writing temp {prefix}audio file: {e}", exc_info=True)
return None
waveform_in = audio_data["waveform"]
sample_rate_in = audio_data["sample_rate"]
waveform_cpu = waveform_in.cpu()
if waveform_cpu.shape[0] > 1:
print(f"ChatterboxVC: {prefix}Audio has batch size {waveform_cpu.shape[0]}, using first item.")
current_waveform = waveform_cpu[0]
if current_waveform.shape[0] > 1: # If C > 1 (stereo or more)
print(f"ChatterboxVC: {prefix}Audio has {current_waveform.shape[0]} channels, converting to mono by averaging.")
current_waveform = torch.mean(current_waveform, dim=0)
else: # If C == 1
current_waveform = current_waveform.squeeze(0)
processed_audio = current_waveform.numpy().astype(np.float32)
temp_file_path = None
try:
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp_wav:
temp_file_path = tmp_wav.name
sf.write(temp_file_path, processed_audio, sample_rate_in)
#print(f"ChatterboxVC: Saved {prefix}audio to temp file: {temp_file_path}")
return temp_file_path
except Exception as e:
print(f"ChatterboxVC: Error writing temp {prefix}audio file: {e}")
if temp_file_path and os.path.exists(temp_file_path):
os.remove(temp_file_path)
return None
def convert_voice(self, model_pack_name, source_audio, device, target_voice_audio=None):
def convert_voice(self, model_pack_name, source_audio, n_timesteps, temperature, flow_cfg_scale, use_watermark, target_voice_audio=None):
if source_audio is None or source_audio.get("waveform") is None or source_audio["waveform"].numel() == 0:
print("ChatterboxVC: No source audio provided, returning silent audio.")
dummy_sr = 24000
silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
logger.warning("No source audio provided, returning silent audio.")
dummy_sr = 24000; silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
return ({"waveform": silent_waveform.unsqueeze(0), "sample_rate": dummy_sr},)
try:
vc_model = get_cached_chatterbox_vc_model(model_pack_name, device_str=device)
except Exception as e:
print(f"ChatterboxVC: Error loading/downloading VC model pack '{model_pack_name}': {e}")
dummy_sr = 24000
silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
cache_key = model_pack_name
if cache_key not in CHATTERBOX_PATCHER_CACHE:
load_device = mm.get_torch_device()
logger.info(f"Creating Chatterbox ModelPatcher for {model_pack_name} on device {load_device}")
model_wrapper = ChatterboxModelWrapper(model_pack_name)
patcher = ChatterboxPatcher(
model=model_wrapper,
load_device=load_device,
offload_device=mm.unet_offload_device(),
size=int(1.5 * 1024**3)
)
CHATTERBOX_PATCHER_CACHE[cache_key] = patcher
patcher = CHATTERBOX_PATCHER_CACHE[cache_key]
mm.load_model_gpu(patcher)
vc_model = patcher.model.vc_model
if vc_model is None:
logger.error("VC model failed to load. Please check logs for download or loading errors.")
dummy_sr = 24000; silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
return ({"waveform": silent_waveform.unsqueeze(0), "sample_rate": dummy_sr},)
source_audio_path_temp = None
target_voice_path_temp = None
is_perth_installed = not getattr(perth, '_is_mock', False)
if use_watermark and not is_perth_installed:
logger.warning("Watermarking is enabled, but 'resemble-perth' is not installed. Output will not be watermarked.")
original_watermarker = vc_model.watermarker
if not use_watermark:
class TmpDummyWatermarker:
def apply_watermark(self, wav, sample_rate): return wav
vc_model.watermarker = TmpDummyWatermarker()
if is_perth_installed: logger.info("Watermarking disabled by user.")
source_audio_path_temp = None; target_voice_path_temp = None
pbar = ProgressBar(n_timesteps)
try:
source_audio_path_temp = self._save_audio_to_temp_file(source_audio, prefix="Source ")
if not source_audio_path_temp:
raise ValueError("Failed to process source audio.")
if target_voice_audio is not None and \
target_voice_audio.get("waveform") is not None and \
target_voice_audio["waveform"].numel() > 0:
if not source_audio_path_temp: raise ValueError("Failed to process source audio.")
if target_voice_audio and target_voice_audio.get("waveform") is not None and target_voice_audio["waveform"].numel() > 0:
target_voice_path_temp = self._save_audio_to_temp_file(target_voice_audio, prefix="Target ")
#print(f"ChatterboxVC: Using target voice from temp file: {target_voice_path_temp}")
else:
print("ChatterboxVC: No target voice audio provided or it's empty. Using default reference from model pack if available.")
# ChatterboxVC.generate expects file paths
converted_wav_tensor = vc_model.generate(
audio=source_audio_path_temp,
target_voice_path=target_voice_path_temp # This will be None if no target_voice_audio was provided or saving it failed
) # Expected output: (1, num_samples)
target_voice_path=target_voice_path_temp,
n_timesteps=n_timesteps,
pbar=pbar,
temperature=temperature,
flow_cfg_scale=flow_cfg_scale
)
except Exception as e:
print(f"ChatterboxVC: Error during voice conversion: {e}")
dummy_sr = 24000
silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
logger.error(f"Error during voice conversion: {e}", exc_info=True)
dummy_sr = 24000; silent_waveform = torch.zeros((1, dummy_sr), dtype=torch.float32, device="cpu")
return ({"waveform": silent_waveform.unsqueeze(0), "sample_rate": dummy_sr},)
finally:
if source_audio_path_temp and os.path.exists(source_audio_path_temp):
try:
os.remove(source_audio_path_temp)
except Exception as e:
print(f"ChatterboxVC: Error removing temp source audio file {source_audio_path_temp}: {e}")
if target_voice_path_temp and os.path.exists(target_voice_path_temp):
try:
os.remove(target_voice_path_temp)
except Exception as e:
print(f"ChatterboxVC: Error removing temp target audio file {target_voice_path_temp}: {e}")
vc_model.watermarker = original_watermarker
if source_audio_path_temp and os.path.exists(source_audio_path_temp): os.remove(source_audio_path_temp)
if target_voice_path_temp and os.path.exists(target_voice_path_temp): os.remove(target_voice_path_temp)
# ComfyUI AUDIO format: {"waveform": tensor (B, C, T), "sample_rate": int}
# ChatterboxVC output: tensor (1, T)
vc_wav_tensor_comfy = converted_wav_tensor.cpu().unsqueeze(0) # (1, T) -> (1, 1, T)
return ({"waveform": vc_wav_tensor_comfy, "sample_rate": vc_model.sr},)
vc_wav_tensor_comfy = converted_wav_tensor.cpu().unsqueeze(0)
return ({"waveform": vc_wav_tensor_comfy, "sample_rate": vc_model.sr},)
+7
View File
@@ -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
View File
+10
View File
@@ -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"
})
+105 -47
View File
@@ -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
+22 -18
View File
@@ -14,18 +14,9 @@
import threading
import torch
import torch.nn.functional as F
from tqdm import tqdm
from .matcha.flow_matching import BASECFM
from omegaconf import OmegaConf
CFM_PARAMS = OmegaConf.create({
"sigma_min": 1e-06,
"solver": "euler",
"t_scheduler": "cosine",
"training_cfg_rate": 0.2,
"inference_cfg_rate": 0.7,
"reg_loss_type": "l1"
})
from .configs import CFM_PARAMS
class ConditionalCFM(BASECFM):
@@ -45,7 +36,7 @@ class ConditionalCFM(BASECFM):
self.lock = threading.Lock()
@torch.inference_mode()
def forward(self, mu, mask, n_timesteps, temperature=1.0, spks=None, cond=None, prompt_len=0, flow_cache=torch.zeros(1, 80, 0, 2)):
def forward(self, mu, mask, n_timesteps, temperature=1.0, spks=None, cond=None, prompt_len=0, flow_cache=torch.zeros(1, 80, 0, 2), pbar=None, flow_cfg_scale=None):
"""Forward diffusion
Args:
@@ -58,11 +49,15 @@ class ConditionalCFM(BASECFM):
spks (torch.Tensor, optional): speaker ids. Defaults to None.
shape: (batch_size, spk_emb_dim)
cond: Not used but kept for future purposes
pbar: ComfyUI ProgressBar instance
Returns:
sample: generated mel-spectrogram
shape: (batch_size, n_feats, mel_timesteps)
"""
if flow_cfg_scale is not None:
self.inference_cfg_rate = flow_cfg_scale
z = torch.randn_like(mu).to(mu.device).to(mu.dtype) * temperature
cache_size = flow_cache.shape[2]
@@ -77,9 +72,10 @@ class ConditionalCFM(BASECFM):
t_span = torch.linspace(0, 1, n_timesteps + 1, device=mu.device, dtype=mu.dtype)
if self.t_scheduler == 'cosine':
t_span = 1 - torch.cos(t_span * 0.5 * torch.pi)
return self.solve_euler(z, t_span=t_span, mu=mu, mask=mask, spks=spks, cond=cond), flow_cache
return self.solve_euler(z, t_span=t_span, mu=mu, mask=mask, spks=spks, cond=cond, pbar=pbar), flow_cache
def solve_euler(self, x, t_span, mu, mask, spks, cond):
def solve_euler(self, x, t_span, mu, mask, spks, cond, pbar=None):
"""
Fixed euler solver for ODEs.
Args:
@@ -93,6 +89,7 @@ class ConditionalCFM(BASECFM):
spks (torch.Tensor, optional): speaker ids. Defaults to None.
shape: (batch_size, spk_emb_dim)
cond: Not used but kept for future purposes
pbar: ComfyUI ProgressBar instance
"""
t, _, dt = t_span[0], t_span[-1], t_span[1] - t_span[0]
t = t.unsqueeze(dim=0)
@@ -100,7 +97,8 @@ class ConditionalCFM(BASECFM):
# I am storing this because I can later plot it by putting a debugger here and saving it to a file
# Or in future might add like a return_all_steps flag
sol = []
iterator = tqdm(range(1, len(t_span)), desc="Voice Converting", dynamic_ncols=True)
# Do not use concat, it may cause memory format changed and trt infer with wrong results!
x_in = torch.zeros([2, 80, x.size(2)], device=x.device, dtype=x.dtype)
mask_in = torch.zeros([2, 1, x.size(2)], device=x.device, dtype=x.dtype)
@@ -108,7 +106,7 @@ class ConditionalCFM(BASECFM):
t_in = torch.zeros([2], device=x.device, dtype=x.dtype)
spks_in = torch.zeros([2, 80], device=x.device, dtype=x.dtype)
cond_in = torch.zeros([2, 80, x.size(2)], device=x.device, dtype=x.dtype)
for step in range(1, len(t_span)):
for step in iterator:
# Classifier-Free Guidance inference introduced in VoiceBox
x_in[:] = x
mask_in[:] = mask
@@ -129,6 +127,9 @@ class ConditionalCFM(BASECFM):
sol.append(x)
if step < len(t_span) - 1:
dt = t_span[step + 1] - t
if pbar:
pbar.update(1)
return sol[-1].float()
@@ -201,7 +202,7 @@ class CausalConditionalCFM(ConditionalCFM):
self.rand_noise = torch.randn([1, 80, 50 * 300])
@torch.inference_mode()
def forward(self, mu, mask, n_timesteps, temperature=1.0, spks=None, cond=None):
def forward(self, mu, mask, n_timesteps, temperature=1.0, spks=None, cond=None, pbar=None, flow_cfg_scale=None, **kwargs):
"""Forward diffusion
Args:
@@ -220,9 +221,12 @@ class CausalConditionalCFM(ConditionalCFM):
shape: (batch_size, n_feats, mel_timesteps)
"""
if flow_cfg_scale is not None:
self.inference_cfg_rate = flow_cfg_scale
z = self.rand_noise[:, :, :mu.size(2)].to(mu.device).to(mu.dtype) * temperature
# fix prompt and overlap part mu and z
t_span = torch.linspace(0, 1, n_timesteps + 1, device=mu.device, dtype=mu.dtype)
if self.t_scheduler == 'cosine':
t_span = 1 - torch.cos(t_span * 0.5 * torch.pi)
return self.solve_euler(z, t_span=t_span, mu=mu, mask=mask, spks=spks, cond=cond), None
return self.solve_euler(z, t_span=t_span, mu=mu, mask=mask, spks=spks, cond=cond, pbar=pbar), None
@@ -2,15 +2,21 @@ from typing import Any, Dict, Optional
import torch
import torch.nn as nn
from diffusers.models.attention import (
from diffusers.models.activations import (
GEGLU,
GELU,
AdaLayerNorm,
AdaLayerNormZero,
ApproximateGELU,
)
from diffusers.models.normalization import (
AdaLayerNorm,
AdaLayerNormZero,
)
from diffusers.models.attention_processor import Attention
from diffusers.models.lora import LoRACompatibleLinear
# deprecated LoRACompatibleLinear
# from diffusers.models.lora import LoRACompatibleLinear
from diffusers.utils.torch_utils import maybe_allow_in_graph
@@ -45,7 +51,9 @@ class SnakeBeta(nn.Module):
"""
super().__init__()
self.in_features = out_features if isinstance(out_features, list) else [out_features]
self.proj = LoRACompatibleLinear(in_features, out_features)
# Switched to standard torch.nn.Linear
self.proj = nn.Linear(in_features, out_features)
# initialize alpha
self.alpha_logscale = alpha_logscale
@@ -123,7 +131,9 @@ class FeedForward(nn.Module):
# project dropout
self.net.append(nn.Dropout(dropout))
# project out
self.net.append(LoRACompatibleLinear(inner_dim, dim_out))
# Switched to standard torch.nn.Linear
self.net.append(nn.Linear(inner_dim, dim_out))
# FF as used in Vision Transformer, MLP-Mixer, etc. have a final dropout
if final_dropout:
self.net.append(nn.Dropout(dropout))
+59 -15
View File
@@ -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.
+19 -11
View File
@@ -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,
}
}
+35 -8
View File
@@ -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):
+39 -33
View File
@@ -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
+4
View File
@@ -0,0 +1,4 @@
class AttrDict(dict):
def __init__(self, *args, **kwargs):
super(AttrDict, self).__init__(*args, **kwargs)
self.__dict__ = self
@@ -259,7 +259,7 @@ class VoiceEncoder(nn.Module):
"""
if sample_rate != self.hp.sample_rate:
wavs = [
librosa.resample(wav, orig_sr=sample_rate, target_sr=self.hp.sample_rate, res_type="kaiser_fast")
librosa.resample(wav, orig_sr=sample_rate, target_sr=self.hp.sample_rate, res_type="kaiser_best")
for wav in wavs
]
+69 -11
View File
@@ -1,3 +1,4 @@
import logging
from dataclasses import dataclass
from pathlib import Path
@@ -6,14 +7,16 @@ import torch
import perth
import torch.nn.functional as F
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
from .models.t3 import T3
from .models.s3tokenizer import S3_SR, drop_invalid_tokens
from .models.s3tokenizer import S3_SR, S3_TOKEN_RATE, drop_invalid_tokens
from .models.s3gen import S3GEN_SR, S3Gen
from .models.tokenizers import EnTokenizer
from .models.voice_encoder import VoiceEncoder
from .models.t3.modules.cond_enc import T3Cond
logger = logging.getLogger(__name__)
REPO_ID = "ResembleAI/chatterbox"
@@ -96,9 +99,27 @@ class Conditionals:
@classmethod
def load(cls, fpath, map_location="cpu"):
if isinstance(map_location, str):
map_location = torch.device(map_location)
kwargs = torch.load(fpath, map_location=map_location, weights_only=True)
return cls(T3Cond(**kwargs['t3']), kwargs['gen'])
# helper function for padding
def _pad_wav_to_40ms_multiple(wav: torch.Tensor, sr: int) -> torch.Tensor:
"""
Pads a waveform to be a multiple of 40ms to prevent rounding errors between
the mel spectrogram (20ms hop) and the speech tokenizer (40ms hop).
"""
S3_TOKEN_DURATION_S = 1 / S3_TOKEN_RATE # 0.04 seconds
samples_per_token = int(sr * S3_TOKEN_DURATION_S)
current_samples = wav.shape[-1]
remainder = current_samples % samples_per_token
if remainder != 0:
padding_needed = samples_per_token - remainder
padded_wav = F.pad(wav, (0, padding_needed))
return padded_wav
return wav
class ChatterboxTTS:
ENC_COND_LEN = 6 * S3_SR
@@ -126,14 +147,20 @@ class ChatterboxTTS:
def from_local(cls, ckpt_dir, device) -> 'ChatterboxTTS':
ckpt_dir = Path(ckpt_dir)
# Always load to CPU first for non-CUDA devices to handle CUDA-saved models
if device in ["cpu", "mps"]:
map_location = torch.device('cpu')
else:
map_location = None
ve = VoiceEncoder()
ve.load_state_dict(
torch.load(ckpt_dir / "ve.pt")
load_file(ckpt_dir / "ve.safetensors")
)
ve.to(device).eval()
t3 = T3()
t3_state = torch.load(ckpt_dir / "t3_cfg.pt")
t3_state = load_file(ckpt_dir / "t3_cfg.safetensors")
if "model" in t3_state.keys():
t3_state = t3_state["model"][0]
t3.load_state_dict(t3_state)
@@ -141,7 +168,7 @@ class ChatterboxTTS:
s3gen = S3Gen()
s3gen.load_state_dict(
torch.load(ckpt_dir / "s3gen.pt")
load_file(ckpt_dir / "s3gen.safetensors"), strict=False
)
s3gen.to(device).eval()
@@ -151,13 +178,21 @@ class ChatterboxTTS:
conds = None
if (builtin_voice := ckpt_dir / "conds.pt").exists():
conds = Conditionals.load(builtin_voice).to(device)
conds = Conditionals.load(builtin_voice, map_location=map_location).to(device)
return cls(t3, s3gen, ve, tokenizer, device, conds=conds)
@classmethod
def from_pretrained(cls, device) -> 'ChatterboxTTS':
for fpath in ["ve.pt", "t3_cfg.pt", "s3gen.pt", "tokenizer.json", "conds.pt"]:
# Check if MPS is available on macOS
if device == "mps" and not torch.backends.mps.is_available():
if not torch.backends.mps.is_built():
logger.warning("MPS not available because the current PyTorch install was not built with MPS enabled.")
else:
logger.warning("MPS not available because the current MacOS version is not 12.3+ and/or you do not have an MPS-enabled device on this machine.")
device = "cpu"
for fpath in ["ve.safetensors", "t3_cfg.safetensors", "s3gen.safetensors", "tokenizer.json", "conds.pt"]:
local_path = hf_hub_download(repo_id=REPO_ID, filename=fpath)
return cls.from_local(Path(local_path).parent, device)
@@ -165,11 +200,18 @@ class ChatterboxTTS:
def prepare_conditionals(self, wav_fpath, exaggeration=0.5):
## Load reference wav
s3gen_ref_wav, _sr = librosa.load(wav_fpath, sr=S3GEN_SR)
# Convert to tensor and pad to a 40ms boundary
s3gen_ref_wav = torch.from_numpy(s3gen_ref_wav).float().unsqueeze(0)
s3gen_ref_wav = _pad_wav_to_40ms_multiple(s3gen_ref_wav, S3GEN_SR)
# Now convert back to numpy for librosa, or use torch audio for resampling
s3gen_ref_wav_np = s3gen_ref_wav.squeeze(0).numpy()
ref_16k_wav = librosa.resample(s3gen_ref_wav, orig_sr=S3GEN_SR, target_sr=S3_SR)
ref_16k_wav = librosa.resample(s3gen_ref_wav_np, orig_sr=S3GEN_SR, target_sr=S3_SR)
s3gen_ref_wav = s3gen_ref_wav[:self.DEC_COND_LEN]
s3gen_ref_dict = self.s3gen.embed_ref(s3gen_ref_wav, S3GEN_SR, device=self.device)
s3gen_ref_wav_np = s3gen_ref_wav_np[:self.DEC_COND_LEN]
s3gen_ref_dict = self.s3gen.embed_ref(s3gen_ref_wav_np, S3GEN_SR, device=self.device)
# Speech cond prompt tokens
if plen := self.t3.hp.speech_cond_prompt_len:
@@ -191,10 +233,16 @@ class ChatterboxTTS:
def generate(
self,
text,
repetition_penalty=1.2,
min_p=0.05,
top_p=1.0,
audio_prompt_path=None,
exaggeration=0.5,
cfg_weight=0.5,
temperature=0.8,
pbar=None,
max_new_tokens=1000,
flow_cfg_scale=0.7
):
if audio_prompt_path:
self.prepare_conditionals(audio_prompt_path, exaggeration=exaggeration)
@@ -213,7 +261,9 @@ class ChatterboxTTS:
# Norm and tokenize text
text = punc_norm(text)
text_tokens = self.tokenizer.text_to_tokens(text).to(self.device)
text_tokens = torch.cat([text_tokens, text_tokens], dim=0) # Need two seqs for CFG
if cfg_weight > 0.0:
text_tokens = torch.cat([text_tokens, text_tokens], dim=0) # Need two seqs for CFG
sot = self.t3.hp.start_text_token
eot = self.t3.hp.stop_text_token
@@ -224,20 +274,28 @@ class ChatterboxTTS:
speech_tokens = self.t3.inference(
t3_cond=self.conds.t3,
text_tokens=text_tokens,
max_new_tokens=1000, # TODO: use the value in config
max_new_tokens=max_new_tokens,
temperature=temperature,
cfg_weight=cfg_weight,
repetition_penalty=repetition_penalty,
min_p=min_p,
top_p=top_p,
pbar=pbar
)
# Extract only the conditional batch.
speech_tokens = speech_tokens[0]
# TODO: output becomes 1D
speech_tokens = drop_invalid_tokens(speech_tokens)
speech_tokens = speech_tokens[speech_tokens < 6561]
speech_tokens = speech_tokens.to(self.device)
wav, _ = self.s3gen.inference(
speech_tokens=speech_tokens,
ref_dict=self.conds.gen,
flow_cfg_scale=flow_cfg_scale
)
wav = wav.squeeze(0).detach().cpu().numpy()
watermarked_wav = self.watermarker.apply_watermark(wav, sample_rate=self.sr)
+53 -7
View File
@@ -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)