Initial commit
This commit is contained in:
@@ -0,0 +1,37 @@
|
||||
infer:
|
||||
n_steps: 32
|
||||
cfg: 3
|
||||
|
||||
audio:
|
||||
hop_size: 480
|
||||
sample_rate: 24000
|
||||
max_length: 36000
|
||||
n_fft: 1920
|
||||
num_mels: 128
|
||||
win_size: 1920
|
||||
fmin: 0
|
||||
fmax: 12000
|
||||
mel_var: 8.14
|
||||
mel_mean: -4.92
|
||||
|
||||
model:
|
||||
encoder:
|
||||
vocab_size: 3000
|
||||
text_dim: 512
|
||||
pitch_dim: 512
|
||||
type_dim: 512
|
||||
f0_bin: 361
|
||||
f0_dim: 512
|
||||
num_layers: 4
|
||||
|
||||
flow_matching:
|
||||
mel_dim: 128
|
||||
hidden_size: 1024
|
||||
num_layers: 22
|
||||
num_heads: 16
|
||||
cfg_drop_prob: 0.2
|
||||
use_embedding: False
|
||||
cond_codebook_size: 512
|
||||
cond_scale_factor: 1
|
||||
sigma: 1e-5
|
||||
time_scheduler: cos
|
||||
@@ -0,0 +1,46 @@
|
||||
import torch.nn as nn
|
||||
import torch
|
||||
|
||||
|
||||
class GRN(nn.Module):
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.gamma = nn.Parameter(torch.zeros(1, 1, dim))
|
||||
self.beta = nn.Parameter(torch.zeros(1, 1, dim))
|
||||
|
||||
def forward(self, x):
|
||||
Gx = torch.norm(x, p=2, dim=1, keepdim=True)
|
||||
Nx = Gx / (Gx.mean(dim=-1, keepdim=True) + 1e-6)
|
||||
return self.gamma * (x * Nx) + self.beta + x
|
||||
|
||||
|
||||
# ref: https://github.com/SWivid/F5-TTS/blob/main/src/f5_tts/model/modules.py#L247
|
||||
class ConvNeXtV2Block(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
intermediate_dim: int,
|
||||
dilation: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
padding = (dilation * (7 - 1)) // 2
|
||||
self.dwconv = nn.Conv1d(
|
||||
dim, dim, kernel_size=7, padding=padding, groups=dim, dilation=dilation
|
||||
) # depthwise conv
|
||||
self.norm = nn.LayerNorm(dim, eps=1e-6)
|
||||
self.pwconv1 = nn.Linear(dim, intermediate_dim) # pointwise/1x1 convs, implemented with linear layers
|
||||
self.act = nn.GELU()
|
||||
self.grn = GRN(intermediate_dim)
|
||||
self.pwconv2 = nn.Linear(intermediate_dim, dim)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
residual = x
|
||||
x = x.transpose(1, 2) # b n d -> b d n
|
||||
x = self.dwconv(x)
|
||||
x = x.transpose(1, 2) # b d n -> b n d
|
||||
x = self.norm(x)
|
||||
x = self.pwconv1(x)
|
||||
x = self.act(x)
|
||||
x = self.grn(x)
|
||||
x = self.pwconv2(x)
|
||||
return residual + x
|
||||
@@ -0,0 +1,29 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from soulxsinger.models.modules.flow_matching import FlowMatchingTransformer
|
||||
|
||||
|
||||
class CFMDecoder(nn.Module):
|
||||
def __init__(self, config):
|
||||
super(CFMDecoder, self).__init__()
|
||||
self.model = FlowMatchingTransformer(cfg=config, **config)
|
||||
|
||||
def forward(self, mel, x_mask, decoder_inp, is_prompt):
|
||||
outputs = self.model(mel, x_mask, decoder_inp, is_prompt)
|
||||
|
||||
noise, x, flow_pred, final_mask, prompt_len = outputs["output"]
|
||||
return noise, x, flow_pred, final_mask, prompt_len
|
||||
|
||||
def reverse_diffusion(self, pt_mel, pt_decoder_inp, gt_decoder_inp, n_timesteps=32, cfg=1):
|
||||
diffusion_cond = torch.cat([pt_decoder_inp, gt_decoder_inp], dim=1)
|
||||
diffusion_cond_emb = self.model.cond_emb(diffusion_cond)
|
||||
diffusion_prompt = pt_mel
|
||||
|
||||
generated = self.model.reverse_diffusion(
|
||||
diffusion_cond_emb,
|
||||
diffusion_prompt,
|
||||
n_timesteps=n_timesteps,
|
||||
cfg=cfg
|
||||
)
|
||||
return generated
|
||||
@@ -0,0 +1,445 @@
|
||||
# https://github.com/open-mmlab/Amphion/blob/main/models/svc/flow_matching_transformer/fmt_model.py
|
||||
|
||||
# Copyright (c) 2023 Amphion.
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
import torch.nn as nn
|
||||
import math
|
||||
from .llama import DiffLlama
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class FlowMatchingTransformer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
mel_dim=100,
|
||||
hidden_size=1024,
|
||||
num_layers=12,
|
||||
num_heads=16,
|
||||
cfg_drop_prob=0.2,
|
||||
use_embedding=True,
|
||||
cond_codebook_size=1024,
|
||||
cond_scale_factor=1,
|
||||
sigma=1e-5,
|
||||
time_scheduler="linear",
|
||||
cfg=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.cfg = cfg
|
||||
|
||||
if cfg is not None:
|
||||
mel_dim = getattr(cfg, "mel_dim", mel_dim)
|
||||
hidden_size = getattr(cfg, "hidden_size", hidden_size)
|
||||
num_layers = getattr(cfg, "num_layers", num_layers)
|
||||
num_heads = getattr(cfg, "num_heads", num_heads)
|
||||
cfg_drop_prob = getattr(cfg, "cfg_drop_prob", cfg_drop_prob)
|
||||
cond_codebook_size = getattr(cfg, "cond_codebook_size", cond_codebook_size)
|
||||
time_scheduler = getattr(cfg, "time_scheduler", time_scheduler)
|
||||
sigma = getattr(cfg, "sigma", sigma)
|
||||
cond_scale_factor = getattr(cfg, "cond_scale_factor", cond_scale_factor)
|
||||
|
||||
self.mel_dim = mel_dim
|
||||
self.hidden_size = hidden_size
|
||||
self.num_layers = num_layers
|
||||
self.num_heads = num_heads
|
||||
self.cfg_drop_prob = cfg_drop_prob
|
||||
self.cond_codebook_size = cond_codebook_size
|
||||
self.time_scheduler = time_scheduler
|
||||
self.sigma = sigma
|
||||
self.cond_scale_factor = cond_scale_factor
|
||||
|
||||
if use_embedding:
|
||||
self.cond_emb = nn.Embedding(cond_codebook_size, self.hidden_size)
|
||||
else:
|
||||
self.cond_emb = nn.Linear(cond_codebook_size, self.hidden_size)
|
||||
|
||||
if cond_scale_factor != 1:
|
||||
self.do_resampling = True
|
||||
assert np.log2(cond_scale_factor).is_integer()
|
||||
|
||||
up_layers = []
|
||||
for _ in range(int(np.log2(cond_scale_factor))):
|
||||
up_layers.extend(
|
||||
[
|
||||
nn.ConvTranspose1d(
|
||||
hidden_size, hidden_size, kernel_size=4, stride=2, padding=1
|
||||
),
|
||||
nn.GELU(),
|
||||
]
|
||||
)
|
||||
self.resampling_layers = nn.Sequential(*up_layers)
|
||||
else:
|
||||
self.do_resampling = False
|
||||
|
||||
### REPA: Use the Wav2Vec2Bert features to align. ###
|
||||
self.use_repa = "repa" in cfg
|
||||
self.repa_layer_index = None
|
||||
if self.use_repa:
|
||||
self.repa_layer_index = cfg.repa.layer_index
|
||||
|
||||
self.repa_mlp_layer = nn.Sequential(
|
||||
nn.Linear(hidden_size, hidden_size * 4),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size * 4, cfg.repa.output_dim),
|
||||
)
|
||||
|
||||
### CTC: Use the ASR loss ###
|
||||
self.use_ctc = "ctc" in cfg
|
||||
self.ctc_layer_index = None
|
||||
if self.use_ctc:
|
||||
self.ctc_layer_index = cfg.ctc.layer_index
|
||||
|
||||
self.ctc_mlp_layer = nn.Sequential(
|
||||
nn.Linear(hidden_size, hidden_size * 4),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size * 4, cfg.ctc.output_dim),
|
||||
)
|
||||
|
||||
self.reset_parameters()
|
||||
|
||||
self.diff_estimator = DiffLlama(
|
||||
mel_dim=mel_dim,
|
||||
hidden_size=hidden_size,
|
||||
num_heads=num_heads,
|
||||
num_layers=num_layers,
|
||||
)
|
||||
|
||||
self.sigma = sigma
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_diffusion(self, x, t, is_prompt=None):
|
||||
"""
|
||||
x: (B, T, mel_dim)
|
||||
t: (B,)
|
||||
"""
|
||||
new_t = t
|
||||
t = t.unsqueeze(-1).unsqueeze(-1)
|
||||
z = torch.randn(
|
||||
x.shape, dtype=x.dtype, device=x.device, requires_grad=False
|
||||
) # (B, T, mel_dim)
|
||||
|
||||
# get prompt len
|
||||
if torch.rand(1) <= self.cfg_drop_prob:
|
||||
prompt_len = torch.zeros(x.shape[0]).to(x)
|
||||
is_prompt = torch.zeros_like(x[:, :, 0])
|
||||
else:
|
||||
if is_prompt is None:
|
||||
prompt_len = torch.randint(
|
||||
min(x.shape[1] // 4, 5), int(x.shape[1] * 0.4), (x.shape[0],)
|
||||
).to(
|
||||
x.device
|
||||
) # (B,)
|
||||
|
||||
# get is_prompt
|
||||
is_prompt = torch.zeros_like(x[:, :, 0]) # (B, T)
|
||||
col_indices = (
|
||||
torch.arange(is_prompt.shape[1])
|
||||
.repeat(is_prompt.shape[0], 1)
|
||||
.to(prompt_len)
|
||||
) # (B, T)
|
||||
is_prompt[col_indices < prompt_len.unsqueeze(1)] = 1 # (B, T) 1 if prompt
|
||||
else:
|
||||
prompt_len = is_prompt.sum(dim=1) # (B,)
|
||||
|
||||
mask = torch.ones_like(x[:, :, 0]) # mask if 1, not mask if 0
|
||||
mask[is_prompt.bool()] = 0
|
||||
mask = mask[:, :, None]
|
||||
|
||||
# flow matching: xt = (1 - (1 - sigma) * t) * x0 + t * x; where x0 ~ N(0, 1), x is a sample
|
||||
# flow gt: x - (1 - sigma) * x0 = x - (1 - sigma) * noise
|
||||
xt = ((1 - (1 - self.sigma) * t) * z + t * x) * mask + x * (1 - mask)
|
||||
|
||||
return xt, z, new_t, prompt_len, mask
|
||||
|
||||
def loss_t(
|
||||
self,
|
||||
x,
|
||||
x_mask,
|
||||
t,
|
||||
cond=None,
|
||||
is_prompt=None
|
||||
):
|
||||
xt, z, new_t, prompt_len, mask = self.forward_diffusion(x, t, is_prompt)
|
||||
|
||||
noise = z
|
||||
|
||||
# drop all condition for cfg, so if prompt_len is 0, we also drop cond
|
||||
if cond is not None:
|
||||
cond = cond * torch.where(
|
||||
prompt_len > 0,
|
||||
torch.ones_like(prompt_len),
|
||||
torch.zeros_like(prompt_len),
|
||||
).to(cond.device).unsqueeze(-1).unsqueeze(-1)
|
||||
|
||||
dit_output = self.diff_estimator(xt, new_t, cond, x_mask, return_dict=True)
|
||||
flow_pred = dit_output["output"] # (B, T, mel_dim)
|
||||
|
||||
# final mask used for loss calculation
|
||||
final_mask = mask * x_mask[..., None] # (B, T, 1)
|
||||
|
||||
results = {"output": (noise, x, flow_pred, final_mask, prompt_len)}
|
||||
|
||||
if self.use_repa:
|
||||
repa_hidden_states = dit_output["hidden_states"][
|
||||
self.repa_layer_index
|
||||
] # (B, T, hidden_size)
|
||||
|
||||
repa_pred = self.repa_mlp_layer(repa_hidden_states) # (B, T, repa_dim)
|
||||
results["repa"] = repa_pred
|
||||
|
||||
if self.use_ctc:
|
||||
ctc_hidden_states = dit_output["hidden_states"][
|
||||
self.ctc_layer_index
|
||||
] # (B, T, hidden_size)
|
||||
ctc_pred = self.ctc_mlp_layer(ctc_hidden_states) # (B, T, ctc_dim)
|
||||
results["ctc"] = ctc_pred
|
||||
|
||||
return results
|
||||
|
||||
def compute_loss(self, x, x_mask, cond=None, is_prompt=None):
|
||||
# x0: (B, T, num_quantizer)
|
||||
# x_mask: (B, T) mask is 0 for padding
|
||||
t = torch.rand(x.shape[0], device=x.device, requires_grad=False)
|
||||
t = torch.clamp(t, 1e-5, 1.0)
|
||||
# from CosyVoice: considering the generation process at the beginning is harder than follows, we involve a cosine scheduler for the timestep t
|
||||
if self.time_scheduler == "cos":
|
||||
t = 1 - torch.cos(t * math.pi * 0.5)
|
||||
else:
|
||||
pass
|
||||
return self.loss_t(x, x_mask, t, cond, is_prompt)
|
||||
|
||||
def reset_parameters(self):
|
||||
def _reset_parameters(m):
|
||||
if isinstance(m, nn.MultiheadAttention):
|
||||
if m._qkv_same_embed_dim:
|
||||
nn.init.normal_(m.in_proj_weight, std=0.02)
|
||||
else:
|
||||
nn.init.normal_(m.q_proj_weight, std=0.02)
|
||||
nn.init.normal_(m.k_proj_weight, std=0.02)
|
||||
nn.init.normal_(m.v_proj_weight, std=0.02)
|
||||
|
||||
if m.in_proj_bias is not None:
|
||||
nn.init.constant_(m.in_proj_bias, 0.0)
|
||||
nn.init.constant_(m.out_proj.bias, 0.0)
|
||||
if m.bias_k is not None:
|
||||
nn.init.xavier_normal_(m.bias_k)
|
||||
if m.bias_v is not None:
|
||||
nn.init.xavier_normal_(m.bias_v)
|
||||
|
||||
elif (
|
||||
isinstance(m, nn.Conv1d)
|
||||
or isinstance(m, nn.ConvTranspose1d)
|
||||
or isinstance(m, nn.Conv2d)
|
||||
or isinstance(m, nn.ConvTranspose2d)
|
||||
):
|
||||
m.weight.data.normal_(0.0, 0.02)
|
||||
|
||||
elif isinstance(m, nn.Linear):
|
||||
m.weight.data.normal_(mean=0.0, std=0.02)
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
elif isinstance(m, nn.Embedding):
|
||||
m.weight.data.normal_(mean=0.0, std=0.02)
|
||||
if m.padding_idx is not None:
|
||||
m.weight.data[m.padding_idx].zero_()
|
||||
|
||||
self.apply(_reset_parameters)
|
||||
|
||||
@torch.no_grad()
|
||||
def reverse_diffusion(
|
||||
self,
|
||||
cond,
|
||||
prompt,
|
||||
x_mask=None,
|
||||
prompt_mask=None,
|
||||
n_timesteps=10,
|
||||
cfg=1.0,
|
||||
rescale_cfg=0.75,
|
||||
):
|
||||
h = 1.0 / n_timesteps
|
||||
prompt_len = prompt.shape[1]
|
||||
target_len = cond.shape[1] - prompt_len
|
||||
|
||||
if x_mask == None:
|
||||
x_mask = torch.ones(cond.shape[0], target_len).to(cond.device) # (B, T)
|
||||
if prompt_mask == None:
|
||||
prompt_mask = torch.ones(cond.shape[0], prompt_len).to(
|
||||
cond.device
|
||||
) # (B, prompt_len)
|
||||
xt_mask = torch.cat([prompt_mask, x_mask], dim=1)
|
||||
z = torch.randn(
|
||||
(cond.shape[0], target_len, self.mel_dim),
|
||||
dtype=cond.dtype,
|
||||
device=cond.device,
|
||||
requires_grad=False,
|
||||
)
|
||||
xt = z
|
||||
|
||||
# t from 0 to 1: x0 = z ~ N(0, 1)
|
||||
for i in range(n_timesteps):
|
||||
xt_input = torch.cat([prompt, xt], dim=1)
|
||||
t = (0 + (i + 0.5) * h) * torch.ones(
|
||||
z.shape[0], dtype=z.dtype, device=z.device
|
||||
)
|
||||
flow_pred = self.diff_estimator(xt_input, t, cond, xt_mask)
|
||||
flow_pred = flow_pred[:, prompt_len:, :]
|
||||
|
||||
# cfg
|
||||
if cfg > 0:
|
||||
uncond_flow_pred = self.diff_estimator(
|
||||
xt, t, torch.zeros_like(cond)[:, : xt.shape[1], :], x_mask
|
||||
)
|
||||
pos_flow_pred_std = flow_pred.std()
|
||||
flow_pred_cfg = flow_pred + cfg * (flow_pred - uncond_flow_pred)
|
||||
rescale_flow_pred = (
|
||||
flow_pred_cfg * pos_flow_pred_std / flow_pred_cfg.std()
|
||||
)
|
||||
flow_pred = (
|
||||
rescale_cfg * rescale_flow_pred + (1 - rescale_cfg) * flow_pred_cfg
|
||||
)
|
||||
|
||||
dxt = flow_pred * h
|
||||
xt = xt + dxt
|
||||
|
||||
return xt
|
||||
|
||||
@torch.no_grad()
|
||||
def reverse_diffusion_v2(
|
||||
self,
|
||||
cond,
|
||||
prompt,
|
||||
x_mask=None,
|
||||
prompt_mask=None,
|
||||
n_timesteps=10,
|
||||
cfg=1.0,
|
||||
rescale_cfg=0.75,
|
||||
):
|
||||
h = 1.0 / n_timesteps
|
||||
prompt_len = prompt.shape[1]
|
||||
target_len = cond.shape[1] - prompt_len * 2
|
||||
|
||||
if x_mask == None:
|
||||
x_mask = torch.ones(cond.shape[0], target_len).to(cond.device) # (B, T)
|
||||
if prompt_mask == None:
|
||||
prompt_mask = torch.ones(cond.shape[0], prompt_len).to(
|
||||
cond.device
|
||||
) # (B, prompt_len)
|
||||
xt_mask = torch.cat([prompt_mask, x_mask, prompt_mask], dim=1)
|
||||
z = torch.randn(
|
||||
(cond.shape[0], target_len, self.mel_dim),
|
||||
dtype=cond.dtype,
|
||||
device=cond.device,
|
||||
requires_grad=False,
|
||||
)
|
||||
xt = z
|
||||
|
||||
# t from 0 to 1: x0 = z ~ N(0, 1)
|
||||
for i in range(n_timesteps):
|
||||
xt_input = torch.cat([prompt, xt, prompt], dim=1)
|
||||
t = (0 + (i + 0.5) * h) * torch.ones(
|
||||
z.shape[0], dtype=z.dtype, device=z.device
|
||||
)
|
||||
flow_pred = self.diff_estimator(xt_input, t, cond, xt_mask)
|
||||
flow_pred = flow_pred[:, prompt_len:-prompt_len, :]
|
||||
|
||||
# cfg
|
||||
if cfg > 0:
|
||||
uncond_flow_pred = self.diff_estimator(
|
||||
xt, t, torch.zeros_like(cond)[:, : xt.shape[1], :], x_mask
|
||||
)
|
||||
pos_flow_pred_std = flow_pred.std()
|
||||
flow_pred_cfg = flow_pred + cfg * (flow_pred - uncond_flow_pred)
|
||||
rescale_flow_pred = (
|
||||
flow_pred_cfg * pos_flow_pred_std / flow_pred_cfg.std()
|
||||
)
|
||||
flow_pred = (
|
||||
rescale_cfg * rescale_flow_pred + (1 - rescale_cfg) * flow_pred_cfg
|
||||
)
|
||||
|
||||
dxt = flow_pred * h
|
||||
xt = xt + dxt
|
||||
|
||||
return xt
|
||||
|
||||
def forward(self, x, x_mask, cond_code, is_prompt=None):
|
||||
"""
|
||||
Args:
|
||||
x: (B, T, mel_dim)
|
||||
x_mask: (B, T)
|
||||
cond_code: (B, T), Note that cond_code might be not at 50Hz!
|
||||
"""
|
||||
T = x.shape[1]
|
||||
|
||||
cond = self.cond_emb(cond_code) # (B, T, hidden_size)
|
||||
if self.do_resampling:
|
||||
# Align to the frame rate of Mels
|
||||
cond = self.resampling_layers(cond.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
# print("cond_code: {}, after resampling: {}".format(cond_code.shape, cond.shape))
|
||||
|
||||
if cond.shape[1] >= T: # Check time dimension
|
||||
cond = cond[:, :T, :]
|
||||
else:
|
||||
padding_frames = T - cond.shape[1]
|
||||
last_frame = cond[:, -1:, :]
|
||||
padding = last_frame.repeat(1, padding_frames, 1)
|
||||
cond = torch.cat([cond, padding], dim=1)
|
||||
|
||||
return self.compute_loss(x, x_mask, cond, is_prompt)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
model_cfg = {
|
||||
"mel_dim": 128,
|
||||
"hidden_size": 256,
|
||||
"num_layers": 8,
|
||||
"num_heads": 8,
|
||||
"cfg_drop_prob": 0.2,
|
||||
"use_embedding": False,
|
||||
"cond_codebook_size": 256,
|
||||
"cond_scale_factor": 1,
|
||||
"sigma": 1e-5,
|
||||
"time_scheduler": "cos",
|
||||
}
|
||||
|
||||
device = "cuda"
|
||||
x = torch.randn(2, 100, 128).to(device)
|
||||
x_mask = torch.ones(2, 100).to(device)
|
||||
# cond_code = torch.randint(0, 16384, (2, 25)).to(device)
|
||||
cond_code = torch.randn(2, 100, 256).to(device)
|
||||
|
||||
model = FlowMatchingTransformer(cfg=model_cfg, **model_cfg).to(device)
|
||||
outputs = model(x, x_mask, cond_code)
|
||||
print(outputs)
|
||||
|
||||
noise, x, flow_pred, final_mask, prompt_len = outputs["output"]
|
||||
final_mask = final_mask.squeeze(-1)
|
||||
|
||||
flow_gt = x - (1 - 1e-5) * noise
|
||||
|
||||
# [B, n_frames, D]
|
||||
diff_loss = F.l1_loss(
|
||||
flow_pred, flow_gt, reduction="none"
|
||||
).float() * final_mask.unsqueeze(-1)
|
||||
diff_loss = torch.mean(diff_loss, dim=2).sum() / final_mask.sum()
|
||||
|
||||
print("diff_loss:", diff_loss.item())
|
||||
|
||||
|
||||
diffusion_cond = torch.randn(2, 150, 256).to(device)
|
||||
diffusion_cond_emb = model.cond_emb(diffusion_cond)
|
||||
diffusion_prompt = torch.randn(2, 50, 128).to(device)
|
||||
n_timesteps = 32
|
||||
|
||||
generated = model.reverse_diffusion(
|
||||
diffusion_cond_emb,
|
||||
diffusion_prompt,
|
||||
n_timesteps=n_timesteps
|
||||
)
|
||||
print("generated:", generated.shape)
|
||||
@@ -0,0 +1,392 @@
|
||||
from transformers import LlamaConfig, LlamaModel
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from typing import List, Optional, Tuple, Union
|
||||
import math
|
||||
|
||||
from transformers.models.llama.modeling_llama import LlamaDecoderLayer
|
||||
from transformers.models.llama.modeling_llama import BaseModelOutputWithPast
|
||||
|
||||
|
||||
# sinusoidal positional encoding
|
||||
class SinusoidalPosEmb(nn.Module):
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
|
||||
def forward(self, x):
|
||||
device = x.device
|
||||
half_dim = self.dim // 2
|
||||
emb = math.log(10000) / (half_dim - 1)
|
||||
emb = torch.exp(torch.arange(half_dim, device=device) * -emb)
|
||||
emb = x[:, None] * emb[None, :] * 1.0
|
||||
emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
|
||||
return emb
|
||||
|
||||
|
||||
class LlamaAdaptiveRMSNorm(nn.Module):
|
||||
def __init__(self, hidden_size=1024, eps=1e-6, dim_cond=1024):
|
||||
super().__init__()
|
||||
self.to_weight = nn.Linear(dim_cond, hidden_size)
|
||||
nn.init.zeros_(self.to_weight.weight)
|
||||
nn.init.ones_(self.to_weight.bias)
|
||||
self.variance_epsilon = eps
|
||||
self._is_hf_initialized = True # disable automatic init
|
||||
|
||||
def forward(self, hidden_states, cond_embedding):
|
||||
input_dtype = hidden_states.dtype
|
||||
variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True)
|
||||
hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
|
||||
|
||||
weight = self.to_weight(cond_embedding)
|
||||
if len(weight.shape) == 2:
|
||||
weight = weight.unsqueeze(1)
|
||||
|
||||
return (weight * hidden_states).to(input_dtype)
|
||||
|
||||
|
||||
class LlamaNARDecoderLayer(LlamaDecoderLayer):
|
||||
def __init__(self, config: LlamaConfig, layer_idx: int):
|
||||
"""Override to adaptive layer norm"""
|
||||
super().__init__(config, layer_idx) # init attention, mlp, etc.
|
||||
self.input_layernorm = LlamaAdaptiveRMSNorm(
|
||||
config.hidden_size, eps=config.rms_norm_eps, dim_cond=config.hidden_size
|
||||
)
|
||||
self.post_attention_layernorm = LlamaAdaptiveRMSNorm(
|
||||
config.hidden_size, eps=config.rms_norm_eps, dim_cond=config.hidden_size
|
||||
)
|
||||
|
||||
# add `cond` in forward function
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
cond_embedding: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
past_key_value: Optional[Tuple[torch.Tensor]] = None,
|
||||
output_attentions: Optional[bool] = False,
|
||||
use_cache: Optional[bool] = False,
|
||||
) -> Tuple[
|
||||
torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]
|
||||
]:
|
||||
"""
|
||||
Args:
|
||||
hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`
|
||||
attention_mask (`torch.FloatTensor`, *optional*): attention mask of size
|
||||
`(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.
|
||||
output_attentions (`bool`, *optional*):
|
||||
Whether or not to return the attentions tensors of all attention layers. See `attentions` under
|
||||
returned tensors for more detail.
|
||||
use_cache (`bool`, *optional*):
|
||||
If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding
|
||||
(see `past_key_values`).
|
||||
past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states
|
||||
"""
|
||||
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.input_layernorm(
|
||||
hidden_states, cond_embedding=cond_embedding
|
||||
)
|
||||
|
||||
# Self Attention
|
||||
hidden_states, self_attn_weights, present_key_value = self.self_attn(
|
||||
hidden_states=hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
position_ids=position_ids,
|
||||
past_key_value=past_key_value,
|
||||
output_attentions=output_attentions,
|
||||
use_cache=use_cache,
|
||||
)
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
# Fully Connected
|
||||
residual = hidden_states
|
||||
hidden_states = self.post_attention_layernorm(
|
||||
hidden_states, cond_embedding=cond_embedding
|
||||
)
|
||||
hidden_states = self.mlp(hidden_states)
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
outputs = (hidden_states,)
|
||||
|
||||
if output_attentions:
|
||||
outputs += (self_attn_weights,)
|
||||
|
||||
if use_cache:
|
||||
outputs += (present_key_value,)
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
class DiffLlama(LlamaModel):
|
||||
def __init__(
|
||||
self,
|
||||
mel_dim=100,
|
||||
hidden_size=1024,
|
||||
num_heads=16,
|
||||
num_layers=16,
|
||||
dropout=0.1,
|
||||
ffn_dropout=0.1,
|
||||
attention_dropout=0.0,
|
||||
config=LlamaConfig(0, 256, 1024, 1, 1),
|
||||
):
|
||||
super().__init__(config)
|
||||
|
||||
self.layers = nn.ModuleList(
|
||||
[
|
||||
LlamaNARDecoderLayer(
|
||||
LlamaConfig(
|
||||
hidden_size=hidden_size,
|
||||
num_attention_heads=num_heads,
|
||||
max_position_embeddings=4096,
|
||||
intermediate_size=hidden_size * 4,
|
||||
),
|
||||
layer_idx=i,
|
||||
)
|
||||
for i in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.norm = LlamaAdaptiveRMSNorm(hidden_size, dim_cond=hidden_size)
|
||||
|
||||
self.diff_step_embedding = SinusoidalPosEmb(hidden_size)
|
||||
self.diff_step_mlp = nn.Sequential(
|
||||
nn.Linear(hidden_size, hidden_size * 4),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size * 4, hidden_size),
|
||||
)
|
||||
|
||||
self.cond_mlp = nn.Sequential(
|
||||
nn.Linear(hidden_size, hidden_size * 4),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size * 4, hidden_size),
|
||||
)
|
||||
|
||||
self.mel_mlp = nn.Sequential(
|
||||
nn.Linear(mel_dim, hidden_size * 4),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size * 4, hidden_size),
|
||||
)
|
||||
|
||||
self.mel_out_mlp = nn.Sequential(
|
||||
nn.Linear(hidden_size, hidden_size * 4),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size * 4, mel_dim),
|
||||
)
|
||||
|
||||
for layer in self.layers:
|
||||
layer.input_layernorm = LlamaAdaptiveRMSNorm(
|
||||
hidden_size, dim_cond=hidden_size
|
||||
)
|
||||
layer.post_attention_layernorm = LlamaAdaptiveRMSNorm(
|
||||
hidden_size, dim_cond=hidden_size
|
||||
)
|
||||
|
||||
self.embed_tokens = None
|
||||
|
||||
self.post_init()
|
||||
|
||||
# self.reset_parameters()
|
||||
|
||||
def _prepare_decoder_attention_mask(
|
||||
self, attention_mask, input_shape, inputs_embeds, past_key_values_length
|
||||
):
|
||||
# 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: Optional[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
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
diffusion_step,
|
||||
cond,
|
||||
x_mask,
|
||||
input_ids: torch.LongTensor = None, # [num_quant, B, T]
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
past_key_values: Optional[List[torch.FloatTensor]] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
use_cache: Optional[bool] = None,
|
||||
output_attentions: Optional[bool] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
return_dict: Optional[bool] = False,
|
||||
) -> Union[Tuple, BaseModelOutputWithPast]:
|
||||
|
||||
# retrieve some shape info
|
||||
batch_size, seq_length, _ = x.shape
|
||||
|
||||
# condtion mlp
|
||||
cond_embedding = self.cond_mlp(cond) # (B, T, C)
|
||||
|
||||
# condition mel
|
||||
x = self.mel_mlp(x)
|
||||
|
||||
# diffusion step embedding
|
||||
diffusion_step = self.diff_step_embedding(diffusion_step).to(x.device)
|
||||
diffusion_step = self.diff_step_mlp(diffusion_step) # (B, C)
|
||||
x = x + cond_embedding
|
||||
|
||||
inputs_embeds = x
|
||||
attention_mask = x_mask
|
||||
|
||||
output_attentions = (
|
||||
output_attentions
|
||||
if output_attentions is not None
|
||||
else self.config.output_attentions
|
||||
)
|
||||
output_hidden_states = (
|
||||
output_hidden_states
|
||||
if output_hidden_states is not None
|
||||
else self.config.output_hidden_states
|
||||
)
|
||||
use_cache = use_cache if use_cache is not None else self.config.use_cache
|
||||
|
||||
seq_length_with_past = seq_length
|
||||
past_key_values_length = 0
|
||||
|
||||
if past_key_values is not None:
|
||||
past_key_values_length = past_key_values[0][0].shape[2]
|
||||
seq_length_with_past = seq_length_with_past + past_key_values_length
|
||||
|
||||
if position_ids is None:
|
||||
device = input_ids.device if input_ids is not None else inputs_embeds.device
|
||||
position_ids = torch.arange(
|
||||
past_key_values_length,
|
||||
seq_length + past_key_values_length,
|
||||
dtype=torch.long,
|
||||
device=device,
|
||||
)
|
||||
position_ids = position_ids.unsqueeze(0).view(-1, seq_length)
|
||||
else:
|
||||
position_ids = position_ids.view(-1, seq_length).long()
|
||||
|
||||
# embed positions
|
||||
if attention_mask is None:
|
||||
attention_mask = torch.ones(
|
||||
(batch_size, seq_length_with_past),
|
||||
dtype=torch.bool,
|
||||
device=inputs_embeds.device,
|
||||
)
|
||||
attention_mask = self._prepare_decoder_attention_mask(
|
||||
attention_mask,
|
||||
(batch_size, seq_length),
|
||||
inputs_embeds,
|
||||
past_key_values_length,
|
||||
)
|
||||
|
||||
hidden_states = inputs_embeds
|
||||
|
||||
if self.gradient_checkpointing and self.training:
|
||||
if use_cache:
|
||||
use_cache = False
|
||||
|
||||
# decoder layers
|
||||
all_hidden_states = () if output_hidden_states else None
|
||||
all_self_attns = () if output_attentions else None
|
||||
next_decoder_cache = () if use_cache else None
|
||||
|
||||
all_layer_hidden_states = []
|
||||
|
||||
for idx, decoder_layer in enumerate(self.layers):
|
||||
if output_hidden_states:
|
||||
all_hidden_states += (hidden_states,)
|
||||
|
||||
past_key_value = (
|
||||
past_key_values[idx] if past_key_values is not None else None
|
||||
)
|
||||
|
||||
if self.gradient_checkpointing and self.training:
|
||||
raise NotImplementedError
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
# None for past_key_value
|
||||
return module(*inputs, output_attentions, None)
|
||||
|
||||
return custom_forward
|
||||
|
||||
layer_outputs = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(decoder_layer),
|
||||
hidden_states,
|
||||
attention_mask,
|
||||
position_ids,
|
||||
None,
|
||||
)
|
||||
else:
|
||||
layer_outputs = decoder_layer(
|
||||
hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
position_ids=position_ids,
|
||||
past_key_value=past_key_value,
|
||||
output_attentions=output_attentions,
|
||||
use_cache=use_cache,
|
||||
cond_embedding=diffusion_step,
|
||||
)
|
||||
|
||||
hidden_states = layer_outputs[0]
|
||||
all_layer_hidden_states.append(hidden_states.clone())
|
||||
|
||||
if use_cache:
|
||||
next_decoder_cache += (layer_outputs[2 if output_attentions else 1],)
|
||||
|
||||
if output_attentions:
|
||||
all_self_attns += (layer_outputs[1],)
|
||||
|
||||
hidden_states = self.norm(hidden_states, cond_embedding=diffusion_step)
|
||||
|
||||
# add hidden states from the last decoder layer
|
||||
if output_hidden_states:
|
||||
all_hidden_states += (hidden_states,)
|
||||
|
||||
next_cache = next_decoder_cache if use_cache else None
|
||||
|
||||
hidden_states = self.mel_out_mlp(hidden_states)
|
||||
|
||||
# if not return_dict:
|
||||
# return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None)
|
||||
# return BaseModelOutputWithPast(
|
||||
# last_hidden_state=hidden_states,
|
||||
# past_key_values=next_cache,
|
||||
# hidden_states=all_hidden_states,
|
||||
# attentions=all_self_attns,
|
||||
# )
|
||||
if return_dict:
|
||||
return {
|
||||
"output": hidden_states,
|
||||
"hidden_states": all_layer_hidden_states,
|
||||
}
|
||||
|
||||
return hidden_states
|
||||
@@ -0,0 +1,151 @@
|
||||
import torch
|
||||
import math
|
||||
import numpy as np
|
||||
from librosa.filters import mel as librosa_mel_fn
|
||||
import torch.nn as nn
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
|
||||
def dynamic_range_compression(x, C=1, clip_val=1e-5):
|
||||
return np.log(np.clip(x, a_min=clip_val, a_max=None) * C)
|
||||
|
||||
|
||||
def dynamic_range_decompression(x, C=1):
|
||||
return np.exp(x) / C
|
||||
|
||||
|
||||
def dynamic_range_compression_torch(x, C=1, clip_val=1e-5):
|
||||
return torch.log(torch.clamp(x, min=clip_val) * C)
|
||||
|
||||
|
||||
def dynamic_range_decompression_torch(x, C=1):
|
||||
return torch.exp(x) / C
|
||||
|
||||
|
||||
def spectral_normalize_torch(magnitudes):
|
||||
output = dynamic_range_compression_torch(magnitudes)
|
||||
return output
|
||||
|
||||
|
||||
def spectral_de_normalize_torch(magnitudes):
|
||||
output = dynamic_range_decompression_torch(magnitudes)
|
||||
return output
|
||||
|
||||
|
||||
class MelSpectrogram(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
n_fft,
|
||||
num_mels,
|
||||
sampling_rate,
|
||||
hop_size,
|
||||
win_size,
|
||||
fmin,
|
||||
fmax,
|
||||
center=False,
|
||||
):
|
||||
super(MelSpectrogram, self).__init__()
|
||||
self.n_fft = n_fft
|
||||
self.hop_size = hop_size
|
||||
self.win_size = win_size
|
||||
self.sampling_rate = sampling_rate
|
||||
self.num_mels = num_mels
|
||||
self.fmin = fmin
|
||||
self.fmax = fmax
|
||||
self.center = center
|
||||
|
||||
mel_basis = {}
|
||||
hann_window = {}
|
||||
|
||||
mel = librosa_mel_fn(
|
||||
sr=sampling_rate, n_fft=n_fft, n_mels=num_mels, fmin=fmin, fmax=fmax
|
||||
)
|
||||
mel_basis = torch.from_numpy(mel).float()
|
||||
hann_window = torch.hann_window(win_size)
|
||||
|
||||
self.register_buffer("mel_basis", mel_basis)
|
||||
self.register_buffer("hann_window", hann_window)
|
||||
|
||||
def forward(self, y):
|
||||
y = torch.nn.functional.pad(
|
||||
y.unsqueeze(1),
|
||||
(
|
||||
int((self.n_fft - self.hop_size) / 2),
|
||||
int((self.n_fft - self.hop_size) / 2),
|
||||
),
|
||||
mode="reflect",
|
||||
)
|
||||
y = y.squeeze(1)
|
||||
spec = torch.stft(
|
||||
y,
|
||||
self.n_fft,
|
||||
hop_length=self.hop_size,
|
||||
win_length=self.win_size,
|
||||
window=self.hann_window,
|
||||
center=self.center,
|
||||
pad_mode="reflect",
|
||||
normalized=False,
|
||||
onesided=True,
|
||||
return_complex=True,
|
||||
)
|
||||
spec = torch.view_as_real(spec)
|
||||
|
||||
spec = torch.sqrt(spec.pow(2).sum(-1) + (1e-9))
|
||||
|
||||
spec = torch.matmul(self.mel_basis, spec)
|
||||
spec = spectral_normalize_torch(spec)
|
||||
|
||||
return spec
|
||||
|
||||
|
||||
def load_mel_spectrogram():
|
||||
return load_mel_spectrogram_from_cfg(None)
|
||||
|
||||
|
||||
def _get_from_mapping(cfg: Any, key: str, default: Any = None) -> Any:
|
||||
"""Safely read a field from a dict/OmegaConf-like object."""
|
||||
if cfg is None:
|
||||
return default
|
||||
if isinstance(cfg, dict):
|
||||
return cfg.get(key, default)
|
||||
return getattr(cfg, key, default)
|
||||
|
||||
|
||||
def load_mel_spectrogram_from_cfg(audio_cfg: Optional[Any] = None) -> MelSpectrogram:
|
||||
"""Build MelSpectrogram from `audio_config`-like config.
|
||||
|
||||
Expected keys (either in dict or Hydra/OmegaConf object):
|
||||
- hop_size, sample_rate (or sampling_rate), n_fft, num_mels, win_size, fmin, fmax
|
||||
"""
|
||||
# Defaults keep current behavior.
|
||||
mel_cfg: Dict[str, Any] = {
|
||||
"hop_size": _get_from_mapping(audio_cfg, "hop_size", 480),
|
||||
"sampling_rate": _get_from_mapping(
|
||||
audio_cfg,
|
||||
"sampling_rate",
|
||||
_get_from_mapping(audio_cfg, "sample_rate", 24000),
|
||||
),
|
||||
"n_fft": _get_from_mapping(audio_cfg, "n_fft", 1920),
|
||||
"num_mels": _get_from_mapping(audio_cfg, "num_mels", 128),
|
||||
"win_size": _get_from_mapping(audio_cfg, "win_size", 1920),
|
||||
"fmin": _get_from_mapping(audio_cfg, "fmin", 0),
|
||||
"fmax": _get_from_mapping(audio_cfg, "fmax", 12000),
|
||||
}
|
||||
|
||||
mel_model = MelSpectrogram(**mel_cfg)
|
||||
mel_model.eval()
|
||||
return mel_model
|
||||
|
||||
|
||||
class MelSpectrogramEncoder(nn.Module):
|
||||
def __init__(self, audio_config: dict | None = None):
|
||||
super(MelSpectrogramEncoder, self).__init__()
|
||||
self.model = load_mel_spectrogram_from_cfg(audio_config)
|
||||
audio_config = audio_config or {}
|
||||
self.mel_mean = audio_config.get("mel_mean", -4.92)
|
||||
self.mel_var = audio_config.get("mel_var", 8.14)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.model(x).transpose(1, 2)
|
||||
x = (x - self.mel_mean) / math.sqrt(self.mel_var)
|
||||
return x
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,187 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import math
|
||||
import numpy as np
|
||||
from typing import Optional, Dict, Any, List
|
||||
|
||||
from soulxsinger.models.modules.vocoder import Vocoder
|
||||
from soulxsinger.models.modules.decoder import CFMDecoder
|
||||
from soulxsinger.models.modules.convnext import ConvNeXtV2Block
|
||||
from soulxsinger.models.modules.mel_transform import MelSpectrogramEncoder
|
||||
|
||||
|
||||
class SoulXSinger(nn.Module):
|
||||
"""
|
||||
SoulXSinger model.
|
||||
"""
|
||||
def __init__(self, config: Dict):
|
||||
super(SoulXSinger, self).__init__()
|
||||
audio_cfg = config.audio
|
||||
enc_cfg = config.model.encoder
|
||||
cfm_cfg = config.model.flow_matching
|
||||
|
||||
self.note_text_encoder = nn.Embedding(enc_cfg["vocab_size"], enc_cfg["text_dim"])
|
||||
self.note_pitch_encoder = nn.Embedding(256, enc_cfg["pitch_dim"])
|
||||
self.note_type_encoder = nn.Embedding(256, enc_cfg["type_dim"])
|
||||
self.f0_encoder = nn.Embedding(enc_cfg["f0_bin"], enc_cfg["f0_dim"])
|
||||
|
||||
self.preflow = nn.Sequential(
|
||||
*[ConvNeXtV2Block(enc_cfg["text_dim"], enc_cfg["text_dim"] * 2) for _ in range(enc_cfg["num_layers"])]
|
||||
)
|
||||
self.cfm_decoder = CFMDecoder(cfm_cfg)
|
||||
|
||||
if audio_cfg is None and isinstance(enc_cfg, dict):
|
||||
audio_cfg = enc_cfg.get("audio_config")
|
||||
self.mel = MelSpectrogramEncoder(audio_cfg)
|
||||
self.vocoder = Vocoder()
|
||||
|
||||
@staticmethod
|
||||
def expand_states(h, mel2token):
|
||||
"""
|
||||
Expand the states to the mel-scale.
|
||||
args:
|
||||
h: states, shape: [B, T, H]
|
||||
mel2token: mel2token, shape: [B, F]
|
||||
returns:
|
||||
h: expanded states, shape: [B, F, H]
|
||||
"""
|
||||
try:
|
||||
assert mel2token.max() <= h.size(1) - 1
|
||||
except:
|
||||
print(f"Warning: mel2token.max() ({mel2token.max()}) is greater than h.size(1) - 1 ({h.size(1) - 1})")
|
||||
mel2token = torch.clamp(mel2token, 0, h.size(1)-1)
|
||||
mel2token_ = mel2token[..., None].repeat([1, 1, h.shape[-1]])
|
||||
h = torch.gather(h, 1, mel2token_) # [B, T, H]
|
||||
return h
|
||||
|
||||
@staticmethod
|
||||
def f0_to_coarse(f0, f0_bin=361, f0_min=32.7031956625, f0_shift=0):
|
||||
"""
|
||||
Convert continuous F0 values to discrete F0 bins (SIL and C1 - B6, 361 bins).
|
||||
args:
|
||||
f0: continuous F0 values
|
||||
f0_bin: number of F0 bins
|
||||
f0_min: minimum F0 value
|
||||
f0_shift: shift value for F0 bins
|
||||
returns:
|
||||
f0_coarse: discrete F0 bins
|
||||
"""
|
||||
is_torch = isinstance(f0, torch.Tensor)
|
||||
uv_mask = f0 <= 0
|
||||
|
||||
if is_torch:
|
||||
f0_safe = torch.maximum(f0, torch.tensor(f0_min))
|
||||
f0_cents = 1200 * torch.log2(f0_safe / f0_min)
|
||||
else:
|
||||
f0_safe = np.maximum(f0, f0_min)
|
||||
f0_cents = 1200 * np.log2(f0_safe / f0_min)
|
||||
|
||||
f0_coarse = (f0_cents / 20) + 1
|
||||
|
||||
if is_torch:
|
||||
f0_coarse = torch.round(f0_coarse).long()
|
||||
f0_coarse = torch.clamp(f0_coarse, min=1, max=f0_bin - 1)
|
||||
else:
|
||||
f0_coarse = np.rint(f0_coarse).astype(int)
|
||||
f0_coarse = np.clip(f0_coarse, 1, f0_bin - 1)
|
||||
|
||||
f0_coarse[uv_mask] = 0
|
||||
|
||||
if f0_shift != 0:
|
||||
if is_torch:
|
||||
voiced = f0_coarse > 0
|
||||
if voiced.any():
|
||||
shifted = f0_coarse[voiced] + f0_shift
|
||||
f0_coarse[voiced] = torch.clamp(shifted, 1, f0_bin - 1)
|
||||
else:
|
||||
voiced = f0_coarse > 0
|
||||
if np.any(voiced):
|
||||
shifted = f0_coarse[voiced] + f0_shift
|
||||
f0_coarse[voiced] = np.clip(shifted, 1, f0_bin - 1)
|
||||
|
||||
return f0_coarse
|
||||
|
||||
def infer(self, meta: dict, auto_shift=False, pitch_shift=0, n_steps=32, cfg=3, control="melody"):
|
||||
|
||||
gt_note_text = meta['target']['phoneme']
|
||||
gt_mel2note = meta['target']['mel2note']
|
||||
gt_note_type = meta['target']['note_type']
|
||||
|
||||
pt_wav = meta['prompt']['waveform']
|
||||
pt_note_text = meta['prompt']['phoneme']
|
||||
pt_mel2note = meta['prompt']['mel2note']
|
||||
pt_note_type = meta['prompt']['note_type']
|
||||
|
||||
if control == "score":
|
||||
gt_note_pitch = meta['target']['note_pitch']
|
||||
pt_note_pitch = meta['prompt']['note_pitch']
|
||||
gt_f0 = None
|
||||
pt_f0 = None
|
||||
elif control == "melody":
|
||||
gt_f0 = meta['target']['f0']
|
||||
pt_f0 = meta['prompt']['f0']
|
||||
gt_note_pitch = None
|
||||
pt_note_pitch = None
|
||||
else:
|
||||
raise ValueError(f"Unknown control mode: {control}")
|
||||
|
||||
# calculate auto pitch shift
|
||||
if auto_shift and pitch_shift == 0:
|
||||
if gt_note_pitch != None and pt_note_pitch != None:
|
||||
gt_median = torch.median(gt_note_pitch[gt_note_pitch >= 1])
|
||||
pt_median = torch.median(pt_note_pitch[pt_note_pitch >= 1])
|
||||
f0_shift = torch.round(pt_median - gt_median).int().item()
|
||||
elif gt_f0 != None and pt_f0 != None:
|
||||
gt_f0_median = torch.median(gt_f0[gt_f0 > 0])
|
||||
pt_f0_median = torch.median(pt_f0[pt_f0 > 0])
|
||||
f0_shift = torch.round(torch.log2(pt_f0_median / gt_f0_median) * 1200 / 100).int().item()
|
||||
else:
|
||||
print("Warning: pitch_shift is True but note_pitch or f0 is None. Set f0_shift to 0.")
|
||||
f0_shift = 0
|
||||
else:
|
||||
f0_shift = 0
|
||||
|
||||
if gt_f0 is None or pt_f0 is None:
|
||||
gt_f0, pt_f0 = torch.zeros_like(gt_mel2note).float().to(gt_mel2note.device), torch.zeros_like(pt_mel2note).float().to(pt_mel2note.device)
|
||||
if gt_note_pitch is None or pt_note_pitch is None:
|
||||
gt_note_pitch, pt_note_pitch = torch.zeros_like(gt_note_type).int().to(gt_note_type.device), torch.zeros_like(pt_note_type).int().to(pt_note_type.device)
|
||||
|
||||
# convert prompt waveform to mel spectrogram
|
||||
pt_mel = self.mel(pt_wav)
|
||||
|
||||
len_prompt = pt_note_pitch.shape[1]
|
||||
len_prompt_mel = pt_f0.shape[1]
|
||||
|
||||
note_pitch = torch.cat([pt_note_pitch, gt_note_pitch], 1)
|
||||
note_text = torch.cat([pt_note_text, gt_note_text], 1)
|
||||
note_type = torch.cat([pt_note_type, gt_note_type], 1)
|
||||
mel2note = torch.cat([pt_mel2note, gt_mel2note + len_prompt], 1)
|
||||
|
||||
f0_course_pt = self.f0_to_coarse(pt_f0)
|
||||
f0_course_gt = self.f0_to_coarse(gt_f0, f0_shift=f0_shift * 5)
|
||||
f0_course = torch.cat([f0_course_pt, f0_course_gt], 1)
|
||||
|
||||
note_pitch[note_pitch > 0] = note_pitch[note_pitch > 0] + f0_shift
|
||||
note_pitch = torch.clamp(note_pitch, 0, 255)
|
||||
|
||||
features = self.note_pitch_encoder(note_pitch) + self.note_type_encoder(note_type) + self.note_text_encoder(note_text)
|
||||
|
||||
features = self.preflow(features)
|
||||
features = self.expand_states(features, mel2note)
|
||||
features = features + self.f0_encoder(f0_course)
|
||||
|
||||
gt_decoder_inp = features[:, len_prompt_mel:, :]
|
||||
pt_decoder_inp = features[:, :len_prompt_mel, :]
|
||||
|
||||
generated_mel = self.cfm_decoder.reverse_diffusion(
|
||||
pt_mel,
|
||||
pt_decoder_inp,
|
||||
gt_decoder_inp,
|
||||
n_timesteps=n_steps,
|
||||
cfg=cfg
|
||||
)
|
||||
|
||||
generated_audio = self.vocoder(generated_mel.transpose(1, 2)[0:1, ...])
|
||||
|
||||
return generated_audio
|
||||
@@ -0,0 +1,25 @@
|
||||
import torch
|
||||
import torchaudio
|
||||
|
||||
|
||||
def load_wav(wav_path: str, sample_rate: int):
|
||||
"""Load wav file and resample to target sample rate.
|
||||
|
||||
Args:
|
||||
wav_path (str): Path to wav file.
|
||||
sample_rate (int): Target sample rate.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Waveform tensor with shape (1, T).
|
||||
"""
|
||||
waveform, sr = torchaudio.load(wav_path)
|
||||
|
||||
if sr != sample_rate:
|
||||
waveform = torchaudio.functional.resample(waveform, sr, sample_rate)
|
||||
|
||||
if len(waveform.shape) > 1 and waveform.shape[0] > 1:
|
||||
waveform = torch.mean(waveform, dim=0, keepdim=True)
|
||||
|
||||
return waveform
|
||||
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
import json
|
||||
import torch
|
||||
import numpy as np
|
||||
import torchaudio
|
||||
from typing import List
|
||||
|
||||
from soulxsinger.utils.audio_utils import load_wav
|
||||
|
||||
|
||||
class DataProcessor:
|
||||
"""Data processor for SoulX-Singer
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
hop_size: int,
|
||||
sample_rate: int,
|
||||
phoneset_path: str = 'soulxsinger/utils/phoneme/phone_set.json',
|
||||
device: str = 'cuda',
|
||||
prompt_append_duration: float = 0.5):
|
||||
"""Initialize data processor.
|
||||
|
||||
Args:
|
||||
hop_size (int): Hop size in samples.
|
||||
sample_rate (int): Sample rate in Hz.
|
||||
phoneset_path (str): Path to phoneme set JSON file.
|
||||
device (str): Device to use for tensor operations.
|
||||
prompt_append_duration (float): Duration to append to prompt in seconds.
|
||||
"""
|
||||
self.hop_size = hop_size
|
||||
self.sample_rate = sample_rate
|
||||
self.device = device
|
||||
self.prompt_append_duration = prompt_append_duration
|
||||
self.prompt_append_length = int(prompt_append_duration * sample_rate / hop_size)
|
||||
self.load_phoneme_id_map(phoneset_path)
|
||||
|
||||
def load_phoneme_id_map(self, phoneset_path: str):
|
||||
with open(phoneset_path, "r", encoding='utf-8') as f:
|
||||
phoneset = json.load(f)
|
||||
self.phone2idx = {ph: idx for idx, ph in enumerate(phoneset)}
|
||||
|
||||
def merge_phoneme(self, meta):
|
||||
merged_items = []
|
||||
|
||||
duration = [float(x) for x in meta["duration"].split()]
|
||||
phoneme = [str(x).replace("<AP>", "<SP>") for i, x in enumerate(meta["phoneme"].split())]
|
||||
note_pitch = [int(x) for x in meta["note_pitch"].split()]
|
||||
note_type = [int(x) if phoneme[i] != "<SP>" else 1 for i, x in enumerate(meta["note_type"].split())]
|
||||
|
||||
for i in range(len(phoneme)):
|
||||
if i > 0 and phoneme[i] == phoneme[i - 1] == "<SP>" and note_type[i] == note_type[i - 1] and note_pitch[i] == note_pitch[i - 1]:
|
||||
merged_items[-1][1] += duration[i]
|
||||
else:
|
||||
merged_items.append([phoneme[i], duration[i], note_pitch[i], note_type[i]])
|
||||
|
||||
single_frame_duration = self.hop_size / self.sample_rate
|
||||
meta['phoneme'] = [x[0] for x in merged_items]
|
||||
meta['duration'] = [x[1] for x in merged_items]
|
||||
meta['note_pitch'] = [x[2] for x in merged_items]
|
||||
meta['note_type'] = [x[3] for x in merged_items]
|
||||
|
||||
return meta
|
||||
|
||||
def preprocess(
|
||||
self,
|
||||
note_duration: List[float],
|
||||
phonemes: List[str],
|
||||
note_pitch: List[int],
|
||||
note_type: List[int],
|
||||
):
|
||||
"""
|
||||
Insert <BOW> and <EOW> for each note.
|
||||
Get aligned indices for each frame.
|
||||
|
||||
Args:
|
||||
note_duration: Duration of each note in seconds
|
||||
phonemes: Phoneme sequence for each note
|
||||
note_pitch: Pitch value for each note
|
||||
note_type: Type value for each note
|
||||
|
||||
"""
|
||||
sample_rate = self.sample_rate
|
||||
hop_size = self.hop_size
|
||||
duration = sum(note_duration) * sample_rate / hop_size
|
||||
mel2note = torch.zeros(int(duration), dtype=torch.long)
|
||||
|
||||
ph_locations = [] # idx at mel scale and length
|
||||
new_phonemes = []
|
||||
dur_sum = 0
|
||||
|
||||
note2origin = []
|
||||
|
||||
for ph_idx in range(len(phonemes)):
|
||||
dur = int(np.round(dur_sum * sample_rate / hop_size))
|
||||
dur = min(dur, len(mel2note) - 1)
|
||||
new_phonemes.append("<BOW>")
|
||||
note2origin.append(ph_idx)
|
||||
if phonemes[ph_idx][:3] == "en_":
|
||||
en_phs = ['en_' + x for x in phonemes[ph_idx][3:].split('-')] + ['<SEP>'] # <sep> between en words in one note
|
||||
ph_locations.append([dur, max(1, len(en_phs))])
|
||||
new_phonemes.extend(en_phs)
|
||||
note2origin.extend([ph_idx] * len(en_phs))
|
||||
else:
|
||||
ph_locations.append([dur, 1])
|
||||
new_phonemes.append(phonemes[ph_idx])
|
||||
note2origin.append(ph_idx)
|
||||
new_phonemes.append("<EOW>")
|
||||
note2origin.append(ph_idx)
|
||||
dur_sum += note_duration[ph_idx]
|
||||
|
||||
ph_idx = 1
|
||||
for idx, (i, j) in enumerate(ph_locations):
|
||||
next_phoneme_start = ph_locations[idx + 1][0] if idx < len(ph_locations) - 1 else len(mel2note)
|
||||
if i >= len(mel2note) or i + j > len(mel2note):
|
||||
break
|
||||
if i < len(mel2note) and mel2note[i] > 0:
|
||||
# print(f"warning: overlap of {idx}: {mel2note[i]}")
|
||||
while i < len(mel2note) and mel2note[i] > 0:
|
||||
i += 1
|
||||
mel2note[i] = ph_idx
|
||||
k = i + 1
|
||||
while k + j < next_phoneme_start:
|
||||
mel2note[k : k + j] = torch.arange(ph_idx, ph_idx + j) + 1
|
||||
k += j
|
||||
mel2note[next_phoneme_start - 1] = ph_idx + j + 1
|
||||
ph_idx += j + 2 # <BOW> + ph repeats + <EOW>
|
||||
|
||||
new_phonemes = ["<PAD>"] + new_phonemes
|
||||
new_note_pitch = [0] + [note_pitch[k] for k in note2origin]
|
||||
new_note_type = [1] + [note_type[k] for k in note2origin]
|
||||
|
||||
return {
|
||||
"phoneme": torch.tensor([self.phone2idx[x] for x in new_phonemes], device=self.device).unsqueeze(0),
|
||||
"note_pitch": torch.tensor(new_note_pitch, device=self.device).unsqueeze(0),
|
||||
"note_type": torch.tensor(new_note_type, device=self.device).unsqueeze(0),
|
||||
"mel2note": mel2note.clone().detach().to(self.device).unsqueeze(0),
|
||||
}
|
||||
|
||||
def process(
|
||||
self,
|
||||
meta: dict,
|
||||
wav_path: str = None
|
||||
):
|
||||
|
||||
meta = self.merge_phoneme(meta)
|
||||
|
||||
item = self.preprocess(
|
||||
meta["duration"],
|
||||
meta["phoneme"],
|
||||
meta["note_pitch"],
|
||||
meta["note_type"],
|
||||
)
|
||||
|
||||
f0 = torch.tensor([float(x) for x in meta["f0"].split()])
|
||||
min_frame = min(item["mel2note"].shape[1], f0.shape[0])
|
||||
item['f0'] = f0[:min_frame].unsqueeze(0).float().to(self.device)
|
||||
item["mel2note"] = item["mel2note"][:, :min_frame]
|
||||
|
||||
if wav_path is not None:
|
||||
waveform = load_wav(wav_path, self.sample_rate)
|
||||
item["waveform"] = waveform.to(self.device)[:, :min_frame * self.hop_size]
|
||||
|
||||
return item
|
||||
|
||||
|
||||
# test
|
||||
if __name__ == "__main__":
|
||||
import json
|
||||
with open("example/metadata/zh_prompt.json", "r", encoding="utf-8") as f:
|
||||
meta = json.load(f)
|
||||
if isinstance(meta, list):
|
||||
meta = meta[0]
|
||||
processor = DataProcessor(hop_size=480, sample_rate=24000)
|
||||
item = processor.process(meta, "example/audio/zh_prompt.wav")
|
||||
print(item.keys())
|
||||
@@ -0,0 +1,77 @@
|
||||
|
||||
"""
|
||||
Description:
|
||||
This script contains a collection of functions designed to handle various
|
||||
file reading and writing operations. It provides utilities to read from files,
|
||||
write data to files, and perform file manipulation tasks.
|
||||
"""
|
||||
|
||||
import os
|
||||
import json
|
||||
|
||||
from tqdm import tqdm
|
||||
from typing import List, Dict
|
||||
from pathlib import Path
|
||||
from omegaconf import OmegaConf, DictConfig
|
||||
|
||||
|
||||
def write_jsonl(metadata: List[dict], file_path: Path):
|
||||
"""Writes a list of dictionaries to a JSONL file.
|
||||
|
||||
Args:
|
||||
metadata : List[dict]
|
||||
A list of dictionaries, each representing a piece of meta.
|
||||
file_path : Path
|
||||
The file path to save the JSONL file
|
||||
|
||||
This function writes each dictionary in the list to a new line in the specified file.
|
||||
"""
|
||||
with open(file_path, "w", encoding="utf-8") as f:
|
||||
for meta in tqdm(metadata, desc="writing jsonl"):
|
||||
# Convert dictionary to JSON string and write it to the file with a newline
|
||||
json_str = json.dumps(meta, ensure_ascii=False) + "\n"
|
||||
f.write(json_str)
|
||||
print(f"jsonl saved to {file_path}")
|
||||
|
||||
|
||||
def read_jsonl(file_path: Path) -> List[dict]:
|
||||
"""
|
||||
Reads a JSONL file and returns a list of dictionaries.
|
||||
|
||||
Args:
|
||||
file_path : Path
|
||||
The path to the JSONL file to be read.
|
||||
|
||||
Returns:
|
||||
List[dict]
|
||||
A list of dictionaries parsed from each line of the JSONL file.
|
||||
"""
|
||||
metadata = []
|
||||
# Open the file for reading
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
# Split the file into lines
|
||||
lines = f.read().splitlines()
|
||||
# Process each line
|
||||
for line in lines:
|
||||
# Convert JSON string back to dictionary and append to list
|
||||
meta = json.loads(line)
|
||||
metadata.append(meta)
|
||||
# Return the list of metadata
|
||||
return metadata
|
||||
|
||||
|
||||
def load_config(config_path: Path) -> DictConfig:
|
||||
"""Loads a configuration file and optionally merges it with a base configuration.
|
||||
|
||||
Args:
|
||||
config_path (Path): Path to the configuration file.
|
||||
"""
|
||||
# Load the initial configuration from the given path
|
||||
config = OmegaConf.load(config_path)
|
||||
|
||||
# Check if there is a base configuration specified and merge if necessary
|
||||
if config.get("base_config", None) is not None:
|
||||
base_config = OmegaConf.load(config["base_config"])
|
||||
config = OmegaConf.merge(base_config, config)
|
||||
|
||||
return config
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,142 @@
|
||||
# https://github.com/gwx314/TechSinger/blob/main/utils/audio/pitch/utils.py
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def to_lf0(f0):
|
||||
f0[f0 < 1.0e-5] = 1.0e-6
|
||||
lf0 = f0.log() if isinstance(f0, torch.Tensor) else np.log(f0)
|
||||
lf0[f0 < 1.0e-5] = - 1.0E+10
|
||||
return lf0
|
||||
|
||||
|
||||
def to_f0(lf0):
|
||||
f0 = np.where(lf0 <= 0, 0.0, np.exp(lf0))
|
||||
return f0.flatten()
|
||||
|
||||
|
||||
def f0_to_coarse_mel(f0, f0_bin=256, f0_max=900.0, f0_min=50.0, f0_shift=0):
|
||||
f0_mel_min = 1127 * np.log(1 + f0_min / 700)
|
||||
f0_mel_max = 1127 * np.log(1 + f0_max / 700)
|
||||
is_torch = isinstance(f0, torch.Tensor)
|
||||
f0_mel = 1127 * (1 + f0 / 700).log() if is_torch else 1127 * np.log(1 + f0 / 700)
|
||||
f0_mel[f0_mel > 0] = (f0_mel[f0_mel > 0] - f0_mel_min) * (f0_bin - 2) / (f0_mel_max - f0_mel_min) + 1
|
||||
|
||||
f0_mel[f0_mel <= 1] = 1
|
||||
f0_mel[f0_mel > f0_bin - 1] = f0_bin - 1
|
||||
f0_coarse = (f0_mel + 0.5).long() if is_torch else np.rint(f0_mel).astype(int)
|
||||
|
||||
if f0_shift != 0:
|
||||
if f0_shift > 0:
|
||||
f0_shift = min(f0_shift, f0_bin - 1 - f0_coarse[f0_coarse > 1].max().item())
|
||||
else:
|
||||
f0_shift = max(f0_shift, 1 - f0_coarse[f0_coarse > 1].min().item())
|
||||
|
||||
f0_coarse[f0_coarse > 1] = f0_coarse[f0_coarse > 1] + f0_shift
|
||||
|
||||
assert f0_coarse.max() <= 255 and f0_coarse.min() >= 1, (f0_coarse.max(), f0_coarse.min(), f0.min(), f0.max())
|
||||
return f0_coarse
|
||||
|
||||
|
||||
def coarse_to_f0_mel(f0_coarse, f0_bin=256, f0_max=900.0, f0_min=50.0):
|
||||
f0_mel_min = 1127 * np.log(1 + f0_min / 700)
|
||||
f0_mel_max = 1127 * np.log(1 + f0_max / 700)
|
||||
uv = f0_coarse == 1
|
||||
f0 = f0_mel_min + (f0_coarse - 1) * (f0_mel_max - f0_mel_min) / (f0_bin - 2)
|
||||
f0 = ((f0 / 1127).exp() - 1) * 700
|
||||
f0[uv] = 0
|
||||
return f0
|
||||
|
||||
CONST_C1_FREQ = 32.7031956625 # C1 frequency in Hz
|
||||
CONST_B6_FREQ = 1975.53320502 # B6 frequency in Hz
|
||||
|
||||
def f0_to_coarse_midi(f0, f0_bin=361, f0_max=CONST_B6_FREQ, f0_min=CONST_C1_FREQ, f0_shift=0):
|
||||
is_torch = isinstance(f0, torch.Tensor)
|
||||
uv_mask = f0 <= 0
|
||||
|
||||
if is_torch:
|
||||
f0_safe = torch.maximum(f0, torch.tensor(f0_min))
|
||||
f0_cents = 1200 * torch.log2(f0_safe / f0_min)
|
||||
else:
|
||||
f0_safe = np.maximum(f0, f0_min)
|
||||
f0_cents = 1200 * np.log2(f0_safe / f0_min)
|
||||
|
||||
f0_coarse = (f0_cents / 20) + 1
|
||||
|
||||
if is_torch:
|
||||
f0_coarse = torch.round(f0_coarse).long()
|
||||
f0_coarse = torch.clamp(f0_coarse, min=1, max=f0_bin - 1)
|
||||
else:
|
||||
f0_coarse = np.rint(f0_coarse).astype(int)
|
||||
f0_coarse = np.clip(f0_coarse, 1, f0_bin - 1)
|
||||
|
||||
f0_coarse[uv_mask] = 0
|
||||
|
||||
if f0_shift != 0:
|
||||
if is_torch:
|
||||
voiced = f0_coarse > 0
|
||||
if voiced.any():
|
||||
shifted = f0_coarse[voiced] + f0_shift
|
||||
f0_coarse[voiced] = torch.clamp(shifted, 1, f0_bin - 1)
|
||||
else:
|
||||
voiced = f0_coarse > 0
|
||||
if np.any(voiced):
|
||||
shifted = f0_coarse[voiced] + f0_shift
|
||||
f0_coarse[voiced] = np.clip(shifted, 1, f0_bin - 1)
|
||||
|
||||
return f0_coarse
|
||||
|
||||
|
||||
def coarse_to_f0_midi(f0_coarse, f0_bin=361, f0_max=CONST_B6_FREQ, f0_min=CONST_C1_FREQ):
|
||||
|
||||
uv_mask = f0_coarse == 0
|
||||
cents = (f0_coarse - 1) * 20
|
||||
f0 = f0_min * (2 ** (cents / 1200))
|
||||
f0[uv_mask] = 0
|
||||
|
||||
return f0
|
||||
|
||||
|
||||
def norm_f0(f0, uv, pitch_norm='log', f0_mean=400, f0_std=100):
|
||||
is_torch = isinstance(f0, torch.Tensor)
|
||||
if pitch_norm == 'standard':
|
||||
f0 = (f0 - f0_mean) / f0_std
|
||||
if pitch_norm == 'log':
|
||||
f0 = torch.log2(f0 + 1e-8) if is_torch else np.log2(f0 + 1e-8)
|
||||
if uv is not None:
|
||||
f0[uv > 0] = 0
|
||||
return f0
|
||||
|
||||
|
||||
def norm_interp_f0(f0, pitch_norm='log', f0_mean=None, f0_std=None):
|
||||
is_torch = isinstance(f0, torch.Tensor)
|
||||
if is_torch:
|
||||
device = f0.device
|
||||
f0 = f0.data.cpu().numpy()
|
||||
uv = f0 == 0
|
||||
f0 = norm_f0(f0, uv, pitch_norm, f0_mean, f0_std)
|
||||
if sum(uv) == len(f0):
|
||||
f0[uv] = 0
|
||||
elif sum(uv) > 0:
|
||||
f0[uv] = np.interp(np.where(uv)[0], np.where(~uv)[0], f0[~uv])
|
||||
if is_torch:
|
||||
uv = torch.FloatTensor(uv)
|
||||
f0 = torch.FloatTensor(f0)
|
||||
f0 = f0.to(device)
|
||||
uv = uv.to(device)
|
||||
return f0, uv
|
||||
|
||||
|
||||
def denorm_f0(f0, uv, pitch_norm='log', f0_mean=400, f0_std=100, pitch_padding=None, min=50, max=900):
|
||||
is_torch = isinstance(f0, torch.Tensor)
|
||||
if pitch_norm == 'standard':
|
||||
f0 = f0 * f0_std + f0_mean
|
||||
if pitch_norm == 'log':
|
||||
f0 = 2 ** f0
|
||||
f0 = f0.clamp(min=min, max=max) if is_torch else np.clip(f0, a_min=min, a_max=max)
|
||||
if uv is not None:
|
||||
f0[uv > 0] = 0
|
||||
if pitch_padding is not None:
|
||||
f0[pitch_padding] = 0
|
||||
return f0
|
||||
Reference in New Issue
Block a user