Initial commit
This commit is contained in:
@@ -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
Reference in New Issue
Block a user