MCPcopy Create free account
hub / github.com/KeepTryingTo/Pytorch-GAN / loss_function

Method loss_function

PyTorch-VAE/models/swae.py:110–127  ·  view source on GitHub ↗
(self,
                      *args,
                      **kwargs)

Source from the content-addressed store, hash-verified

108 return [self.decode(z), input, z]
109
110 def loss_function(self,
111 *args,
112 **kwargs) -> dict:
113 recons = args[0]
114 input = args[1]
115 z = args[2]
116
117 batch_size = input.size(0)
118 bias_corr = batch_size * (batch_size - 1)
119 reg_weight = self.reg_weight / bias_corr
120
121 recons_loss_l2 = F.mse_loss(recons, input)
122 recons_loss_l1 = F.l1_loss(recons, input)
123
124 swd_loss = self.compute_swd(z, self.p, reg_weight)
125
126 loss = recons_loss_l2 + recons_loss_l1 + swd_loss
127 return {'loss': loss, 'Reconstruction_Loss':(recons_loss_l2 + recons_loss_l1), 'SWD': swd_loss}
128
129 def get_random_projections(self, latent_dim: int, num_samples: int) -> Tensor:
130 """

Callers

nothing calls this directly

Calls 1

compute_swdMethod · 0.95

Tested by

no test coverage detected