pulsatrix
Loading...
Searching...
No Matches
kl_divergence_loss.hpp
Go to the documentation of this file.
1
5#pragma once
6
8#include "pulsatrix/reparameterize.hpp" // ReparamGrad -- the shared (grad_mu, grad_log_sigma) pair
10
11namespace pulsatrix {
12
31public:
36 explicit KLDivergenceLoss(DeviceBackend* backend);
37
49 [[nodiscard]] float forward(const Tensor& mu, const Tensor& log_sigma);
50
63 [[nodiscard]] ReparamGrad backward() const;
64
65private:
66 DeviceBackend* backend_;
67 Tensor last_mu_;
68 Tensor last_log_sigma_;
69 bool has_forwarded_ = false;
70};
71
72} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
Closed-form KL(N(mu, sigma^2) || N(0, I)) for a diagonal Gaussian posterior (Kingma & Welling 2013,...
Definition kl_divergence_loss.hpp:30
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,...
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Abstract interface isolating vendor-specific memory/compute operations from Tensor/ComputationGraph.
Definition acquisition_functions.hpp:16
VAE reparameterization trick: z = mu + exp(log_sigma) * epsilon.
The (grad_mu, grad_log_sigma) pair both VAE building blocks produce.
Definition reparameterize.hpp:22
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).