pulsatrix
Loading...
Searching...
No Matches
pulsatrix::KLDivergenceLoss Class Reference

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...

#include <kl_divergence_loss.hpp>

Public Member Functions

 KLDivergenceLoss (DeviceBackend *backend)
 Constructs a KL divergence loss.
 
float forward (const Tensor &mu, const Tensor &log_sigma)
 Computes the KL term and caches mu/log_sigma for backward().
 
ReparamGrad backward () const
 Gradients of the KL term w.r.t. its two inputs: grad_mu[b,d] = mu[b,d] / N grad_log_sigma[b,d] = (exp(2*log_sigma[b,d]) - 1) / N (N is the batch size; the 0.5 factor cancels against the derivative of the squared/exponential terms).
 

Detailed Description

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.

Note
Not a Module subclass, for exactly MSELoss's reason: losses are the seed point relevance/gradient propagation starts from, not something a propagate_relevance rule is defined for. Same shape as MSELoss/CrossEntropyLoss – forward() computes and caches, backward() consumes the cache.
backward() takes no incoming gradient: like MSELoss, this loss is a graph root.
Pairs with MSELoss (Gaussian decoder reconstruction term) to form the full VAE objective; no separate reconstruction loss is built for VAE, MSELoss is reused directly.

Constructor & Destructor Documentation

◆ KLDivergenceLoss()

pulsatrix::KLDivergenceLoss::KLDivergenceLoss ( DeviceBackend *  backend)
explicit

Constructs a KL divergence loss.

Parameters
backendBackend to allocate/compute through. Not owned; must outlive this loss.

Member Function Documentation

◆ backward()

ReparamGrad pulsatrix::KLDivergenceLoss::backward ( ) const

Gradients of the KL term w.r.t. its two inputs: grad_mu[b,d] = mu[b,d] / N grad_log_sigma[b,d] = (exp(2*log_sigma[b,d]) - 1) / N (N is the batch size; the 0.5 factor cancels against the derivative of the squared/exponential terms).

Returns
Both gradients, each the shape of the mu passed to forward().
Exceptions
std::logic_errorif forward() has never been called – uses the cached mu/log_sigma.
Note
Device-generic: runs on Cpu, Cuda or Hip tensors (GPU-native-kernels Mission 1b); inputs must share one device.

◆ forward()

float pulsatrix::KLDivergenceLoss::forward ( const Tensor &  mu,
const Tensor &  log_sigma 
)

Computes the KL term and caches mu/log_sigma for backward().

Parameters
muPosterior mean, shape (N, latent_dim).
log_sigmaPosterior log-standard-deviation. Must match mu's shape.
Returns
The scalar KL value.
Exceptions
std::invalid_argumentif mu/log_sigma shapes don't match, or either isn't rank-2 (N, latent_dim) – external boundary, same classification as MSELoss::forward's shape check.
Note
Device-generic: runs on Cpu, Cuda or Hip tensors (GPU-native-kernels Mission 1b); inputs must share one device.

The documentation for this class was generated from the following file: