(self,
*args,
**kwargs)
| 127 | return [self.decode(z), input, mu, log_var] |
| 128 | |
| 129 | def loss_function(self, |
| 130 | *args, |
| 131 | **kwargs) -> dict: |
| 132 | self.num_iter += 1 |
| 133 | recons = args[0] |
| 134 | input = args[1] |
| 135 | mu = args[2] |
| 136 | log_var = args[3] |
| 137 | kld_weight = kwargs['M_N'] # Account for the minibatch samples from the dataset |
| 138 | |
| 139 | recons_loss =F.mse_loss(recons, input) |
| 140 | |
| 141 | kld_loss = torch.mean(-0.5 * torch.sum(1 + log_var - mu ** 2 - log_var.exp(), dim = 1), dim = 0) |
| 142 | |
| 143 | if self.loss_type == 'H': # https://openreview.net/forum?id=Sy2fzU9gl |
| 144 | loss = recons_loss + self.beta * kld_weight * kld_loss |
| 145 | elif self.loss_type == 'B': # https://arxiv.org/pdf/1804.03599.pdf |
| 146 | self.C_max = self.C_max.to(input.device) |
| 147 | C = torch.clamp(self.C_max/self.C_stop_iter * self.num_iter, 0, self.C_max.data[0]) |
| 148 | loss = recons_loss + self.gamma * kld_weight* (kld_loss - C).abs() |
| 149 | else: |
| 150 | raise ValueError('Undefined loss type.') |
| 151 | |
| 152 | return {'loss': loss, 'Reconstruction_Loss':recons_loss, 'KLD':kld_loss} |
| 153 | |
| 154 | def sample(self, |
| 155 | num_samples:int, |
nothing calls this directly
no outgoing calls
no test coverage detected