This commit is contained in:
billwuhao
2025-05-14 18:38:07 +08:00
parent f5abbf60e7
commit 95ca7b438d
83 changed files with 2768 additions and 1688 deletions
+207 -316
View File
@@ -1,141 +1,25 @@
import os
import time
import random
import torch
import torchaudio
from einops import rearrange
import sys
import os
import json
from easydict import EasyDict
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 (
from infer_utils import (
decode_audio,
get_lrc_token,
get_negative_style_prompt,
get_reference_latent,
CNENTokenizer,
get_audio_style_prompt,
get_text_style_prompt,
prepare_model,
eval_song,
)
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(),)
import folder_paths
models_dir = folder_paths.models_dir
model_path = os.path.join(models_dir, "TTS")
models = ["cfm_model.pt", "cfm_full_model.pt"]
def set_all_seeds(seed):
# import random
# import numpy as np
@@ -153,6 +37,123 @@ def set_all_seeds(seed):
# torch.backends.cudnn.benchmark = False # 关闭优化(牺牲速度换取确定性)
import folder_paths
cache_dir = folder_paths.get_temp_directory()
import tempfile
from typing import Optional
def cache_audio_tensor(
cache_dir,
audio_tensor: torch.Tensor,
sample_rate: int,
filename_prefix: str = "cached_audio_",
audio_format: Optional[str] = ".wav"
) -> str:
try:
with tempfile.NamedTemporaryFile(
prefix=filename_prefix,
suffix=audio_format,
dir=cache_dir,
delete=False
) as tmp_file:
temp_filepath = tmp_file.name
torchaudio.save(temp_filepath, audio_tensor, sample_rate)
return temp_filepath
except Exception as e:
raise Exception(f"Error caching audio tensor: {e}")
def inference(
cfm_model,
vae_model,
eval_model,
eval_muq,
cond,
text,
duration,
style_prompt,
negative_style_prompt,
steps,
cfg_strength,
sway_sampling_coef,
start_time,
# file_type,
vocal_flag,
odeint_method,
pred_frames,
batch_infer_num,
chunked=True,
):
with torch.inference_mode():
latents, _ = 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,
sway_sampling_coef=sway_sampling_coef,
start_time=start_time,
vocal_flag=vocal_flag,
odeint_method=odeint_method,
latent_pred_segments=pred_frames,
batch_infer_num=batch_infer_num
)
outputs = []
for latent in latents:
latent = latent.to(torch.float32)
latent = latent.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)")
outputs.append(output)
if batch_infer_num > 1:
generated_song = eval_song(eval_model, eval_muq, outputs)
else:
generated_song = outputs[0]
output_tensor = generated_song.to(torch.float32).div(torch.max(torch.abs(output))).clamp(-1, 1).cpu()
return output_tensor
node_dir = os.path.dirname(os.path.abspath(__file__))
folder = f'{node_dir}/diffrhythm/example'
files = [f for f in os.listdir(folder) if os.path.isfile(os.path.join(folder, f))]
selected = random.choice(files)
with open(os.path.join(folder, selected), 'r', encoding='utf-8') as f:
lyrics = f.read()
class MultiLineLyricsDR:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"lyrics": ("STRING", {
"multiline": True,
"default": lyrics}),
},
}
CATEGORY = "🎤MW/MW-DiffRhythm"
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("lyrics",)
FUNCTION = "lyricsgen"
def lyricsgen(self, lyrics: str):
return (lyrics.strip(),)
class DiffRhythmRun:
def __init__(self):
device = "cpu"
@@ -161,29 +162,35 @@ class DiffRhythmRun:
elif torch.backends.mps.is_available():
device = "mps"
self.device = device
self.cfm = None
self.vae = None
self.muq = None
self.tokenizer = None
self.eval_model = None
self.eval_muq = None
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": (models, {"default": "cfm_full_model.pt"}),
"model": (["cfm_model_v1_2.pt", "cfm_model.pt", "cfm_full_model.pt"], {"default": "cfm_model_v1_2.pt"}),
"style_prompt": ("STRING", {
"multiline": True,
"default": ""}),
"default": "Indie folk ballad, coming-of-age themes, acoustic guitar picking with harmonica interludes"}),
},
"optional": {
"lyrics_prompt": ("STRING", {"forceInput": True}),
"style_audio": ("AUDIO", ),
"chunked": ("BOOLEAN", {"default": False, "tooltip": "Whether to use chunked decoding."}),
"lyrics_or_edit_lyrics": ("STRING", {"forceInput": True}),
"style_audio_or_edit_song": ("AUDIO", ),
# "chunked": ("BOOLEAN", {"default": False, "tooltip": "Whether to use chunked decoding."}),
"unload_model": ("BOOLEAN", {"default": True}),
"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}),
"quality_or_speed":(["quality", "speed"], {"default": "speed"}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
"edit": ("BOOLEAN", {"default": False}),
"edit_segments": ("STRING", {"default":"[-1, 20], [60, -1]", "multiline": True}),
},
}
@@ -194,63 +201,94 @@ class DiffRhythmRun:
def diffrhythmgen(
self,
edit,
model: str,
style_prompt: str,
lyrics_prompt: str = "",
style_audio: str = None,
chunked: bool = False,
style_prompt: str = None,
lyrics_or_edit_lyrics: str = "",
style_audio_or_edit_song = None,
edit_segments: str = "",
chunked: bool = True,
odeint_method: str = "euler",
steps: int = 30,
cfg: int = 4,
quality_or_speed: str = "speed",
unload_model: bool = False,
seed: int = 0):
if seed != 0:
set_all_seeds(seed)
if model == "cfm_model.pt":
if model == "cfm_model.pt" or model == "cfm_model_v1_2.pt":
max_frames = 2048
elif model == "cfm_full_model.pt":
else:
max_frames = 6144
if self.cfm is None:
self.cfm, self.tokenizer, self.muq, self.vae = self.prepare_model(model, self.device)
self.cfm, self.tokenizer, self.muq, self.vae, self.eval_model, self.eval_muq = prepare_model(max_frames, self.device, model)
lrc_prompt, start_time = get_lrc_token(max_frames, lyrics_prompt, self.tokenizer, self.device)
batch_infer_num = 1 if quality_or_speed == "speed" else 5
lyrics = lyrics_or_edit_lyrics.strip()
vocal_flag = False
if style_audio:
prompt, vocal_flag = self.get_audio_style_prompt(self.muq, style_audio)
elif style_prompt:
prompt = self.get_text_style_prompt(self.muq, style_prompt)
if style_audio_or_edit_song is not None:
style_audio_path = cache_audio_tensor(cache_dir,
style_audio_or_edit_song["waveform"].squeeze(0),
style_audio_or_edit_song["sample_rate"],
filename_prefix="style_audio_")
prompt, vocal_flag = get_audio_style_prompt(self.muq, style_audio_path)
print("Provided style_audio, style_prompt will be ineffective")
else:
raise ValueError("Style prompt or style audio must be provided")
assert style_prompt.strip(), "One of style_audio and style_prompt must be provided"
prompt = get_text_style_prompt(self.muq, style_prompt)
edit_song_path = None
if edit:
if style_audio_or_edit_song is not None:
edit_song_path = style_audio_path
prompt, vocal_flag = get_audio_style_prompt(self.muq, edit_song_path)
assert edit_song_path and lyrics and edit_segments.strip(), "edit song, edit lyrics, edit segments must be provided"
edit_segments = "["+edit_segments+"]"
else:
edit_segments = None
lrc_prompt, start_time = get_lrc_token(max_frames, lyrics.strip(), self.tokenizer, self.device)
negative_style_prompt = get_negative_style_prompt(self.device)
latent_prompt = get_reference_latent(self.device, max_frames)
latent_prompt, pred_frames = get_reference_latent(self.device,
max_frames,
edit,
pred_segments=edit_segments,
ref_song=edit_song_path,
vae_model=self.vae)
sway_sampling_coef = -1 if steps < 32 else None
try:
generated_song = inference(
cfm_model=self.cfm,
vae_model=self.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)
s_t = time.time()
generated_songs = inference(
cfm_model=self.cfm,
vae_model=self.vae,
eval_model=self.eval_model,
eval_muq=self.eval_muq,
odeint_method=odeint_method,
vocal_flag=vocal_flag,
sway_sampling_coef=sway_sampling_coef,
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,
start_time=start_time,
pred_frames=pred_frames,
batch_infer_num=batch_infer_num
)
e_t = time.time() - s_t
print(f"inference cost {e_t:.2f} seconds")
audio_tensor = generated_songs[0].unsqueeze(0).unsqueeze(0)
if unload_model:
import gc
@@ -258,167 +296,20 @@ class DiffRhythmRun:
self.muq = None
self.vae = None
self.tokenizer = None
self.eval_model = None
self.eval_muq = None
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):
# prepare tokenizer
try:
tokenizer = CNENTokenizer()
except Exception as e:
raise
from huggingface_hub import snapshot_download
# prepare cfm model
if model == "cfm_full_model.pt":
dit_ckpt_path = f"{model_path}/DiffRhythm/cfm_full_model.pt"
dit_config_path = f"{model_path}/DiffRhythm/config.json"
if not os.path.exists(dit_ckpt_path):
snapshot_download(repo_id="ASLP-lab/DiffRhythm-full",
local_dir=f"{model_path}/DiffRhythm")
elif model == "cfm_model.pt":
dit_ckpt_path = f"{model_path}/DiffRhythm/cfm_model.pt"
dit_config_path = f"{model_path}/DiffRhythm/config.json"
if not os.path.exists(dit_ckpt_path):
snapshot_download(repo_id="ASLP-lab/DiffRhythm-base",
local_dir=f"{model_path}/DiffRhythm")
vae_ckpt_path = f"{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"{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:
main_model_dir = f"{model_path}/DiffRhythm/MuQ-MuLan-large"
local_audio_model_dir = f"{model_path}/DiffRhythm/MuQ-large-msd-iter"
local_text_model_dir = f"{model_path}/DiffRhythm/xlm-roberta-base"
config_path = f"{main_model_dir}/config.json"
with open(config_path, 'r') as f:
config_dict = json.load(f)
config_dict['audio_model']['name'] = local_audio_model_dir
config_dict['text_model']['name'] = local_text_model_dir
config_obj = EasyDict(config_dict)
muq = MuQMuLan(config=config_obj, hf_hub_cache_dir=None)
weights_path = f"{main_model_dir}/pytorch_model.bin"
try:
state_dict = torch.load(weights_path, map_location='cpu')
# Adjust loading based on how weights are saved (e.g., remove prefixes if needed)
muq.load_state_dict(state_dict, strict=False) # Use strict=False initially
except FileNotFoundError:
raise FileNotFoundError(f"Weights file not found at {weights_path}")
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
return (cfm, tokenizer, muq, vae)
from MWAudioRecorderDR import AudioRecorderDR
NODE_CLASS_MAPPINGS = {
"DiffRhythmRun": DiffRhythmRun,
"MultiLinePrompt": MultiLinePrompt,
"AudioRecorderDR": AudioRecorderDR
"MultiLineLyricsDR": MultiLineLyricsDR
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DiffRhythmRun": "DiffRhythm Run",
"MultiLinePrompt": "Multi Line Prompt",
"AudioRecorderDR": "MW Audio Recorder"
"MultiLineLyricsDR": "MultiLine Lyrics"
}
-141
View File
@@ -1,141 +0,0 @@
import numpy as np
import torch
import time
import librosa
import sounddevice as sd
from scipy import ndimage
from comfy.utils import ProgressBar
class AudioRecorderDR:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
# Recording duration
"record_sec": ("INT", {
"default": 5,
"min": 1,
"step": 1
}),
"sample_rate": (["16000", "44100", "48000"], {
"default": "48000"
}),
"n_fft": ("INT", {
"default": 2048,
"min": 512,
"max": 4096,
"step": 512
}),
"sensitivity": ("FLOAT", {
"default": 1.2,
"min": 0.1,
"max": 3.0,
"step": 0.1
}),
"smooth": ("INT", {
"default": 1,
"min": 5,
"max": 7,
"step": 2
}),
"seed": ("INT", {
"default": 0,
"min": 0,
"max": 0xFFFFFFFFFFFFFFFF
}),
},
"optional": { # 可选参数
"enable": ("BOOLEAN", {"default": True}),
},
}
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("audio",)
FUNCTION = "record_and_clean"
CATEGORY = "🎤MW/MW-DiffRhythm"
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) # Increase the mask value for smoother edges
def record_and_clean(
self,
trigger: bool,
record_sec: int,
n_fft: int,
sensitivity: float,
smooth: int,
sample_rate: str,
seed: int
):
if not trigger:
return (None,)
sr = int(sample_rate)
final_audio = None
try:
noise_clip = None
# Main recording
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()
# Auto noise detection
if noise_clip is None:
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]
# Noise reduction
noise_profile = self._calc_noise_profile(noise_clip, n_fft)
spec = self._stft(audio, n_fft)
# Multi-step processing
mask = np.ones_like(spec) # Initial mask
for _ in range(2): # Dual processing loop
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
# Phase reconstruction
processed = self._istft(spec * mask, n_fft)
# Dynamic gain normalization
peak = np.max(np.abs(processed))
processed = processed * (0.99 / peak) if peak > 0 else processed
# Format conversion
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,)
+20 -10
View File
@@ -4,26 +4,29 @@
快速而简单的端到端全长歌曲生成.
![](https://github.com/billwuhao/ComfyUI_DiffRhythm/blob/master/images/2025-03-12_23-49-32.png)
![](https://github.com/billwuhao/ComfyUI_DiffRhythm/blob/master/images/2025-05-13_01-51-00.png)
## 📣 更新
[2025-05-13]⚒️: 支持 DiffRhythm v1.2 版本, 质量更好, 可编辑歌词. 目前发布 95 秒长度歌曲模型, 全长歌曲发布将即时更新. **注意**: 版本代码更新, 之前的模型生成质量可能会受到影响. 如果尝试之前的版本, 请退回到 v2.2.0 之前版本.
[2025-04-26]⚒️: 改为手动选择下载 muq 模型.
[2025-03-21]⚒️: 代码重构, 超快生成速度, 4分45秒音乐, 20秒不到生成, 1分35秒音乐, 7秒不到生成. 增加更多可调参数, 畅玩更自由. 可选是否卸载模型.
[2025-03-16]⚒️: 发布版本 v2.0.0. 支持全长音乐生成, 4 分钟仅需 62 秒.
![](https://github.com/billwuhao/ComfyUI_DiffRhythm/blob/master/images/2025-03-16_03-53-48.png)
下载模型放到 `ComfyUI\models\TTS\DiffRhythm` 文件夹下:
- [DiffRhythm-full](https://huggingface.co/ASLP-lab/DiffRhythm-full) 模型重命名为 `cfm_full_model.pt`, 同时下载 comfig.json 放到一起.
- [DiffRhythm-full](https://huggingface.co/ASLP-lab/DiffRhythm-full) 模型重命名为 `cfm_full_model.pt`.
[2025-03-13]⚒️: 发布版本 v1.0.0.
- 所有参数均是可选的, 不提供任何参数随机生成音乐.
## 使用
- 自动生成歌曲, 自动添加双语歌词字幕:
![](https://github.com/billwuhao/ComfyUI_DiffRhythm/blob/master/images/2025-05-14_16-33-54.png)
## 安装
@@ -43,14 +46,19 @@ pip install -r requirements.txt
结构如下:
![](https://github.com/billwuhao/ComfyUI_DiffRhythm/blob/master/images/2025-03-13_00-08-51.png)
![](https://github.com/billwuhao/ComfyUI_DiffRhythm/blob/master/images/2025-05-13_01-54-13.png)
```
.
| cfm_model_v1_2.pt
│ cfm_full_model.pt
│ cfm_model.pt
│ config.json
│ vae_model.pt
|
├─eval-model
│ eval.yaml
│ eval.safetensors
│
├─MuQ-large-msd-iter
│ config.json
@@ -69,11 +77,13 @@ pip install -r requirements.txt
```
手动下载地址:
https://huggingface.co/ASLP-lab/DiffRhythm-1_2/blob/main/cfm_model.pt 重命名: `cfm_model_v1_2.pt`
https://huggingface.co/spaces/ASLP-lab/DiffRhythm/tree/main/pretrained
https://huggingface.co/ASLP-lab/DiffRhythm-full/tree/main
https://huggingface.co/ASLP-lab/DiffRhythm-base/blob/main/cfm_model.pt
https://huggingface.co/ASLP-lab/DiffRhythm-vae/blob/main/vae_model.pt
https://huggingface.co/OpenMuQ/MuQ-MuLan-large/tree/main
https://huggingface.co/OpenMuQ/MuQ-large-msd-iter/tree/main → `.safetensors`: (https://huggingface.co/OpenMuQ/MuQ-large-msd-iter/blob/refs%2Fpr%2F1/model.safetensors)
https://huggingface.co/OpenMuQ/MuQ-large-msd-iter/tree/main 要下载 `.safetensors` 格式: (https://huggingface.co/OpenMuQ/MuQ-large-msd-iter/blob/refs%2Fpr%2F1/model.safetensors)
https://huggingface.co/FacebookAI/xlm-roberta-base/tree/main
## 环境配置
@@ -82,7 +92,7 @@ Windows 系统做如下配置.
下载安装最新版 [espeak-ng](https://github.com/espeak-ng/espeak-ng/releases/tag/1.52.0)
添加环境变量 `PHONEMIZER_ESPEAK_LIBRARY` 到系统中, 值是你安装的 espeak-ng 软件中 `libespeak-ng.dll` 文件的路径, 例如: `C:\Program Files\eSpeak NG\libespeak-ng.dll`.
添加系统环境变量 `PHONEMIZER_ESPEAK_LIBRARY`, 值是你安装的 espeak-ng 软件中 `libespeak-ng.dll` 文件的路径, 例如: `C:\Program Files\eSpeak NG\libespeak-ng.dll`.
Linux 系统下, 需要安装 `espeak-ng` 软件包. 执行如下命令安装:
@@ -96,4 +106,4 @@ Linux 系统下, 需要安装 `espeak-ng` 软件包. 执行如下命令安装:
[DiffRhythm](https://github.com/ASLP-lab/DiffRhythm)
感谢 DiffRhythm 团队的卓越的工作, 目前最强开源 音乐/歌曲 生成模型👍.
感谢 DiffRhythm 团队的卓越的工作👍.
+34 -21
View File
@@ -1,28 +1,32 @@
[中文](README-CN.md) | [English](README.md)
[中文](README-CN.md) | [English](README.md)
# DiffRhythm Node for ComfyUI
# DiffRhythm Nodes for ComfyUI
Blazingly Fast and Embarrassingly Simple End-to-End Full-Length Song Generation.
Fast and easy end-to-end full-length song generation.
![](https://github.com/billwuhao/ComfyUI_DiffRhythm/blob/master/images/2025-03-12_23-49-32.png)
![](https://github.com/billwuhao/ComfyUI_DiffRhythm/blob/master/images/2025-05-13_01-51-00.png)
## 📣 update
## 📣 Updates
[2025-04-26]⚒️: Change to manually selecting to download the `muq` model.
[2025-05-13]⚒️: Supports DiffRhythm v1.2, better quality, editable lyrics. Currently released a 95-second song model, full-length song release will be updated promptly. **Note**: The version code has been updated, and the generation quality of previous models may be affected. If you want to try the previous version, please revert to the version before v2.2.0.
[2025-03-21] ⚒️: Code refactored, ultra-fast generation speed. 4 minutes 45 seconds of music generated in less than 20 seconds, 1 minute 35 seconds of music generated in less than 7 seconds. Added more tunable parameters for more creative freedom. Optional model unloading.
[2025-04-26]⚒️: Changed to manually download the muq model.
[2025-03-21]⚒️: Code refactoring, super fast generation speed, 4 minutes 45 seconds of music generated in less than 20 seconds, 1 minute 35 seconds of music generated in less than 7 seconds. Added more adjustable parameters for more freedom. Option to uninstall the model.
[2025-03-16]⚒️: Released version v2.0.0. Supports full-length music generation, 4 minutes only takes 62 seconds.
![](https://github.com/billwuhao/ComfyUI_DiffRhythm/blob/master/images/2025-03-16_03-53-48.png)
Download the model and place it in the `ComfyUI\models\TTS\DiffRhythm` folder:
- [DiffRhythm-full](https://huggingface.co/ASLP-lab/DiffRhythm-full) Rename the model to `cfm_full_model.pt`, and also download `comfig.json` and put it together.
- [DiffRhythm-full](https://huggingface.co/ASLP-lab/DiffRhythm-full) rename the model to `cfm_full_model.pt`.
[2025-03-13]⚒️: Release version v1.0.0.
[2025-03-13]⚒️: Released version v1.0.0.
- All parameters are optional; you can generate random music without providing any parameters.
## Usage
- Automatically generate song and add bilingual lyrics subtitles:
![](https://github.com/billwuhao/ComfyUI_DiffRhythm/blob/master/images/2025-05-14_16-33-54.png)
## Installation
@@ -42,14 +46,19 @@ The model needs to be manually downloaded to the `ComfyUI\models\TTS\DiffRhythm`
The structure is as follows:
![](https://github.com/billwuhao/ComfyUI_DiffRhythm/blob/master/images/2025-03-13_00-08-51.png)
![](https://github.com/billwuhao/ComfyUI_DiffRhythm/blob/master/images/2025-05-13_01-54-13.png)
```
.
| cfm_model_v1_2.pt
│ cfm_full_model.pt
│ cfm_model.pt
│ config.json
│ vae_model.pt
|
├─eval-model
│ eval.yaml
│ eval.safetensors
│
├─MuQ-large-msd-iter
│ config.json
@@ -67,31 +76,35 @@ The structure is as follows:
tokenizer_config.json
```
https://huggingface.co/ASLP-lab/DiffRhythm-full/tree/main → `cfm_full_model.pt` and `config.json`
Manual download links:
https://huggingface.co/ASLP-lab/DiffRhythm-1_2/blob/main/cfm_model.pt → `cfm_model_v1_2.pt`
https://huggingface.co/spaces/ASLP-lab/DiffRhythm/tree/main/pretrained
https://huggingface.co/ASLP-lab/DiffRhythm-full/tree/main
https://huggingface.co/ASLP-lab/DiffRhythm-base/blob/main/cfm_model.pt
https://huggingface.co/ASLP-lab/DiffRhythm-vae/blob/main/vae_model.pt
https://huggingface.co/OpenMuQ/MuQ-MuLan-large/tree/main
https://huggingface.co/OpenMuQ/MuQ-large-msd-iter/tree/main → `.safetensors`: (https://huggingface.co/OpenMuQ/MuQ-large-msd-iter/blob/refs%2Fpr%2F1/model.safetensors)
https://huggingface.co/OpenMuQ/MuQ-large-msd-iter/tree/main → `.safetensors`: (https://huggingface.co/OpenMuQ/MuQ-large-msd-iter/blob/refs%2Fpr%2F1/model.safetensors)
https://huggingface.co/FacebookAI/xlm-roberta-base/tree/main
## Environment Configuration
- Configure the following on Windows systems:
For Windows systems, configure as follows:
Download and install the latest version of [espeak-ng](https://github.com/espeak-ng/espeak-ng/releases/tag/1.52.0)
Add the environment variable `PHONEMIZER_ESPEAK_LIBRARY` to your system. The value should be the path to the `libespeak-ng.dll` file in your espeak-ng installation, for example: `C:\Program Files\eSpeak NG\libespeak-ng.dll`.
Add the system environment variable `PHONEMIZER_ESPEAK_LIBRARY`, the value is the path to the `libespeak-ng.dll` file in your espeak-ng installation, for example: `C:\Program Files\eSpeak NG\libespeak-ng.dll`.
- On Linux systems, you need to install the `espeak-ng` package. Execute the following command to install:
For Linux systems, you need to install the `espeak-ng` package. Execute the following command to install:
`apt-get -qq -y install espeak-ng`
It should support Mac, but has not been tested.
Mac is supported, but untested.
Enjoy the music! 🎶
Enjoy the music🎶
## Acknowledgements
[DiffRhythm](https://github.com/ASLP-lab/DiffRhythm)
Thanks to the DiffRhythm team for their excellent work. Currently the strongest open-source music/song generation model 👍.
Thanks to the DiffRhythm team for their excellent work👍.
File diff suppressed because it is too large Load Diff
+9
View File
@@ -0,0 +1,9 @@
from diffrhythm.LangSegment.LangSegment import LangSegment,getTexts,classify,getCounts,printList,setfilters,getfilters,setPriorityThreshold,getPriorityThreshold,setEnablePreview,getEnablePreview,setKeepPinyin,getKeepPinyin,setLangMerge,getLangMerge
# release
__version__ = '0.3.5'
# develop
__develop__ = 'dev-0.0.1'
+327
View File
@@ -0,0 +1,327 @@
# Copyright (c) 2021 PaddlePaddle Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# Digital processing from GPT_SoVITS num.py (thanks)
"""
Rules to verbalize numbers into Chinese characters.
https://zh.wikipedia.org/wiki/中文数字#現代中文
"""
import re
from collections import OrderedDict
from typing import List
DIGITS = {str(i): tran for i, tran in enumerate('零一二三四五六七八九')}
UNITS = OrderedDict({
1: '十',
2: '百',
3: '千',
4: '万',
8: '亿',
})
COM_QUANTIFIERS = '(处|台|架|枚|趟|幅|平|方|堵|间|床|株|批|项|例|列|篇|栋|注|亩|封|艘|把|目|套|段|人|所|朵|匹|张|座|回|场|尾|条|个|首|阙|阵|网|炮|顶|丘|棵|只|支|袭|辆|挑|担|颗|壳|窠|曲|墙|群|腔|砣|座|客|贯|扎|捆|刀|令|打|手|罗|坡|山|岭|江|溪|钟|队|单|双|对|出|口|头|脚|板|跳|枝|件|贴|针|线|管|名|位|身|堂|课|本|页|家|户|层|丝|毫|厘|分|钱|两|斤|担|铢|石|钧|锱|忽|(千|毫|微)克|毫|厘|(公)分|分|寸|尺|丈|里|寻|常|铺|程|(千|分|厘|毫|微)米|米|撮|勺|合|升|斗|石|盘|碗|碟|叠|桶|笼|盆|盒|杯|钟|斛|锅|簋|篮|盘|桶|罐|瓶|壶|卮|盏|箩|箱|煲|啖|袋|钵|年|月|日|季|刻|时|周|天|秒|分|小时|旬|纪|岁|世|更|夜|春|夏|秋|冬|代|伏|辈|丸|泡|粒|颗|幢|堆|条|根|支|道|面|片|张|颗|块|元|(亿|千万|百万|万|千|百)|(亿|千万|百万|万|千|百|美|)元|(亿|千万|百万|万|千|百|十|)吨|(亿|千万|百万|万|千|百|)块|角|毛|分)'
# 分数表达式
RE_FRAC = re.compile(r'(-?)(\d+)/(\d+)')
def replace_frac(match) -> str:
"""
Args:
match (re.Match)
Returns:
str
"""
sign = match.group(1)
nominator = match.group(2)
denominator = match.group(3)
sign: str = "负" if sign else ""
nominator: str = num2str(nominator)
denominator: str = num2str(denominator)
result = f"{sign}{denominator}分之{nominator}"
return result
# 百分数表达式
RE_PERCENTAGE = re.compile(r'(-?)(\d+(\.\d+)?)%')
def replace_percentage(match) -> str:
"""
Args:
match (re.Match)
Returns:
str
"""
sign = match.group(1)
percent = match.group(2)
sign: str = "负" if sign else ""
percent: str = num2str(percent)
result = f"{sign}百分之{percent}"
return result
# 整数表达式
# 带负号的整数 -10
RE_INTEGER = re.compile(r'(-)' r'(\d+)')
def replace_negative_num(match) -> str:
"""
Args:
match (re.Match)
Returns:
str
"""
sign = match.group(1)
number = match.group(2)
sign: str = "负" if sign else ""
number: str = num2str(number)
result = f"{sign}{number}"
return result
# 编号-无符号整形
# 00078
RE_DEFAULT_NUM = re.compile(r'\d{3}\d*')
def replace_default_num(match):
"""
Args:
match (re.Match)
Returns:
str
"""
number = match.group(0)
return verbalize_digit(number, alt_one=True)
# 加减乘除
# RE_ASMD = re.compile(
# r'((-?)((\d+)(\.\d+)?)|(\.(\d+)))([\+\-\×÷=])((-?)((\d+)(\.\d+)?)|(\.(\d+)))')
RE_ASMD = re.compile(
r'((-?)((\d+)(\.\d+)?[⁰¹²³⁴⁵⁶⁷⁸⁹ˣʸⁿ]*)|(\.\d+[⁰¹²³⁴⁵⁶⁷⁸⁹ˣʸⁿ]*)|([A-Za-z][⁰¹²³⁴⁵⁶⁷⁸⁹ˣʸⁿ]*))([\+\-\×÷=])((-?)((\d+)(\.\d+)?[⁰¹²³⁴⁵⁶⁷⁸⁹ˣʸⁿ]*)|(\.\d+[⁰¹²³⁴⁵⁶⁷⁸⁹ˣʸⁿ]*)|([A-Za-z][⁰¹²³⁴⁵⁶⁷⁸⁹ˣʸⁿ]*))')
asmd_map = {
'+': '加',
'-': '减',
'×': '乘',
'÷': '除',
'=': '等于'
}
def replace_asmd(match) -> str:
"""
Args:
match (re.Match)
Returns:
str
"""
result = match.group(1) + asmd_map[match.group(8)] + match.group(9)
return result
# 次方专项
RE_POWER = re.compile(r'[⁰¹²³⁴⁵⁶⁷⁸⁹ˣʸⁿ]+')
power_map = {
'⁰': '0',
'¹': '1',
'²': '2',
'³': '3',
'⁴': '4',
'⁵': '5',
'⁶': '6',
'⁷': '7',
'⁸': '8',
'⁹': '9',
'ˣ': 'x',
'ʸ': 'y',
'ⁿ': 'n'
}
def replace_power(match) -> str:
"""
Args:
match (re.Match)
Returns:
str
"""
power_num = ""
for m in match.group(0):
power_num += power_map[m]
result = "的" + power_num + "次方"
return result
# 数字表达式
# 纯小数
RE_DECIMAL_NUM = re.compile(r'(-?)((\d+)(\.\d+))' r'|(\.(\d+))')
# 正整数 + 量词
RE_POSITIVE_QUANTIFIERS = re.compile(r"(\d+)([多余几\+])?" + COM_QUANTIFIERS)
RE_NUMBER = re.compile(r'(-?)((\d+)(\.\d+)?)' r'|(\.(\d+))')
def replace_positive_quantifier(match) -> str:
"""
Args:
match (re.Match)
Returns:
str
"""
number = match.group(1)
match_2 = match.group(2)
if match_2 == "+":
match_2 = "多"
match_2: str = match_2 if match_2 else ""
quantifiers: str = match.group(3)
number: str = num2str(number)
result = f"{number}{match_2}{quantifiers}"
return result
def replace_number(match) -> str:
"""
Args:
match (re.Match)
Returns:
str
"""
sign = match.group(1)
number = match.group(2)
pure_decimal = match.group(5)
if pure_decimal:
result = num2str(pure_decimal)
else:
sign: str = "负" if sign else ""
number: str = num2str(number)
result = f"{sign}{number}"
return result
# 范围表达式
# match.group(1) and match.group(8) are copy from RE_NUMBER
RE_RANGE = re.compile(
r"""
(?<![\d\+\-\×÷=]) # 使用反向前瞻以确保数字范围之前没有其他数字和操作符
((-?)((\d+)(\.\d+)?)) # 匹配范围起始的负数或正数(整数或小数)
[-~] # 匹配范围分隔符
((-?)((\d+)(\.\d+)?)) # 匹配范围结束的负数或正数(整数或小数)
(?![\d\+\-\×÷=]) # 使用正向前瞻以确保数字范围之后没有其他数字和操作符
""", re.VERBOSE)
def replace_range(match) -> str:
"""
Args:
match (re.Match)
Returns:
str
"""
first, second = match.group(1), match.group(6)
first = RE_NUMBER.sub(replace_number, first)
second = RE_NUMBER.sub(replace_number, second)
result = f"{first}到{second}"
return result
# ~至表达式
RE_TO_RANGE = re.compile(
r'((-?)((\d+)(\.\d+)?)|(\.(\d+)))(%|°C|℃|度|摄氏度|cm2|cm²|cm3|cm³|cm|db|ds|kg|km|m2|m²|m³|m3|ml|m|mm|s)[~]((-?)((\d+)(\.\d+)?)|(\.(\d+)))(%|°C|℃|度|摄氏度|cm2|cm²|cm3|cm³|cm|db|ds|kg|km|m2|m²|m³|m3|ml|m|mm|s)')
def replace_to_range(match) -> str:
"""
Args:
match (re.Match)
Returns:
str
"""
result = match.group(0).replace('~', '至')
return result
def _get_value(value_string: str, use_zero: bool=True) -> List[str]:
stripped = value_string.lstrip('0')
if len(stripped) == 0:
return []
elif len(stripped) == 1:
if use_zero and len(stripped) < len(value_string):
return [DIGITS['0'], DIGITS[stripped]]
else:
return [DIGITS[stripped]]
else:
largest_unit = next(
power for power in reversed(UNITS.keys()) if power < len(stripped))
first_part = value_string[:-largest_unit]
second_part = value_string[-largest_unit:]
return _get_value(first_part) + [UNITS[largest_unit]] + _get_value(
second_part)
def verbalize_cardinal(value_string: str) -> str:
if not value_string:
return ''
# 000 -> '零' , 0 -> '零'
value_string = value_string.lstrip('0')
if len(value_string) == 0:
return DIGITS['0']
result_symbols = _get_value(value_string)
# verbalized number starting with '一十*' is abbreviated as `十*`
if len(result_symbols) >= 2 and result_symbols[0] == DIGITS[
'1'] and result_symbols[1] == UNITS[1]:
result_symbols = result_symbols[1:]
return ''.join(result_symbols)
def verbalize_digit(value_string: str, alt_one=False) -> str:
result_symbols = [DIGITS[digit] for digit in value_string]
result = ''.join(result_symbols)
if alt_one:
result = result.replace("一", "幺")
return result
def num2str(value_string: str) -> str:
integer_decimal = value_string.split('.')
if len(integer_decimal) == 1:
integer = integer_decimal[0]
decimal = ''
elif len(integer_decimal) == 2:
integer, decimal = integer_decimal
else:
raise ValueError(
f"The value string: '${value_string}' has more than one point in it."
)
result = verbalize_cardinal(integer)
decimal = decimal.rstrip('0')
if decimal:
# '.22' is verbalized as '零点二二'
# '3.20' is verbalized as '三点二
result = result if result else "零"
result += '点' + verbalize_digit(decimal)
return result
if __name__ == "__main__":
text = ""
text = num2str(text)
print(text)
pass
@@ -9,7 +9,7 @@ mixed_precision: fp16
num_machines: 1
rdzv_backend: static
same_network: true
num_processes: 8
num_processes: 1
tpu_env: []
tpu_use_cluster: false
tpu_use_sudo: false
@@ -1,12 +1,12 @@
# Copyright (c) 2025 ASLP-LAB
# 2025 Huakang Chen (huakang@mail.nwpu.edu.cn)
#
# Licensed under the Stability AI License (the "License");
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://huggingface.co/stabilityai/stable-audio-open-1.0/blob/main/LICENSE.md
#
# http://www.apache.org/licenses/LICENSE-2.0
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
@@ -80,7 +80,8 @@ class DiffusionDataset(torch.utils.data.Dataset):
else:
raise
lrc_with_time = lrc_with_time[:-1] if len(lrc_with_time) >= 1 else lrc_with_time # drop last, can be empty
if self.max_frames == 2048:
lrc_with_time = lrc_with_time[:-1] if len(lrc_with_time) >= 1 else lrc_with_time # drop last, can be empty
lrc = torch.zeros((self.max_frames,), dtype=torch.long)
+23
View File
@@ -0,0 +1,23 @@
[00:13.55]亲爱的 人间爱情就像一朵花
[00:20.29]想要种鲜花
[00:22.51]手上难免沾泥巴
[00:27.90]大红喜字铺开乡间田野
[00:32.01]让我们住下
[00:34.96]听微风吟唱吧
[00:38.44]用光阴发新芽
[00:41.04]裂缝钻出的新绿
[00:43.97]正叩响整个盛夏
[00:47.15]候鸟衔来远方的彩霞
[00:52.03]荆棘里开出的花
[00:55.67]每片花瓣都是回答
[00:58.56]伤疤酿成琥珀时差
[01:01.42]折射出星空的密码
[01:04.36]蝴蝶吻过褪皮的痂
[01:07.28]绽放成盔甲
[01:10.11]牵你的手啊
[01:12.90]在岁月种下鲜花
[01:16.49]在寻常烟火人家
[01:20.05]用漫长慢慢地回答
[01:25.23]云会升起 雨会落下
[01:28.86]不用催促啊
[01:33.25]耐心的人啊才可以看见童话
+19
View File
@@ -0,0 +1,19 @@
[00:12.56]后来
[00:14.58]我总算学会了
[00:17.81]如何去爱
[00:19.37]可惜你 早已远去
[00:22.21]消失在人海
[00:25.42]Through tears I understood too late
[00:32.11]Some souls slip away beyond fate's gate
[00:39.86]White petals of gardenia
[00:46.32]Drift onto my blue pleated skirt
[00:51.95]I love you your whisper fell
[00:58.96]I bowed my head
[01:00.97]Breathing your lingering scent
[01:05.40]That eternal summer night
[01:09.35]Seventeen years young
[01:12.20]When your lips found mine
[01:18.23]Now through passing years
[01:21.79]Each wistful sigh
[01:24.91]Recalls starlight in your eyes
[01:30.91]Why did young love's promise
+16
View File
@@ -0,0 +1,16 @@
[00:18.23]让我掉下眼泪的
[00:21.80]不止昨夜的酒
[00:26.06]让我依依不舍的
[00:29.91]不止你的温柔
[00:33.79]余路还要走多久
[00:37.92]你攥着我的手
[00:41.81]让我感到为难的
[00:45.78]是挣扎的自由
[00:51.81]分别总是在九月
[00:55.66]回忆是思念的愁
[00:59.77]深秋嫩绿的垂柳
[01:03.55]亲吻着我额头
[01:07.55]在那座阴雨的小城里
[01:11.47]我从未忘记你
[01:15.37]成都 带不走的 只有你
[01:23.51]和我在成都的街头走一走
+44
View File
@@ -0,0 +1,44 @@
[00:18.23]让我掉下眼泪的
[00:21.80]不止昨夜的酒
[00:26.06]让我依依不舍的
[00:29.91]不止你的温柔
[00:33.79]余路还要走多久
[00:37.92]你攥着我的手
[00:41.81]让我感到为难的
[00:45.78]是挣扎的自由
[00:51.81]分别总是在九月
[00:55.66]回忆是思念的愁
[00:59.77]深秋嫩绿的垂柳
[01:03.55]亲吻着我额头
[01:07.55]在那座阴雨的小城里
[01:11.47]我从未忘记你
[01:15.37]成都 带不走的 只有你
[01:23.51]和我在成都的街头走一走
[01:31.38]直到所有的灯都熄灭了也不停留
[01:39.23]你会挽着我的衣袖
[01:43.12]我会把手揣进裤兜
[01:47.13]走到玉林路的尽头
[01:50.97]坐在小酒馆的门口
[02:30.76]分别总是在九月
[02:34.46]回忆是思念的愁
[02:38.48]深秋嫩绿的垂柳
[02:42.38]亲吻着我额头
[02:46.40]在那座阴雨的小城里
[02:50.35]我从未忘记你
[02:54.21]成都 带不走的 只有你
[03:02.18]和我在成都的街头走一走
[03:10.02]直到所有的灯都熄灭了也不停留
[03:18.16]你会挽着我的衣袖
[03:21.90]我会把手揣进裤兜
[03:25.80]走到玉林路的尽头
[03:29.81]坐在小酒馆的门口
[03:37.97]和我在成都的街头走一走
[03:45.76]直到所有的灯都熄灭了也不停留
[03:53.72]和我在成都的街头走一走
[04:01.44]直到所有的灯都熄灭了也不停留
[04:09.68]你会挽着我的衣袖
[04:13.35]我会把手揣进裤兜
[04:17.41]走到玉林路的尽头
[04:21.27]走过小酒馆的门口
[04:35.78]和我在成都的街头走一走
[04:43.18]直到所有的灯都熄灭了也不停留
+26
View File
@@ -0,0 +1,26 @@
[00:10.00]Moonlight spills through broken blinds
[00:13.20]Your shadow dances on the dashboard shrine
[00:16.85]Neon ghosts in gasoline rain
[00:20.40]I hear your laughter down the midnight train
[00:24.15]Static whispers through frayed wires
[00:27.65]Guitar strings hum our cathedral choirs
[00:31.30]Flicker screens show reruns of June
[00:34.90]I'm drowning in this mercury lagoon
[00:38.55]Electric veins pulse through concrete skies
[00:42.10]Your name echoes in the hollow where my heartbeat lies
[00:45.75]We're satellites trapped in parallel light
[00:49.25]Burning through the atmosphere of endless night
[01:00.00]Dusty vinyl spins reverse
[01:03.45]Our polaroid timeline bleeds through the verse
[01:07.10]Telescope aimed at dead stars
[01:10.65]Still tracing constellations through prison bars
[01:14.30]Electric veins pulse through concrete skies
[01:17.85]Your name echoes in the hollow where my heartbeat lies
[01:21.50]We're satellites trapped in parallel light
[01:25.05]Burning through the atmosphere of endless night
[02:10.00]Clockwork gears grind moonbeams to rust
[02:13.50]Our fingerprint smudged by interstellar dust
[02:17.15]Velvet thunder rolls through my veins
[02:20.70]Chasing phantom trains through solar plane
[02:24.35]Electric veins pulse through concrete skies
[02:27.90]Your name echoes in the hollow where my heartbeat lies
+57
View File
@@ -0,0 +1,57 @@
[00:00.52]Abracadabra abracadabra
[00:03.97]Ha
[00:04.66]Abracadabra abracadabra
[00:12.02]Yeah
[00:15.80]Pay the toll to the angels
[00:19.08]Drawin' circles in the clouds
[00:23.31]Keep your mind on the distance
[00:26.67]When the devil turns around
[00:30.95]Hold me in your heart tonight
[00:34.11]In the magic of the dark moonlight
[00:38.44]Save me from this empty fight
[00:43.83]In the game of life
[00:45.84]Like a poem said by a lady in red
[00:49.45]You hear the last few words of your life
[00:53.15]With a haunting dance now you're both in a trance
[00:56.90]It's time to cast your spell on the night
[01:01.40]Abracadabra ama-ooh-na-na
[01:04.88]Abracadabra porta-ooh-ga-ga
[01:08.92]Abracadabra abra-ooh-na-na
[01:12.30]In her tongue she's sayin'
[01:14.76]Death or love tonight
[01:18.61]Abracadabra abracadabra
[01:22.18]Abracadabra abracadabra
[01:26.08]Feel the beat under your feet
[01:27.82]The floor's on fire
[01:29.90]Abracadabra abracadabra
[01:33.78]Choose the road on the west side
[01:37.09]As the dust flies watch it burn
[01:41.45]Don't waste time on feeling
[01:44.64]Your depression won't return
[01:49.15]Hold me in your heart tonight
[01:52.21]In the magic of the dark moonlight
[01:56.54]Save me from this empty fight
[02:01.77]In the game of life
[02:03.94]Like a poem said by a lady in red
[02:07.52]You hear the last few words of your life
[02:11.19]With a haunting dance now you're both in a trance
[02:14.95]It's time to cast your spell on the night
[02:19.53]Abracadabra ama-ooh-na-na
[02:22.71]Abracadabra porta-ooh-ga-ga
[02:26.94]Abracadabra abra-ooh-na-na
[02:30.42]In her tongue she's sayin'
[02:32.83]Death or love tonight
[02:36.55]Abracadabra abracadabra
[02:40.27]Abracadabra abracadabra
[02:44.19]Feel the beat under your feet
[02:46.14]The floor's on fire
[02:47.95]Abracadabra abracadabra
[02:51.17]Phantom of the dance floor come to me
[02:58.46]Sing for me a sinful melody
[03:06.51]Ah-ah-ah-ah-ah ah-ah ah-ah
[03:13.76]Ah-ah-ah-ah-ah ah-ah ah-ah
[03:22.39]Abracadabra ama-ooh-na-na
[03:25.66]Abracadabra porta-ooh-ga-ga
[03:29.87]Abracadabra abra-ooh-na-na
[03:33.16]In her tongue she's sayin'
[03:35.55]Death or love tonight
View File
@@ -3,18 +3,17 @@
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
from g2p.g2p import cleaners
from diffrhythm.g2p.g2p import cleaners
from tokenizers import Tokenizer
from g2p.g2p.text_tokenizers import TextTokenizer
import LangSegment
from diffrhythm.g2p.g2p.text_tokenizers import TextTokenizer
from diffrhythm.LangSegment import LangSegment
import json
import os
import re
current_path = os.path.dirname(os.path.abspath(__file__))
class PhonemeBpeTokenizer:
def __init__(self, vacab_path=f"{current_path}/vocab.json"):
import os
def __init__(self, vacab_path=f"{os.path.dirname(os.path.abspath(__file__))}/vocab.json"):
self.lang2backend = {
"zh": "cmn",
"en": "en-us",
@@ -25,7 +24,7 @@ class PhonemeBpeTokenizer:
self.text_tokenizers = {}
self.int_text_tokenizers()
with open(vacab_path, "r", encoding="utf-8") as f:
with open(vacab_path, "r", encoding='utf-8') as f:
json_data = f.read()
data = json.loads(json_data)
self.vocab = data["vocab"]
@@ -116,7 +116,7 @@ class BertPolyPredict:
self.polydataset = PolyDataset
options = SessionOptions() # initialize session options
options.graph_optimization_level = GraphOptimizationLevel.ORT_ENABLE_ALL
# print(os.path.join(bert_model, "poly_bert_model.onnx"))
print(os.path.join(bert_model, "poly_bert_model.onnx"))
self.session = InferenceSession(
os.path.join(bert_model, "poly_bert_model.onnx"),
sess_options=options,
@@ -4,11 +4,11 @@
# LICENSE file in the root directory of this source tree.
import re
from g2p.g2p.mandarin import chinese_to_ipa
from g2p.g2p.english import english_to_ipa
from g2p.g2p.french import french_to_ipa
from g2p.g2p.korean import korean_to_ipa
from g2p.g2p.german import german_to_ipa
from diffrhythm.g2p.g2p.mandarin import chinese_to_ipa
from diffrhythm.g2p.g2p.english import english_to_ipa
from diffrhythm.g2p.g2p.french import french_to_ipa
from diffrhythm.g2p.g2p.korean import korean_to_ipa
from diffrhythm.g2p.g2p.german import german_to_ipa
def cjekfd_cleaners(text, sentence, language, text_tokenizers):
@@ -8,20 +8,18 @@ import jieba
import cn2an
from pypinyin import lazy_pinyin, BOPOMOFO
from typing import List
from g2p.g2p.chinese_model_g2p import BertPolyPredict
from g2p.utils.front_utils import *
from diffrhythm.g2p.g2p.chinese_model_g2p import BertPolyPredict
from diffrhythm.g2p.utils.front_utils import *
import os
# from g2pw import G2PWConverter
current_path = os.path.dirname(os.path.abspath(__file__))
# set blank level, {0:"none",1:"char", 2:"word"}
BLANK_LEVEL = 0
# conv = G2PWConverter(style='pinyin', enable_non_tradional_chinese=True)
resource_path = os.path.dirname(current_path)
resource_path = f"{os.path.dirname(os.path.dirname(os.path.abspath(__file__)))}"
poly_all_class_path = os.path.join(
resource_path, "sources", "g2p_chinese_model", "polychar.txt"
)
@@ -185,7 +183,7 @@ must_not_er_words = {"女儿", "老儿", "男儿", "少儿", "小儿"}
word_pinyin_dict = {}
with open(
rf"{resource_path}/sources/chinese_lexicon.txt", "r", encoding="utf-8"
f"{resource_path}/sources/chinese_lexicon.txt", "r", encoding="utf-8"
) as fread:
txt_list = fread.readlines()
for txt in txt_list:
@@ -195,7 +193,7 @@ with open(
pinyin_2_bopomofo_dict = {}
with open(
rf"{resource_path}/sources/pinyin_2_bpmf.txt", "r", encoding="utf-8"
f"{resource_path}/sources/pinyin_2_bpmf.txt", "r", encoding="utf-8"
) as fread:
txt_list = fread.readlines()
for txt in txt_list:
@@ -214,7 +212,7 @@ tone_dict = {
bopomofos2pinyin_dict = {}
with open(
rf"{resource_path}/sources/bpmf_2_pinyin.txt", "r", encoding="utf-8"
f"{resource_path}/sources/bpmf_2_pinyin.txt", "r", encoding="utf-8"
) as fread:
txt_list = fread.readlines()
for txt in txt_list:
@@ -6,8 +6,8 @@
import os
import sys
from g2p.g2p import PhonemeBpeTokenizer
from g2p.utils.g2p import phonemizer_g2p
from diffrhythm.g2p.g2p import PhonemeBpeTokenizer
from diffrhythm.g2p.utils.g2p import phonemizer_g2p
import tqdm
from typing import List
import json
@@ -114,9 +114,7 @@ def chn_eng_g2p(text: str):
text_tokenizer = PhonemeBpeTokenizer()
current_path = os.path.dirname(os.path.abspath(__file__))
with open(f"{current_path}/g2p/vocab.json", "r", encoding="utf-8") as f:
with open(f"{os.path.dirname(os.path.abspath(__file__))}/g2p/vocab.json", "r", encoding='utf-8') as f:
json_data = f.read()
data = json.loads(json_data)
vocab = data["vocab"]
View File
@@ -27,11 +27,6 @@ phonemizer_en = EspeakBackend(
)
# phonemizer_en.separator = separator
# phonemizer_ja = EspeakBackend(
# "ja", preserve_punctuation=False, with_stress=False, language_switch="remove-flags"
# )
# phonemizer_ja.separator = separator
phonemizer_ko = EspeakBackend(
"ko", preserve_punctuation=False, with_stress=False, language_switch="remove-flags"
)
@@ -53,15 +48,13 @@ phonemizer_de = EspeakBackend(
lang2backend = {
"zh": phonemizer_zh,
# "ja": phonemizer_ja,
"en": phonemizer_en,
"fr": phonemizer_fr,
"ko": phonemizer_ko,
"de": phonemizer_de,
}
current_path = os.path.dirname(os.path.abspath(__file__))
with open(f"{current_path}/mls_en.json", "r", encoding="utf-8") as f:
with open(f"{os.path.dirname(os.path.abspath(__file__))}/mls_en.json", "r", encoding='utf-8') as f:
json_data = f.read()
token = json.loads(json_data)
+6
View File
@@ -0,0 +1,6 @@
from diffrhythm.model.cfm import CFM
from diffrhythm.model.dit import DiT
from diffrhythm.model.trainer import Trainer
__all__ = ["CFM", "DiT", "Trainer"]
+71 -58
View File
@@ -1,10 +1,22 @@
"""
ein notation:
b - batch
n - sequence
nt - text sequence
nw - raw wave length
d - dimension
# Copyright (c) 2025 ASLP-LAB
# 2025 Ziqian Ning (ningziqian@mail.nwpu.edu.cn)
# 2025 Huakang Chen (huakang@mail.nwpu.edu.cn)
# 2025 Guobin Ma (guobin.ma@mail.nwpu.edu.cn)
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
# http://www.apache.org/licenses/LICENSE-2.0
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
""" This implementation is adapted from github repo:
https://github.com/SWivid/F5-TTS.
"""
from __future__ import annotations
@@ -19,9 +31,7 @@ from torch.nn.utils.rnn import pad_sequence
from torchdiffeq import odeint
from model.modules import MelSpec
from model.utils import (
default,
from diffrhythm.model.utils import (
exists,
list_str_to_idx,
list_str_to_tensor,
@@ -29,12 +39,25 @@ from model.utils import (
mask_from_frac_lengths,
)
def custom_mask_from_start_end_indices(seq_len: int["b"], start: int["b"], end: int["b"], device, max_seq_len): # noqa: F722 F821
def custom_mask_from_start_end_indices(
seq_len: int["b"], # noqa: F821
latent_pred_segments,
device,
max_seq_len
):
max_seq_len = max_seq_len
seq = torch.arange(max_seq_len, device=device).long()
start_mask = seq[None, :] >= start[:, None]
end_mask = seq[None, :] < end[:, None]
return start_mask & end_mask
res_mask = torch.zeros(max_seq_len, device=device, dtype=torch.bool)
for start, end in latent_pred_segments:
start = start.unsqueeze(0)
end = end.unsqueeze(0)
start_mask = seq[None, :] >= start[:, None]
end_mask = seq[None, :] < end[:, None]
res_mask = res_mask | (start_mask & end_mask)
return res_mask
class CFM(nn.Module):
def __init__(
@@ -42,7 +65,7 @@ class CFM(nn.Module):
transformer: nn.Module,
sigma=0.0,
odeint_kwargs: dict = dict(
method="euler" # 'midpoint'
method="euler"
),
odeint_options: dict = dict(
min_step=0.05
@@ -54,7 +77,7 @@ class CFM(nn.Module):
num_channels=None,
frac_lengths_mask: tuple[float, float] = (0.7, 1.0),
vocab_char_map: dict[str:int] | None = None,
use_style_prompt: bool = False
max_frames=2048
):
super().__init__()
@@ -83,8 +106,8 @@ class CFM(nn.Module):
# vocab map for tokenization
self.vocab_char_map = vocab_char_map
self.use_style_prompt = use_style_prompt
self.max_frames = max_frames
@property
def device(self):
@@ -112,10 +135,10 @@ class CFM(nn.Module):
t_inter=0.1,
edit_mask=None,
start_time=None,
latent_pred_start_frame=0,
latent_pred_end_frame=2048,
latent_pred_segments=None,
vocal_flag=False,
odeint_method="euler"
odeint_method="euler",
batch_infer_num=5
):
self.eval()
@@ -125,7 +148,6 @@ class CFM(nn.Module):
cond = cond.half()
# raw wave
if cond.shape[1] > duration:
cond = cond[:, :duration, :]
@@ -139,7 +161,6 @@ class CFM(nn.Module):
lens = torch.full((batch,), cond_seq_len, device=device, dtype=torch.long)
# text
if isinstance(text, list):
if exists(self.vocab_char_map):
text = list_str_to_idx(text, self.vocab_char_map).to(device)
@@ -147,26 +168,18 @@ class CFM(nn.Module):
text = list_str_to_tensor(text).to(device)
assert text.shape[0] == batch
if exists(text):
text_lens = (text != -1).sum(dim=-1)
# duration
cond_mask = lens_to_mask(lens)
if edit_mask is not None:
cond_mask = cond_mask & edit_mask
latent_pred_start_frame = torch.tensor([latent_pred_start_frame]).to(cond.device)
latent_pred_end_frame = duration
latent_pred_end_frame = torch.tensor([latent_pred_end_frame]).to(cond.device)
fixed_span_mask = custom_mask_from_start_end_indices(cond_seq_len, latent_pred_start_frame, latent_pred_end_frame, device=cond.device, max_seq_len=duration)
latent_pred_segments = torch.tensor(latent_pred_segments).to(cond.device)
fixed_span_mask = custom_mask_from_start_end_indices(cond_seq_len, latent_pred_segments, device=cond.device, max_seq_len=duration)
fixed_span_mask = fixed_span_mask.unsqueeze(-1)
step_cond = torch.where(fixed_span_mask, torch.zeros_like(cond), cond)
if isinstance(duration, int):
duration = torch.full((batch,), duration, device=device, dtype=torch.long)
duration = torch.full((batch_infer_num,), duration, device=device, dtype=torch.long)
duration = duration.clamp(max=max_duration)
max_duration = duration.amax()
@@ -175,7 +188,6 @@ class CFM(nn.Module):
if duplicate_test:
test_cond = F.pad(cond, (0, 0, cond_seq_len, max_duration - 2 * cond_seq_len), value=0.0)
if batch > 1:
mask = lens_to_mask(duration)
else: # save memory and speed up, as single inference need no mask currently
@@ -184,32 +196,33 @@ class CFM(nn.Module):
# test for no ref audio
if no_ref_audio:
cond = torch.zeros_like(cond)
start_time_embed, positive_text_embed, positive_text_residuals = self.transformer.forward_timestep_invariant(text, step_cond.shape[1], drop_text=False, start_time=start_time)
_, negative_text_embed, negative_text_residuals = self.transformer.forward_timestep_invariant(text, step_cond.shape[1], drop_text=True, start_time=start_time)
if vocal_flag:
style_prompt = negative_style_prompt
negative_style_prompt = torch.zeros_like(style_prompt)
text_embed = torch.cat([positive_text_embed, negative_text_embed], 0)
text_residuals = [torch.cat([a, b], 0) for a, b in zip(positive_text_residuals, negative_text_residuals)]
step_cond = torch.cat([step_cond, step_cond], 0)
style_prompt = torch.cat([style_prompt, negative_style_prompt], 0)
start_time_embed = torch.cat([start_time_embed, start_time_embed], 0)
cond = cond.repeat(batch_infer_num, 1, 1)
step_cond = step_cond.repeat(batch_infer_num, 1, 1)
text = text.repeat(batch_infer_num, 1)
style_prompt = style_prompt.repeat(batch_infer_num, 1)
negative_style_prompt = negative_style_prompt.repeat(batch_infer_num, 1)
start_time = start_time.repeat(batch_infer_num)
fixed_span_mask = fixed_span_mask.repeat(batch_infer_num, 1, 1)
def fn(t, x):
x = torch.cat([x, x], 0)
# predict flow
pred = self.transformer(
x=x, text_embed=text_embed, text_residuals=text_residuals, cond=step_cond, time=t,
drop_audio_cond=True, drop_prompt=False, style_prompt=style_prompt, start_time=start_time_embed
x=x, cond=step_cond, text=text, time=t, drop_audio_cond=False, drop_text=False, drop_prompt=False,
style_prompt=style_prompt, start_time=start_time
)
if cfg_strength < 1e-5:
return pred
positive_pred, negative_pred = pred.chunk(2, 0)
cfg_pred = positive_pred + (positive_pred - negative_pred) * cfg_strength
return cfg_pred
null_pred = self.transformer(
x=x, cond=step_cond, text=text, time=t, drop_audio_cond=True, drop_text=True, drop_prompt=False,
style_prompt=negative_style_prompt, start_time=start_time
)
return pred + (pred - null_pred) * cfg_strength
# noise input
# to make sure batch inference result is same with different batch size, and for sure single inference
@@ -228,7 +241,7 @@ class CFM(nn.Module):
t_start = t_inter
y0 = (1 - t_start) * y0 + t_start * test_cond
steps = int(steps * (1 - t_start))
t = torch.linspace(t_start, 1, steps, device=self.device, dtype=step_cond.dtype)
if sway_sampling_coef is not None:
t = t + sway_sampling_coef * (torch.cos(torch.pi / 2 * t) - 1 + t)
@@ -243,6 +256,7 @@ class CFM(nn.Module):
out = out.permute(0, 2, 1)
out = vocoder(out)
out = torch.chunk(out, batch_infer_num, dim=0)
return out, trajectory
def forward(
@@ -267,11 +281,10 @@ class CFM(nn.Module):
# get a random span to mask out for training conditionally
frac_lengths = torch.zeros((batch,), device=self.device).float().uniform_(*self.frac_lengths_mask)
rand_span_mask = mask_from_frac_lengths(lens, frac_lengths)
rand_span_mask = mask_from_frac_lengths(lens, frac_lengths, self.max_frames)
if exists(mask):
rand_span_mask = mask
# rand_span_mask &= mask
# mel is x1
x1 = inp
@@ -301,7 +314,7 @@ class CFM(nn.Module):
# adding mask will use more memory, thus also need to adjust batchsampler with scaled down threshold for long sequences
pred = self.transformer(
x=φ, cond=cond, text=text, time=time, drop_audio_cond=drop_audio_cond, drop_text=drop_text, drop_prompt=drop_prompt,
style_prompt=style_prompt, style_prompt_lens=style_prompt_lens, grad_ckpt=grad_ckpt, start_time=start_time
style_prompt=style_prompt, start_time=start_time
)
# flow matching loss
@@ -309,4 +322,4 @@ class CFM(nn.Module):
loss = loss[rand_span_mask]
return loss.mean(), cond, pred
+49 -29
View File
@@ -1,10 +1,22 @@
"""
ein notation:
b - batch
n - sequence
nt - text sequence
nw - raw wave length
d - dimension
# Copyright (c) 2025 ASLP-LAB
# 2025 Ziqian Ning (ningziqian@mail.nwpu.edu.cn)
# 2025 Huakang Chen (huakang@mail.nwpu.edu.cn)
# 2025 Yuepeng Jiang (Jiangyp@mail.nwpu.edu.cn)
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
# http://www.apache.org/licenses/LICENSE-2.0
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
""" This implementation is adapted from github repo:
https://github.com/SWivid/F5-TTS.
"""
from __future__ import annotations
@@ -12,22 +24,19 @@ from __future__ import annotations
import torch
from torch import nn
import torch
import torch.nn.functional as F
from transformers.models.llama.modeling_llama import LlamaDecoderLayer, LlamaRotaryEmbedding
from transformers.models.llama import LlamaConfig
from torch.utils.checkpoint import checkpoint
from model.modules import (
from diffrhythm.model.modules import (
TimestepEmbedding,
ConvNeXtV2Block,
ConvPositionEmbedding,
DiTBlock,
AdaLayerNormZero_Final,
precompute_freqs_cis,
get_pos_embed_indices,
_prepare_decoder_attention_mask,
)
# from liger_kernel.transformers import apply_liger_kernel_to_llama
# apply_liger_kernel_to_llama()
# Text embedding
class TextEmbedding(nn.Module):
@@ -77,7 +86,6 @@ class InputEmbedding(nn.Module):
def forward(self, x: float["b n d"], cond: float["b n d"], text_embed: float["b n d"], style_emb, time_emb, drop_audio_cond=False): # noqa: F722
if drop_audio_cond: # cfg for cond audio
cond = torch.zeros_like(cond)
style_emb = style_emb.unsqueeze(1).repeat(1, x.shape[1], 1)
time_emb = time_emb.unsqueeze(1).repeat(1, x.shape[1], 1)
x = self.proj(torch.cat((x, cond, text_embed, style_emb, time_emb), dim=-1))
@@ -85,9 +93,7 @@ class InputEmbedding(nn.Module):
return x
# Transformer backbone using DiT blocks
# Transformer backbone using Llama blocks
class DiT(nn.Module):
def __init__(
self,
@@ -103,26 +109,25 @@ class DiT(nn.Module):
text_dim=None,
conv_layers=0,
long_skip_connection=False,
use_style_prompt=False,
max_pos=2048,
max_frames=2048
):
super().__init__()
self.max_frames = max_frames
cond_dim = 512
self.time_embed = TimestepEmbedding(cond_dim)
self.start_time_embed = TimestepEmbedding(cond_dim)
if text_dim is None:
text_dim = mel_dim
self.text_embed = TextEmbedding(text_num_embeds, text_dim, conv_layers=conv_layers, max_pos=max_pos)
self.text_embed = TextEmbedding(text_num_embeds, text_dim, conv_layers=conv_layers, max_pos=self.max_frames)
self.input_embed = InputEmbedding(mel_dim, text_dim, dim, cond_dim=cond_dim)
self.dim = dim
self.depth = depth
llama_config = LlamaConfig(hidden_size=dim, intermediate_size=dim * ff_mult, hidden_act='silu', max_position_embeddings=max_pos)
llama_config = LlamaConfig(hidden_size=dim, intermediate_size=dim * ff_mult, hidden_act='silu', max_position_embeddings=self.max_frames)
llama_config._attn_implementation = 'sdpa'
self.transformer_blocks = nn.ModuleList(
[LlamaDecoderLayer(llama_config, layer_idx=i) for i in range(depth)]
)
@@ -144,7 +149,6 @@ class DiT(nn.Module):
self.norm_out = AdaLayerNormZero_Final(dim, cond_dim) # final modulation
self.proj_out = nn.Linear(dim, mel_dim)
def forward_timestep_invariant(self, text, seq_len, drop_text, start_time):
s_t = self.start_time_embed(start_time)
text_embed = self.text_embed(text, seq_len, drop_text=drop_text)
@@ -158,21 +162,25 @@ class DiT(nn.Module):
def forward(
self,
x: float["b n d"], # nosied input audio # noqa: F722
text_embed: int["b nt"], # text # noqa: F722
text_residuals,
cond: float["b n d"], # masked cond audio # noqa: F722
text: int["b nt"], # text # noqa: F722
time: float["b"] | float[""], # time step # noqa: F821 F722
drop_audio_cond, # cfg for cond audio
drop_text, # cfg for text
drop_prompt=False,
style_prompt=None, # [b d t]
start_time=None,
):
batch, seq_len = x.shape[0], x.shape[1]
if time.ndim == 0:
time = time.repeat(batch)
# t: conditioning time, c: context (text + masked cond audio), x: noised input audio
t = self.time_embed(time)
c = t + start_time
s_t = self.start_time_embed(start_time)
c = t + s_t
text_embed = self.text_embed(text, seq_len, drop_text=drop_text)
if drop_prompt:
style_prompt = torch.zeros_like(style_prompt)
@@ -187,11 +195,22 @@ class DiT(nn.Module):
pos_ids = torch.arange(x.shape[1], device=x.device)
pos_ids = pos_ids.unsqueeze(0).repeat(x.shape[0], 1)
rotary_embed = self.rotary_emb(x, pos_ids)
attention_mask = torch.ones(
(batch, seq_len),
dtype=torch.bool,
device=x.device,
)
attention_mask = _prepare_decoder_attention_mask(
attention_mask,
(batch, seq_len),
x,
)
for i, block in enumerate(self.transformer_blocks):
x, *_ = block(x, position_embeddings=rotary_embed)
x, *_ = block(x, attention_mask=attention_mask, position_embeddings=rotary_embed)
if i < self.depth // 2:
x = x + text_residuals[i]
x = x + self.text_fusion_linears[i](text_embed)
if self.long_skip_connection is not None:
x = self.long_skip_connection(torch.cat((x, residual), dim=-1))
@@ -200,3 +219,4 @@ class DiT(nn.Module):
output = self.proj_out(x)
return output
@@ -1,19 +1,10 @@
# Copyright (c) 2025 ASLP-LAB
#
# Licensed under the Stability AI License (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://huggingface.co/stabilityai/stable-audio-open-1.0/blob/main/LICENSE.md
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
""" This implementation is adapted from github repo:
https://github.com/SWivid/F5-TTS.
"""
ein notation:
b - batch
n - sequence
nt - text sequence
nw - raw wave length
d - dimension
"""
from __future__ import annotations
@@ -45,6 +36,7 @@ class FiLMLayer(nn.Module):
gamma, beta = torch.chunk(self.film(c.unsqueeze(2)), chunks=2, dim=1)
gamma = gamma.transpose(1, 2)
beta = beta.transpose(1, 2)
# print(gamma.shape, beta.shape)
return gamma * x + beta
# raw wav to mel spec
@@ -617,3 +609,44 @@ class TimestepEmbedding(nn.Module):
time_hidden = time_hidden.to(timestep.dtype)
time = self.time_mlp(time_hidden) # b d
return time
# attention mask realated
def _prepare_decoder_attention_mask(attention_mask, input_shape, inputs_embeds):
# create noncausal mask
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
combined_attention_mask = None
def _expand_mask(
mask: torch.Tensor, dtype: torch.dtype, tgt_len: int = None
):
"""
Expands attention_mask from `[bsz, seq_len]` to `[bsz, 1, tgt_seq_len, src_seq_len]`.
"""
bsz, src_len = mask.size()
tgt_len = tgt_len if tgt_len is not None else src_len
expanded_mask = (
mask[:, None, None, :].expand(bsz, 1, tgt_len, src_len).to(dtype)
)
inverted_mask = 1.0 - expanded_mask
return inverted_mask.masked_fill(
inverted_mask.to(torch.bool), torch.finfo(dtype).min
)
if attention_mask is not None:
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
expanded_attn_mask = _expand_mask(
attention_mask, inputs_embeds.dtype, tgt_len=input_shape[-1]
).to(inputs_embeds.device)
combined_attention_mask = (
expanded_attn_mask
if combined_attention_mask is None
else expanded_attn_mask + combined_attention_mask
)
return combined_attention_mask
@@ -2,12 +2,12 @@
# 2025 Ziqian Ning (ningziqian@mail.nwpu.edu.cn)
# 2025 Huakang Chen (huakang@mail.nwpu.edu.cn)
#
# Licensed under the Stability AI License (the "License");
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://huggingface.co/stabilityai/stable-audio-open-1.0/blob/main/LICENSE.md
#
# http://www.apache.org/licenses/LICENSE-2.0
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
@@ -31,15 +31,14 @@ from torch.optim.lr_scheduler import LinearLR, SequentialLR, ConstantLR
from accelerate import Accelerator
from accelerate.utils import DistributedDataParallelKwargs
from dr_dataset.dataset import DiffusionDataset
from diffrhythm.dataset.dataset import DiffusionDataset
from torch.utils.data import DataLoader
from ema_pytorch import EMA
from model import CFM
from model.utils import exists, default
from diffrhythm.model import CFM
from diffrhythm.model.utils import exists, default
class Trainer:
def __init__(
+5 -23
View File
@@ -1,21 +1,3 @@
# Copyright (c) 2025 ASLP-LAB
#
# Licensed under the Stability AI License (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://huggingface.co/stabilityai/stable-audio-open-1.0/blob/main/LICENSE.md
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
""" This implementation is adapted from github repo:
https://github.com/SWivid/F5-TTS.
"""
from __future__ import annotations
import os
@@ -62,15 +44,15 @@ def lens_to_mask(t: int["b"], length: int | None = None) -> bool["b n"]: # noqa
return seq[None, :] < t[:, None]
def mask_from_start_end_indices(seq_len: int["b"], start: int["b"], end: int["b"]): # noqa: F722 F821
max_seq_len = 2048
def mask_from_start_end_indices(seq_len: int["b"], start: int["b"], end: int["b"], max_frames): # noqa: F722 F821
max_seq_len = max_frames
seq = torch.arange(max_seq_len, device=start.device).long()
start_mask = seq[None, :] >= start[:, None]
end_mask = seq[None, :] < end[:, None]
return start_mask & end_mask
def mask_from_frac_lengths(seq_len: int["b"], frac_lengths: float["b"]): # noqa: F722 F821
def mask_from_frac_lengths(seq_len: int["b"], frac_lengths: float["b"], max_frames): # noqa: F722 F821
lengths = (frac_lengths * seq_len).long()
max_start = seq_len - lengths
@@ -78,7 +60,7 @@ def mask_from_frac_lengths(seq_len: int["b"], frac_lengths: float["b"]): # noqa
start = (max_start * rand).long().clamp(min=0)
end = start + lengths
return mask_from_start_end_indices(seq_len, start, end)
return mask_from_start_end_indices(seq_len, start, end, max_frames)
def maybe_masked_mean(t: float["b n d"], mask: bool["b n"] = None) -> float["b d"]: # noqa: F722
@@ -197,4 +179,4 @@ def repetition_found(text, length=2, tolerance=10):
for pattern, count in pattern_count.items():
if count > tolerance:
return True
return False
return False
View File
+77
View File
@@ -0,0 +1,77 @@
# Copyright (c) 2025 ASLP-LAB
# 2025 Ziqian Ning (ningziqian@mail.nwpu.edu.cn)
# 2025 Huakang Chen (huakang@mail.nwpu.edu.cn)
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
# http://www.apache.org/licenses/LICENSE-2.0
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from importlib.resources import files
from diffrhythm.model import CFM, DiT, Trainer
from prefigure.prefigure import get_all_args
import json
import os
os.environ['OMP_NUM_THREADS']="1"
os.environ['MKL_NUM_THREADS']="1"
def main():
args = get_all_args("config/default.ini")
with open(args.model_config) as f:
model_config = json.load(f)
if model_config["model_type"] == "diffrhythm":
wandb_resume_id = None
model_cls = DiT
model = CFM(
transformer=model_cls(**model_config["model"], max_frames=args.max_frames),
num_channels=model_config["model"]['mel_dim'],
audio_drop_prob=args.audio_drop_prob,
cond_drop_prob=args.cond_drop_prob,
style_drop_prob=args.style_drop_prob,
lrc_drop_prob=args.lrc_drop_prob,
max_frames=args.max_frames
)
total_params = sum(p.numel() for p in model.parameters())
print(f"Total parameters: {total_params}")
trainer = Trainer(
model,
args,
args.epochs,
args.learning_rate,
num_warmup_updates=args.num_warmup_updates,
save_per_updates=args.save_per_updates,
checkpoint_path=f"ckpts/{args.exp_name}",
grad_accumulation_steps=args.grad_accumulation_steps,
max_grad_norm=args.max_grad_norm,
wandb_project="diffrhythm-test",
wandb_run_name=args.exp_name,
wandb_resume_id=wandb_resume_id,
last_per_steps=args.last_per_steps,
bnb_optimizer=False,
reset_lr=args.reset_lr,
batch_size=args.batch_size,
grad_ckpt=args.grad_ckpt
)
trainer.train(
resumable_with_seed=args.resumable_with_seed, # seed for shuffling dataset
)
if __name__ == "__main__":
main()
View File
-200
View File
@@ -1,200 +0,0 @@
import torch
import random
import json
import os
import numpy as np
node_dir = os.path.dirname(os.path.abspath(__file__))
def decode_audio(
latents: torch.Tensor,
vae_model: torch.nn.Module,
chunked: bool = False,
overlap: int = 32,
chunk_size: int = 128
):
downsampling_ratio = 2048
io_channels = 2
if not chunked:
try:
output = vae_model.decode_export(latents)
return output
except Exception as e:
raise
else:
# Chunked decoding logic
hop_size = chunk_size - overlap
total_size = latents.shape[2]
batch_size = latents.shape[0]
chunks = []
i = 0
for i in range(0, total_size - chunk_size + 1, hop_size):
chunk = latents[:, :, i : i + chunk_size]
chunks.append(chunk)
if i + chunk_size != total_size:
# Final chunk
chunk = latents[:, :, -chunk_size:]
chunks.append(chunk)
chunks = torch.stack(chunks)
num_chunks = chunks.shape[0]
# samples_per_latent is just the downsampling ratio
samples_per_latent = downsampling_ratio
# Create an empty waveform, we will populate it with chunks as decode them
y_size = total_size * samples_per_latent
y_final = torch.zeros((batch_size, io_channels, y_size)).to(latents.device)
for i in range(num_chunks):
x_chunk = chunks[i, :]
try:
y_chunk = vae_model.decode_export(x_chunk)
except Exception as e:
raise
# figure out where to put the audio along the time domain
if i == num_chunks - 1:
# final chunk always goes at the end
t_end = y_size
t_start = t_end - y_chunk.shape[2]
else:
t_start = i * hop_size * samples_per_latent
t_end = t_start + chunk_size * samples_per_latent
# remove the edges of the overlaps
ol = (overlap // 2) * samples_per_latent
chunk_start = 0
chunk_end = y_chunk.shape[2]
if i > 0:
# no overlap for the start of the first chunk
t_start += ol
chunk_start += ol
if i < num_chunks - 1:
# no overlap for the end of the last chunk
t_end -= ol
chunk_end -= ol
# paste the chunked audio into our y_final output audio
y_final[:, :, t_start:t_end] = y_chunk[:, :, chunk_start:chunk_end]
return y_final
def get_reference_latent(device: torch.device, max_frames: int):
return torch.zeros(1, max_frames, 64).to(device)
def get_negative_style_prompt(device: torch.device):
file_path = f"{node_dir}/vocal.npy"
try:
vocal_style = np.load(file_path)
except Exception as e:
raise
vocal_style = torch.from_numpy(vocal_style).to(device) # [1, 512]
return vocal_style.half()
def parse_lyrics(lyrics: str):
lyrics_with_time = []
lyrics = lyrics.strip()
for line in lyrics.split("\n"):
try:
time, lyric = line[1:9], line[10:]
mins, secs = time.split(":")
secs = int(mins) * 60 + float(secs)
lyrics_with_time.append((secs, lyric.strip()))
except ValueError:
continue
return lyrics_with_time
class CNENTokenizer:
def __init__(self):
vocab_path = f"{node_dir}/g2p/g2p/vocab.json"
try:
with open(vocab_path, "r", encoding="utf-8") as file:
self.phone2id: dict = json.load(file)["vocab"]
except Exception as e:
raise
self.id2phone = {v: k for k, v in self.phone2id.items()}
try:
from g2p.g2p_generation import chn_eng_g2p
self.tokenizer = chn_eng_g2p
except Exception as e:
raise
def encode(self, text: str):
try:
phone, token = self.tokenizer(text)
return [x + 1 for x in token]
except Exception as e:
print(f"Text encoding failed: {str(e)}")
raise
def decode(self, token: list):
try:
return "|".join([self.id2phone[x - 1] for x in token])
except Exception as e:
raise
def get_lrc_token(
max_frames: int,
text: str,
tokenizer: CNENTokenizer,
device: torch.device
):
# Audio processing parameters
lyrics_shift = 0
sampling_rate = 44100
downsample_rate = 2048
max_secs = max_frames / (sampling_rate / downsample_rate)
# Token configuration
comma_token_id = 1
period_token_id = 2
lrc_with_time = parse_lyrics(text)
modified_lrc_with_time = []
for i in range(len(lrc_with_time)):
time, line = lrc_with_time[i]
try:
line_token = tokenizer.encode(line)
modified_lrc_with_time.append((time, line_token))
except Exception as e:
raise
lrc_with_time = modified_lrc_with_time
lrc_with_time = [
(time_start, line)
for (time_start, line) in lrc_with_time
if time_start < max_secs
]
# lrc_with_time = lrc_with_time[:-1] if len(lrc_with_time) >= 1 else lrc_with_time
normalized_start_time = 0.0
lrc = torch.zeros((max_frames,), dtype=torch.long)
tokens_count = 0
last_end_pos = 0
for time_start, line in lrc_with_time:
tokens = [
token if token != period_token_id else comma_token_id for token in line
] + [period_token_id]
tokens = torch.tensor(tokens, dtype=torch.long)
num_tokens = tokens.shape[0]
gt_frame_start = int(time_start * sampling_rate / downsample_rate)
frame_shift = random.randint(int(lyrics_shift), int(lyrics_shift))
frame_start = max(gt_frame_start - frame_shift, last_end_pos)
frame_len = min(num_tokens, max_frames - frame_start)
lrc[frame_start : frame_start + frame_len] = tokens[:frame_len]
tokens_count += num_tokens
last_end_pos = frame_start + frame_len
lrc_emb = lrc.unsqueeze(0).to(device)
normalized_start_time = torch.tensor(normalized_start_time).unsqueeze(0).to(device)
if device == "cuda":
normalized_start_time = normalized_start_time.half()
else:
normalized_start_time = normalized_start_time.float()
return lrc_emb, normalized_start_time
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
-816
View File
@@ -1,816 +0,0 @@
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
import io, re, os, sys, time, argparse, pdb, json
from io import StringIO
from typing import Optional
import numpy as np
import traceback
import pyopenjtalk
from pykakasi import kakasi
punctuation = [",", ".", "!", "?", ":", ";", "'", "…"]
jp_xphone2ipa = [
" a a",
" i i",
" u ɯ",
" e e",
" o o",
" a: aː",
" i: iː",
" u: ɯː",
" e: eː",
" o: oː",
" k k",
" s s",
" t t",
" n n",
" h ç",
" f ɸ",
" m m",
" y j",
" r ɾ",
" w ɰᵝ",
" N ɴ",
" g g",
" j d ʑ",
" z z",
" d d",
" b b",
" p p",
" q q",
" v v",
" : :",
" by b j",
" ch t ɕ",
" dy d e j",
" ty t e j",
" gy g j",
" gw g ɯ",
" hy ç j",
" ky k j",
" kw k ɯ",
" my m j",
" ny n j",
" py p j",
" ry ɾ j",
" sh ɕ",
" ts t s ɯ",
]
_mora_list_minimum: list[tuple[str, Optional[str], str]] = [
("ヴォ", "v", "o"),
("ヴェ", "v", "e"),
("ヴィ", "v", "i"),
("ヴァ", "v", "a"),
("ヴ", "v", "u"),
("ン", None, "N"),
("ワ", "w", "a"),
("ロ", "r", "o"),
("レ", "r", "e"),
("ル", "r", "u"),
("リョ", "ry", "o"),
("リュ", "ry", "u"),
("リャ", "ry", "a"),
("リェ", "ry", "e"),
("リ", "r", "i"),
("ラ", "r", "a"),
("ヨ", "y", "o"),
("ユ", "y", "u"),
("ヤ", "y", "a"),
("モ", "m", "o"),
("メ", "m", "e"),
("ム", "m", "u"),
("ミョ", "my", "o"),
("ミュ", "my", "u"),
("ミャ", "my", "a"),
("ミェ", "my", "e"),
("ミ", "m", "i"),
("マ", "m", "a"),
("ポ", "p", "o"),
("ボ", "b", "o"),
("ホ", "h", "o"),
("ペ", "p", "e"),
("ベ", "b", "e"),
("ヘ", "h", "e"),
("プ", "p", "u"),
("ブ", "b", "u"),
("フォ", "f", "o"),
("フェ", "f", "e"),
("フィ", "f", "i"),
("ファ", "f", "a"),
("フ", "f", "u"),
("ピョ", "py", "o"),
("ピュ", "py", "u"),
("ピャ", "py", "a"),
("ピェ", "py", "e"),
("ピ", "p", "i"),
("ビョ", "by", "o"),
("ビュ", "by", "u"),
("ビャ", "by", "a"),
("ビェ", "by", "e"),
("ビ", "b", "i"),
("ヒョ", "hy", "o"),
("ヒュ", "hy", "u"),
("ヒャ", "hy", "a"),
("ヒェ", "hy", "e"),
("ヒ", "h", "i"),
("パ", "p", "a"),
("バ", "b", "a"),
("ハ", "h", "a"),
("ノ", "n", "o"),
("ネ", "n", "e"),
("ヌ", "n", "u"),
("ニョ", "ny", "o"),
("ニュ", "ny", "u"),
("ニャ", "ny", "a"),
("ニェ", "ny", "e"),
("ニ", "n", "i"),
("ナ", "n", "a"),
("ドゥ", "d", "u"),
("ド", "d", "o"),
("トゥ", "t", "u"),
("ト", "t", "o"),
("デョ", "dy", "o"),
("デュ", "dy", "u"),
("デャ", "dy", "a"),
# ("デェ", "dy", "e"),
("ディ", "d", "i"),
("デ", "d", "e"),
("テョ", "ty", "o"),
("テュ", "ty", "u"),
("テャ", "ty", "a"),
("ティ", "t", "i"),
("テ", "t", "e"),
("ツォ", "ts", "o"),
("ツェ", "ts", "e"),
("ツィ", "ts", "i"),
("ツァ", "ts", "a"),
("ツ", "ts", "u"),
("ッ", None, "q"), # 「cl」から「q」に変更
("チョ", "ch", "o"),
("チュ", "ch", "u"),
("チャ", "ch", "a"),
("チェ", "ch", "e"),
("チ", "ch", "i"),
("ダ", "d", "a"),
("タ", "t", "a"),
("ゾ", "z", "o"),
("ソ", "s", "o"),
("ゼ", "z", "e"),
("セ", "s", "e"),
("ズィ", "z", "i"),
("ズ", "z", "u"),
("スィ", "s", "i"),
("ス", "s", "u"),
("ジョ", "j", "o"),
("ジュ", "j", "u"),
("ジャ", "j", "a"),
("ジェ", "j", "e"),
("ジ", "j", "i"),
("ショ", "sh", "o"),
("シュ", "sh", "u"),
("シャ", "sh", "a"),
("シェ", "sh", "e"),
("シ", "sh", "i"),
("ザ", "z", "a"),
("サ", "s", "a"),
("ゴ", "g", "o"),
("コ", "k", "o"),
("ゲ", "g", "e"),
("ケ", "k", "e"),
("グヮ", "gw", "a"),
("グ", "g", "u"),
("クヮ", "kw", "a"),
("ク", "k", "u"),
("ギョ", "gy", "o"),
("ギュ", "gy", "u"),
("ギャ", "gy", "a"),
("ギェ", "gy", "e"),
("ギ", "g", "i"),
("キョ", "ky", "o"),
("キュ", "ky", "u"),
("キャ", "ky", "a"),
("キェ", "ky", "e"),
("キ", "k", "i"),
("ガ", "g", "a"),
("カ", "k", "a"),
("オ", None, "o"),
("エ", None, "e"),
("ウォ", "w", "o"),
("ウェ", "w", "e"),
("ウィ", "w", "i"),
("ウ", None, "u"),
("イェ", "y", "e"),
("イ", None, "i"),
("ア", None, "a"),
]
_mora_list_additional: list[tuple[str, Optional[str], str]] = [
("ヴョ", "by", "o"),
("ヴュ", "by", "u"),
("ヴャ", "by", "a"),
("ヲ", None, "o"),
("ヱ", None, "e"),
("ヰ", None, "i"),
("ヮ", "w", "a"),
("ョ", "y", "o"),
("ュ", "y", "u"),
("ヅ", "z", "u"),
("ヂ", "j", "i"),
("ヶ", "k", "e"),
("ャ", "y", "a"),
("ォ", None, "o"),
("ェ", None, "e"),
("ゥ", None, "u"),
("ィ", None, "i"),
("ァ", None, "a"),
]
# 例: "vo" -> "ヴォ", "a" -> "ア"
mora_phonemes_to_mora_kata: dict[str, str] = {
(consonant or "") + vowel: kana for [kana, consonant, vowel] in _mora_list_minimum
}
# 例: "ヴォ" -> ("v", "o"), "ア" -> (None, "a")
mora_kata_to_mora_phonemes: dict[str, tuple[Optional[str], str]] = {
kana: (consonant, vowel)
for [kana, consonant, vowel] in _mora_list_minimum + _mora_list_additional
}
# 正規化で記号を変換するための辞書
rep_map = {
":": ":",
";": ";",
",": ",",
"。": ".",
"!": "!",
"?": "?",
"\n": ".",
".": ".",
"⋯": "…",
"···": "…",
"・・・": "…",
"·": ",",
"・": ",",
"•": ",",
"、": ",",
"$": ".",
# "“": "'",
# "”": "'",
# '"': "'",
"‘": "'",
"’": "'",
# "(": "'",
# ")": "'",
# "(": "'",
# ")": "'",
# "《": "'",
# "》": "'",
# "【": "'",
# "】": "'",
# "[": "'",
# "]": "'",
# "——": "-",
# "−": "-",
# "-": "-",
# "『": "'",
# "』": "'",
# "〈": "'",
# "〉": "'",
# "«": "'",
# "»": "'",
# # "~": "-", # これは長音記号「ー」として扱うよう変更
# # "~": "-", # これは長音記号「ー」として扱うよう変更
# "「": "'",
# "」": "'",
}
def _numeric_feature_by_regex(regex, s):
match = re.search(regex, s)
if match is None:
return -50
return int(match.group(1))
def replace_punctuation(text: str) -> str:
"""句読点等を「.」「,」「!」「?」「'」「-」に正規化し、OpenJTalkで読みが取得できるもののみ残す:
漢字・平仮名・カタカナ、アルファベット、ギリシャ文字
"""
pattern = re.compile("|".join(re.escape(p) for p in rep_map.keys()))
# print("before: ", text)
# 句読点を辞書で置換
replaced_text = pattern.sub(lambda x: rep_map[x.group()], text)
replaced_text = re.sub(
# ↓ ひらがな、カタカナ、漢字
r"[^\u3040-\u309F\u30A0-\u30FF\u4E00-\u9FFF\u3400-\u4DBF\u3005"
# ↓ 半角アルファベット(大文字と小文字)
+ r"\u0041-\u005A\u0061-\u007A"
# ↓ 全角アルファベット(大文字と小文字)
+ r"\uFF21-\uFF3A\uFF41-\uFF5A"
# ↓ ギリシャ文字
+ r"\u0370-\u03FF\u1F00-\u1FFF"
# ↓ "!", "?", "…", ",", ".", "'", "-", 但し`…`はすでに`...`に変換されている
+ "".join(punctuation) + r"]+",
# 上述以外の文字を削除
"",
replaced_text,
)
# print("after: ", replaced_text)
return replaced_text
def fix_phone_tone(phone_tone_list: list[tuple[str, int]]) -> list[tuple[str, int]]:
"""
`phone_tone_list`のtone(アクセントの値)を0か1の範囲に修正する。
例: [(a, 0), (i, -1), (u, -1)] → [(a, 1), (i, 0), (u, 0)]
"""
tone_values = set(tone for _, tone in phone_tone_list)
if len(tone_values) == 1:
assert tone_values == {0}, tone_values
return phone_tone_list
elif len(tone_values) == 2:
if tone_values == {0, 1}:
return phone_tone_list
elif tone_values == {-1, 0}:
return [
(letter, 0 if tone == -1 else 1) for letter, tone in phone_tone_list
]
else:
raise ValueError(f"Unexpected tone values: {tone_values}")
else:
raise ValueError(f"Unexpected tone values: {tone_values}")
def fix_phone_tone_wplen(phone_tone_list, word_phone_length_list):
phones = []
tones = []
w_p_len = []
p_len = len(phone_tone_list)
idx = 0
w_idx = 0
while idx < p_len:
offset = 0
if phone_tone_list[idx] == "▁":
w_p_len.append(w_idx + 1)
curr_w_p_len = word_phone_length_list[w_idx]
for i in range(curr_w_p_len):
p, t = phone_tone_list[idx]
if p == ":" and len(phones) > 0:
if phones[-1][-1] != ":":
phones[-1] += ":"
offset -= 1
else:
phones.append(p)
tones.append(str(t))
idx += 1
if idx >= p_len:
break
w_p_len.append(curr_w_p_len + offset)
w_idx += 1
# print(w_p_len)
return phones, tones, w_p_len
def g2phone_tone_wo_punct(prosodies) -> list[tuple[str, int]]:
"""
テキストに対して、音素とアクセント(0か1)のペアのリストを返す。
ただし「!」「.」「?」等の非音素記号(punctuation)は全て消える(ポーズ記号も残さない)。
非音素記号を含める処理は`align_tones()`で行われる。
また「っ」は「cl」でなく「q」に変換される(「ん」は「N」のまま)。
例: "こんにちは、世界ー。。元気?!" →
[('k', 0), ('o', 0), ('N', 1), ('n', 1), ('i', 1), ('ch', 1), ('i', 1), ('w', 1), ('a', 1), ('s', 1), ('e', 1), ('k', 0), ('a', 0), ('i', 0), ('i', 0), ('g', 1), ('e', 1), ('N', 0), ('k', 0), ('i', 0)]
"""
result: list[tuple[str, int]] = []
current_phrase: list[tuple[str, int]] = []
current_tone = 0
last_accent = ""
for i, letter in enumerate(prosodies):
# 特殊記号の処理
# 文頭記号、無視する
if letter == "^":
assert i == 0, "Unexpected ^"
# アクセント句の終わりに来る記号
elif letter in ("$", "?", "_", "#"):
# 保持しているフレーズを、アクセント数値を0-1に修正し結果に追加
result.extend(fix_phone_tone(current_phrase))
# 末尾に来る終了記号、無視(文中の疑問文は`_`になる)
if letter in ("$", "?"):
assert i == len(prosodies) - 1, f"Unexpected {letter}"
# あとは"_"(ポーズ)と"#"(アクセント句の境界)のみ
# これらは残さず、次のアクセント句に備える。
current_phrase = []
# 0を基準点にしてそこから上昇・下降する(負の場合は上の`fix_phone_tone`で直る)
current_tone = 0
last_accent = ""
# アクセント上昇記号
elif letter == "[":
if last_accent != letter:
current_tone = current_tone + 1
last_accent = letter
# アクセント下降記号
elif letter == "]":
if last_accent != letter:
current_tone = current_tone - 1
last_accent = letter
# それ以外は通常の音素
else:
if letter == "cl": # 「っ」の処理
letter = "q"
current_phrase.append((letter, current_tone))
return result
def handle_long(sep_phonemes: list[list[str]]) -> list[list[str]]:
for i in range(len(sep_phonemes)):
if sep_phonemes[i][0] == "ー":
# sep_phonemes[i][0] = sep_phonemes[i - 1][-1]
sep_phonemes[i][0] = ":"
if "ー" in sep_phonemes[i]:
for j in range(len(sep_phonemes[i])):
if sep_phonemes[i][j] == "ー":
# sep_phonemes[i][j] = sep_phonemes[i][j - 1][-1]
sep_phonemes[i][j] = ":"
return sep_phonemes
def handle_long_word(sep_phonemes: list[list[str]]) -> list[list[str]]:
res = []
for i in range(len(sep_phonemes)):
if sep_phonemes[i][0] == "ー":
sep_phonemes[i][0] = sep_phonemes[i - 1][-1]
# sep_phonemes[i][0] = ':'
if "ー" in sep_phonemes[i]:
for j in range(len(sep_phonemes[i])):
if sep_phonemes[i][j] == "ー":
sep_phonemes[i][j] = sep_phonemes[i][j - 1][-1]
# sep_phonemes[i][j] = ':'
res.append(sep_phonemes[i])
res.append("▁")
return res
def align_tones(
phones_with_punct: list[str], phone_tone_list: list[tuple[str, int]]
) -> list[tuple[str, int]]:
"""
例:
…私は、、そう思う。
phones_with_punct:
[".", ".", ".", "w", "a", "t", "a", "sh", "i", "w", "a", ",", ",", "s", "o", "o", "o", "m", "o", "u", "."]
phone_tone_list:
[("w", 0), ("a", 0), ("t", 1), ("a", 1), ("sh", 1), ("i", 1), ("w", 1), ("a", 1), ("s", 0), ("o", 0), ("o", 1), ("o", 1), ("m", 1), ("o", 1), ("u", 0))]
Return:
[(".", 0), (".", 0), (".", 0), ("w", 0), ("a", 0), ("t", 1), ("a", 1), ("sh", 1), ("i", 1), ("w", 1), ("a", 1), (",", 0), (",", 0), ("s", 0), ("o", 0), ("o", 1), ("o", 1), ("m", 1), ("o", 1), ("u", 0), (".", 0)]
"""
result: list[tuple[str, int]] = []
tone_index = 0
for phone in phones_with_punct:
if tone_index >= len(phone_tone_list):
# 余ったpunctuationがある場合 → (punctuation, 0)を追加
result.append((phone, 0))
elif phone == phone_tone_list[tone_index][0]:
# phone_tone_listの現在の音素と一致する場合 → toneをそこから取得、(phone, tone)を追加
result.append((phone, phone_tone_list[tone_index][1]))
# 探すindexを1つ進める
tone_index += 1
elif phone in punctuation or phone == "▁":
# phoneがpunctuationの場合 → (phone, 0)を追加
result.append((phone, 0))
else:
print(f"phones: {phones_with_punct}")
print(f"phone_tone_list: {phone_tone_list}")
print(f"result: {result}")
print(f"tone_index: {tone_index}")
print(f"phone: {phone}")
raise ValueError(f"Unexpected phone: {phone}")
return result
def kata2phoneme_list(text: str) -> list[str]:
"""
原則カタカナの`text`を受け取り、それをそのままいじらずに音素記号のリストに変換。
注意点:
- punctuationが来た場合(punctuationが1文字の場合がありうる)、処理せず1文字のリストを返す
- 冒頭に続く「ー」はそのまま「ー」のままにする(`handle_long()`で処理される)
- 文中の「ー」は前の音素記号の最後の音素記号に変換される。
例:
`ーーソーナノカーー` → ["ー", "ー", "s", "o", "o", "n", "a", "n", "o", "k", "a", "a", "a"]
`?` → ["?"]
"""
if text in punctuation:
return [text]
# `text`がカタカナ(`ー`含む)のみからなるかどうかをチェック
if re.fullmatch(r"[\u30A0-\u30FF]+", text) is None:
raise ValueError(f"Input must be katakana only: {text}")
sorted_keys = sorted(mora_kata_to_mora_phonemes.keys(), key=len, reverse=True)
pattern = "|".join(map(re.escape, sorted_keys))
def mora2phonemes(mora: str) -> str:
cosonant, vowel = mora_kata_to_mora_phonemes[mora]
if cosonant is None:
return f" {vowel}"
return f" {cosonant} {vowel}"
spaced_phonemes = re.sub(pattern, lambda m: mora2phonemes(m.group()), text)
# 長音記号「ー」の処理
long_pattern = r"(\w)(ー*)"
long_replacement = lambda m: m.group(1) + (" " + m.group(1)) * len(m.group(2))
spaced_phonemes = re.sub(long_pattern, long_replacement, spaced_phonemes)
# spaced_phonemes += ' ▁'
return spaced_phonemes.strip().split(" ")
def frontend2phoneme(labels, drop_unvoiced_vowels=False):
N = len(labels)
phones = []
for n in range(N):
lab_curr = labels[n]
# print(lab_curr)
# current phoneme
p3 = re.search(r"\-(.*?)\+", lab_curr).group(1)
# deal unvoiced vowels as normal vowels
if drop_unvoiced_vowels and p3 in "AEIOU":
p3 = p3.lower()
# deal with sil at the beginning and the end of text
if p3 == "sil":
# assert n == 0 or n == N - 1
# if n == 0:
# phones.append("^")
# elif n == N - 1:
# # check question form or not
# e3 = _numeric_feature_by_regex(r"!(\d+)_", lab_curr)
# if e3 == 0:
# phones.append("$")
# elif e3 == 1:
# phones.append("?")
continue
elif p3 == "pau":
phones.append("_")
continue
else:
phones.append(p3)
# accent type and position info (forward or backward)
a1 = _numeric_feature_by_regex(r"/A:([0-9\-]+)\+", lab_curr)
a2 = _numeric_feature_by_regex(r"\+(\d+)\+", lab_curr)
a3 = _numeric_feature_by_regex(r"\+(\d+)/", lab_curr)
# number of mora in accent phrase
f1 = _numeric_feature_by_regex(r"/F:(\d+)_", lab_curr)
a2_next = _numeric_feature_by_regex(r"\+(\d+)\+", labels[n + 1])
# accent phrase border
# print(p3, a1, a2, a3, f1, a2_next, lab_curr)
if a3 == 1 and a2_next == 1 and p3 in "aeiouAEIOUNcl":
phones.append("#")
# pitch falling
elif a1 == 0 and a2_next == a2 + 1 and a2 != f1:
phones.append("]")
# pitch rising
elif a2 == 1 and a2_next == 2:
phones.append("[")
# phones = ' '.join(phones)
return phones
class JapanesePhoneConverter(object):
def __init__(self, lexicon_path=None, ipa_dict_path=None):
# lexicon_lines = open(lexicon_path, 'r', encoding='utf-8').readlines()
# self.lexicon = {}
# self.single_dict = {}
# self.double_dict = {}
# for curr_line in lexicon_lines:
# k,v = curr_line.strip().split('+',1)
# self.lexicon[k] = v
# if len(k) == 2:
# self.double_dict[k] = v
# elif len(k) == 1:
# self.single_dict[k] = v
self.ipa_dict = {}
for curr_line in jp_xphone2ipa:
k, v = curr_line.strip().split(" ", 1)
self.ipa_dict[k] = re.sub("\s", "", v)
# kakasi1 = kakasi()
# kakasi1.setMode("H","K")
# kakasi1.setMode("J","K")
# kakasi1.setMode("r","Hepburn")
self.japan_JH2K = kakasi()
self.table = {ord(f): ord(t) for f, t in zip("67", "_¯")}
def text2sep_kata(self, parsed) -> tuple[list[str], list[str]]:
"""
`text_normalize`で正規化済みの`norm_text`を受け取り、それを単語分割し、
分割された単語リストとその読み(カタカナor記号1文字)のリストのタプルを返す。
単語分割結果は、`g2p()`の`word2ph`で1文字あたりに割り振る音素記号の数を決めるために使う。
例:
`私はそう思う!って感じ?` →
["私", "は", "そう", "思う", "!", "って", "感じ", "?"], ["ワタシ", "ワ", "ソー", "オモウ", "!", "ッテ", "カンジ", "?"]
"""
# parsed: OpenJTalkの解析結果
sep_text: list[str] = []
sep_kata: list[str] = []
fix_parsed = []
i = 0
while i <= len(parsed) - 1:
# word: 実際の単語の文字列
# yomi: その読み、但し無声化サインの`’`は除去
# print(parsed)
yomi = parsed[i]["pron"]
tmp_parsed = parsed[i]
if i != len(parsed) - 1 and parsed[i + 1]["string"] in [
"々",
"ゝ",
"ヽ",
"ゞ",
"ヾ",
"゛",
]:
word = parsed[i]["string"] + parsed[i + 1]["string"]
i += 1
else:
word = parsed[i]["string"]
word, yomi = replace_punctuation(word), yomi.replace("’", "")
"""
ここで`yomi`の取りうる値は以下の通りのはず。
- `word`が通常単語 → 通常の読み(カタカナ)
(カタカナからなり、長音記号も含みうる、`アー` 等)
- `word`が`ー` から始まる → `ーラー` や `ーーー` など
- `word`が句読点や空白等 → `、`
- `word`が`?` → `?`(全角になる)
他にも`word`が読めないキリル文字アラビア文字等が来ると`、`になるが、正規化でこの場合は起きないはず。
また元のコードでは`yomi`が空白の場合の処理があったが、これは起きないはず。
処理すべきは`yomi`が`、`の場合のみのはず。
"""
assert yomi != "", f"Empty yomi: {word}"
if yomi == "、":
# wordは正規化されているので、`.`, `,`, `!`, `'`, `-`のいずれか
if word not in (
".",
",",
"!",
"'",
"-",
"?",
":",
";",
"…",
"",
):
# ここはpyopenjtalkが読めない文字等のときに起こる
#print(
# "{}Cannot read:{}, yomi:{}, new_word:{};".format(
# parsed, word, yomi, self.japan_JH2K.convert(word)[0]["kana"]
# )
#)
# raise ValueError(word)
word = self.japan_JH2K.convert(word)[0]["kana"]
# print(word, self.japan_JH2K.convert(word)[0]['kana'], kata2phoneme_list(self.japan_JH2K.convert(word)[0]['kana']))
tmp_parsed["pron"] = word
# yomi = "-"
# word = ','
# yomiは元の記号のままに変更
# else:
# parsed[i]['pron'] = parsed[i]["string"]
yomi = word
elif yomi == "?":
assert word == "?", f"yomi `?` comes from: {word}"
yomi = "?"
if word == "":
i += 1
continue
sep_text.append(word)
sep_kata.append(yomi)
# print(word, yomi, parts)
fix_parsed.append(tmp_parsed)
i += 1
# print(sep_text, sep_kata)
return sep_text, sep_kata, fix_parsed
def getSentencePhone(self, sentence, blank_mode=True, phoneme_mode=False):
# print("origin:", sentence)
words = []
words_phone_len = []
short_char_flag = False
output_duration_flag = []
output_before_sil_flag = []
normed_text = []
sentence = sentence.strip().strip("'")
sentence = re.sub(r"\s+", "", sentence)
output_res = []
failed_words = []
last_long_pause = 4
last_word = None
frontend_text = pyopenjtalk.run_frontend(sentence)
# print("frontend_text: ", frontend_text)
try:
frontend_text = pyopenjtalk.estimate_accent(frontend_text)
except:
pass
# print("estimate_accent: ", frontend_text)
# sep_text: 単語単位の単語のリスト
# sep_kata: 単語単位の単語のカタカナ読みのリスト
sep_text, sep_kata, frontend_text = self.text2sep_kata(frontend_text)
# print("sep_text: ", sep_text)
# print("sep_kata: ", sep_kata)
# print("frontend_text: ", frontend_text)
# sep_phonemes: 各単語ごとの音素のリストのリスト
sep_phonemes = handle_long_word([kata2phoneme_list(i) for i in sep_kata])
# print("sep_phonemes: ", sep_phonemes)
pron_text = [x["pron"].strip().replace("’", "") for x in frontend_text]
# pdb.set_trace()
prosodys = pyopenjtalk.make_label(frontend_text)
prosodys = frontend2phoneme(prosodys, drop_unvoiced_vowels=True)
# print("prosodys: ", ' '.join(prosodys))
# print("pron_text: ", pron_text)
normed_text = [x["string"].strip() for x in frontend_text]
# punctuationがすべて消えた、音素とアクセントのタプルのリスト
phone_tone_list_wo_punct = g2phone_tone_wo_punct(prosodys)
# print("phone_tone_list_wo_punct: ", phone_tone_list_wo_punct)
# phone_w_punct: sep_phonemesを結合した、punctuationを元のまま保持した音素列
phone_w_punct: list[str] = []
w_p_len = []
for i in sep_phonemes:
phone_w_punct += i
w_p_len.append(len(i))
phone_w_punct = phone_w_punct[:-1]
# punctuation無しのアクセント情報を使って、punctuationを含めたアクセント情報を作る
# print("phone_w_punct: ", phone_w_punct)
# print("phone_tone_list_wo_punct: ", phone_tone_list_wo_punct)
phone_tone_list = align_tones(phone_w_punct, phone_tone_list_wo_punct)
jp_item = {}
jp_p = ""
jp_t = ""
# mye rye pye bye nye
# je she
# print(phone_tone_list)
for p, t in phone_tone_list:
if p in self.ipa_dict:
curr_p = self.ipa_dict[p]
jp_p += curr_p
jp_t += str(t + 6) * len(curr_p)
elif p in punctuation:
jp_p += p
jp_t += "0"
elif p == "▁":
jp_p += p
jp_t += " "
else:
print(p, t)
jp_p += "|"
jp_t += "0"
# return phones, tones, w_p_len
jp_p = jp_p.replace("▁", " ")
jp_t = jp_t.translate(self.table)
jp_l = ""
for t in jp_t:
if t == " ":
jp_l += " "
else:
jp_l += "2"
# print(jp_p)
# print(jp_t)
# print(jp_l)
# print(len(jp_p_len), sum(w_p_len), len(jp_p), sum(jp_p_len))
assert len(jp_p) == len(jp_t) and len(jp_p) == len(jp_l)
jp_item["jp_p"] = jp_p.replace("| |", "|").rstrip("|")
jp_item["jp_t"] = jp_t
jp_item["jp_l"] = jp_l
jp_item["jp_normed_text"] = " ".join(normed_text)
jp_item["jp_pron_text"] = " ".join(pron_text)
# jp_item['jp_ruoma'] = sep_phonemes
# print(len(normed_text), len(sep_phonemes))
# print(normed_text)
return jp_item
jpc = JapanesePhoneConverter()
def japanese_to_ipa(text, text_tokenizer):
# phonemes = text_tokenizer(text)
if type(text) == str:
return jpc.getSentencePhone(text)["jp_p"]
else:
result_ph = []
for t in text:
result_ph.append(jpc.getSentencePhone(t)["jp_p"])
return result_ph
Binary file not shown.

Before

Width:  |  Height:  |  Size: 82 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 4.8 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 119 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 3.8 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 363 KiB

+551
View File
@@ -0,0 +1,551 @@
# Copyright (c) 2025 ASLP-LAB
# 2025 Huakang Chen (huakang@mail.nwpu.edu.cn)
# 2025 Guobin Ma (guobin.ma@gmail.com)
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
# http://www.apache.org/licenses/LICENSE-2.0
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import torch
import librosa
import torchaudio
import random
import json
from muq import MuQMuLan, MuQ
from omegaconf import OmegaConf
from safetensors.torch import load_file
from hydra.utils import instantiate
from mutagen.mp3 import MP3
import os
import numpy as np
from sys import path
path.append(os.getcwd())
from diffrhythm.model import DiT, CFM
node_dir = os.path.dirname(os.path.abspath(__file__))
import folder_paths
models_dir = folder_paths.models_dir
model_path = os.path.join(models_dir, "TTS")
def vae_sample(mean, scale):
stdev = torch.nn.functional.softplus(scale) + 1e-4
var = stdev * stdev
logvar = torch.log(var)
latents = torch.randn_like(mean) * stdev + mean
kl = (mean * mean + var - logvar - 1).sum(1).mean()
return latents, kl
def normalize_audio(y, target_dbfs=0):
max_amplitude = torch.max(torch.abs(y))
target_amplitude = 10.0**(target_dbfs / 20.0)
scale_factor = target_amplitude / max_amplitude
normalized_audio = y * scale_factor
return normalized_audio
def set_audio_channels(audio, target_channels):
if target_channels == 1:
# Convert to mono
audio = audio.mean(1, keepdim=True)
elif target_channels == 2:
# Convert to stereo
if audio.shape[1] == 1:
audio = audio.repeat(1, 2, 1)
elif audio.shape[1] > 2:
audio = audio[:, :2, :]
return audio
class PadCrop(torch.nn.Module):
def __init__(self, n_samples, randomize=True):
super().__init__()
self.n_samples = n_samples
self.randomize = randomize
def __call__(self, signal):
n, s = signal.shape
start = 0 if (not self.randomize) else torch.randint(0, max(0, s - self.n_samples) + 1, []).item()
end = start + self.n_samples
output = signal.new_zeros([n, self.n_samples])
output[:, :min(s, self.n_samples)] = signal[:, start:end]
return output
def prepare_audio(audio, in_sr, target_sr, target_length, target_channels, device):
audio = audio.to(device)
if in_sr != target_sr:
resample_tf = torchaudio.transforms.Resample(in_sr, target_sr).to(device)
audio = resample_tf(audio)
if target_length is None:
target_length = audio.shape[-1]
audio = PadCrop(target_length, randomize=False)(audio)
# Add batch dimension
if audio.dim() == 1:
audio = audio.unsqueeze(0).unsqueeze(0)
elif audio.dim() == 2:
audio = audio.unsqueeze(0)
audio = set_audio_channels(audio, target_channels)
return audio
def decode_audio(latents, vae_model, chunked=False, overlap=32, chunk_size=128):
downsampling_ratio = 2048
io_channels = 2
if not chunked:
return vae_model.decode_export(latents)
else:
# chunked decoding
hop_size = chunk_size - overlap
total_size = latents.shape[2]
batch_size = latents.shape[0]
chunks = []
i = 0
for i in range(0, total_size - chunk_size + 1, hop_size):
chunk = latents[:, :, i : i + chunk_size]
chunks.append(chunk)
if i + chunk_size != total_size:
# Final chunk
chunk = latents[:, :, -chunk_size:]
chunks.append(chunk)
chunks = torch.stack(chunks)
num_chunks = chunks.shape[0]
# samples_per_latent is just the downsampling ratio
samples_per_latent = downsampling_ratio
# Create an empty waveform, we will populate it with chunks as decode them
y_size = total_size * samples_per_latent
y_final = torch.zeros((batch_size, io_channels, y_size)).to(latents.device)
for i in range(num_chunks):
x_chunk = chunks[i, :]
# decode the chunk
y_chunk = vae_model.decode_export(x_chunk)
# figure out where to put the audio along the time domain
if i == num_chunks - 1:
# final chunk always goes at the end
t_end = y_size
t_start = t_end - y_chunk.shape[2]
else:
t_start = i * hop_size * samples_per_latent
t_end = t_start + chunk_size * samples_per_latent
# remove the edges of the overlaps
ol = (overlap // 2) * samples_per_latent
chunk_start = 0
chunk_end = y_chunk.shape[2]
if i > 0:
# no overlap for the start of the first chunk
t_start += ol
chunk_start += ol
if i < num_chunks - 1:
# no overlap for the end of the last chunk
t_end -= ol
chunk_end -= ol
# paste the chunked audio into our y_final output audio
y_final[:, :, t_start:t_end] = y_chunk[:, :, chunk_start:chunk_end]
return y_final
def encode_audio(audio, vae_model, chunked=False, overlap=32, chunk_size=128):
downsampling_ratio = 2048
latent_dim = 128
if not chunked:
# default behavior. Encode the entire audio in parallel
return vae_model.encode_export(audio)
else:
# CHUNKED ENCODING
# samples_per_latent is just the downsampling ratio (which is also the upsampling ratio)
samples_per_latent = downsampling_ratio
total_size = audio.shape[2] # in samples
batch_size = audio.shape[0]
chunk_size *= samples_per_latent # converting metric in latents to samples
overlap *= samples_per_latent # converting metric in latents to samples
hop_size = chunk_size - overlap
chunks = []
for i in range(0, total_size - chunk_size + 1, hop_size):
chunk = audio[:,:,i:i+chunk_size]
chunks.append(chunk)
if i+chunk_size != total_size:
# Final chunk
chunk = audio[:,:,-chunk_size:]
chunks.append(chunk)
chunks = torch.stack(chunks)
num_chunks = chunks.shape[0]
# Note: y_size might be a different value from the latent length used in diffusion training
# because we can encode audio of varying lengths
# However, the audio should've been padded to a multiple of samples_per_latent by now.
y_size = total_size // samples_per_latent
# Create an empty latent, we will populate it with chunks as we encode them
y_final = torch.zeros((batch_size,latent_dim,y_size)).to(audio.device)
for i in range(num_chunks):
x_chunk = chunks[i,:]
# encode the chunk
y_chunk = vae_model.encode_export(x_chunk)
# figure out where to put the audio along the time domain
if i == num_chunks-1:
# final chunk always goes at the end
t_end = y_size
t_start = t_end - y_chunk.shape[2]
else:
t_start = i * hop_size // samples_per_latent
t_end = t_start + chunk_size // samples_per_latent
# remove the edges of the overlaps
ol = overlap//samples_per_latent//2
chunk_start = 0
chunk_end = y_chunk.shape[2]
if i > 0:
# no overlap for the start of the first chunk
t_start += ol
chunk_start += ol
if i < num_chunks-1:
# no overlap for the end of the last chunk
t_end -= ol
chunk_end -= ol
# paste the chunked audio into our y_final output audio
y_final[:,:,t_start:t_end] = y_chunk[:,:,chunk_start:chunk_end]
return y_final
def prepare_model(max_frames, device, model_name):
# prepare cfm model
dit_ckpt_path = os.path.join(model_path, "DiffRhythm", model_name)
dit_config_path = f"{node_dir}/diffrhythm/config/diffrhythm-1b.json"
with open(dit_config_path) as f:
model_config = json.load(f)
dit_model_cls = DiT
cfm = CFM(
transformer=dit_model_cls(**model_config["model"], max_frames=max_frames),
num_channels=model_config["model"]['mel_dim'],
)
cfm = cfm.to(device)
cfm = load_checkpoint(cfm, dit_ckpt_path, device=device, use_ema=False)
# prepare tokenizer
tokenizer = CNENTokenizer()
# prepare muq model
try:
from easydict import EasyDict
main_model_dir = f"{model_path}/DiffRhythm/MuQ-MuLan-large"
local_audio_model_dir = f"{model_path}/DiffRhythm/MuQ-large-msd-iter"
local_text_model_dir = f"{model_path}/DiffRhythm/xlm-roberta-base"
config_path = f"{main_model_dir}/config.json"
with open(config_path, 'r') as f:
config_dict = json.load(f)
config_dict['audio_model']['name'] = local_audio_model_dir
config_dict['text_model']['name'] = local_text_model_dir
config_obj = EasyDict(config_dict)
muq = MuQMuLan(config=config_obj, hf_hub_cache_dir=None)
weights_path = f"{main_model_dir}/pytorch_model.bin"
try:
state_dict = torch.load(weights_path, map_location='cpu')
# Adjust loading based on how weights are saved (e.g., remove prefixes if needed)
muq.load_state_dict(state_dict, strict=False) # Use strict=False initially
except FileNotFoundError:
raise FileNotFoundError(f"Weights file not found at {weights_path}")
except Exception as e:
raise
muq = muq.to(device).eval()
# prepare vae
vae_ckpt_path = f"{model_path}/DiffRhythm/vae_model.pt"
vae = torch.jit.load(vae_ckpt_path, map_location="cpu").to(device)
# prepare eval model
train_config = OmegaConf.load(f"{model_path}/DiffRhythm/eval-model/eval.yaml")
checkpoint_path = f"{model_path}/DiffRhythm/eval-model/eval.safetensors"
eval_model = instantiate(train_config.generator).to(device).eval()
state_dict = load_file(checkpoint_path, device="cpu")
eval_model.load_state_dict(state_dict)
eval_muq = MuQ.from_pretrained(f"{model_path}/DiffRhythm/MuQ-large-msd-iter")
eval_muq = eval_muq.to(device).eval()
return cfm, tokenizer, muq, vae, eval_model, eval_muq
# for song edit, will be added in the future
def get_reference_latent(device, max_frames, edit, pred_segments, ref_song, vae_model):
sampling_rate = 44100
downsample_rate = 2048
io_channels = 2
if edit:
input_audio, in_sr = torchaudio.load(ref_song)
input_audio = prepare_audio(input_audio, in_sr=in_sr, target_sr=sampling_rate, target_length=None, target_channels=io_channels, device=device)
input_audio = normalize_audio(input_audio, -6)
with torch.no_grad():
latent = encode_audio(input_audio, vae_model, chunked=True) # [b d t]
mean, scale = latent.chunk(2, dim=1)
prompt, _ = vae_sample(mean, scale)
prompt = prompt.transpose(1, 2) # [b t d]
prompt = prompt[:,:max_frames,:] if prompt.shape[1] >= max_frames else torch.nn.functional.pad(prompt, (0, 0, 0, max_frames - prompt.shape[1]), mode="constant", value=0)
pred_segments = json.loads(pred_segments)
# import pdb; pdb.set_trace()
pred_frames = []
for st, et in pred_segments:
sf = 0 if st == -1 else int(st * sampling_rate / downsample_rate)
# if st == -1:
# sf = 0
# else:
# sf = int(st * sampling_rate / downsample_rate )
ef = max_frames if et == -1 else int(et * sampling_rate / downsample_rate)
# if et == -1:
# ef = max_frames
# else:
# ef = int(et * sampling_rate / downsample_rate )
pred_frames.append((sf, ef))
# import pdb; pdb.set_trace()
return prompt, pred_frames
else:
prompt = torch.zeros(1, max_frames, 64).to(device)
pred_frames = [(0, max_frames)]
return prompt, pred_frames
def get_negative_style_prompt(device):
file_path = f"{node_dir}/diffrhythm/vocal.npy"
vocal_stlye = np.load(file_path)
vocal_stlye = torch.from_numpy(vocal_stlye).to(device) # [1, 512]
vocal_stlye = vocal_stlye.half()
return vocal_stlye
@torch.no_grad()
def eval_song(eval_model, eval_muq, songs):
resampled_songs = [torchaudio.functional.resample(song.mean(dim=0, keepdim=True), 44100, 24000) for song in songs]
ssl_list = []
for i in range(len(resampled_songs)):
output = eval_muq(resampled_songs[i], output_hidden_states=True)
muq_ssl = output["hidden_states"][6]
ssl_list.append(muq_ssl.squeeze(0))
ssl = torch.stack(ssl_list)
scores_g = eval_model(ssl)
score = torch.mean(scores_g, dim=1)
idx = score.argmax(dim=0)
return songs[idx]
@torch.no_grad()
def get_audio_style_prompt(model, wav_path):
vocal_flag = False
mulan = model
audio, _ = librosa.load(wav_path, sr=24000)
audio_len = librosa.get_duration(y=audio, sr=24000)
if audio_len <= 1:
vocal_flag = True
if audio_len > 10:
start_time = int(audio_len // 2 - 5)
wav = audio[start_time*24000:(start_time+10)*24000]
else:
wav = audio
wav = torch.tensor(wav).unsqueeze(0).to(model.device)
with torch.no_grad():
audio_emb = mulan(wavs = wav) # [1, 512]
audio_emb = audio_emb.half()
return audio_emb, vocal_flag
@torch.no_grad()
def get_text_style_prompt(model, text_prompt):
mulan = model
with torch.no_grad():
text_emb = mulan(texts = text_prompt) # [1, 512]
text_emb = text_emb.half()
return text_emb
@torch.no_grad()
def get_style_prompt(model, wav_path=None, prompt=None):
mulan = model
if prompt is not None:
return mulan(texts=prompt).half()
ext = os.path.splitext(wav_path)[-1].lower()
if ext == ".mp3":
meta = MP3(wav_path)
audio_len = meta.info.length
elif ext in [".wav", ".flac"]:
audio_len = librosa.get_duration(path=wav_path)
else:
raise ValueError("Unsupported file format: {}".format(ext))
if audio_len < 10:
print(
f"Warning: The audio file {wav_path} is too short ({audio_len:.2f} seconds). Expected at least 10 seconds."
)
assert audio_len >= 10
mid_time = audio_len // 2
start_time = mid_time - 5
wav, _ = librosa.load(wav_path, sr=24000, offset=start_time, duration=10)
wav = torch.tensor(wav).unsqueeze(0).to(model.device)
with torch.no_grad():
audio_emb = mulan(wavs=wav) # [1, 512]
audio_emb = audio_emb
audio_emb = audio_emb.half()
return audio_emb
def parse_lyrics(lyrics: str):
lyrics_with_time = []
lyrics = lyrics.strip()
for line in lyrics.split("\n"):
try:
time, lyric = line[1:9], line[10:]
lyric = lyric.strip()
mins, secs = time.split(":")
secs = int(mins) * 60 + float(secs)
lyrics_with_time.append((secs, lyric))
except:
continue
return lyrics_with_time
class CNENTokenizer:
def __init__(self):
with open(f"{node_dir}/diffrhythm/g2p/g2p/vocab.json", "r", encoding='utf-8') as file:
self.phone2id: dict = json.load(file)["vocab"]
self.id2phone = {v: k for (k, v) in self.phone2id.items()}
from diffrhythm.g2p.g2p_generation import chn_eng_g2p
self.tokenizer = chn_eng_g2p
def encode(self, text):
phone, token = self.tokenizer(text)
token = [x + 1 for x in token]
return token
def decode(self, token):
return "|".join([self.id2phone[x - 1] for x in token])
def get_lrc_token(max_frames, text, tokenizer, device):
lyrics_shift = 0
sampling_rate = 44100
downsample_rate = 2048
max_secs = max_frames / (sampling_rate / downsample_rate)
comma_token_id = 1
period_token_id = 2
lrc_with_time = parse_lyrics(text)
modified_lrc_with_time = []
for i in range(len(lrc_with_time)):
time, line = lrc_with_time[i]
line_token = tokenizer.encode(line)
modified_lrc_with_time.append((time, line_token))
lrc_with_time = modified_lrc_with_time
lrc_with_time = [
(time_start, line)
for (time_start, line) in lrc_with_time
if time_start < max_secs
]
if max_frames == 2048:
lrc_with_time = lrc_with_time[:-1] if len(lrc_with_time) >= 1 else lrc_with_time
normalized_start_time = 0.0
lrc = torch.zeros((max_frames,), dtype=torch.long)
tokens_count = 0
last_end_pos = 0
for time_start, line in lrc_with_time:
tokens = [
token if token != period_token_id else comma_token_id for token in line
] + [period_token_id]
tokens = torch.tensor(tokens, dtype=torch.long)
num_tokens = tokens.shape[0]
gt_frame_start = int(time_start * sampling_rate / downsample_rate)
frame_shift = random.randint(int(-lyrics_shift), int(lyrics_shift))
frame_start = max(gt_frame_start - frame_shift, last_end_pos)
frame_len = min(num_tokens, max_frames - frame_start)
lrc[frame_start : frame_start + frame_len] = tokens[:frame_len]
tokens_count += num_tokens
last_end_pos = frame_start + frame_len
lrc_emb = lrc.unsqueeze(0).to(device)
normalized_start_time = torch.tensor(normalized_start_time).unsqueeze(0).to(device)
normalized_start_time = normalized_start_time.half()
return lrc_emb, normalized_start_time
def load_checkpoint(model, ckpt_path, device, use_ema=True):
model = model.half()
ckpt_type = ckpt_path.split(".")[-1]
if ckpt_type == "safetensors":
from safetensors.torch import load_file
checkpoint = load_file(ckpt_path)
else:
checkpoint = torch.load(ckpt_path, weights_only=True)
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)
return model.to(device)
-6
View File
@@ -1,6 +0,0 @@
from model.cfm import CFM
from model.dit import DiT
from model.trainer import Trainer
__all__ = ["CFM", "DiT", "Trainer"]
View File
+65
View File
@@ -0,0 +1,65 @@
import numpy as np
import torch
import torch.nn as nn
class Generator(nn.Module):
def __init__(self,
in_features,
ffd_hidden_size,
num_classes,
attn_layer_num,
):
super(Generator, self).__init__()
self.attn = nn.ModuleList(
[
nn.MultiheadAttention(
embed_dim=in_features,
num_heads=8,
dropout=0.2,
batch_first=True,
)
for _ in range(attn_layer_num)
]
)
self.ffd = nn.Sequential(
nn.Linear(in_features, ffd_hidden_size),
nn.ReLU(),
nn.Linear(ffd_hidden_size, in_features)
)
self.dropout = nn.Dropout(0.2)
self.fc = nn.Linear(in_features * 2, num_classes)
self.proj = nn.Tanh()
def forward(self, ssl_feature, judge_id=None):
'''
ssl_feature: [B, T, D]
output: [B, num_classes]
'''
B, T, D = ssl_feature.shape
ssl_feature = self.ffd(ssl_feature)
tmp_ssl_feature = ssl_feature
for attn in self.attn:
tmp_ssl_feature, _ = attn(tmp_ssl_feature, tmp_ssl_feature, tmp_ssl_feature)
ssl_feature = self.dropout(torch.concat([torch.mean(tmp_ssl_feature, dim=1), torch.max(ssl_feature, dim=1)[0]], dim=1)) # B, 2D
x = self.fc(ssl_feature) # B, num_classes
x = self.proj(x) * 2.0 + 3
return x
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "diffrhythm_mw"
description = "Blazingly Fast and Embarrassingly Simple End-to-End Full-Length Song Generation. A node for ComfyUI."
version = "2.1.6"
version = "2.2.0"
license = {file = "LICENSE"}
dependencies = ["# accelerate==1.4.0", "# torchdiffeq==0.2.5", "# torchaudio==2.6.0", "# transformers==4.49.0", "# librosa==0.10.2.post1", "# pyarrow==19.0.1", "# pandas==2.2.3", "# bitsandbytes", "# jieba==0.42.1", "# cn2an==0.5.23", "# pypinyin==0.53.0", "# onnxruntime", "LangSegment", "x-transformers", "pylance", "ema-pytorch", "prefigure", "muq", "mutagen", "pyopenjtalk", "pykakasi", "Unidecode", "phonemizer"]
+6 -2
View File
@@ -9,7 +9,6 @@ jieba
cn2an
pypinyin
onnxruntime
LangSegment
x-transformers
pylance
ema-pytorch
@@ -19,4 +18,9 @@ mutagen
pyopenjtalk
pykakasi
Unidecode
phonemizer
phonemizer
inflect
py3langid
easydict
hydra
omegaconf
File diff suppressed because one or more lines are too long