Added my good ole pattern loss. God I love that thing, conv transpose pattern instantly wiped from vae
This commit is contained in:
@@ -1,4 +1,5 @@
|
||||
import torch
|
||||
from .llvae import LosslessLatentEncoder
|
||||
|
||||
|
||||
def total_variation(image):
|
||||
@@ -45,3 +46,40 @@ def get_gradient_penalty(critic, real, fake, device):
|
||||
gradient_penalty = ((gradient_norm - 1) ** 2).mean()
|
||||
return gradient_penalty
|
||||
|
||||
|
||||
class PatternLoss(torch.nn.Module):
|
||||
def __init__(self, pattern_size=4, dtype=torch.float32):
|
||||
super().__init__()
|
||||
self.pattern_size = pattern_size
|
||||
self.llvae_encoder = LosslessLatentEncoder(3, pattern_size, dtype=dtype)
|
||||
|
||||
def forward(self, pred, target):
|
||||
pred_latents = self.llvae_encoder(pred)
|
||||
target_latents = self.llvae_encoder(target)
|
||||
|
||||
matrix_pixels = self.pattern_size * self.pattern_size
|
||||
|
||||
color_chans = pred_latents.shape[1] // 3
|
||||
# pytorch
|
||||
r_chans, g_chans, b_chans = torch.split(pred_latents, [color_chans, color_chans, color_chans], 1)
|
||||
r_chans_target, g_chans_target, b_chans_target = torch.split(target_latents, [color_chans, color_chans, color_chans], 1)
|
||||
|
||||
def separated_chan_loss(latent_chan):
|
||||
nonlocal matrix_pixels
|
||||
chan_mean = torch.mean(latent_chan, dim=[1, 2, 3])
|
||||
chan_splits = torch.split(latent_chan, [1 for i in range(matrix_pixels)], 1)
|
||||
chan_loss = None
|
||||
for chan in chan_splits:
|
||||
this_mean = torch.mean(chan, dim=[1, 2, 3])
|
||||
this_chan_loss = torch.abs(this_mean - chan_mean)
|
||||
if chan_loss is None:
|
||||
chan_loss = this_chan_loss
|
||||
else:
|
||||
chan_loss = chan_loss + this_chan_loss
|
||||
chan_loss = chan_loss * (1 / matrix_pixels)
|
||||
return chan_loss
|
||||
|
||||
r_chan_loss = torch.abs(separated_chan_loss(r_chans) - separated_chan_loss(r_chans_target))
|
||||
g_chan_loss = torch.abs(separated_chan_loss(g_chans) - separated_chan_loss(g_chans_target))
|
||||
b_chan_loss = torch.abs(separated_chan_loss(b_chans) - separated_chan_loss(b_chans_target))
|
||||
return (r_chan_loss + g_chan_loss + b_chan_loss) * 0.3333
|
||||
|
||||
Reference in New Issue
Block a user