This commit is contained in:
billwuhao
2025-03-18 23:26:01 +08:00
parent 8e736d0c2d
commit 02a55fa18c
15 changed files with 1001 additions and 2 deletions
+22
View File
@@ -0,0 +1,22 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- master
- main
paths:
- "pyproject.toml"
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+510
View File
@@ -0,0 +1,510 @@
from dataclasses import dataclass
from typing import List, Tuple
import ast
import silentcipher
import torch
import torchaudio
import os
from huggingface_hub import hf_hub_download
from .models import Model, ModelArgs
from moshi.models import loaders
from tokenizers.processors import TemplateProcessing
from transformers import AutoTokenizer
import folder_paths
models_dir = folder_paths.models_dir
class AddWatermark:
if torch.backends.mps.is_available():
device = "mps"
elif torch.cuda.is_available():
device = "cuda"
else:
device = "cpu"
@classmethod
def INPUT_TYPES(s):
return {"required": {
"audio": ("AUDIO",),
"add_watermark": ("BOOLEAN", {"default": False, "tooltip": "Add watermark or not."}),
"key": ("STRING", {"default": "[212, 211, 146, 56, 201]", "tooltip": "List of integers such as [212, 211, 146, 56, 201]"}),
},
# "optional": {
# "check_watermark": ("BOOLEAN", {"default": False, "tooltip": "Check if the audio contains watermark."}),
# }
}
CATEGORY = "MW_CSM"
RETURN_TYPES = ("AUDIO", "STRING")
RETURN_NAMES = ("audio", "watermark")
FUNCTION = "watermarkgen"
def watermarkgen(self, audio, add_watermark, key):
watermarker = self.load_watermarker(device=self.device)
audio_array, sample_rate = self.load_audio(audio)
# 确保 audio_array 在正确的设备上
audio_array = audio_array.to(self.device)
if add_watermark:
key = self._parse_key(key)
audio_array, sample_rate = self.watermark(watermarker, audio_array, sample_rate, key)
watermark = self.verify(watermarker, audio_array, sample_rate)
# 返回前将音频数据移回 CPU
return ({"waveform": audio_array.unsqueeze(0).unsqueeze(0).cpu(), "sample_rate": sample_rate}, watermark)
@torch.inference_mode()
def watermark(self,
watermarker: silentcipher.server.Model,
audio_array: torch.Tensor,
sample_rate: int,
watermark_key: list[int],
) -> tuple[torch.Tensor, int]:
# 确保音频是单声道
if len(audio_array.shape) > 1 and audio_array.shape[0] > 1:
audio_array = audio_array.mean(dim=0)
audio_array = audio_array.to(self.device)
audio_array_44khz = torchaudio.functional.resample(
audio_array,
orig_freq=sample_rate,
new_freq=44100
).to(self.device)
# 确保音频形状正确 (应为一维张量)
if len(audio_array_44khz.shape) != 1:
audio_array_44khz = audio_array_44khz.reshape(-1)
try:
# 增加水印强度,降低message_sdr值使水印更明显
encoded, _ = watermarker.encode_wav(audio_array_44khz, 44100, watermark_key, calc_sdr=False, message_sdr=30)
verify_result = watermarker.decode_wav(encoded, 44100, phase_shift_decoding=True)
if not verify_result["status"]:
encoded, _ = watermarker.encode_wav(audio_array_44khz, 44100, watermark_key, calc_sdr=False, message_sdr=25)
verify_result = watermarker.decode_wav(encoded, 44100, phase_shift_decoding=True)
except Exception as e:
return audio_array, sample_rate
# 如果需要,重采样回原始采样率
output_sample_rate = min(44100, sample_rate)
if output_sample_rate != 44100:
encoded = torchaudio.functional.resample(
encoded,
orig_freq=44100,
new_freq=output_sample_rate
).to(self.device)
return encoded, output_sample_rate
@torch.inference_mode()
def verify(self,
watermarker: silentcipher.server.Model,
watermarked_audio: torch.Tensor,
sample_rate: int,
) -> str:
if len(watermarked_audio.shape) > 1 and watermarked_audio.shape[0] > 1:
watermarked_audio = watermarked_audio.mean(dim=0)
if sample_rate != 44100:
watermarked_audio_44khz = torchaudio.functional.resample(
watermarked_audio,
orig_freq=sample_rate,
new_freq=44100
).to(self.device)
else:
watermarked_audio_44khz = watermarked_audio.to(self.device)
if len(watermarked_audio_44khz.shape) != 1:
watermarked_audio_44khz = watermarked_audio_44khz.reshape(-1)
# 尝试不同的解码参数
# 1. 使用相位偏移解码
result_phase = watermarker.decode_wav(watermarked_audio_44khz, 44100, phase_shift_decoding=True)
# 2. 不使用相位偏移解码
result_no_phase = watermarker.decode_wav(watermarked_audio_44khz, 44100, phase_shift_decoding=False)
# 使用两种方法中任一种成功的结果
if result_phase["status"]:
watermark = "Watermarked:" + str(result_phase["messages"][0])
elif result_no_phase["status"]:
watermark = "Watermarked:" + str(result_no_phase["messages"][0])
else:
watermark = "No watermarked"
return watermark
def load_watermarker(self, device: str = "cuda") -> silentcipher.server.Model:
ckpt_path = os.path.join(models_dir, "TTS", "SilentCipher", "44_1_khz", "73999_iteration")
config_path = os.path.join(models_dir, ckpt_path, "hparams.yaml")
model = silentcipher.get_model(
model_type="44.1k",
ckpt_path=ckpt_path,
config_path=config_path,
device=device,
)
return model
def _parse_key(self, key_string):
"""Helper function to safely parse the key."""
try:
key = ast.literal_eval(key_string)
return key
except (ValueError, SyntaxError) as e:
raise
def load_audio(self, audio) -> tuple[torch.Tensor, int]:
waveform = audio["waveform"].squeeze(0)
audio_array = waveform.mean(dim=0)
sample_rate = audio["sample_rate"]
return audio_array, int(sample_rate)
@dataclass
class Segment:
speaker: int
text: str
# (num_samples,), sample_rate = 24_000
audio: torch.Tensor
SEGMENTS = []
SPEAKERS = []
class Generator:
# 添加类变量用于缓存
_cached_llama3_tokenizer = None
_cached_mimi = None
def __init__(
self,
model: Model,
device: str = "cuda",
):
self._model = model
self._model.setup_caches(1)
self.device = device
self._text_tokenizer = self.load_llama3_tokenizer()
mimi = self.load_mimi()
self._audio_tokenizer = mimi
self.sample_rate = mimi.sample_rate
def load_llama3_tokenizer(self):
"""
https://github.com/huggingface/transformers/issues/22794#issuecomment-2092623992
"""
if Generator._cached_llama3_tokenizer is not None:
return Generator._cached_llama3_tokenizer
llama_path = os.path.join(models_dir, "LLM", "Llama-3.2-1B")
tokenizer = AutoTokenizer.from_pretrained(llama_path)
bos = tokenizer.bos_token
eos = tokenizer.eos_token
tokenizer._tokenizer.post_processor = TemplateProcessing(
single=f"{bos}:0 $A:0 {eos}:0",
pair=f"{bos}:0 $A:0 {eos}:0 {bos}:1 $B:1 {eos}:1",
special_tokens=[(f"{bos}", tokenizer.bos_token_id), (f"{eos}", tokenizer.eos_token_id)],
)
Generator._cached_llama3_tokenizer = tokenizer
return tokenizer
def load_mimi(self):
if Generator._cached_mimi is not None:
return Generator._cached_mimi
mimi_path = os.path.join(models_dir, "TTS", "moshiko-pytorch-bf16", loaders.MIMI_NAME)
mimi = loaders.get_mimi(mimi_path, device=self.device)
mimi.set_num_codebooks(32)
Generator._cached_mimi = mimi
return mimi
def _tokenize_text_segment(self, text: str, speaker: int) -> Tuple[torch.Tensor, torch.Tensor]:
frame_tokens = []
frame_masks = []
text_tokens = self._text_tokenizer.encode(f"[{speaker}]{text}")
text_frame = torch.zeros(len(text_tokens), 33).long()
text_frame_mask = torch.zeros(len(text_tokens), 33).bool()
text_frame[:, -1] = torch.tensor(text_tokens)
text_frame_mask[:, -1] = True
frame_tokens.append(text_frame.to(self.device))
frame_masks.append(text_frame_mask.to(self.device))
return torch.cat(frame_tokens, dim=0), torch.cat(frame_masks, dim=0)
def _tokenize_audio(self, audio: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
frame_tokens = []
frame_masks = []
# (K, T)
audio = audio.to(self.device)
audio_tokens = self._audio_tokenizer.encode(audio.unsqueeze(0).unsqueeze(0))[0]
# add EOS frame
eos_frame = torch.zeros(audio_tokens.size(0), 1).to(self.device)
audio_tokens = torch.cat([audio_tokens, eos_frame], dim=1)
audio_frame = torch.zeros(audio_tokens.size(1), 33).long().to(self.device)
audio_frame_mask = torch.zeros(audio_tokens.size(1), 33).bool().to(self.device)
audio_frame[:, :-1] = audio_tokens.transpose(0, 1)
audio_frame_mask[:, :-1] = True
frame_tokens.append(audio_frame)
frame_masks.append(audio_frame_mask)
return torch.cat(frame_tokens, dim=0), torch.cat(frame_masks, dim=0)
def _tokenize_segment(self, segment: Segment) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Returns:
(seq_len, 33), (seq_len, 33)
"""
text_tokens, text_masks = self._tokenize_text_segment(segment.text, segment.speaker)
audio_tokens, audio_masks = self._tokenize_audio(segment.audio)
return torch.cat([text_tokens, audio_tokens], dim=0), torch.cat([text_masks, audio_masks], dim=0)
@torch.inference_mode()
def generate(
self,
text: str,
speaker: int,
context: List[Segment],
max_audio_length_ms: float = 90_000,
temperature: float = 0.9,
topk: int = 50,
) -> torch.Tensor:
self._model.reset_caches()
max_audio_frames = int(max_audio_length_ms / 80)
tokens, tokens_mask = [], []
for segment in context:
segment_tokens, segment_tokens_mask = self._tokenize_segment(segment)
tokens.append(segment_tokens)
tokens_mask.append(segment_tokens_mask)
gen_segment_tokens, gen_segment_tokens_mask = self._tokenize_text_segment(text, speaker)
tokens.append(gen_segment_tokens)
tokens_mask.append(gen_segment_tokens_mask)
prompt_tokens = torch.cat(tokens, dim=0).long().to(self.device)
prompt_tokens_mask = torch.cat(tokens_mask, dim=0).bool().to(self.device)
samples = []
curr_tokens = prompt_tokens.unsqueeze(0)
curr_tokens_mask = prompt_tokens_mask.unsqueeze(0)
curr_pos = torch.arange(0, prompt_tokens.size(0)).unsqueeze(0).long().to(self.device)
max_seq_len = 2048 - max_audio_frames
if curr_tokens.size(1) >= max_seq_len:
raise ValueError(f"Inputs too long, must be below max_seq_len - max_audio_frames: {max_seq_len}")
for _ in range(max_audio_frames):
sample = self._model.generate_frame(curr_tokens, curr_tokens_mask, curr_pos, temperature, topk)
if torch.all(sample == 0):
break # eos
samples.append(sample)
curr_tokens = torch.cat([sample, torch.zeros(1, 1).long().to(self.device)], dim=1).unsqueeze(1)
curr_tokens_mask = torch.cat(
[torch.ones_like(sample).bool(), torch.zeros(1, 1).bool().to(self.device)], dim=1
).unsqueeze(1)
curr_pos = curr_pos[:, -1:] + 1
audio = self._audio_tokenizer.decode(torch.stack(samples).permute(1, 2, 0)).squeeze(0).squeeze(0)
return audio
class MultiLinePromptCSM:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"multi_line_prompt": ("STRING", {
"multiline": True,
"default": ""}),
},
}
CATEGORY = "MW_CSM"
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("prompt",)
FUNCTION = "promptgen"
def promptgen(self, multi_line_prompt: str):
return (multi_line_prompt.strip(),)
class CSMDialogRun:
# 添加类变量用于缓存
_cached_csm_1b = None
_cached_generator = None
if torch.backends.mps.is_available():
device = "mps"
elif torch.cuda.is_available():
device = "cuda"
else:
device = "cpu"
@classmethod
def INPUT_TYPES(s):
return {"required": {
"text": ("STRING",),
"unload_speakers": ("BOOLEAN",{ "default": False}),
},
"optional": {
"prompt0": ("STRING",),
"prompt1": ("STRING",),
"prompt2": ("STRING",),
"prompt3": ("STRING",),
"audio0": ("AUDIO",),
"audio1": ("AUDIO",),
"audio2": ("AUDIO",),
"audio3": ("AUDIO",),
"who_will_speak": ("INT", {"default": 0, "min": 0, "max": 9, "step": 1,}),
}
}
CATEGORY = "MW_CSM"
RETURN_TYPES = ("AUDIO", "STRING")
RETURN_NAMES = ("audio", "prompt")
FUNCTION = "run"
def run(self, text, unload_speakers, prompt0="", prompt1="", prompt2="", prompt3="", audio0=None, audio1=None, audio2=None, audio3=None, who_will_speak=1):
generator = self.load_csm_1b()
global SEGMENTS, SPEAKERS
if unload_speakers:
SEGMENTS.clear()
SPEAKERS.clear()
segments = []
for i in range(4):
prompt = locals()[f"prompt{i}"]
audio = locals()[f"audio{i}"]
# print(f"prompt{i}: {prompt}→audio{i}")
if audio is not None:
audio_tensor = audio["waveform"].squeeze(0).mean(dim=0)
sample_rate = int(audio["sample_rate"])
# print(f"audio{i} sample_rate: {sample_rate}")
audio_tensor = torchaudio.functional.resample(
audio_tensor.squeeze(0),
orig_freq=sample_rate,
new_freq=generator.sample_rate
)
else:
audio_tensor = None
speaker, prompt = self.get_speaker_text(prompt)
segment = self.get_segment(speaker, prompt, audio_tensor)
if segment is not None:
SEGMENTS.append(segment)
SPEAKERS.append(speaker)
if SEGMENTS:
SEGMENTS.extend(segments)
if len(SEGMENTS) > 6:
SEGMENTS = SEGMENTS[-6:]
SPEAKERS = SPEAKERS[-6:]
if who_will_speak not in SPEAKERS:
raise ValueError(f"The speaker {who_will_speak} not found or cleared, up to 6 recent speakers saved.")
# print(f"使用 {len(segments)} 个上下文段落生成音频,说话者: {will_speaker}")
audio = generator.generate(
text=text,
speaker=who_will_speak,
context=SEGMENTS,
max_audio_length_ms=10_000,
)
out_prompt = str(who_will_speak) + ": " + text
else:
# print(f"无上下文生成音频,说话者: 0")
audio = generator.generate(
text=text,
speaker=0,
context=[],
max_audio_length_ms=10_000,
)
out_prompt = "0: " + text
return ({"waveform": audio.unsqueeze(0).unsqueeze(0).cpu(), "sample_rate": generator.sample_rate}, out_prompt)
def get_speaker_text(self, text):
import re
if text.strip() != "":
st = [i.strip() for i in re.split('[::]', text, 1)]
if len(st) == 2:
speaker, prompt = st
return int(speaker), prompt
else:
raise ValueError("Invalid text format")
else:
return None, None
def get_segment(self, speaker, text, audio):
if speaker is not None:
if audio is not None:
return Segment(speaker=speaker, text=text, audio=audio)
else:
raise ValueError(f"{text}: Audio is required")
else:
return None
def load_csm_1b(self) -> Generator:
if CSMDialogRun._cached_generator is not None:
return CSMDialogRun._cached_generator
if CSMDialogRun._cached_csm_1b is None:
csm_1b_path = os.path.join(models_dir, "TTS", "csm-1b")
config_path = os.path.join(csm_1b_path, "config.json")
import json
with open(config_path, 'r') as f:
config = json.load(f)["args"]
configs = ModelArgs(backbone_flavor = config["backbone_flavor"],
decoder_flavor = config["decoder_flavor"],
text_vocab_size = config["text_vocab_size"],
audio_vocab_size = config["audio_vocab_size"],
audio_num_codebooks = config["audio_num_codebooks"])
model = Model.from_pretrained(csm_1b_path, config=configs)
model.to(device=self.device, dtype=torch.bfloat16)
CSMDialogRun._cached_csm_1b = model
generator = Generator(CSMDialogRun._cached_csm_1b, device=self.device)
CSMDialogRun._cached_generator = generator
return generator
from .MWAudioRecorderCSM import AudioRecorderCSM
NODE_CLASS_MAPPINGS = {
"AddWatermark": AddWatermark,
"CSMDialogRun": CSMDialogRun,
"MultiLinePromptCSM": MultiLinePromptCSM,
"AudioRecorderCSM": AudioRecorderCSM,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"AddWatermark": "Add Watermark",
"CSMDialogRun": "CSM Dialog Run",
"MultiLinePromptCSM": "Multi Line Prompt",
"AudioRecorderCSM": "MW Audio Recorder",
}
+138
View File
@@ -0,0 +1,138 @@
import numpy as np
import torch
import time
import librosa
import sounddevice as sd
from scipy import ndimage
from comfy.utils import ProgressBar
class AudioRecorderCSM:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
# 触发控制
"trigger": ("BOOLEAN", {"default": False}),
# 录音时长
"record_sec": ("INT", {
"default": 5,
"min": 1,
"max": 60,
"step": 1 # 整数秒递增
}),
"sample_rate": (["16000", "44100", "48000"], { # 限定标准采样率
"default": "48000"
}),
"n_fft": ("INT", { # 限定为2的幂次方
"default": 2048,
"min": 512,
"max": 4096,
"step": 512 # 只能选择512,1024,1536,2048,...4096
}),
"sensitivity": ("FLOAT", { # 灵敏度精确控制
"default": 1.2,
"min": 0.5,
"max": 3.0,
"step": 0.1 # 0.1步进
}),
"smooth": ("INT", { # 确保为奇数
"default": 5,
"min": 1,
"max": 11,
"step": 2 # 生成1,3,5,7,9,11
}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
},
"optional": {
"interlocutor": ("AUDIO",),
},
}
RETURN_TYPES = ("AUDIO", "AUDIO")
RETURN_NAMES = ("audio", "interlocutor")
FUNCTION = "record_and_clean"
CATEGORY = "MW_CSM"
def _stft(self, y, n_fft):
hop = n_fft // 4
return librosa.stft(y, n_fft=n_fft, hop_length=hop, win_length=n_fft)
def _istft(self, spec, n_fft):
hop = n_fft // 4
return librosa.istft(spec, hop_length=hop, win_length=n_fft)
def _calc_noise_profile(self, noise_clip, n_fft):
noise_spec = self._stft(noise_clip, n_fft)
return {
'mean': np.mean(np.abs(noise_spec), axis=1, keepdims=True),
'std': np.std(np.abs(noise_spec), axis=1, keepdims=True)
}
def _spectral_gate(self, spec, noise_profile, sensitivity):
threshold = noise_profile['mean'] + sensitivity * noise_profile['std']
return np.where(np.abs(spec) > threshold, spec, 0)
def _smooth_mask(self, mask, kernel_size):
smoothed = ndimage.uniform_filter(mask, size=(kernel_size, kernel_size))
return np.clip(smoothed * 1.2, 0, 1) # 增强边缘保留
def record_and_clean(self, trigger, record_sec, n_fft, sensitivity, smooth, sample_rate, interlocutor=None, seed=0):
if not trigger:
if interlocutor is not None:
return (None, interlocutor)
raise ValueError("No trigger received: Recording not opened.")
sr = int(sample_rate)
final_audio = None
try:
noise_clip = None
# 主录音
# print(f"开始主录音 {record_sec}秒...")
main_rec = sd.rec(int(record_sec * sr), samplerate=sr, channels=1, dtype='float32')
pb = ProgressBar(record_sec)
for _ in range(record_sec * 2):
time.sleep(0.5)
pb.update(0.5)
sd.wait()
audio = main_rec.flatten()
# 自动噪声检测
if noise_clip is None:
# print("自动检测静默段作为噪声参考...")
energy = librosa.feature.rms(y=audio, frame_length=n_fft, hop_length=n_fft//4)
min_idx = np.argmin(energy)
start = min_idx * (n_fft//4)
noise_clip = audio[start:start + n_fft*2]
# 降噪处理
# print("进行频谱降噪...")
noise_profile = self._calc_noise_profile(noise_clip, n_fft)
spec = self._stft(audio, n_fft)
# 多步骤处理
mask = np.ones_like(spec) # 初始掩膜
for _ in range(2): # 双重处理循环
cleaned_spec = self._spectral_gate(spec, noise_profile, sensitivity)
mask = np.where(np.abs(cleaned_spec) > 0, 1, 0)
mask = self._smooth_mask(mask, smooth//2+1)
spec = spec * mask
# 相位恢复重建
processed = self._istft(spec * mask, n_fft)
# 动态增益归一化
peak = np.max(np.abs(processed))
processed = processed * (0.99 / peak) if peak > 0 else processed
# 格式转换
waveform = torch.from_numpy(processed).float().unsqueeze(0).unsqueeze(0)
final_audio = {"waveform": waveform, "sample_rate": sr}
except Exception as e:
print(f"Recording/processing failed: {str(e)}")
raise
return (final_audio, interlocutor)
+49
View File
@@ -0,0 +1,49 @@
[中文](README-CN.md)|[English](README.md)
# ComfyUI Node for CSM
CSM(Conversational Speech Model 会话语音模型), 多人会话, 克隆声音, 然后可根据会话中语音的情绪变化等, 生成相应情绪变化的语音的模型. 遗憾的是目前只有英文可用. 该节点, 暂时支持同时 10 人会话.
![](https://github.com/billwuhao/ComfyUI_CSM/blob/master/images/2025-03-18_15-38-45.png)
可以将录音节点穿插在其中, 进行多人会话.
![](https://github.com/billwuhao/ComfyUI_CSM/blob/master/images/2025-03-18_19-01-15.png)
还支持音频水印检测(自动检测水印) 和 音频添加加密水印.
![](https://github.com/billwuhao/ComfyUI_CSM/blob/master/images/2025-03-18_14-43-49.png)
## 📣 更新
[2025-03-18]⚒️: 发布版本 v1.0.0.
节点使用方法详解示例工作流: [example_workflows](https://github.com/billwuhao/ComfyUI_CSM/blob/master/example_workflows)
## 安装
```
cd ComfyUI/custom_nodes
git clone https://github.com/billwuhao/ComfyUI_CSM.git
cd ComfyUI_CSM
pip install -r requirements.txt
# python_embeded
./python_embeded/python.exe -m pip install -r requirements.txt
```
## 模型下载
- [csm-1b](https://huggingface.co/sesame/csm-1b/tree/main): `config.json` 和 `model.safetensors` 下载放到 `ComfyUI/models/TTS/csm-1b` 目录下.
- [moshiko-pytorch-bf16](https://huggingface.co/kyutai/moshiko-pytorch-bf16/tree/main): `tokenizer-e351c8d8-checkpoint125.safetensors` 下载放到 `ComfyUI/models/TTS/moshiko-pytorch-bf16` 目录下.
- [SilentCipher](https://huggingface.co/Sony/SilentCipher/tree/main/44_1_khz/73999_iteration): 全部模型下载放到 `ComfyUI\models\TTS\SilentCipher\44_1_khz\73999_iteration` 目录下.
- [Llama-3.2-1B](https://huggingface.co/meta-llama/Llama-3.2-1B/tree/main): 除了 `original` 目录, 其他全部下载放到 `ComfyUI\models\LLM\Llama-3.2-1B` 目录下.
## 鸣谢
[csm](https://github.com/SesameAILabs/csm)
感谢 SesameAILabs 团队的卓越的工作👍.
+49 -2
View File
@@ -1,2 +1,49 @@
# ComfyUI_CSM
ComfyUI node of Conversational Speech Model (CSM).
[中文](README-CN.md)|[English](README.md)
# ComfyUI Node for CSM
CSM (Conversational Speech Model), a model that supports multi-person conversations, voice cloning, and generates speech with corresponding emotional changes based on the emotional changes in the conversation. Unfortunately, it is currently only available in English. This node temporarily supports simultaneous conversations with up to 10 people.
![](https://github.com/billwuhao/ComfyUI_CSM/blob/master/images/2025-03-18_15-38-45.png)
Recording nodes can be interspersed within to create multi-person conversations.
![](https://github.com/billwuhao/ComfyUI_CSM/blob/master/images/2025-03-18_19-01-15.png)
It also supports audio watermark detection (automatic watermark detection) and audio adding encrypted watermarks.
![](https://github.com/billwuhao/ComfyUI_CSM/blob/master/images/2025-03-18_14-43-49.png)
## 📣 Updates
[2025-03-18]⚒️: Released version v1.0.0.
Detailed node usage example workflows: [example_workflows](https://github.com/billwuhao/ComfyUI_CSM/blob/master/example_workflows)
## Installation
```
cd ComfyUI/custom_nodes
git clone https://github.com/billwuhao/ComfyUI_CSM.git
cd ComfyUI_CSM
pip install -r requirements.txt
# python_embeded
./python_embeded/python.exe -m pip install -r requirements.txt
```
## Model Download
- [csm-1b](https://huggingface.co/sesame/csm-1b/tree/main): Download `config.json` and `model.safetensors` and place them in the `ComfyUI/models/TTS/csm-1b` directory.
- [moshiko-pytorch-bf16](https://huggingface.co/kyutai/moshiko-pytorch-bf16/tree/main): Download `tokenizer-e351c8d8-checkpoint125.safetensors` and place it in the `ComfyUI/models/TTS/moshiko-pytorch-bf16` directory.
- [SilentCipher](https://huggingface.co/Sony/SilentCipher/tree/main/44_1_khz/73999_iteration): Download all models and place them in the `ComfyUI\models\TTS\SilentCipher\44_1_khz\73999_iteration` directory.
- [Llama-3.2-1B](https://huggingface.co/meta-llama/Llama-3.2-1B/tree/main): Download everything except the `original` directory and place it in the `ComfyUI\models\LLM\Llama-3.2-1B` directory.
## Acknowledgements
[csm](https://github.com/SesameAILabs/csm)
Thanks to the SesameAILabs team for their excellent work 👍.
+3
View File
@@ -0,0 +1,3 @@
from .CSMNode import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -0,0 +1 @@
{"last_node_id":15,"last_link_id":15,"nodes":[{"id":2,"type":"LoadAudio","pos":[-1.9409359693527222,-13.219098091125488],"size":[315,136],"flags":{},"order":0,"mode":4,"inputs":[],"outputs":[{"name":"AUDIO","localized_name":"AUDIO","type":"AUDIO","links":[9],"slot_index":0}],"properties":{"cnr_id":"comfy-core","ver":"0.3.26","Node name for S&R":"LoadAudio"},"widgets_values":["ComfyUI_00001_.flac","",""]},{"id":5,"type":"AddWatermark","pos":[371.0498352050781,-3.9667160511016846],"size":[315,102],"flags":{},"order":2,"mode":4,"inputs":[{"name":"audio","localized_name":"audio","type":"AUDIO","link":9}],"outputs":[{"name":"audio","localized_name":"audio","type":"AUDIO","links":[5],"slot_index":0},{"name":"watermark","localized_name":"watermark","type":"STRING","links":[6]}],"properties":{"aux_id":"billwuhao/ComfyUI_CSM","ver":"8e736d0c2d6a2c9b165b88c048e5110e2c731f59","Node name for S&R":"AddWatermark"},"widgets_values":[true,"[0, 211, 146, 56, 201]"]},{"id":4,"type":"PreviewAudio","pos":[749.3675537109375,-7.144290447235107],"size":[315,88],"flags":{},"order":4,"mode":4,"inputs":[{"name":"audio","localized_name":"audio","type":"AUDIO","link":5}],"outputs":[],"properties":{"cnr_id":"comfy-core","ver":"0.3.26","Node name for S&R":"PreviewAudio"},"widgets_values":[""]},{"id":3,"type":"easy showAnything","pos":[762.9191284179688,143.41648864746094],"size":[316.5421142578125,88],"flags":{},"order":5,"mode":4,"inputs":[{"name":"anything","localized_name":"anything","type":"*","shape":7,"link":6}],"outputs":[{"name":"output","localized_name":"output","type":"*","links":null}],"properties":{"cnr_id":"comfyui-easy-use","ver":"f888e3d75dc09c1fc8045bf429e699ed468eda1e","Node name for S&R":"easy showAnything"},"widgets_values":["Watermarked:[0, 211, 146, 56, 201]"]},{"id":14,"type":"CSMDialogRun","pos":[359.74139404296875,275.2856140136719],"size":[315,258],"flags":{},"order":3,"mode":0,"inputs":[{"name":"audio0","localized_name":"audio0","type":"AUDIO","shape":7,"link":null},{"name":"audio1","localized_name":"audio1","type":"AUDIO","shape":7,"link":null},{"name":"audio2","localized_name":"audio2","type":"AUDIO","shape":7,"link":null},{"name":"audio3","localized_name":"audio3","type":"AUDIO","shape":7,"link":null},{"name":"text","type":"STRING","widget":{"name":"text"},"link":15}],"outputs":[{"name":"audio","localized_name":"audio","type":"AUDIO","links":[14]}],"properties":{"aux_id":"billwuhao/ComfyUI_CSM","ver":"8e736d0c2d6a2c9b165b88c048e5110e2c731f59","Node name for S&R":"CSMDialogRun"},"widgets_values":["","","","","",1]},{"id":15,"type":"MultiLinePromptCSM","pos":[-100.63268280029297,330.7996520996094],"size":[400,200],"flags":{},"order":1,"mode":0,"inputs":[],"outputs":[{"name":"prompt","localized_name":"prompt","type":"STRING","links":[15],"slot_index":0}],"properties":{"aux_id":"billwuhao/ComfyUI_CSM","ver":"8e736d0c2d6a2c9b165b88c048e5110e2c731f59","Node name for S&R":"MultiLinePromptCSM"},"widgets_values":["Hello, what are you doing"]},{"id":9,"type":"PreviewAudio","pos":[732.0776977539062,319.2107849121094],"size":[315,88],"flags":{},"order":6,"mode":0,"inputs":[{"name":"audio","localized_name":"audio","type":"AUDIO","link":14}],"outputs":[],"properties":{"cnr_id":"comfy-core","ver":"0.3.26","Node name for S&R":"PreviewAudio"},"widgets_values":[""]}],"links":[[5,5,0,4,0,"AUDIO"],[6,5,1,3,0,"*"],[9,2,0,5,0,"AUDIO"],[14,14,0,9,0,"AUDIO"],[15,15,0,14,4,"STRING"]],"groups":[],"config":{},"extra":{"ds":{"scale":0.9090909090909091,"offset":[505.47957687377937,121.62480659484862]},"ue_links":[]},"version":0.4}
Binary file not shown.

After

Width:  |  Height:  |  Size: 67 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 218 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 211 KiB

+203
View File
@@ -0,0 +1,203 @@
from dataclasses import dataclass
import torch
import torch.nn as nn
import torchtune
from huggingface_hub import PyTorchModelHubMixin
from torchtune.models import llama3_2
def llama3_2_1B() -> torchtune.modules.transformer.TransformerDecoder:
return llama3_2.llama3_2(
vocab_size=128_256,
num_layers=16,
num_heads=32,
num_kv_heads=8,
embed_dim=2048,
max_seq_len=2048,
intermediate_dim=8192,
attn_dropout=0.0,
norm_eps=1e-5,
rope_base=500_000,
scale_factor=32,
)
def llama3_2_100M() -> torchtune.modules.transformer.TransformerDecoder:
return llama3_2.llama3_2(
vocab_size=128_256,
num_layers=4,
num_heads=8,
num_kv_heads=2,
embed_dim=1024,
max_seq_len=2048,
intermediate_dim=8192,
attn_dropout=0.0,
norm_eps=1e-5,
rope_base=500_000,
scale_factor=32,
)
FLAVORS = {
"llama-1B": llama3_2_1B,
"llama-100M": llama3_2_100M,
}
def _prepare_transformer(model):
embed_dim = model.tok_embeddings.embedding_dim
model.tok_embeddings = nn.Identity()
model.output = nn.Identity()
return model, embed_dim
def _create_causal_mask(seq_len: int, device: torch.device):
return torch.tril(torch.ones(seq_len, seq_len, dtype=torch.bool, device=device))
def _index_causal_mask(mask: torch.Tensor, input_pos: torch.Tensor):
"""
Args:
mask: (max_seq_len, max_seq_len)
input_pos: (batch_size, seq_len)
Returns:
(batch_size, seq_len, max_seq_len)
"""
r = mask[input_pos, :]
return r
def _multinomial_sample_one_no_sync(probs): # Does multinomial sampling without a cuda synchronization
q = torch.empty_like(probs).exponential_(1)
return torch.argmax(probs / q, dim=-1, keepdim=True).to(dtype=torch.int)
def sample_topk(logits: torch.Tensor, topk: int, temperature: float):
logits = logits / temperature
filter_value: float = -float("Inf")
indices_to_remove = logits < torch.topk(logits, topk)[0][..., -1, None]
scores_processed = logits.masked_fill(indices_to_remove, filter_value)
scores_processed = torch.nn.functional.log_softmax(scores_processed, dim=-1)
probs = torch.nn.functional.softmax(scores_processed, dim=-1)
sample_token = _multinomial_sample_one_no_sync(probs)
return sample_token
@dataclass
class ModelArgs:
backbone_flavor: str
decoder_flavor: str
text_vocab_size: int
audio_vocab_size: int
audio_num_codebooks: int
class Model(
nn.Module,
PyTorchModelHubMixin,
repo_url="https://github.com/SesameAILabs/csm",
pipeline_tag="text-to-speech",
license="apache-2.0",
):
def __init__(self, config: ModelArgs):
super().__init__()
self.config = config
self.backbone, backbone_dim = _prepare_transformer(FLAVORS[config.backbone_flavor]())
self.decoder, decoder_dim = _prepare_transformer(FLAVORS[config.decoder_flavor]())
self.text_embeddings = nn.Embedding(config.text_vocab_size, backbone_dim)
self.audio_embeddings = nn.Embedding(config.audio_vocab_size * config.audio_num_codebooks, backbone_dim)
self.projection = nn.Linear(backbone_dim, decoder_dim, bias=False)
self.codebook0_head = nn.Linear(backbone_dim, config.audio_vocab_size, bias=False)
self.audio_head = nn.Parameter(torch.empty(config.audio_num_codebooks - 1, decoder_dim, config.audio_vocab_size))
def setup_caches(self, max_batch_size: int) -> torch.Tensor:
"""Setup KV caches and return a causal mask."""
dtype = next(self.parameters()).dtype
device = next(self.parameters()).device
with device:
self.backbone.setup_caches(max_batch_size, dtype)
self.decoder.setup_caches(max_batch_size, dtype, decoder_max_seq_len=self.config.audio_num_codebooks)
self.register_buffer("backbone_causal_mask", _create_causal_mask(self.backbone.max_seq_len, device))
self.register_buffer("decoder_causal_mask", _create_causal_mask(self.config.audio_num_codebooks, device))
def generate_frame(
self,
tokens: torch.Tensor,
tokens_mask: torch.Tensor,
input_pos: torch.Tensor,
temperature: float,
topk: int,
) -> torch.Tensor:
"""
Args:
tokens: (batch_size, seq_len, audio_num_codebooks+1)
tokens_mask: (batch_size, seq_len, audio_num_codebooks+1)
input_pos: (batch_size, seq_len) positions for each token
mask: (batch_size, seq_len, max_seq_len)
Returns:
(batch_size, audio_num_codebooks) sampled tokens
"""
dtype = next(self.parameters()).dtype
b, s, _ = tokens.size()
assert self.backbone.caches_are_enabled(), "backbone caches are not enabled"
curr_backbone_mask = _index_causal_mask(self.backbone_causal_mask, input_pos)
embeds = self._embed_tokens(tokens)
masked_embeds = embeds * tokens_mask.unsqueeze(-1)
h = masked_embeds.sum(dim=2)
h = self.backbone(h, input_pos=input_pos, mask=curr_backbone_mask).to(dtype=dtype)
last_h = h[:, -1, :]
c0_logits = self.codebook0_head(last_h)
c0_sample = sample_topk(c0_logits, topk, temperature)
c0_embed = self._embed_audio(0, c0_sample)
curr_h = torch.cat([last_h.unsqueeze(1), c0_embed], dim=1)
curr_sample = c0_sample.clone()
curr_pos = torch.arange(0, curr_h.size(1), device=curr_h.device).unsqueeze(0).repeat(curr_h.size(0), 1)
# Decoder caches must be reset every frame.
self.decoder.reset_caches()
for i in range(1, self.config.audio_num_codebooks):
curr_decoder_mask = _index_causal_mask(self.decoder_causal_mask, curr_pos)
decoder_h = self.decoder(self.projection(curr_h), input_pos=curr_pos, mask=curr_decoder_mask).to(
dtype=dtype
)
ci_logits = torch.mm(decoder_h[:, -1, :], self.audio_head[i - 1])
ci_sample = sample_topk(ci_logits, topk, temperature)
ci_embed = self._embed_audio(i, ci_sample)
curr_h = ci_embed
curr_sample = torch.cat([curr_sample, ci_sample], dim=1)
curr_pos = curr_pos[:, -1:] + 1
return curr_sample
def reset_caches(self):
self.backbone.reset_caches()
self.decoder.reset_caches()
def _embed_audio(self, codebook: int, tokens: torch.Tensor) -> torch.Tensor:
return self.audio_embeddings(tokens + codebook * self.config.audio_vocab_size)
def _embed_tokens(self, tokens: torch.Tensor) -> torch.Tensor:
text_embeds = self.text_embeddings(tokens[:, :, -1]).unsqueeze(-2)
audio_tokens = tokens[:, :, :-1] + (
self.config.audio_vocab_size * torch.arange(self.config.audio_num_codebooks, device=tokens.device)
)
audio_embeds = self.audio_embeddings(audio_tokens.view(-1)).reshape(
tokens.size(0), tokens.size(1), self.config.audio_num_codebooks, -1
)
return torch.cat([audio_embeds, text_embeds], dim=-2)
+14
View File
@@ -0,0 +1,14 @@
[project]
name = "csm_mw"
description = "ComfyUI node of Conversational Speech Model (CSM)."
version = "1.0.0"
license = {file = "LICENSE"}
[project.urls]
Repository = "https://github.com/billwuhao/ComfyUI_CSM"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "mw"
DisplayName = "ComfyUI_CSM"
Icon = "https://github.com/billwuhao/aiart.website/blob/master/hb.png"
+10
View File
@@ -0,0 +1,10 @@
torch>=2.4.0
torchaudio>=2.4.0
tokenizers>=0.21.0
transformers>=4.49.0
huggingface_hub>=0.28.1
moshi>=0.2.2
torchtune>=0.4.0
torchao>=0.9.0
librosa
silentcipher @ git+https://github.com/SesameAILabs/silentcipher@master