v1.2
This commit is contained in:
+207
-316
@@ -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"
|
||||
}
|
||||
@@ -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
@@ -4,26 +4,29 @@
|
||||
|
||||
快速而简单的端到端全长歌曲生成.
|
||||
|
||||

|
||||
|
||||

|
||||
|
||||
## 📣 更新
|
||||
|
||||
[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 秒.
|
||||
|
||||

|
||||
|
||||
下载模型放到 `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.
|
||||
|
||||
- 所有参数均是可选的, 不提供任何参数随机生成音乐.
|
||||
## 使用
|
||||
|
||||
- 自动生成歌曲, 自动添加双语歌词字幕:
|
||||
|
||||

|
||||
|
||||
## 安装
|
||||
|
||||
@@ -43,14 +46,19 @@ pip install -r requirements.txt
|
||||
|
||||
结构如下:
|
||||
|
||||

|
||||

|
||||
|
||||
```
|
||||
.
|
||||
| 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 团队的卓越的工作👍.
|
||||
@@ -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.
|
||||
|
||||

|
||||

|
||||
|
||||
## 📣 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.
|
||||
|
||||

|
||||
|
||||
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:
|
||||
|
||||

|
||||
|
||||
## Installation
|
||||
|
||||
@@ -42,14 +46,19 @@ The model needs to be manually downloaded to the `ComfyUI\models\TTS\DiffRhythm`
|
||||
|
||||
The structure is as follows:
|
||||
|
||||

|
||||

|
||||
|
||||
```
|
||||
.
|
||||
| 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
@@ -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'
|
||||
@@ -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)
|
||||
|
||||
@@ -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]耐心的人啊才可以看见童话
|
||||
@@ -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
|
||||
@@ -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]和我在成都的街头走一走
|
||||
@@ -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]直到所有的灯都熄灭了也不停留
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"]
|
||||
@@ -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)
|
||||
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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__(
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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.
@@ -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
@@ -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)
|
||||
@@ -1,6 +0,0 @@
|
||||
from model.cfm import CFM
|
||||
from model.dit import DiT
|
||||
from model.trainer import Trainer
|
||||
|
||||
|
||||
__all__ = ["CFM", "DiT", "Trainer"]
|
||||
@@ -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
@@ -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
@@ -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
Reference in New Issue
Block a user