init (not checked bat script yet)

This commit is contained in:
litagin02
2023-12-27 05:37:46 +09:00
parent 11a1e7e80d
commit 58fab45b84
235 changed files with 2831 additions and 773509 deletions

157
models.py
View File

@@ -1,18 +1,21 @@
import math
import warnings
import torch
from torch import nn
from torch.nn import Conv1d, Conv2d, ConvTranspose1d
from torch.nn import functional as F
import attentions
import commons
import modules
import attentions
import monotonic_align
from commons import get_padding, init_weights
from text import num_languages, num_tones, symbols
from torch.nn import Conv1d, ConvTranspose1d, Conv2d
from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm
from commons import init_weights, get_padding
from text import symbols, num_tones, num_languages
with warnings.catch_warnings():
warnings.simplefilter("ignore")
from torch.nn.utils import remove_weight_norm, spectral_norm, weight_norm
class DurationDiscriminator(nn.Module): # vits2
@@ -38,22 +41,33 @@ class DurationDiscriminator(nn.Module): # vits2
self.norm_2 = modules.LayerNorm(filter_channels)
self.dur_proj = nn.Conv1d(1, filter_channels, 1)
self.LSTM = nn.LSTM(
2 * filter_channels, filter_channels, batch_first=True, bidirectional=True
self.pre_out_conv_1 = nn.Conv1d(
2 * filter_channels, filter_channels, kernel_size, padding=kernel_size // 2
)
self.pre_out_norm_1 = modules.LayerNorm(filter_channels)
self.pre_out_conv_2 = nn.Conv1d(
filter_channels, filter_channels, kernel_size, padding=kernel_size // 2
)
self.pre_out_norm_2 = modules.LayerNorm(filter_channels)
if gin_channels != 0:
self.cond = nn.Conv1d(gin_channels, in_channels, 1)
self.output_layer = nn.Sequential(
nn.Linear(2 * filter_channels, 1), nn.Sigmoid()
)
self.output_layer = nn.Sequential(nn.Linear(filter_channels, 1), nn.Sigmoid())
def forward_probability(self, x, dur):
def forward_probability(self, x, x_mask, dur, g=None):
dur = self.dur_proj(dur)
x = torch.cat([x, dur], dim=1)
x = self.pre_out_conv_1(x * x_mask)
x = torch.relu(x)
x = self.pre_out_norm_1(x)
x = self.drop(x)
x = self.pre_out_conv_2(x * x_mask)
x = torch.relu(x)
x = self.pre_out_norm_2(x)
x = self.drop(x)
x = x * x_mask
x = x.transpose(1, 2)
x, _ = self.LSTM(x)
output_prob = self.output_layer(x)
return output_prob
@@ -73,7 +87,7 @@ class DurationDiscriminator(nn.Module): # vits2
output_probs = []
for dur in [dur_r, dur_hat]:
output_prob = self.forward_probability(x, dur)
output_prob = self.forward_probability(x, x_mask, dur, g)
output_probs.append(output_prob)
return output_probs
@@ -299,37 +313,6 @@ class DurationPredictor(nn.Module):
return x * x_mask
class Bottleneck(nn.Sequential):
def __init__(self, in_dim, hidden_dim):
c_fc1 = nn.Linear(in_dim, hidden_dim, bias=False)
c_fc2 = nn.Linear(in_dim, hidden_dim, bias=False)
super().__init__(*[c_fc1, c_fc2])
class Block(nn.Module):
def __init__(self, in_dim, hidden_dim) -> None:
super().__init__()
self.norm = nn.LayerNorm(in_dim)
self.mlp = MLP(in_dim, hidden_dim)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x + self.mlp(self.norm(x))
return x
class MLP(nn.Module):
def __init__(self, in_dim, hidden_dim):
super().__init__()
self.c_fc1 = nn.Linear(in_dim, hidden_dim, bias=False)
self.c_fc2 = nn.Linear(in_dim, hidden_dim, bias=False)
self.c_proj = nn.Linear(hidden_dim, in_dim, bias=False)
def forward(self, x: torch.Tensor):
x = F.silu(self.c_fc1(x)) * self.c_fc2(x)
x = self.c_proj(x)
return x
class TextEncoder(nn.Module):
def __init__(
self,
@@ -341,6 +324,7 @@ class TextEncoder(nn.Module):
n_layers,
kernel_size,
p_dropout,
n_speakers,
gin_channels=0,
):
super().__init__()
@@ -362,6 +346,7 @@ class TextEncoder(nn.Module):
self.bert_proj = nn.Conv1d(1024, hidden_channels, 1)
self.ja_bert_proj = nn.Conv1d(1024, hidden_channels, 1)
self.en_bert_proj = nn.Conv1d(1024, hidden_channels, 1)
self.style_proj = nn.Linear(256, hidden_channels)
self.encoder = attentions.Encoder(
hidden_channels,
@@ -374,10 +359,24 @@ class TextEncoder(nn.Module):
)
self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
def forward(self, x, x_lengths, tone, language, bert, ja_bert, en_bert, g=None):
def forward(
self,
x,
x_lengths,
tone,
language,
bert,
ja_bert,
en_bert,
style_vec,
sid,
g=None,
):
bert_emb = self.bert_proj(bert).transpose(1, 2)
ja_bert_emb = self.ja_bert_proj(ja_bert).transpose(1, 2)
en_bert_emb = self.en_bert_proj(en_bert).transpose(1, 2)
style_emb = self.style_proj(style_vec.unsqueeze(1))
x = (
self.emb(x)
+ self.tone_emb(tone)
@@ -385,6 +384,7 @@ class TextEncoder(nn.Module):
+ bert_emb
+ ja_bert_emb
+ en_bert_emb
+ style_emb
) * math.sqrt(
self.hidden_channels
) # [b, t, h]
@@ -700,55 +700,6 @@ class MultiPeriodDiscriminator(torch.nn.Module):
return y_d_rs, y_d_gs, fmap_rs, fmap_gs
class WavLMDiscriminator(nn.Module):
"""docstring for Discriminator."""
def __init__(
self, slm_hidden=768, slm_layers=13, initial_channel=64, use_spectral_norm=False
):
super(WavLMDiscriminator, self).__init__()
norm_f = weight_norm if use_spectral_norm == False else spectral_norm
self.pre = norm_f(
Conv1d(slm_hidden * slm_layers, initial_channel, 1, 1, padding=0)
)
self.convs = nn.ModuleList(
[
norm_f(
nn.Conv1d(
initial_channel, initial_channel * 2, kernel_size=5, padding=2
)
),
norm_f(
nn.Conv1d(
initial_channel * 2,
initial_channel * 4,
kernel_size=5,
padding=2,
)
),
norm_f(
nn.Conv1d(initial_channel * 4, initial_channel * 4, 5, 1, padding=2)
),
]
)
self.conv_post = norm_f(Conv1d(initial_channel * 4, 1, 3, 1, padding=1))
def forward(self, x):
x = self.pre(x)
fmap = []
for l in self.convs:
x = l(x)
x = F.leaky_relu(x, modules.LRELU_SLOPE)
fmap.append(x)
x = self.conv_post(x)
x = torch.flatten(x, 1, -1)
return x
class ReferenceEncoder(nn.Module):
"""
inputs --- [N, Ty/r, n_mels*r] mels
@@ -838,7 +789,7 @@ class SynthesizerTrn(nn.Module):
n_layers_trans_flow=4,
flow_share_parameter=False,
use_transformer_flow=True,
**kwargs
**kwargs,
):
super().__init__()
self.n_vocab = n_vocab
@@ -879,6 +830,7 @@ class SynthesizerTrn(nn.Module):
n_layers,
kernel_size,
p_dropout,
self.n_speakers,
gin_channels=self.enc_gin_channels,
)
self.dec = Generator(
@@ -946,13 +898,14 @@ class SynthesizerTrn(nn.Module):
bert,
ja_bert,
en_bert,
style_vec,
):
if self.n_speakers > 0:
g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]
else:
g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1)
x, m_p, logs_p, x_mask = self.enc_p(
x, x_lengths, tone, language, bert, ja_bert, en_bert, g=g
x, x_lengths, tone, language, bert, ja_bert, en_bert, style_vec, sid, g=g
)
z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)
z_p = self.flow(z, y_mask, g=g)
@@ -995,11 +948,11 @@ class SynthesizerTrn(nn.Module):
logw_ = torch.log(w + 1e-6) * x_mask
logw = self.dp(x, x_mask, g=g)
logw_sdp = self.sdp(x, x_mask, g=g, reverse=True, noise_scale=1.0)
# logw_sdp = self.sdp(x, x_mask, g=g, reverse=True, noise_scale=1.0)
l_length_dp = torch.sum((logw - logw_) ** 2, [1, 2]) / torch.sum(
x_mask
) # for averaging
l_length_sdp += torch.sum((logw_sdp - logw_) ** 2, [1, 2]) / torch.sum(x_mask)
# l_length_sdp += torch.sum((logw_sdp - logw_) ** 2, [1, 2]) / torch.sum(x_mask)
l_length = l_length_dp + l_length_sdp
@@ -1019,8 +972,7 @@ class SynthesizerTrn(nn.Module):
x_mask,
y_mask,
(z, z_p, m_p, logs_p, m_q, logs_q),
(x, logw, logw_, logw_sdp),
g,
(x, logw, logw_),
)
def infer(
@@ -1033,6 +985,7 @@ class SynthesizerTrn(nn.Module):
bert,
ja_bert,
en_bert,
style_vec,
noise_scale=0.667,
length_scale=1,
noise_scale_w=0.8,
@@ -1047,7 +1000,7 @@ class SynthesizerTrn(nn.Module):
else:
g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1)
x, m_p, logs_p, x_mask = self.enc_p(
x, x_lengths, tone, language, bert, ja_bert, en_bert, g=g
x, x_lengths, tone, language, bert, ja_bert, en_bert, style_vec, sid, g=g
)
logw = self.sdp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w) * (
sdp_ratio