Skip to content
AITroveRead. Build. Understand.
Make this comfortable

Variational latent sampling, KL accounting and collapse

Last updated: 7 Oct 20265 min read
tutorial
AdvancedBy AITrove Editorial

A variational autoencoder predicts a distribution for each latent coordinate and balances reconstruction against a KL penalty; neither term should be hidden in an opaque total.

Parameterize a distribution

The encoder produces a mean and log variance for every latent coordinate. Draw standard noise and form a sample as mean plus exp(0.5 × log variance) times noise. This separates the random draw from trainable parameters, allowing gradients through the sample path. Store log variance rather than a raw standard deviation to avoid requiring a positive unconstrained network output. Numerical range still matters: unbounded log variance can overflow an exponential. Loss contracts should record reductions and dtype.

Keep both losses on named scales

A reconstruction term measures the decoded output against the observed input under a chosen likelihood or proxy loss. For a diagonal Gaussian posterior against a unit Gaussian prior, the per-example KL is minus one half times the sum over latent coordinates of one plus log variance minus squared mean minus variance. Average both per-example sums over the batch before combining them. Changing reconstruction from a sum over features to a mean per feature changes its relative weight against KL even if a coefficient stays unchanged.

Watch for latent collapse

If the decoder can reconstruct without using the sampled code, the posterior can move close to the prior and KL can approach zero. That can look like a tidy optimization trace while the latent carries little input information. Track reconstruction, KL, latent means and an intervention that swaps latent codes between inputs. A scheduled KL weight or reduced decoder capacity may help in some tasks, but choose it using held-out behavior rather than treating a nonzero KL as proof of useful representation.

Separate generation from anomaly scoring

Sampling from the prior can generate new outputs, while scoring an observed telemetry window requires a stable procedure. A single random latent sample injects score noise; use the posterior mean or average a fixed number of samples and report the compute and variance. A VAE reconstruction score is not calibrated fault probability. Compare it to a deterministic autoencoder and a simple per-feature baseline on the same held-out incidents. Threshold calibration is still required.

Guard numerical and deployment state

Check finite mean, log variance, KL, reconstruction and gradients before stepping. Save model weights, preprocessing, latent policy, random-seed policy when sampling, reduction convention and threshold. At serving, switching from sampled latents to means changes score distribution, so a threshold calibrated under one policy cannot be reused silently under the other. Test a fixed batch after reload and track both score components through model revisions.

Implementation

python
import torch
from torch import nn
from torch.nn import functional as functional

class SensorVariationalEncoder(nn.Module):
    def __init__(self):
        super().__init__()
        self.shared = nn.Sequential(nn.Linear(6, 9), nn.ReLU())
        self.mean_head = nn.Linear(9, 3)
        self.log_variance_head = nn.Linear(9, 3)
        self.decoder = nn.Linear(3, 6)

    def forward(self, normalized_readings: torch.Tensor):
        hidden = self.shared(normalized_readings)
        latent_mean = self.mean_head(hidden)
        latent_log_variance = self.log_variance_head(hidden).clamp(-12, 8)
        latent_std = torch.exp(0.5 * latent_log_variance)
        latent_sample = latent_mean + latent_std * torch.randn_like(latent_std)
        return self.decoder(latent_sample), latent_mean, latent_log_variance

torch.manual_seed(67)
model = SensorVariationalEncoder()
normal_readings = torch.rand(4, 6)
reconstructed, latent_mean, log_variance = model(normal_readings)
reconstruction_per_case = functional.mse_loss(
    reconstructed, normal_readings, reduction="none").sum(dim=1)
kl_per_case = -0.5 * (1 + log_variance - latent_mean.square()
                      - log_variance.exp()).sum(dim=1)
reconstruction_loss = reconstruction_per_case.mean()
kl_loss = kl_per_case.mean()
total_loss = reconstruction_loss + 0.37 * kl_loss
total_loss.backward()
assert torch.isfinite(total_loss)
assert kl_per_case.shape == (4,)

Performance and operating cost

Encoder and decoder matrix products dominate work; sampling and diagonal KL add O(batch × latent width) arithmetic and storage. Averaging several latent samples multiplies decoder work by the sample count. Clamping log variance prevents some overflow but changes gradients at the clamp bounds, so log clamp frequency should be monitored. A small KL number can signal collapse rather than computational efficiency.

Common Mistakes

  • Do not compare KL to a reconstruction term whose reduction convention changed unnoticed.
  • Do not call a near-zero KL proof that the latent representation is good.
  • Do not calibrate an anomaly threshold with sampled latents and serve with posterior means.

Read next

Continue the workflow: GAN mode coverage and near-duplicate audits.

Continue the workflow: Invertible flows and change-of-variables accounting.

ai-data
deep-learning
Storage details