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)
| 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, |
nothing calls this directly
no outgoing calls
no test coverage detected