mirror of
https://github.com/storytold/vits-finetuning.git
synced 2026-10-09 00:09:52 +00:00
980 lines
36 KiB
Python
980 lines
36 KiB
Python
import copy
|
|
import math
|
|
import numpy as np
|
|
import scipy
|
|
import torch
|
|
from torch import nn
|
|
from torch.nn import functional as F
|
|
|
|
from torch.nn import Conv1d, ConvTranspose1d, AvgPool1d, Conv2d
|
|
from torch.nn.utils import weight_norm, remove_weight_norm
|
|
|
|
import commons
|
|
from commons import init_weights, get_padding
|
|
from transforms import piecewise_rational_quadratic_transform
|
|
from typing import Optional, Tuple
|
|
from numba import jit, prange
|
|
from scipy import signal as sig
|
|
|
|
LRELU_SLOPE = 0.1
|
|
|
|
|
|
|
|
|
|
|
|
class TMEncoder(nn.Module):
|
|
def __init__(self, tm_start_size, enc_sizes):
|
|
super().__init__()
|
|
enc_layers = [torch.nn.Linear(tm_start_size,enc_sizes[0])]
|
|
|
|
last_size = enc_sizes[0]
|
|
for ly_size in enc_sizes:
|
|
enc_layers.append(torch.nn.ReLU())
|
|
enc_layers.append(torch.nn.Linear(last_size,ly_size))
|
|
last_size = ly_size
|
|
|
|
self.net = torch.nn.Sequential(*enc_layers)
|
|
|
|
def forward(self, x):
|
|
return self.net(x)
|
|
|
|
|
|
|
|
class CoMBD(torch.nn.Module):
|
|
|
|
def __init__(self, filters, kernels, groups, strides, use_spectral_norm=False):
|
|
super(CoMBD, self).__init__()
|
|
norm_f = weight_norm if use_spectral_norm == False else spectral_norm
|
|
self.convs = nn.ModuleList()
|
|
init_channel = 1
|
|
for i, (f, k, g, s) in enumerate(zip(filters, kernels, groups, strides)):
|
|
self.convs.append(norm_f(Conv1d(init_channel, f, k, s, padding=get_padding(k, 1), groups=g)))
|
|
init_channel = f
|
|
self.conv_post = norm_f(Conv1d(filters[-1], 1, 3, 1, padding=get_padding(3, 1)))
|
|
|
|
def forward(self, x):
|
|
fmap = []
|
|
for l in self.convs:
|
|
x = l(x)
|
|
x = F.leaky_relu(x, 0.1)
|
|
fmap.append(x)
|
|
x = self.conv_post(x)
|
|
#fmap.append(x)
|
|
x = torch.flatten(x, 1, -1)
|
|
return x, fmap
|
|
|
|
|
|
class MDC(torch.nn.Module):
|
|
|
|
def __init__(self, in_channel, channel, kernel, stride, dilations, use_spectral_norm=False):
|
|
super(MDC, self).__init__()
|
|
norm_f = weight_norm if use_spectral_norm == False else spectral_norm
|
|
self.convs = torch.nn.ModuleList()
|
|
self.num_dilations = len(dilations)
|
|
for d in dilations:
|
|
self.convs.append(norm_f(Conv1d(in_channel, channel, kernel, stride=1, padding=get_padding(kernel, d),
|
|
dilation=d)))
|
|
|
|
self.conv_out = norm_f(Conv1d(channel, channel, 3, stride=stride, padding=get_padding(3, 1)))
|
|
|
|
def forward(self, x):
|
|
xs = None
|
|
for l in self.convs:
|
|
if xs is None:
|
|
xs = l(x)
|
|
else:
|
|
xs += l(x)
|
|
|
|
x = xs / self.num_dilations
|
|
|
|
x = self.conv_out(x)
|
|
x = F.leaky_relu(x, 0.1)
|
|
return x
|
|
|
|
|
|
class SubBandDiscriminator(torch.nn.Module):
|
|
|
|
def __init__(self, init_channel, channels, kernel, strides, dilations, use_spectral_norm=False):
|
|
super(SubBandDiscriminator, self).__init__()
|
|
norm_f = weight_norm if use_spectral_norm == False else spectral_norm
|
|
|
|
self.mdcs = torch.nn.ModuleList()
|
|
|
|
for c, s, d in zip(channels, strides, dilations):
|
|
self.mdcs.append(MDC(init_channel, c, kernel, s, d))
|
|
init_channel = c
|
|
self.conv_post = norm_f(Conv1d(init_channel, 1, 3, padding=get_padding(3, 1)))
|
|
|
|
def forward(self, x):
|
|
fmap = []
|
|
|
|
for l in self.mdcs:
|
|
x = l(x)
|
|
fmap.append(x)
|
|
x = self.conv_post(x)
|
|
#fmap.append(x)
|
|
x = torch.flatten(x, 1, -1)
|
|
|
|
return x, fmap
|
|
|
|
|
|
|
|
|
|
# adapted from
|
|
# https://github.com/kan-bayashi/ParallelWaveGAN/tree/master/parallel_wavegan
|
|
class PQMF(torch.nn.Module):
|
|
def __init__(self, N=4, taps=62, cutoff=0.15, beta=9.0):
|
|
super(PQMF, self).__init__()
|
|
|
|
self.N = N
|
|
self.taps = taps
|
|
self.cutoff = cutoff
|
|
self.beta = beta
|
|
|
|
QMF = sig.firwin(taps + 1, cutoff, window=('kaiser', beta))
|
|
H = np.zeros((N, len(QMF)))
|
|
G = np.zeros((N, len(QMF)))
|
|
for k in range(N):
|
|
constant_factor = (2 * k + 1) * (np.pi /
|
|
(2 * N)) * (np.arange(taps + 1) -
|
|
((taps - 1) / 2)) # TODO: (taps - 1) -> taps
|
|
phase = (-1)**k * np.pi / 4
|
|
H[k] = 2 * QMF * np.cos(constant_factor + phase)
|
|
|
|
G[k] = 2 * QMF * np.cos(constant_factor - phase)
|
|
|
|
H = torch.from_numpy(H[:, None, :]).float()
|
|
G = torch.from_numpy(G[None, :, :]).float()
|
|
|
|
self.register_buffer("H", H)
|
|
self.register_buffer("G", G)
|
|
|
|
updown_filter = torch.zeros((N, N, N)).float()
|
|
for k in range(N):
|
|
updown_filter[k, k, 0] = 1.0
|
|
self.register_buffer("updown_filter", updown_filter)
|
|
self.N = N
|
|
|
|
self.pad_fn = torch.nn.ConstantPad1d(taps // 2, 0.0)
|
|
|
|
def forward(self, x):
|
|
return self.analysis(x)
|
|
|
|
def analysis(self, x):
|
|
return F.conv1d(x, self.H, padding=self.taps // 2, stride=self.N)
|
|
|
|
def synthesis(self, x):
|
|
x = F.conv_transpose1d(x,
|
|
self.updown_filter * self.N,
|
|
stride=self.N)
|
|
x = F.conv1d(x, self.G, padding=self.taps // 2)
|
|
return
|
|
|
|
@jit(nopython=True)
|
|
def mas_width1(attn_map):
|
|
"""mas with hardcoded width=1"""
|
|
# assumes mel x text
|
|
opt = np.zeros_like(attn_map)
|
|
attn_map = np.log(attn_map)
|
|
attn_map[0, 1:] = -np.inf
|
|
log_p = np.zeros_like(attn_map)
|
|
log_p[0, :] = attn_map[0, :]
|
|
prev_ind = np.zeros_like(attn_map, dtype=np.int64)
|
|
for i in range(1, attn_map.shape[0]):
|
|
for j in range(attn_map.shape[1]): # for each text dim
|
|
prev_log = log_p[i - 1, j]
|
|
prev_j = j
|
|
|
|
if j - 1 >= 0 and log_p[i - 1, j - 1] >= log_p[i - 1, j]:
|
|
prev_log = log_p[i - 1, j - 1]
|
|
prev_j = j - 1
|
|
|
|
log_p[i, j] = attn_map[i, j] + prev_log
|
|
prev_ind[i, j] = prev_j
|
|
|
|
# now backtrack
|
|
curr_text_idx = attn_map.shape[1] - 1
|
|
for i in range(attn_map.shape[0] - 1, -1, -1):
|
|
opt[i, curr_text_idx] = 1
|
|
curr_text_idx = prev_ind[i, curr_text_idx]
|
|
opt[0, curr_text_idx] = 1
|
|
return opt
|
|
|
|
|
|
@jit(nopython=True)
|
|
def b_mas(b_attn_map, in_lens, out_lens, width=1):
|
|
assert width == 1
|
|
attn_out = np.zeros_like(b_attn_map)
|
|
|
|
for b in prange(b_attn_map.shape[0]):
|
|
out = mas_width1(b_attn_map[b, 0, : out_lens[b], : in_lens[b]])
|
|
attn_out[b, 0, : out_lens[b], : in_lens[b]] = out
|
|
return attn_out
|
|
|
|
|
|
|
|
|
|
class PartialConv1d(torch.nn.Conv1d):
|
|
"""
|
|
Zero padding creates a unique identifier for where the edge of the data is, such that the model can almost always identify
|
|
exactly where it is relative to either edge given a sufficient receptive field. Partial padding goes to some lengths to remove
|
|
this affect.
|
|
"""
|
|
|
|
def __init__(self, *args, **kwargs):
|
|
super(PartialConv1d, self).__init__(*args, **kwargs)
|
|
weight_maskUpdater = torch.ones(1, 1, self.kernel_size[0])
|
|
self.register_buffer("weight_maskUpdater", weight_maskUpdater, persistent=False)
|
|
slide_winsize = torch.tensor(self.weight_maskUpdater.shape[1] * self.weight_maskUpdater.shape[2])
|
|
self.register_buffer("slide_winsize", slide_winsize, persistent=False)
|
|
|
|
if self.bias is not None:
|
|
bias_view = self.bias.view(1, self.out_channels, 1)
|
|
self.register_buffer('bias_view', bias_view, persistent=False)
|
|
# caching part
|
|
self.last_size = (-1, -1, -1)
|
|
|
|
update_mask = torch.ones(1, 1, 1)
|
|
self.register_buffer('update_mask', update_mask, persistent=False)
|
|
mask_ratio = torch.ones(1, 1, 1)
|
|
self.register_buffer('mask_ratio', mask_ratio, persistent=False)
|
|
self.partial: bool = True
|
|
|
|
def calculate_mask(self, input: torch.Tensor, mask_in: Optional[torch.Tensor]):
|
|
with torch.no_grad():
|
|
if mask_in is None:
|
|
mask = torch.ones(1, 1, input.shape[2], dtype=input.dtype, device=input.device)
|
|
else:
|
|
mask = mask_in
|
|
update_mask = F.conv1d(
|
|
mask,
|
|
self.weight_maskUpdater,
|
|
bias=None,
|
|
stride=self.stride,
|
|
padding=self.padding,
|
|
dilation=self.dilation,
|
|
groups=1,
|
|
)
|
|
# for mixed precision training, change 1e-8 to 1e-6
|
|
mask_ratio = self.slide_winsize / (update_mask + 1e-6)
|
|
update_mask = torch.clamp(update_mask, 0, 1)
|
|
mask_ratio = torch.mul(mask_ratio.to(update_mask), update_mask)
|
|
return torch.mul(input, mask), mask_ratio, update_mask
|
|
|
|
def forward_aux(self, input: torch.Tensor, mask_ratio: torch.Tensor, update_mask: torch.Tensor) -> torch.Tensor:
|
|
assert len(input.shape) == 3
|
|
|
|
raw_out = self._conv_forward(input, self.weight, self.bias)
|
|
|
|
if self.bias is not None:
|
|
output = torch.mul(raw_out - self.bias_view, mask_ratio) + self.bias_view
|
|
output = torch.mul(output, update_mask)
|
|
else:
|
|
output = torch.mul(raw_out, mask_ratio)
|
|
|
|
return output
|
|
|
|
@torch.jit.ignore
|
|
def forward_with_cache(self, input: torch.Tensor, mask_in: Optional[torch.Tensor] = None) -> torch.Tensor:
|
|
use_cache = not (torch.jit.is_tracing() or torch.onnx.is_in_onnx_export())
|
|
cache_hit = use_cache and mask_in is None and self.last_size == input.shape
|
|
if cache_hit:
|
|
mask_ratio = self.mask_ratio
|
|
update_mask = self.update_mask
|
|
else:
|
|
input, mask_ratio, update_mask = self.calculate_mask(input, mask_in)
|
|
if use_cache:
|
|
# if a mask is input, or tensor shape changed, update mask ratio
|
|
self.last_size = tuple(input.shape)
|
|
self.update_mask = update_mask
|
|
self.mask_ratio = mask_ratio
|
|
return self.forward_aux(input, mask_ratio, update_mask)
|
|
|
|
def forward_no_cache(self, input: torch.Tensor, mask_in: Optional[torch.Tensor] = None) -> torch.Tensor:
|
|
if self.partial:
|
|
input, mask_ratio, update_mask = self.calculate_mask(input, mask_in)
|
|
return self.forward_aux(input, mask_ratio, update_mask)
|
|
else:
|
|
if mask_in is not None:
|
|
input = torch.mul(input, mask_in)
|
|
return self._conv_forward(input, self.weight, self.bias)
|
|
|
|
def forward(self, input: torch.Tensor, mask_in: Optional[torch.Tensor] = None) -> torch.Tensor:
|
|
if self.partial:
|
|
return self.forward_with_cache(input, mask_in)
|
|
else:
|
|
if mask_in is not None:
|
|
input = torch.mul(input, mask_in)
|
|
return self._conv_forward(input, self.weight, self.bias)
|
|
|
|
|
|
class ConvNorm(torch.nn.Module):
|
|
def __init__(
|
|
self,
|
|
in_channels,
|
|
out_channels,
|
|
kernel_size=1,
|
|
stride=1,
|
|
padding=None,
|
|
dilation=1,
|
|
bias=True,
|
|
w_init_gain='linear',
|
|
init_gain_param=None,
|
|
use_partial_padding: bool = False,
|
|
use_weight_norm: bool = False,
|
|
norm_fn=None,
|
|
):
|
|
super(ConvNorm, self).__init__()
|
|
if padding is None:
|
|
assert kernel_size % 2 == 1
|
|
padding = int(dilation * (kernel_size - 1) / 2)
|
|
self.use_partial_padding: bool = use_partial_padding
|
|
conv = PartialConv1d(
|
|
in_channels,
|
|
out_channels,
|
|
kernel_size=kernel_size,
|
|
stride=stride,
|
|
padding=padding,
|
|
dilation=dilation,
|
|
bias=bias,
|
|
)
|
|
conv.partial = use_partial_padding
|
|
torch.nn.init.xavier_uniform_(conv.weight, gain=torch.nn.init.calculate_gain(w_init_gain,init_gain_param))
|
|
if use_weight_norm:
|
|
conv = torch.nn.utils.weight_norm(conv)
|
|
if norm_fn is not None:
|
|
self.norm = norm_fn(out_channels, affine=True)
|
|
else:
|
|
self.norm = None
|
|
self.conv = conv
|
|
|
|
def forward(self, input: torch.Tensor, mask_in: Optional[torch.Tensor] = None) -> torch.Tensor:
|
|
ret = self.conv(input, mask_in)
|
|
if self.norm is not None:
|
|
ret = self.norm(ret)
|
|
return ret
|
|
|
|
def binarize_attention_parallel(attn, in_lens, out_lens):
|
|
"""For training purposes only. Binarizes attention with MAS.
|
|
These will no longer receive a gradient.
|
|
Args:
|
|
attn: B x 1 x max_mel_len x max_text_len
|
|
"""
|
|
with torch.no_grad():
|
|
attn_cpu = attn.data.cpu().numpy()
|
|
attn_out = b_mas(attn_cpu, in_lens.cpu().numpy(), out_lens.cpu().numpy(), width=1)
|
|
return torch.from_numpy(attn_out).to(attn.device)
|
|
|
|
class AlignmentEncoder(torch.nn.Module):
|
|
"""Module for alignment text and mel spectrogram. """
|
|
|
|
def __init__(
|
|
self, n_mel_channels=80, n_text_channels=512, n_att_channels=80, temperature=0.0005,lr_slope=0.05,
|
|
):
|
|
super().__init__()
|
|
self.temperature = temperature
|
|
self.softmax = torch.nn.Softmax(dim=3)
|
|
self.log_softmax = torch.nn.LogSoftmax(dim=3)
|
|
|
|
self.key_proj = nn.Sequential(
|
|
ConvNorm(n_text_channels, n_text_channels * 2, kernel_size=3, bias=True, w_init_gain='leaky_relu',init_gain_param=lr_slope),
|
|
torch.nn.LeakyReLU(lr_slope),
|
|
ConvNorm(n_text_channels * 2, n_att_channels, kernel_size=1, bias=True),
|
|
)
|
|
|
|
self.query_proj = nn.Sequential(
|
|
ConvNorm(n_mel_channels, n_mel_channels * 2, kernel_size=3, bias=True, w_init_gain='leaky_relu',init_gain_param=lr_slope),
|
|
torch.nn.LeakyReLU(lr_slope),
|
|
ConvNorm(n_mel_channels * 2, n_mel_channels, kernel_size=1, bias=True),
|
|
torch.nn.LeakyReLU(lr_slope),
|
|
ConvNorm(n_mel_channels, n_att_channels, kernel_size=1, bias=True),
|
|
)
|
|
|
|
def get_dist(self, keys, queries, mask=None):
|
|
"""Calculation of distance matrix.
|
|
Args:
|
|
queries (torch.tensor): B x C x T1 tensor (probably going to be mel data).
|
|
keys (torch.tensor): B x C2 x T2 tensor (text data).
|
|
mask (torch.tensor): B x T2 x 1 tensor, binary mask for variable length entries and also can be used
|
|
for ignoring unnecessary elements from keys in the resulting distance matrix (True = mask element, False = leave unchanged).
|
|
Output:
|
|
dist (torch.tensor): B x T1 x T2 tensor.
|
|
"""
|
|
keys_enc = self.key_proj(keys) # B x n_attn_dims x T2
|
|
queries_enc = self.query_proj(queries) # B x n_attn_dims x T1
|
|
attn = (queries_enc[:, :, :, None] - keys_enc[:, :, None]) ** 2 # B x n_attn_dims x T1 x T2
|
|
dist = attn.sum(1, keepdim=True) # B x 1 x T1 x T2
|
|
|
|
if mask is not None:
|
|
dist.data.masked_fill_(mask.permute(0, 2, 1).unsqueeze(2), float("inf"))
|
|
|
|
return dist.squeeze(1)
|
|
|
|
@staticmethod
|
|
def get_durations(attn_soft, text_len, spect_len):
|
|
"""Calculation of durations.
|
|
Args:
|
|
attn_soft (torch.tensor): B x 1 x T1 x T2 tensor.
|
|
text_len (torch.tensor): B tensor, lengths of text.
|
|
spect_len (torch.tensor): B tensor, lengths of mel spectrogram.
|
|
"""
|
|
attn_hard = binarize_attention_parallel(attn_soft, text_len, spect_len)
|
|
durations = attn_hard.sum(2)[:, 0, :]
|
|
assert torch.all(torch.eq(durations.sum(dim=1), spect_len))
|
|
return durations
|
|
|
|
@staticmethod
|
|
def get_mean_dist_by_durations(dist, durations, mask=None):
|
|
"""Select elements from the distance matrix for the given durations and mask and return mean distance.
|
|
Args:
|
|
dist (torch.tensor): B x T1 x T2 tensor.
|
|
durations (torch.tensor): B x T2 tensor. Dim T2 should sum to T1.
|
|
mask (torch.tensor): B x T2 x 1 binary mask for variable length entries and also can be used
|
|
for ignoring unnecessary elements in dist by T2 dim (True = mask element, False = leave unchanged).
|
|
Output:
|
|
mean_dist (torch.tensor): B x 1 tensor.
|
|
"""
|
|
batch_size, t1_size, t2_size = dist.size()
|
|
assert torch.all(torch.eq(durations.sum(dim=1), t1_size))
|
|
|
|
if mask is not None:
|
|
dist = dist.masked_fill(mask.permute(0, 2, 1).unsqueeze(2), 0)
|
|
|
|
# TODO(oktai15): make it more efficient
|
|
mean_dist_by_durations = []
|
|
for dist_idx in range(batch_size):
|
|
mean_dist_by_durations.append(
|
|
torch.mean(
|
|
dist[
|
|
dist_idx,
|
|
torch.arange(t1_size),
|
|
torch.repeat_interleave(torch.arange(t2_size), repeats=durations[dist_idx]),
|
|
]
|
|
)
|
|
)
|
|
|
|
return torch.tensor(mean_dist_by_durations, dtype=dist.dtype, device=dist.device)
|
|
|
|
@staticmethod
|
|
def get_mean_distance_for_word(l2_dists, durs, start_token, num_tokens):
|
|
"""Calculates the mean distance between text and audio embeddings given a range of text tokens.
|
|
Args:
|
|
l2_dists (torch.tensor): L2 distance matrix from Aligner inference. T1 x T2 tensor.
|
|
durs (torch.tensor): List of durations corresponding to each text token. T2 tensor. Should sum to T1.
|
|
start_token (int): Index of the starting token for the word of interest.
|
|
num_tokens (int): Length (in tokens) of the word of interest.
|
|
Output:
|
|
mean_dist_for_word (float): Mean embedding distance between the word indicated and its predicted audio frames.
|
|
"""
|
|
# Need to calculate which audio frame we start on by summing all durations up to the start token's duration
|
|
start_frame = torch.sum(durs[:start_token]).data
|
|
|
|
total_frames = 0
|
|
dist_sum = 0
|
|
|
|
# Loop through each text token
|
|
for token_ind in range(start_token, start_token + num_tokens):
|
|
# Loop through each frame for the given text token
|
|
for frame_ind in range(start_frame, start_frame + durs[token_ind]):
|
|
# Recall that the L2 distance matrix is shape [spec_len, text_len]
|
|
dist_sum += l2_dists[frame_ind, token_ind]
|
|
|
|
# Update total frames so far & the starting frame for the next token
|
|
total_frames += durs[token_ind]
|
|
start_frame += durs[token_ind]
|
|
|
|
return dist_sum / total_frames
|
|
|
|
def forward(self, queries, keys, mask=None, attn_prior=None, conditioning=None):
|
|
"""Forward pass of the aligner encoder.
|
|
Args:
|
|
queries (torch.tensor): B x C x T1 tensor (probably going to be mel data).
|
|
keys (torch.tensor): B x C2 x T2 tensor (text data).
|
|
mask (torch.tensor): B x T2 x 1 tensor, binary mask for variable length entries (True = mask element, False = leave unchanged).
|
|
attn_prior (torch.tensor): prior for attention matrix.
|
|
conditioning (torch.tensor): B x T2 x 1 conditioning embedding
|
|
Output:
|
|
attn (torch.tensor): B x 1 x T1 x T2 attention mask. Final dim T2 should sum to 1.
|
|
attn_logprob (torch.tensor): B x 1 x T1 x T2 log-prob attention mask.
|
|
"""
|
|
if conditioning is not None:
|
|
keys = keys + conditioning
|
|
|
|
keys_enc = self.key_proj(keys) # B x n_attn_dims x T2
|
|
queries_enc = self.query_proj(queries) # B x n_attn_dims x T1
|
|
|
|
# Simplistic Gaussian Isotopic Attention
|
|
attn = (queries_enc[:, :, :, None] - keys_enc[:, :, None]) ** 2 # B x n_attn_dims x T1 x T2
|
|
attn = -self.temperature * attn.sum(1, keepdim=True)
|
|
|
|
if attn_prior is not None:
|
|
attn = self.log_softmax(attn) + torch.log(attn_prior[:, None] + 1e-8)
|
|
|
|
attn_logprob = attn.clone()
|
|
|
|
if mask is not None:
|
|
attn.data.masked_fill_(mask.permute(0, 2, 1).unsqueeze(2), -float("inf"))
|
|
|
|
attn = self.softmax(attn) # softmax along T2
|
|
return attn, attn_logprob
|
|
|
|
|
|
|
|
class LayerNorm(nn.Module):
|
|
def __init__(self, channels, eps=1e-5):
|
|
super().__init__()
|
|
self.channels = channels
|
|
self.eps = eps
|
|
|
|
self.gamma = nn.Parameter(torch.ones(channels))
|
|
self.beta = nn.Parameter(torch.zeros(channels))
|
|
|
|
def forward(self, x):
|
|
x = x.transpose(1, -1)
|
|
x = F.layer_norm(x, (self.channels,), self.gamma, self.beta, self.eps)
|
|
return x.transpose(1, -1)
|
|
|
|
|
|
class ConvReluNorm(nn.Module):
|
|
def __init__(self, in_channels, hidden_channels, out_channels, kernel_size, n_layers, p_dropout):
|
|
super().__init__()
|
|
self.in_channels = in_channels
|
|
self.hidden_channels = hidden_channels
|
|
self.out_channels = out_channels
|
|
self.kernel_size = kernel_size
|
|
self.n_layers = n_layers
|
|
self.p_dropout = p_dropout
|
|
assert n_layers > 1, "Number of layers should be larger than 0."
|
|
|
|
self.conv_layers = nn.ModuleList()
|
|
self.norm_layers = nn.ModuleList()
|
|
self.conv_layers.append(nn.Conv1d(in_channels, hidden_channels, kernel_size, padding=kernel_size//2))
|
|
self.norm_layers.append(LayerNorm(hidden_channels))
|
|
self.relu_drop = nn.Sequential(
|
|
nn.ReLU(),
|
|
nn.Dropout(p_dropout))
|
|
for _ in range(n_layers-1):
|
|
self.conv_layers.append(nn.Conv1d(hidden_channels, hidden_channels, kernel_size, padding=kernel_size//2))
|
|
self.norm_layers.append(LayerNorm(hidden_channels))
|
|
self.proj = nn.Conv1d(hidden_channels, out_channels, 1)
|
|
self.proj.weight.data.zero_()
|
|
self.proj.bias.data.zero_()
|
|
|
|
def forward(self, x, x_mask):
|
|
x_org = x
|
|
for i in range(self.n_layers):
|
|
x = self.conv_layers[i](x * x_mask)
|
|
x = self.norm_layers[i](x)
|
|
x = self.relu_drop(x)
|
|
x = x_org + self.proj(x)
|
|
return x * x_mask
|
|
|
|
|
|
class DDSConv(nn.Module):
|
|
"""
|
|
Dialted and Depth-Separable Convolution
|
|
"""
|
|
def __init__(self, channels, kernel_size, n_layers, p_dropout=0.):
|
|
super().__init__()
|
|
self.channels = channels
|
|
self.kernel_size = kernel_size
|
|
self.n_layers = n_layers
|
|
self.p_dropout = p_dropout
|
|
|
|
self.drop = nn.Dropout(p_dropout)
|
|
self.convs_sep = nn.ModuleList()
|
|
self.convs_1x1 = nn.ModuleList()
|
|
self.norms_1 = nn.ModuleList()
|
|
self.norms_2 = nn.ModuleList()
|
|
for i in range(n_layers):
|
|
dilation = kernel_size ** i
|
|
padding = (kernel_size * dilation - dilation) // 2
|
|
self.convs_sep.append(nn.Conv1d(channels, channels, kernel_size,
|
|
groups=channels, dilation=dilation, padding=padding
|
|
))
|
|
self.convs_1x1.append(nn.Conv1d(channels, channels, 1))
|
|
self.norms_1.append(LayerNorm(channels))
|
|
self.norms_2.append(LayerNorm(channels))
|
|
|
|
def forward(self, x, x_mask, g=None):
|
|
if g is not None:
|
|
x = x + g
|
|
for i in range(self.n_layers):
|
|
y = self.convs_sep[i](x * x_mask)
|
|
y = self.norms_1[i](y)
|
|
y = F.gelu(y)
|
|
y = self.convs_1x1[i](y)
|
|
y = self.norms_2[i](y)
|
|
y = F.gelu(y)
|
|
y = self.drop(y)
|
|
x = x + y
|
|
return x * x_mask
|
|
|
|
|
|
class WN(torch.nn.Module):
|
|
def __init__(self, hidden_channels, kernel_size, dilation_rate, n_layers, gin_channels=0, p_dropout=0):
|
|
super(WN, self).__init__()
|
|
assert(kernel_size % 2 == 1)
|
|
self.hidden_channels =hidden_channels
|
|
self.kernel_size = kernel_size,
|
|
self.dilation_rate = dilation_rate
|
|
self.n_layers = n_layers
|
|
self.gin_channels = gin_channels
|
|
self.p_dropout = p_dropout
|
|
|
|
self.in_layers = torch.nn.ModuleList()
|
|
self.res_skip_layers = torch.nn.ModuleList()
|
|
self.drop = nn.Dropout(p_dropout)
|
|
|
|
if gin_channels != 0:
|
|
cond_layer = torch.nn.Conv1d(gin_channels, 2*hidden_channels*n_layers, 1)
|
|
self.cond_layer = torch.nn.utils.weight_norm(cond_layer, name='weight')
|
|
|
|
for i in range(n_layers):
|
|
dilation = dilation_rate ** i
|
|
padding = int((kernel_size * dilation - dilation) / 2)
|
|
in_layer = torch.nn.Conv1d(hidden_channels, 2*hidden_channels, kernel_size,
|
|
dilation=dilation, padding=padding)
|
|
in_layer = torch.nn.utils.weight_norm(in_layer, name='weight')
|
|
self.in_layers.append(in_layer)
|
|
|
|
# last one is not necessary
|
|
if i < n_layers - 1:
|
|
res_skip_channels = 2 * hidden_channels
|
|
else:
|
|
res_skip_channels = hidden_channels
|
|
|
|
res_skip_layer = torch.nn.Conv1d(hidden_channels, res_skip_channels, 1)
|
|
res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name='weight')
|
|
self.res_skip_layers.append(res_skip_layer)
|
|
|
|
def forward(self, x, x_mask, g=None, **kwargs):
|
|
output = torch.zeros_like(x)
|
|
n_channels_tensor = torch.IntTensor([self.hidden_channels])
|
|
|
|
if g is not None:
|
|
g = self.cond_layer(g)
|
|
|
|
for i in range(self.n_layers):
|
|
x_in = self.in_layers[i](x)
|
|
if g is not None:
|
|
cond_offset = i * 2 * self.hidden_channels
|
|
g_l = g[:,cond_offset:cond_offset+2*self.hidden_channels,:]
|
|
else:
|
|
g_l = torch.zeros_like(x_in)
|
|
|
|
acts = commons.fused_add_tanh_sigmoid_multiply(
|
|
x_in,
|
|
g_l,
|
|
n_channels_tensor)
|
|
acts = self.drop(acts)
|
|
|
|
res_skip_acts = self.res_skip_layers[i](acts)
|
|
if i < self.n_layers - 1:
|
|
res_acts = res_skip_acts[:,:self.hidden_channels,:]
|
|
x = (x + res_acts) * x_mask
|
|
output = output + res_skip_acts[:,self.hidden_channels:,:]
|
|
else:
|
|
output = output + res_skip_acts
|
|
return output * x_mask
|
|
|
|
def remove_weight_norm(self):
|
|
if self.gin_channels != 0:
|
|
torch.nn.utils.remove_weight_norm(self.cond_layer)
|
|
for l in self.in_layers:
|
|
torch.nn.utils.remove_weight_norm(l)
|
|
for l in self.res_skip_layers:
|
|
torch.nn.utils.remove_weight_norm(l)
|
|
|
|
|
|
class ResBlock1(torch.nn.Module):
|
|
def __init__(self, channels, kernel_size=3, dilation=(1, 3, 5)):
|
|
super(ResBlock1, self).__init__()
|
|
self.convs1 = nn.ModuleList([
|
|
weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=dilation[0],
|
|
padding=get_padding(kernel_size, dilation[0]))),
|
|
weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=dilation[1],
|
|
padding=get_padding(kernel_size, dilation[1]))),
|
|
weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=dilation[2],
|
|
padding=get_padding(kernel_size, dilation[2])))
|
|
])
|
|
self.convs1.apply(init_weights)
|
|
|
|
self.convs2 = nn.ModuleList([
|
|
weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=1,
|
|
padding=get_padding(kernel_size, 1))),
|
|
weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=1,
|
|
padding=get_padding(kernel_size, 1))),
|
|
weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=1,
|
|
padding=get_padding(kernel_size, 1)))
|
|
])
|
|
self.convs2.apply(init_weights)
|
|
|
|
def forward(self, x, x_mask=None):
|
|
for c1, c2 in zip(self.convs1, self.convs2):
|
|
xt = F.leaky_relu(x, LRELU_SLOPE)
|
|
if x_mask is not None:
|
|
xt = xt * x_mask
|
|
xt = c1(xt)
|
|
xt = F.leaky_relu(xt, LRELU_SLOPE)
|
|
if x_mask is not None:
|
|
xt = xt * x_mask
|
|
xt = c2(xt)
|
|
x = xt + x
|
|
if x_mask is not None:
|
|
x = x * x_mask
|
|
return x
|
|
|
|
def remove_weight_norm(self):
|
|
for l in self.convs1:
|
|
remove_weight_norm(l)
|
|
for l in self.convs2:
|
|
remove_weight_norm(l)
|
|
|
|
|
|
class ResBlock2(torch.nn.Module):
|
|
def __init__(self, channels, kernel_size=3, dilation=(1, 3)):
|
|
super(ResBlock2, self).__init__()
|
|
self.convs = nn.ModuleList([
|
|
weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=dilation[0],
|
|
padding=get_padding(kernel_size, dilation[0]))),
|
|
weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=dilation[1],
|
|
padding=get_padding(kernel_size, dilation[1])))
|
|
])
|
|
self.convs.apply(init_weights)
|
|
|
|
def forward(self, x, x_mask=None):
|
|
for c in self.convs:
|
|
xt = F.leaky_relu(x, LRELU_SLOPE)
|
|
if x_mask is not None:
|
|
xt = xt * x_mask
|
|
xt = c(xt)
|
|
x = xt + x
|
|
if x_mask is not None:
|
|
x = x * x_mask
|
|
return x
|
|
|
|
def remove_weight_norm(self):
|
|
for l in self.convs:
|
|
remove_weight_norm(l)
|
|
|
|
|
|
class Log(nn.Module):
|
|
def forward(self, x, x_mask, reverse=False, **kwargs):
|
|
if not reverse:
|
|
y = torch.log(torch.clamp_min(x, 1e-5)) * x_mask
|
|
logdet = torch.sum(-y, [1, 2])
|
|
return y, logdet
|
|
else:
|
|
x = torch.exp(x) * x_mask
|
|
return x
|
|
|
|
|
|
class Flip(nn.Module):
|
|
def forward(self, x, *args, reverse=False, **kwargs):
|
|
x = torch.flip(x, [1])
|
|
if not reverse:
|
|
logdet = torch.zeros(x.size(0)).to(dtype=x.dtype, device=x.device)
|
|
return x, logdet
|
|
else:
|
|
return x
|
|
|
|
|
|
class ElementwiseAffine(nn.Module):
|
|
def __init__(self, channels):
|
|
super().__init__()
|
|
self.channels = channels
|
|
self.m = nn.Parameter(torch.zeros(channels,1))
|
|
self.logs = nn.Parameter(torch.zeros(channels,1))
|
|
|
|
def forward(self, x, x_mask, reverse=False, **kwargs):
|
|
if not reverse:
|
|
y = self.m + torch.exp(self.logs) * x
|
|
y = y * x_mask
|
|
logdet = torch.sum(self.logs * x_mask, [1,2])
|
|
return y, logdet
|
|
else:
|
|
x = (x - self.m) * torch.exp(-self.logs) * x_mask
|
|
return x
|
|
|
|
|
|
class ResidualCouplingLayer(nn.Module):
|
|
def __init__(self,
|
|
channels,
|
|
hidden_channels,
|
|
kernel_size,
|
|
dilation_rate,
|
|
n_layers,
|
|
p_dropout=0,
|
|
gin_channels=0,
|
|
mean_only=False):
|
|
assert channels % 2 == 0, "channels should be divisible by 2"
|
|
super().__init__()
|
|
self.channels = channels
|
|
self.hidden_channels = hidden_channels
|
|
self.kernel_size = kernel_size
|
|
self.dilation_rate = dilation_rate
|
|
self.n_layers = n_layers
|
|
self.half_channels = channels // 2
|
|
self.mean_only = mean_only
|
|
|
|
self.pre = nn.Conv1d(self.half_channels, hidden_channels, 1)
|
|
self.enc = WN(hidden_channels, kernel_size, dilation_rate, n_layers, p_dropout=p_dropout, gin_channels=gin_channels)
|
|
self.post = nn.Conv1d(hidden_channels, self.half_channels * (2 - mean_only), 1)
|
|
self.post.weight.data.zero_()
|
|
self.post.bias.data.zero_()
|
|
|
|
def forward(self, x, x_mask, g=None, reverse=False):
|
|
x0, x1 = torch.split(x, [self.half_channels]*2, 1)
|
|
h = self.pre(x0) * x_mask
|
|
h = self.enc(h, x_mask, g=g)
|
|
stats = self.post(h) * x_mask
|
|
if not self.mean_only:
|
|
m, logs = torch.split(stats, [self.half_channels]*2, 1)
|
|
else:
|
|
m = stats
|
|
logs = torch.zeros_like(m)
|
|
|
|
if not reverse:
|
|
x1 = m + x1 * torch.exp(logs) * x_mask
|
|
x = torch.cat([x0, x1], 1)
|
|
logdet = torch.sum(logs, [1,2])
|
|
return x, logdet
|
|
else:
|
|
x1 = (x1 - m) * torch.exp(-logs) * x_mask
|
|
x = torch.cat([x0, x1], 1)
|
|
return x
|
|
|
|
|
|
class ConvFlow(nn.Module):
|
|
def __init__(self, in_channels, filter_channels, kernel_size, n_layers, num_bins=10, tail_bound=5.0):
|
|
super().__init__()
|
|
self.in_channels = in_channels
|
|
self.filter_channels = filter_channels
|
|
self.kernel_size = kernel_size
|
|
self.n_layers = n_layers
|
|
self.num_bins = num_bins
|
|
self.tail_bound = tail_bound
|
|
self.half_channels = in_channels // 2
|
|
|
|
self.pre = nn.Conv1d(self.half_channels, filter_channels, 1)
|
|
self.convs = DDSConv(filter_channels, kernel_size, n_layers, p_dropout=0.)
|
|
self.proj = nn.Conv1d(filter_channels, self.half_channels * (num_bins * 3 - 1), 1)
|
|
self.proj.weight.data.zero_()
|
|
self.proj.bias.data.zero_()
|
|
|
|
def forward(self, x, x_mask, g=None, reverse=False):
|
|
x0, x1 = torch.split(x, [self.half_channels]*2, 1)
|
|
h = self.pre(x0)
|
|
h = self.convs(h, x_mask, g=g)
|
|
h = self.proj(h) * x_mask
|
|
|
|
b, c, t = x0.shape
|
|
h = h.reshape(b, c, -1, t).permute(0, 1, 3, 2) # [b, cx?, t] -> [b, c, t, ?]
|
|
|
|
unnormalized_widths = h[..., :self.num_bins] / math.sqrt(self.filter_channels)
|
|
unnormalized_heights = h[..., self.num_bins:2*self.num_bins] / math.sqrt(self.filter_channels)
|
|
unnormalized_derivatives = h[..., 2 * self.num_bins:]
|
|
|
|
x1, logabsdet = piecewise_rational_quadratic_transform(x1,
|
|
unnormalized_widths,
|
|
unnormalized_heights,
|
|
unnormalized_derivatives,
|
|
inverse=reverse,
|
|
tails='linear',
|
|
tail_bound=self.tail_bound
|
|
)
|
|
|
|
x = torch.cat([x0, x1], 1) * x_mask
|
|
logdet = torch.sum(logabsdet * x_mask, [1,2])
|
|
if not reverse:
|
|
return x, logdet
|
|
else:
|
|
return x
|
|
|
|
|
|
class PositionalEmbedding(nn.Module):
|
|
def __init__(self, demb):
|
|
super(PositionalEmbedding, self).__init__()
|
|
self.demb = demb
|
|
inv_freq = 1 / (10000 ** (torch.arange(0.0, demb, 2.0) / demb))
|
|
self.register_buffer('inv_freq', inv_freq)
|
|
|
|
def forward(self, pos_seq, bsz=None):
|
|
# sinusoid_inp = torch.ger(pos_seq, self.inv_freq)
|
|
sinusoid_inp = torch.matmul(torch.unsqueeze(pos_seq, -1), torch.unsqueeze(self.inv_freq, 0))
|
|
|
|
pos_emb = torch.cat([sinusoid_inp.sin(), sinusoid_inp.cos()], dim=1)
|
|
if bsz is not None:
|
|
return pos_emb[None, :, :].repeat(bsz, 1, 1)
|
|
else:
|
|
return pos_emb[None, :, :]
|
|
|
|
|
|
class SelfAttentionModule(nn.Module):
|
|
"""Self-attention for lm tokens and text. """
|
|
|
|
def __init__(self, n_text_channels=384, n_lm_tokens_channels=128):
|
|
super().__init__()
|
|
|
|
self.text_pos_emb = PositionalEmbedding(n_text_channels)
|
|
self.lm_pos_emb = PositionalEmbedding(n_lm_tokens_channels)
|
|
|
|
self.query_proj = nn.Sequential(
|
|
ConvNorm(n_text_channels, n_text_channels, kernel_size=3, bias=True, w_init_gain='relu'),
|
|
torch.nn.ReLU(),
|
|
ConvNorm(n_text_channels, n_text_channels, kernel_size=1, bias=True),
|
|
)
|
|
|
|
self.key_proj = nn.Sequential(
|
|
ConvNorm(n_lm_tokens_channels, n_text_channels, kernel_size=3, bias=True, w_init_gain='relu'),
|
|
torch.nn.ReLU(),
|
|
ConvNorm(n_text_channels, n_text_channels, kernel_size=1, bias=True),
|
|
)
|
|
|
|
self.value_proj = nn.Sequential(
|
|
ConvNorm(n_lm_tokens_channels, n_text_channels, kernel_size=3, bias=True, w_init_gain='relu'),
|
|
torch.nn.ReLU(),
|
|
ConvNorm(n_text_channels, n_text_channels, kernel_size=1, bias=True),
|
|
)
|
|
|
|
self.scale = math.sqrt(n_text_channels)
|
|
|
|
def forward(self, queries, keys, values, q_mask=None, kv_mask=None):
|
|
"""Forward pass of self-attention.
|
|
Args:
|
|
queries (torch.tensor): B x T1 x C1 tensor
|
|
keys (torch.tensor): B x T2 x C2 tensor
|
|
values (torch.tensor): B x T2 x C2 tensor
|
|
q_mask (torch.tensor): B x T1 tensor, bool mask for variable length entries
|
|
kv_mask (torch.tensor): B x T2 tensor, bool mask for variable length entries
|
|
Output:
|
|
attn_out (torch.tensor): B x T1 x C1 tensor
|
|
"""
|
|
pos_q_seq = torch.arange(queries.size(-2), device=queries.device).to(queries.dtype)
|
|
pos_kv_seq = torch.arange(keys.size(-2), device=queries.device).to(queries.dtype)
|
|
|
|
pos_q_emb = self.text_pos_emb(pos_q_seq)
|
|
pos_kv_emb = self.lm_pos_emb(pos_kv_seq)
|
|
|
|
if q_mask is not None:
|
|
pos_q_emb = pos_q_emb * q_mask.unsqueeze(2)
|
|
|
|
if kv_mask is not None:
|
|
pos_kv_emb = pos_kv_emb * kv_mask.unsqueeze(2)
|
|
|
|
queries = (queries + pos_q_emb).transpose(1, 2)
|
|
keys = (keys + pos_kv_emb).transpose(1, 2)
|
|
values = (values + pos_kv_emb).transpose(1, 2)
|
|
|
|
queries_enc = self.query_proj(queries).transpose(-2, -1) # B x T1 x C1
|
|
keys_enc = self.key_proj(keys) # B x C1 x T2
|
|
values_enc = self.value_proj(values).transpose(-2, -1) # B x T2 x C1
|
|
|
|
scores = torch.matmul(queries_enc, keys_enc) / self.scale # B x T1 x T2
|
|
|
|
if kv_mask is not None:
|
|
scores.masked_fill_(~kv_mask.unsqueeze(-2), -float("inf"))
|
|
|
|
return torch.matmul(torch.softmax(scores, dim=-1), values_enc) # B x T1 x C1
|