Files
billwuhao-ComfyUI_DiffRhythm/DiffRhythmNode.py
T
2025-03-25 18:48:41 +08:00

393 lines
12 KiB
Python

import torch
import torchaudio
from einops import rearrange
import sys
import os
import json
from muq import MuQMuLan
current_dir = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, current_dir)
from model import DiT, CFM
from diffrhythm_utils import (
decode_audio,
get_lrc_token,
get_negative_style_prompt,
get_reference_latent,
CNENTokenizer,
)
def load_checkpoint(
model: torch.nn.Module,
ckpt_path: str,
device: torch.device,
use_ema: bool = True
):
model = model.half()
if device == 'mps':
model = model.float()
ckpt_type = ckpt_path.split(".")[-1]
try:
if ckpt_type == "safetensors":
from safetensors.torch import load_file
checkpoint = load_file(ckpt_path)
else:
checkpoint = torch.load(ckpt_path, weights_only=True)
except Exception as e:
raise
try:
if use_ema:
if ckpt_type == "safetensors":
checkpoint = {"ema_model_state_dict": checkpoint}
checkpoint["model_state_dict"] = {
k.replace("ema_model.", ""): v
for k, v in checkpoint["ema_model_state_dict"].items()
if k not in ["initted", "step"]
}
model.load_state_dict(checkpoint["model_state_dict"], strict=False)
else:
if ckpt_type == "safetensors":
checkpoint = {"model_state_dict": checkpoint}
model.load_state_dict(checkpoint["model_state_dict"], strict=False)
except Exception as e:
raise
return model.to(device)
def inference(
cfm_model,
vae_model,
cond,
text,
duration,
style_prompt,
negative_style_prompt,
steps,
cfg_strength,
start_time,
odeint_method,
sway_sampling_coef=None,
chunked=False,
vocal_flag=False,
):
with torch.inference_mode():
generated, _ = cfm_model.sample(
cond=cond,
text=text,
duration=duration,
style_prompt=style_prompt,
negative_style_prompt=negative_style_prompt,
steps=steps,
cfg_strength=cfg_strength,
start_time=start_time,
odeint_method=odeint_method,
vocal_flag=vocal_flag,
sway_sampling_coef=sway_sampling_coef,
)
generated = generated.to(torch.float32)
latent = generated.transpose(1, 2) # [b d t]
output = decode_audio(latent, vae_model, chunked=chunked)
# Rearrange audio batch to a single sequence
output = rearrange(output, "b d n -> d (b n)")
# Peak normalize, clip, convert to int16, and save to file
output = (
output.to(torch.float32)
.div(torch.max(torch.abs(output)))
.clamp(-1, 1)
.mul(32767)
.to(torch.int16)
.cpu()
)
return output
class MultiLinePrompt:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"multi_line_prompt": ("STRING", {
"multiline": True,
"default": ""}),
},
}
CATEGORY = "🎤MW/MW-DiffRhythm"
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("prompt",)
FUNCTION = "promptgen"
def promptgen(self, multi_line_prompt: str):
return (multi_line_prompt.strip(),)
class DiffRhythmRun:
device = "cpu"
if torch.cuda.is_available():
device = "cuda"
elif torch.backends.mps.is_available():
device = "mps"
node_dir = os.path.dirname(os.path.abspath(__file__))
comfy_path = os.path.dirname(os.path.dirname(node_dir))
model_path = os.path.join(comfy_path, "models", "TTS")
models = ["cfm_model.pt", "cfm_full_model.pt"]
model_cache = {"cfm": None, "muq": None, "vae": None}
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": (cls.models, {"default": "cfm_full_model.pt"}),
"style_prompt": ("STRING", {
"multiline": True,
"default": ""}),
},
"optional": {
"lyrics_prompt": ("STRING",),
"style_audio": ("AUDIO", ),
"chunked": ("BOOLEAN", {"default": False, "tooltip": "Whether to use chunked decoding."}),
"unload_model": ("BOOLEAN", {"default": False}),
"odeint_method": (["euler", "midpoint", "rk4","implicit_adams"], {"default": "euler"}),
"steps": ("INT", {"default": 30, "min": 1, "max": 100, "step": 1}),
"cfg": ("INT", {"default": 4, "min": 1, "max": 10, "step": 1}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
},
}
CATEGORY = "🎤MW/MW-DiffRhythm"
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("audio",)
FUNCTION = "diffrhythmgen"
def diffrhythmgen(
self,
model: str,
style_prompt: str,
lyrics_prompt: str = "",
style_audio: str = None,
chunked: bool = False,
odeint_method: str = "euler",
steps: int = 30,
cfg: int = 4,
unload_model: bool = False,
seed: int = 0):
if seed != 0:
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
if model == "cfm_model.pt":
max_frames = 2048
elif model == "cfm_full_model.pt":
max_frames = 6144
cfm, tokenizer, muq, vae = self.prepare_model(model, self.device, unload_model)
lrc_prompt, start_time = get_lrc_token(max_frames, lyrics_prompt, tokenizer, self.device)
vocal_flag = False
if style_audio:
prompt, vocal_flag = self.get_audio_style_prompt(muq, style_audio)
elif style_prompt:
prompt = self.get_text_style_prompt(muq, style_prompt)
else:
raise ValueError("Style prompt or style audio must be provided")
negative_style_prompt = get_negative_style_prompt(self.device)
latent_prompt = get_reference_latent(self.device, max_frames)
sway_sampling_coef = -1 if steps < 32 else None
try:
generated_song = inference(
cfm_model=cfm,
vae_model=vae,
cond=latent_prompt,
text=lrc_prompt,
duration=max_frames,
style_prompt=prompt,
negative_style_prompt=negative_style_prompt,
steps=steps,
chunked=chunked,
cfg_strength=cfg,
sway_sampling_coef=sway_sampling_coef,
start_time=start_time,
vocal_flag=vocal_flag,
odeint_method=odeint_method,
)
except Exception as e:
raise
audio_tensor = generated_song.unsqueeze(0)
if unload_model:
import gc
del cfm
del tokenizer
del vae
del muq
gc.collect()
torch.cuda.empty_cache()
return ({"waveform": audio_tensor, "sample_rate": 44100},)
def get_audio_style_prompt(self, model, audio):
vocal_flag = False
if audio is None:
return None, vocal_flag
mulan = model
waveform = audio["waveform"]
sample_rate = audio["sample_rate"]
# Ensure waveform has correct shape
if len(waveform.shape) == 3: # [1, channels, samples]
waveform = waveform.squeeze(0)
if waveform.shape[0] > 1: # If stereo, convert to mono
waveform = waveform.mean(0, keepdim=True)
if sample_rate != 24000:
waveform = torchaudio.transforms.Resample(sample_rate, 24000)(waveform)
# Calculate audio length (seconds)
audio_len = waveform.shape[-1] / 24000
if audio_len <= 1:
vocal_flag = True
if audio_len > 10:
start_sample = int((audio_len // 2 - 5) * 24000)
end_sample = start_sample + 10 * 24000
wav_segment = waveform[..., start_sample:end_sample]
else:
wav_segment = waveform
wav = wav_segment.to(model.device)
with torch.no_grad():
audio_emb = mulan(wavs = wav) # [1, 512]
audio_emb = audio_emb.half()
return audio_emb, vocal_flag
def get_text_style_prompt(self, model, text_prompt):
if text_prompt is None:
return None
mulan = model
with torch.no_grad():
text_emb = mulan(texts = text_prompt) # [1, 512]
text_emb = text_emb.half()
return text_emb
def prepare_model(self, model, device, unload_model=False):
# prepare tokenizer
try:
tokenizer = CNENTokenizer()
except Exception as e:
raise
if not unload_model and all(self.model_cache.values()):
return (self.model_cache["cfm"],
tokenizer,
self.model_cache["muq"],
self.model_cache["vae"])
from huggingface_hub import snapshot_download
# prepare cfm model
if model == "cfm_full_model.pt":
dit_ckpt_path = f"{self.model_path}/DiffRhythm/cfm_full_model.pt"
dit_config_path = f"{self.model_path}/DiffRhythm/config.json"
if not os.path.exists(dit_ckpt_path):
snapshot_download(repo_id="ASLP-lab/DiffRhythm-full",
local_dir=f"{self.model_path}/DiffRhythm")
elif model == "cfm_model.pt":
dit_ckpt_path = f"{self.model_path}/DiffRhythm/cfm_model.pt"
dit_config_path = f"{self.node_dir}/config/diffrhythm-1b.json"
if not os.path.exists(dit_ckpt_path):
snapshot_download(repo_id="ASLP-lab/DiffRhythm-base",
local_dir=f"{self.model_path}/DiffRhythm")
vae_ckpt_path = f"{self.model_path}/DiffRhythm/vae_model.pt"
if not os.path.exists(vae_ckpt_path):
snapshot_download(repo_id="ASLP-lab/DiffRhythm-vae",
local_dir=f"{self.model_path}/DiffRhythm",
ignore_patterns=["*safetensors"])
try:
with open(dit_config_path, "r", encoding="utf-8") as f:
model_config = json.load(f)
except Exception as e:
raise
dit_model_cls = DiT
if model == "cfm_model.pt":
cfm = CFM(
transformer=dit_model_cls(**model_config["model"], use_style_prompt=True, max_pos=2048),
num_channels=model_config["model"]["mel_dim"],
)
elif model == "cfm_full_model.pt":
cfm = CFM(
transformer=dit_model_cls(**model_config["model"], use_style_prompt=True, max_pos=6144),
num_channels=model_config["model"]['mel_dim'],
use_style_prompt=True
)
cfm = cfm.to(device)
try:
cfm = load_checkpoint(cfm, dit_ckpt_path, device=device, use_ema=False)
except Exception as e:
raise
# prepare muq model
try:
muq = MuQMuLan.from_pretrained("OpenMuQ/MuQ-MuLan-large", cache_dir=f"{self.model_path}/DiffRhythm")
except Exception as e:
raise
muq = muq.to(device).eval()
# prepare vae
try:
vae = torch.jit.load(vae_ckpt_path, map_location="cpu").to(device)
except Exception as e:
raise
self.model_cache["cfm"] = cfm
self.model_cache["muq"] = muq
self.model_cache["vae"] = vae
return cfm, tokenizer, muq, vae
from MWAudioRecorderDR import AudioRecorderDR
NODE_CLASS_MAPPINGS = {
"DiffRhythmRun": DiffRhythmRun,
"MultiLinePrompt": MultiLinePrompt,
"AudioRecorderDR": AudioRecorderDR
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DiffRhythmRun": "DiffRhythm Run",
"MultiLinePrompt": "Multi Line Prompt",
"AudioRecorderDR": "MW Audio Recorder"
}