ChazzyG/Retrieval-based-Voice-Conversion-WebUI
0
1import torch2from torch.nn import functional as F3 4 5def feature_loss(fmap_r, fmap_g):6 loss = 07 for dr, dg in zip(fmap_r, fmap_g):8 for rl, gl in zip(dr, dg):9 rl = rl.float().detach()10 gl = gl.float()11 loss += torch.mean(torch.abs(rl - gl))12 13 return loss * 214 15 16def discriminator_loss(disc_real_outputs, disc_generated_outputs):17 loss = 018 r_losses = []19 g_losses = []20 for dr, dg in zip(disc_real_outputs, disc_generated_outputs):21 dr = dr.float()22 dg = dg.float()23 r_loss = torch.mean((1 - dr) ** 2)24 g_loss = torch.mean(dg**2)25 loss += r_loss + g_loss26 r_losses.append(r_loss.item())27 g_losses.append(g_loss.item())28 29 return loss, r_losses, g_losses30 31 32def generator_loss(disc_outputs):33 loss = 034 gen_losses = []35 for dg in disc_outputs:36 dg = dg.float()37 l = torch.mean((1 - dg) ** 2)38 gen_losses.append(l)39 loss += l40 41 return loss, gen_losses42 43 44def kl_loss(z_p, logs_q, m_p, logs_p, z_mask):45 """46 z_p, logs_q: [b, h, t_t]47 m_p, logs_p: [b, h, t_t]48 """49 z_p = z_p.float()50 logs_q = logs_q.float()51 m_p = m_p.float()52 logs_p = logs_p.float()53 z_mask = z_mask.float()54 55 kl = logs_p - logs_q - 0.556 kl += 0.5 * ((z_p - m_p) ** 2) * torch.exp(-2.0 * logs_p)57 kl = torch.sum(kl * z_mask)58 l = kl / torch.sum(z_mask)59 return l60 