Team Ai
Apppublic

fighter-programmer/voicegen

sourceHugging Faceupdated 4y agoView on Hugging Face
0likes
modules.py388 linesDownload Raw Back to root
1import math2import torch3from torch import nn4from torch.nn import functional as F5 6from torch.nn import Conv1d7from torch.nn.utils import weight_norm, remove_weight_norm8 9import commons10from commons import init_weights, get_padding11from transforms import piecewise_rational_quadratic_transform12 13 14LRELU_SLOPE = 0.115 16 17class LayerNorm(nn.Module):18  def __init__(self, channels, eps=1e-5):19    super().__init__()20    self.channels = channels21    self.eps = eps22 23    self.gamma = nn.Parameter(torch.ones(channels))24    self.beta = nn.Parameter(torch.zeros(channels))25 26  def forward(self, x):27    x = x.transpose(1, -1)28    x = F.layer_norm(x, (self.channels,), self.gamma, self.beta, self.eps)29    return x.transpose(1, -1)30 31 32class ConvReluNorm(nn.Module):33  def __init__(self, in_channels, hidden_channels, out_channels, kernel_size, n_layers, p_dropout):34    super().__init__()35    self.in_channels = in_channels36    self.hidden_channels = hidden_channels37    self.out_channels = out_channels38    self.kernel_size = kernel_size39    self.n_layers = n_layers40    self.p_dropout = p_dropout41    assert n_layers > 1, "Number of layers should be larger than 0."42 43    self.conv_layers = nn.ModuleList()44    self.norm_layers = nn.ModuleList()45    self.conv_layers.append(nn.Conv1d(in_channels, hidden_channels, kernel_size, padding=kernel_size//2))46    self.norm_layers.append(LayerNorm(hidden_channels))47    self.relu_drop = nn.Sequential(48        nn.ReLU(),49        nn.Dropout(p_dropout))50    for _ in range(n_layers-1):51      self.conv_layers.append(nn.Conv1d(hidden_channels, hidden_channels, kernel_size, padding=kernel_size//2))52      self.norm_layers.append(LayerNorm(hidden_channels))53    self.proj = nn.Conv1d(hidden_channels, out_channels, 1)54    self.proj.weight.data.zero_()55    self.proj.bias.data.zero_()56 57  def forward(self, x, x_mask):58    x_org = x59    for i in range(self.n_layers):60      x = self.conv_layers[i](x * x_mask)61      x = self.norm_layers[i](x)62      x = self.relu_drop(x)63    x = x_org + self.proj(x)64    return x * x_mask65 66 67class DDSConv(nn.Module):68  """69  Dialted and Depth-Separable Convolution70  """71  def __init__(self, channels, kernel_size, n_layers, p_dropout=0.):72    super().__init__()73    self.channels = channels74    self.kernel_size = kernel_size75    self.n_layers = n_layers76    self.p_dropout = p_dropout77 78    self.drop = nn.Dropout(p_dropout)79    self.convs_sep = nn.ModuleList()80    self.convs_1x1 = nn.ModuleList()81    self.norms_1 = nn.ModuleList()82    self.norms_2 = nn.ModuleList()83    for i in range(n_layers):84      dilation = kernel_size ** i85      padding = (kernel_size * dilation - dilation) // 286      self.convs_sep.append(nn.Conv1d(channels, channels, kernel_size, 87          groups=channels, dilation=dilation, padding=padding88      ))89      self.convs_1x1.append(nn.Conv1d(channels, channels, 1))90      self.norms_1.append(LayerNorm(channels))91      self.norms_2.append(LayerNorm(channels))92 93  def forward(self, x, x_mask, g=None):94    if g is not None:95      x = x + g96    for i in range(self.n_layers):97      y = self.convs_sep[i](x * x_mask)98      y = self.norms_1[i](y)99      y = F.gelu(y)100      y = self.convs_1x1[i](y)101      y = self.norms_2[i](y)102      y = F.gelu(y)103      y = self.drop(y)104      x = x + y105    return x * x_mask106 107 108class WN(torch.nn.Module):109  def __init__(self, hidden_channels, kernel_size, dilation_rate, n_layers, gin_channels=0, p_dropout=0):110    super(WN, self).__init__()111    assert(kernel_size % 2 == 1)112    self.hidden_channels =hidden_channels113    self.kernel_size = kernel_size,114    self.dilation_rate = dilation_rate115    self.n_layers = n_layers116    self.gin_channels = gin_channels117    self.p_dropout = p_dropout118 119    self.in_layers = torch.nn.ModuleList()120    self.res_skip_layers = torch.nn.ModuleList()121    self.drop = nn.Dropout(p_dropout)122 123    if gin_channels != 0:124      cond_layer = torch.nn.Conv1d(gin_channels, 2*hidden_channels*n_layers, 1)125      self.cond_layer = torch.nn.utils.weight_norm(cond_layer, name='weight')126 127    for i in range(n_layers):128      dilation = dilation_rate ** i129      padding = int((kernel_size * dilation - dilation) / 2)130      in_layer = torch.nn.Conv1d(hidden_channels, 2*hidden_channels, kernel_size,131                                 dilation=dilation, padding=padding)132      in_layer = torch.nn.utils.weight_norm(in_layer, name='weight')133      self.in_layers.append(in_layer)134 135      # last one is not necessary136      if i < n_layers - 1:137        res_skip_channels = 2 * hidden_channels138      else:139        res_skip_channels = hidden_channels140 141      res_skip_layer = torch.nn.Conv1d(hidden_channels, res_skip_channels, 1)142      res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name='weight')143      self.res_skip_layers.append(res_skip_layer)144 145  def forward(self, x, x_mask, g=None, **kwargs):146    output = torch.zeros_like(x)147    n_channels_tensor = torch.IntTensor([self.hidden_channels])148 149    if g is not None:150      g = self.cond_layer(g)151 152    for i in range(self.n_layers):153      x_in = self.in_layers[i](x)154      if g is not None:155        cond_offset = i * 2 * self.hidden_channels156        g_l = g[:,cond_offset:cond_offset+2*self.hidden_channels,:]157      else:158        g_l = torch.zeros_like(x_in)159 160      acts = commons.fused_add_tanh_sigmoid_multiply(161          x_in,162          g_l,163          n_channels_tensor)164      acts = self.drop(acts)165 166      res_skip_acts = self.res_skip_layers[i](acts)167      if i < self.n_layers - 1:168        res_acts = res_skip_acts[:,:self.hidden_channels,:]169        x = (x + res_acts) * x_mask170        output = output + res_skip_acts[:,self.hidden_channels:,:]171      else:172        output = output + res_skip_acts173    return output * x_mask174 175  def remove_weight_norm(self):176    if self.gin_channels != 0:177      torch.nn.utils.remove_weight_norm(self.cond_layer)178    for l in self.in_layers:179      torch.nn.utils.remove_weight_norm(l)180    for l in self.res_skip_layers:181     torch.nn.utils.remove_weight_norm(l)182 183 184class ResBlock1(torch.nn.Module):185    def __init__(self, channels, kernel_size=3, dilation=(1, 3, 5)):186        super(ResBlock1, self).__init__()187        self.convs1 = nn.ModuleList([188            weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=dilation[0],189                               padding=get_padding(kernel_size, dilation[0]))),190            weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=dilation[1],191                               padding=get_padding(kernel_size, dilation[1]))),192            weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=dilation[2],193                               padding=get_padding(kernel_size, dilation[2])))194        ])195        self.convs1.apply(init_weights)196 197        self.convs2 = nn.ModuleList([198            weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=1,199                               padding=get_padding(kernel_size, 1))),200            weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=1,201                               padding=get_padding(kernel_size, 1))),202            weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=1,203                               padding=get_padding(kernel_size, 1)))204        ])205        self.convs2.apply(init_weights)206 207    def forward(self, x, x_mask=None):208        for c1, c2 in zip(self.convs1, self.convs2):209            xt = F.leaky_relu(x, LRELU_SLOPE)210            if x_mask is not None:211                xt = xt * x_mask212            xt = c1(xt)213            xt = F.leaky_relu(xt, LRELU_SLOPE)214            if x_mask is not None:215                xt = xt * x_mask216            xt = c2(xt)217            x = xt + x218        if x_mask is not None:219            x = x * x_mask220        return x221 222    def remove_weight_norm(self):223        for l in self.convs1:224            remove_weight_norm(l)225        for l in self.convs2:226            remove_weight_norm(l)227 228 229class ResBlock2(torch.nn.Module):230    def __init__(self, channels, kernel_size=3, dilation=(1, 3)):231        super(ResBlock2, self).__init__()232        self.convs = nn.ModuleList([233            weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=dilation[0],234                               padding=get_padding(kernel_size, dilation[0]))),235            weight_norm(Conv1d(channels, channels, kernel_size, 1, dilation=dilation[1],236                               padding=get_padding(kernel_size, dilation[1])))237        ])238        self.convs.apply(init_weights)239 240    def forward(self, x, x_mask=None):241        for c in self.convs:242            xt = F.leaky_relu(x, LRELU_SLOPE)243            if x_mask is not None:244                xt = xt * x_mask245            xt = c(xt)246            x = xt + x247        if x_mask is not None:248            x = x * x_mask249        return x250 251    def remove_weight_norm(self):252        for l in self.convs:253            remove_weight_norm(l)254 255 256class Log(nn.Module):257  def forward(self, x, x_mask, reverse=False, **kwargs):258    if not reverse:259      y = torch.log(torch.clamp_min(x, 1e-5)) * x_mask260      logdet = torch.sum(-y, [1, 2])261      return y, logdet262    else:263      x = torch.exp(x) * x_mask264      return x265    266 267class Flip(nn.Module):268  def forward(self, x, *args, reverse=False, **kwargs):269    x = torch.flip(x, [1])270    if not reverse:271      logdet = torch.zeros(x.size(0)).to(dtype=x.dtype, device=x.device)272      return x, logdet273    else:274      return x275 276 277class ElementwiseAffine(nn.Module):278  def __init__(self, channels):279    super().__init__()280    self.channels = channels281    self.m = nn.Parameter(torch.zeros(channels,1))282    self.logs = nn.Parameter(torch.zeros(channels,1))283 284  def forward(self, x, x_mask, reverse=False, **kwargs):285    if not reverse:286      y = self.m + torch.exp(self.logs) * x287      y = y * x_mask288      logdet = torch.sum(self.logs * x_mask, [1,2])289      return y, logdet290    else:291      x = (x - self.m) * torch.exp(-self.logs) * x_mask292      return x293 294 295class ResidualCouplingLayer(nn.Module):296  def __init__(self,297      channels,298      hidden_channels,299      kernel_size,300      dilation_rate,301      n_layers,302      p_dropout=0,303      gin_channels=0,304      mean_only=False):305    assert channels % 2 == 0, "channels should be divisible by 2"306    super().__init__()307    self.channels = channels308    self.hidden_channels = hidden_channels309    self.kernel_size = kernel_size310    self.dilation_rate = dilation_rate311    self.n_layers = n_layers312    self.half_channels = channels // 2313    self.mean_only = mean_only314 315    self.pre = nn.Conv1d(self.half_channels, hidden_channels, 1)316    self.enc = WN(hidden_channels, kernel_size, dilation_rate, n_layers, p_dropout=p_dropout, gin_channels=gin_channels)317    self.post = nn.Conv1d(hidden_channels, self.half_channels * (2 - mean_only), 1)318    self.post.weight.data.zero_()319    self.post.bias.data.zero_()320 321  def forward(self, x, x_mask, g=None, reverse=False):322    x0, x1 = torch.split(x, [self.half_channels]*2, 1)323    h = self.pre(x0) * x_mask324    h = self.enc(h, x_mask, g=g)325    stats = self.post(h) * x_mask326    if not self.mean_only:327      m, logs = torch.split(stats, [self.half_channels]*2, 1)328    else:329      m = stats330      logs = torch.zeros_like(m)331 332    if not reverse:333      x1 = m + x1 * torch.exp(logs) * x_mask334      x = torch.cat([x0, x1], 1)335      logdet = torch.sum(logs, [1,2])336      return x, logdet337    else:338      x1 = (x1 - m) * torch.exp(-logs) * x_mask339      x = torch.cat([x0, x1], 1)340      return x341 342 343class ConvFlow(nn.Module):344  def __init__(self, in_channels, filter_channels, kernel_size, n_layers, num_bins=10, tail_bound=5.0):345    super().__init__()346    self.in_channels = in_channels347    self.filter_channels = filter_channels348    self.kernel_size = kernel_size349    self.n_layers = n_layers350    self.num_bins = num_bins351    self.tail_bound = tail_bound352    self.half_channels = in_channels // 2353 354    self.pre = nn.Conv1d(self.half_channels, filter_channels, 1)355    self.convs = DDSConv(filter_channels, kernel_size, n_layers, p_dropout=0.)356    self.proj = nn.Conv1d(filter_channels, self.half_channels * (num_bins * 3 - 1), 1)357    self.proj.weight.data.zero_()358    self.proj.bias.data.zero_()359 360  def forward(self, x, x_mask, g=None, reverse=False):361    x0, x1 = torch.split(x, [self.half_channels]*2, 1)362    h = self.pre(x0)363    h = self.convs(h, x_mask, g=g)364    h = self.proj(h) * x_mask365 366    b, c, t = x0.shape367    h = h.reshape(b, c, -1, t).permute(0, 1, 3, 2) # [b, cx?, t] -> [b, c, t, ?]368 369    unnormalized_widths = h[..., :self.num_bins] / math.sqrt(self.filter_channels)370    unnormalized_heights = h[..., self.num_bins:2*self.num_bins] / math.sqrt(self.filter_channels)371    unnormalized_derivatives = h[..., 2 * self.num_bins:]372 373    x1, logabsdet = piecewise_rational_quadratic_transform(x1,374        unnormalized_widths,375        unnormalized_heights,376        unnormalized_derivatives,377        inverse=reverse,378        tails='linear',379        tail_bound=self.tail_bound380    )381 382    x = torch.cat([x0, x1], 1) * x_mask383    logdet = torch.sum(logabsdet * x_mask, [1,2])384    if not reverse:385        return x, logdet386    else:387        return x388