(self,
*args,
**kwargs)
| 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 | """ |
nothing calls this directly
no test coverage detected