v1.0.0
This commit is contained in:
@@ -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
@@ -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",
|
||||
}
|
||||
@@ -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)
|
||||
@@ -0,0 +1,49 @@
|
||||
[中文](README-CN.md)|[English](README.md)
|
||||
|
||||
# ComfyUI Node for CSM
|
||||
|
||||
CSM(Conversational Speech Model 会话语音模型), 多人会话, 克隆声音, 然后可根据会话中语音的情绪变化等, 生成相应情绪变化的语音的模型. 遗憾的是目前只有英文可用. 该节点, 暂时支持同时 10 人会话.
|
||||
|
||||

|
||||
|
||||
可以将录音节点穿插在其中, 进行多人会话.
|
||||
|
||||

|
||||
|
||||
还支持音频水印检测(自动检测水印) 和 音频添加加密水印.
|
||||
|
||||

|
||||
|
||||
## 📣 更新
|
||||
|
||||
[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 团队的卓越的工作👍.
|
||||
@@ -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.
|
||||
|
||||

|
||||
|
||||
Recording nodes can be interspersed within to create multi-person conversations.
|
||||
|
||||

|
||||
|
||||
It also supports audio watermark detection (automatic watermark detection) and audio adding encrypted watermarks.
|
||||
|
||||

|
||||
|
||||
## 📣 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 👍.
|
||||
|
||||
@@ -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 |
@@ -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)
|
||||
@@ -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"
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user