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

Method loss_function

PyTorch-VAE/models/dfcvae.py:163–190  ·  view source on GitHub ↗

Computes the VAE loss function. KL(N(\mu, \sigma), N(0, 1)) = \log \frac{1}{\sigma} + \frac{\sigma^2 + \mu^2}{2} - \frac{1}{2} :param args: :param kwargs: :return:

(self,
                      *args,
                      **kwargs)

Source from the content-addressed store, hash-verified

161 return features
162
163 def loss_function(self,
164 *args,
165 **kwargs) -> dict:
166 """
167 Computes the VAE loss function.
168 KL(N(\mu, \sigma), N(0, 1)) = \log \frac{1}{\sigma} + \frac{\sigma^2 + \mu^2}{2} - \frac{1}{2}
169 :param args:
170 :param kwargs:
171 :return:
172 """
173 recons = args[0]
174 input = args[1]
175 recons_features = args[2]
176 input_features = args[3]
177 mu = args[4]
178 log_var = args[5]
179
180 kld_weight = kwargs['M_N'] # Account for the minibatch samples from the dataset
181 recons_loss =F.mse_loss(recons, input)
182
183 feature_loss = 0.0
184 for (r, i) in zip(recons_features, input_features):
185 feature_loss += F.mse_loss(r, i)
186
187 kld_loss = torch.mean(-0.5 * torch.sum(1 + log_var - mu ** 2 - log_var.exp(), dim = 1), dim = 0)
188
189 loss = self.beta * (recons_loss + feature_loss) + self.alpha * kld_weight * kld_loss
190 return {'loss': loss, 'Reconstruction_Loss':recons_loss, 'KLD':-kld_loss}
191
192 def sample(self,
193 num_samples:int,

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected