CoolFace
Apppublic

goathead777/Zero_Shot_Inference

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
losses.py74 linesDownload Raw Back to module
1import math2 3import torch4from torch.nn import functional as F5 6 7def feature_loss(fmap_r, fmap_g):8    loss = 09    for dr, dg in zip(fmap_r, fmap_g):10        for rl, gl in zip(dr, dg):11            rl = rl.float().detach()12            gl = gl.float()13            loss += torch.mean(torch.abs(rl - gl))14 15    return loss * 216 17 18def discriminator_loss(disc_real_outputs, disc_generated_outputs):19    loss = 020    r_losses = []21    g_losses = []22    for dr, dg in zip(disc_real_outputs, disc_generated_outputs):23        dr = dr.float()24        dg = dg.float()25        r_loss = torch.mean((1 - dr) ** 2)26        g_loss = torch.mean(dg**2)27        loss += r_loss + g_loss28        r_losses.append(r_loss.item())29        g_losses.append(g_loss.item())30 31    return loss, r_losses, g_losses32 33 34def generator_loss(disc_outputs):35    loss = 036    gen_losses = []37    for dg in disc_outputs:38        dg = dg.float()39        l = torch.mean((1 - dg) ** 2)40        gen_losses.append(l)41        loss += l42 43    return loss, gen_losses44 45 46def kl_loss(z_p, logs_q, m_p, logs_p, z_mask):47    """48    z_p, logs_q: [b, h, t_t]49    m_p, logs_p: [b, h, t_t]50    """51    z_p = z_p.float()52    logs_q = logs_q.float()53    m_p = m_p.float()54    logs_p = logs_p.float()55    z_mask = z_mask.float()56 57    kl = logs_p - logs_q - 0.558    kl += 0.5 * ((z_p - m_p) ** 2) * torch.exp(-2.0 * logs_p)59    kl = torch.sum(kl * z_mask)60    l = kl / torch.sum(z_mask)61    return l62 63 64def mle_loss(z, m, logs, logdet, mask):65    l = torch.sum(logs) + 0.5 * torch.sum(66        torch.exp(-2 * logs) * ((z - m) ** 2)67    )  # neg normal likelihood w/o the constant term68    l = l - torch.sum(logdet)  # log jacobian determinant69    l = l / torch.sum(70        torch.ones_like(z) * mask71    )  # averaging across batch, channel and time axes72    l = l + 0.5 * math.log(2 * math.pi)  # add the remaining constant term73    return l74