|
pulsatrix
|
VAE KL-divergence-to-standard-normal loss term. More...
#include "pulsatrix/device_backend.hpp"#include "pulsatrix/reparameterize.hpp"#include "pulsatrix/tensor.hpp"
Go to the source code of this file.
Classes | |
| class | pulsatrix::KLDivergenceLoss |
| Closed-form KL(N(mu, sigma^2) || N(0, I)) for a diagonal Gaussian posterior (Kingma & Welling 2013, arXiv:1312.6114, Appendix B): loss = mean_b( 0.5 * sum_d( mu[b,d]^2 + exp(2*log_sigma[b,d]) - 2*log_sigma[b,d] - 1 ) ) i.e. summed over the latent dimension per example, averaged over the batch – the standard VAE ELBO normalization, matching how the reconstruction term is typically summed-per-example/averaged-over-batch too. More... | |
Namespaces | |
| namespace | pulsatrix |
VAE KL-divergence-to-standard-normal loss term.