Files
2023-05-22 23:09:00 -03:00

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