pulsatrix
Loading...
Searching...
No Matches
kl_divergence_loss.hpp File Reference

VAE KL-divergence-to-standard-normal loss term. More...

Include dependency graph for kl_divergence_loss.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
 

Detailed Description

VAE KL-divergence-to-standard-normal loss term.