recon_x: generating images x: origin images mu: latent mean logvar: latent log variance
(recon_x, x, mu, logvar)
| 78 | |
| 79 | |
| 80 | def loss_function(recon_x, x, mu, logvar): |
| 81 | """ |
| 82 | recon_x: generating images |
| 83 | x: origin images |
| 84 | mu: latent mean |
| 85 | logvar: latent log variance |
| 86 | """ |
| 87 | BCE = reconstruction_function(recon_x, x) # mse loss |
| 88 | # loss = 0.5 * sum(1 + log(sigma^2) - mu^2 - sigma^2) |
| 89 | KLD_element = mu.pow(2).add_(logvar.exp()).mul_(-1).add_(1).add_(logvar) |
| 90 | KLD = torch.sum(KLD_element).mul_(-0.5) |
| 91 | # KL divergence |
| 92 | return BCE + KLD |
| 93 | |
| 94 | |
| 95 | optimizer = optim.Adam(model.parameters(), lr=1e-3) |